昇思25天学习打卡营第8天|模型权重保存与加载

打卡

目录

打卡

模型的两种保存形式

Checkpoint

中间表示IR

模型保存与加载

模型权重保存-例1

模型权重加载-例1

模型权重保存-例2

模型权重加载-例2

模型权重文件的空间占用计算-例


模型的两种保存形式

Checkpoint

权重参数文件

中间表示IR

中间表示(Intermediate Representation,IR)是程序编译过程中介于源语言和目标语言之间的程序表示。MindIR是一种基于图表示的函数式IR,其最核心的目的是服务于自动微分变换

在图模式set_context(mode=GRAPH_MODE)下运行用MindSpore编写的模型时,若配置中设置了set_context(save_graphs=1),运行时会输出一些图编译过程中生成的一些中间文件,我们称为IR文件。

  • ir后缀结尾的IR文件:一种比较直观易懂的以文本格式描述模型结构的文件,可以直接用文本编辑软件查看。

  • dot后缀结尾的IR文件:描述了不同节点间的拓扑关系,可以用graphviz将此文件作为输入生成图片,方便用户直观地查看模型结构。对于算子比较多的模型,推荐使用可视化组件MindSpore Insight对计算图进行可视化。

模型保存与加载

保存流程:

  • 定义模型网络
  • 选择损失函数、优化器等
  • 训练模型、更新模型权重参数
  • 选择1:保存模型权重参数Checkpoint到本地
  • 选择2:保存中间表示IR到本地

加载流程:

  • 定义模型网络
  • 选择1:从本地加载模型权重参数Checkpoint
  • 选择2:保存中间表示IR到本地

模型权重保存-例1

python 复制代码
model = network()
mindspore.save_checkpoint( 
       model,         ## 待保存的对象。数据类型可为 mindspore.nn.Cell 、list或dict。
       "model.ckpt"   ## 模型权重保存路径
     )

模型权重加载-例1

python 复制代码
model = network()
param_dict = mindspore.load_checkpoint("model.ckpt")
param_not_load, _ = mindspore.load_param_into_net(
                                    model, 
                                    param_dict
                                  )
print(param_not_load)  ## param_not_load是未被加载的参数列表,为空时代表所有参数均加载成功。

模型权重保存-例2

MindIR同时保存了Checkpoint和模型结构,因此需要定义输入Tensor来获取输入shape。

python 复制代码
model = network()
inputs = Tensor(np.ones([1, 1, 28, 28]).astype(np.float32))
mindspore.export(model, 
                inputs, 
                file_name="model", 
                file_format="MINDIR"
                )

模型权重加载-例2

python 复制代码
mindspore.set_context(mode=mindspore.GRAPH_MODE)

graph = mindspore.load("model.mindir")
model = nn.GraphCell(graph)
outputs = model(inputs)
print(outputs.shape)

模型权重文件的空间占用计算-例

  • 计算方式:计算模型参数个数;按照每个参数占用的字节数计算所有参数的字节占用;转换字节占用单位为MB或GB等。
  • 对比:查看实际保存的大小,与计算预期占用字节数做对比。

例子如下:可以看到,计算与预期基本一致。MindIR同时保存了Checkpoint和模型结构,参数文件会更大一些。

相关推荐
大模型momo18 分钟前
Spring AI 实战:多 Agent 协作实战 —— 分工拆解复杂旅游行程任务
人工智能·spring·ai·agent·旅游
小程故事多_8035 分钟前
从A2C、TRPO、PPO到GRPO,强化学习策略梯度算法完整演进与大模型落地实战解析
人工智能·算法
冬奇Lab40 分钟前
开源项目第176期:Better Harness — 不审查 diff,审查工作流本身,给 AI 编程 Agent 的五维评估框架
人工智能·开源·agent
冬奇Lab1 小时前
代码库知识库系列(07):混合检索 BM25 + 向量——Q8 还是失败,而且总分退步了
人工智能
ajassi20001 小时前
AI语音智能体开发日记(十一)为智能设备“声”临其境——详解音频资源自动化生成流程
人工智能·ai·ai编程
2601_949499941 小时前
400G组网低功耗优选!芯瑞科技400G VR4 QSFP112光模块赋能智算中心高速互联
大数据·人工智能·科技
GoAI2 小时前
# AI Agent 记忆框架横向对比报告总结
人工智能·大模型·llm·多模态
AI人工智能+2 小时前
一种基于深度学习技术的高精度医疗机构执业许可证识别系统,构建了一套基于深度神经网络的端到端智能识别系统,为医疗行业提
深度学习·ocr·医疗机构执业许可证识别
硅谷秋水2 小时前
EgoSteer:一种基于第一人称视角视频、实现可控灵巧操作的全栈系统
深度学习·机器学习·语言模型·机器人·音视频
李昊哲小课2 小时前
fastapi sse websocket 奶茶店实时订单看板
人工智能·python·websocket·网络协议·fastapi·sse