PyTorch实现类Keras的模型结构与权重分离保存加载方案问询
实现PyTorch中模型结构与权重的分离保存/加载(类似Keras)
PyTorch默认的torch.save()或state_dict保存方式依赖原模型类定义,会出现Can't get attribute 'net'的错误。以下两种方案可以实现类似Keras的“无需原模型定义即可加载”的效果:
方案1:使用TorchScript(最简便)
TorchScript会将模型的计算图结构和权重一起序列化,加载时无需提前定义模型类,直接使用即可。
保存模型
import torch # 假设已训练好模型model,传入一个匹配输入形状的张量用于追踪结构 traced_model = torch.jit.trace(model, torch.randn(1, 输入特征维度)) # 保存为.pt文件 torch.jit.save(traced_model, "bpnn_model.pt")
注意:如果模型包含动态控制流(如if/else分支、循环),请改用
torch.jit.script(model)代替trace,避免结构追踪不完整。
加载模型
import torch # 直接加载模型,无需定义原模型类 trained_model = torch.jit.load("bpnn_model.pt") # 推理时切换到评估模式 trained_model.eval() # 继续训练时切换到训练模式,并重新设置优化器 trained_model.train() optimizer = torch.optim.Adam(trained_model.parameters(), lr=1e-3)
方案2:手动序列化配置+单独保存权重(自定义性强)
通过手动保存模型的结构参数(如层数、维度、激活函数类型),加载时根据配置动态重建模型,再加载权重。
保存配置与权重
import torch import json # 提取模型的关键配置(根据你的模型结构自定义) model_config = { "input_size": 10, "hidden_sizes": [64, 32], "output_size": 2, "activation": "relu" } # 保存配置到JSON文件 with open("bpnn_model_config.json", "w") as f: json.dump(model_config, f) # 单独保存权重 torch.save(model.state_dict(), "bpnn_model_weights.pth")
加载配置与权重
import torch import json # 加载配置文件 with open("bpnn_model_config.json", "r") as f: model_config = json.load(f) # 根据配置动态构建模型 def build_model(config): layers = [] prev_dim = config["input_size"] for dim in config["hidden_sizes"]: layers.append(torch.nn.Linear(prev_dim, dim)) if config["activation"] == "relu": layers.append(torch.nn.ReLU()) # 可扩展支持其他激活函数(如tanh、sigmoid) prev_dim = dim layers.append(torch.nn.Linear(prev_dim, config["output_size"])) return torch.nn.Sequential(*layers) # 重建模型并加载权重 trained_model = build_model(model_config) trained_model.load_state_dict(torch.load("bpnn_model_weights.pth")) # 后续操作:推理或继续训练 trained_model.train() optimizer = torch.optim.SGD(trained_model.parameters(), lr=0.01)
内容的提问来源于stack exchange,提问作者Hopan
相关产品推荐
相关产品推荐

