如何在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
相关产品推荐
相关产品推荐

