You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中无需记忆模型初始化参数即可正确加载模型

PyTorch 自定义模型免参数记录的保存加载方案

你之前使用的仅保存state_dict的方案是官方推荐的安全做法,但确实需要匹配模型初始化参数,以下两种方案可以解决参数记录的问题,同时规避直接序列化整个模型对象的风险:

方案1:打包存储初始化参数与权重(首选方案)

保存时将模型的初始化参数和权重打包到同一个 checkpoint 中,不需要额外维护参数记录,兼容性最好。

保存代码

from torch import nn
import torch

# 训练阶段保存逻辑
checkpoint = {
    "state_dict": model.state_dict(),
    "init_args": {
        "dense1": model.MLP[0].in_features,
        "dense2": model.MLP[0].out_features,
        "dense3": model.MLP[2].out_features
        # 所有自定义的初始化参数都可以追加到该字典中
    }
}
torch.save(checkpoint, checkpoint_model_path)

加载代码

# 推理阶段加载逻辑
checkpoint = torch.load(model_file)
model = myNN(**checkpoint["init_args"])
model.load_state_dict(checkpoint["state_dict"])

方案2:从已有state_dict自动推断层尺寸

如果已经生成了大量仅保存了state_dict的旧 checkpoint,无法重新导出,可以直接从权重张量的形状推导层的输入输出尺寸:

def load_mynn_auto(model_path):
    state_dict = torch.load(model_path)
    # PyTorch Linear层weight的形状固定为 [输出维度, 输入维度]
    dense1 = state_dict["MLP.0.weight"].shape[1]
    dense2 = state_dict["MLP.0.weight"].shape[0]
    dense3 = state_dict["MLP.2.weight"].shape[0]
    model = myNN(dense1=dense1, dense2=dense2, dense3=dense3)
    model.load_state_dict(state_dict)
    return model

# 直接调用即可,不需要提前知道初始化参数
model = load_mynn_auto(model_file)

注意该方案依赖模型层的命名和顺序固定,如果后续修改了模型结构,需要同步调整参数提取逻辑。

方案优势说明

两种方案都保留了仅使用state_dict的安全特性,不会出现直接保存整个模型对象的常见问题:

  • 不存在和myNN类定义强绑定的问题,移动类位置、升级PyTorch版本不会导致加载失败
  • 没有pickle反序列化的恶意代码执行风险

内容的提问来源于stack exchange,提问作者Michael

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.26 13:24:05