• 【深度学习21天学习挑战赛】备忘篇:模型复用——模型的保存与加载


    活动地址:CSDN21天学习挑战赛

    最近一遍学,一遍尝试进行模型的简单应用,需求驱动也是一个好的学习动力。

    那么问题来了,难道我们每次应用模型,都要从头到尾训练一遍,然后再去做识别任务吗?

    当然不是,所以,记录一下简单的模型保存模型加载过程。

    只是抛砖引玉,和给自己记录一下。更多使用,请参考官方文档

    识别手写数字模型为例。

    1、保存模型

    在识别手写数字模型训练之后,保存代码。

    # 保存模型
    model.save('test.h5') #保存到与代码文件同目录,h5:模型文件后缀名
    
    • 1
    • 2

    在这里插入图片描述

    保存后,即可看到,代码同目录下,我们保存的模型文件。
    在这里插入图片描述

    2、加载模型

    在复用模型的地方:

    my_model = tf.keras.models.load_model('test.h5') 
    
    • 1

    即可加载我们保存过的模型。

    加载后,即可使用模型的一些方法了。

    import tensorflow as tf
    from tensorflow.keras import datasets, layers, models
    import matplotlib.pyplot as plt
    # 加载并预处理数据
    (train_images, train_labels), (test_images, test_labels) = datasets.mnist.load_data()
    train_images, test_images = train_images / 255.0, test_images / 255.0
    # 加载模型
    my_model = tf.keras.models.load_model('test.h5') 
    my_model.summary()
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8
    • 9

    在这里插入图片描述

    在这里插入图片描述

    3、保存模型的更多姿势

    上面这种保存方式,是完整保存模式,保存的模型文件包括:

    • architecture 模型的结构
    • weight values 模型的权值
    • training config 模型的配置:即我们通过compile编译模型的一些信息,如优化器,损失函数等
    • optimizer and its state 优化器的状态信息,我们可以接着之前的训练继续训练

    也可以分别保存:

    3.1 保存模型结构

    # 保存模型结构
    with open('model_config.json', 'w') as json_file:
        json_file.write(json_config)
    
    
    # 加载
    with open('model_config.json') as json_file:
        json_config = json_file.read()
    
    • 1
    • 2
    • 3
    • 4
    • 5
    • 6
    • 7
    • 8

    3.2 保存模型权重

    # 保存权重到磁盘,注意,这里是一个 h5 文件哦!
    model.save_weights('my_weights.h5')
     
    # 新模型从磁盘加载权重
    new_model.load_weights('my_weights.h5')
    
    • 1
    • 2
    • 3
    • 4
    • 5

    3.3 保存为SavedModel格式

    另外,也可以将模型保存为tensorflow标准的SavedModel格式,这是对tensorflow对象标准的序列化格式,是官方推荐使用,不同的是,他不是将模型保存为一个单独的文件,而是有几个文件组成。

    # 模型保存,注意:仅仅是多了一个save_format的参数而已
    # 注意:这里的'path_to_saved_model'不再是模型名称,仅仅是一个文件夹,模型会保存在这个文件夹之下
    model.save('mymodel', save_format='tf')
    
    • 1
    • 2
    • 3
     
    # 加载模型,通过指定存放模型的文件夹来加载
    new_model = keras.models.load_model('mymodel')
    
    • 1
    • 2
    • 3

    需要注意的是:这种方式依然会保存模型的所有信息,即“网络结构、权重、配置、优化器状态”四个信息,所以可以接着训练

    这种方法,可以参考大佬文章

  • 相关阅读:
    Maven在开发中的使用及理解
    真嘟假嘟?!这么清晰简单的字符函数和字符串函数!!!
    10.动态路由绑定怎么做
    跳表的实现
    超硬核!华为智慧屏上的家庭相册竟可以自动精准分类?
    java基于ssm+vue+elementui的多用户博客管理系统
    【CocosCreator】利用遮罩Mask实现单边开门效果
    使用 MySQL 日志 | 慢速日志 - Part 3
    Python基础知识从hello world 开始(第二天)
    PyTorch 2.0 重磅发布:一行代码提速 30%
  • 原文地址:https://blog.csdn.net/m0_48300767/article/details/126187008