PyTorch中如何正确保存自定义函数及训练后的参数?
问题描述
用户定义了以下自定义网络相关函数:
def softmax(X): X_exp=torch.exp(X) partition=X_exp.sum(1,keepdim=True) return X_exp/partition def net(X): return softmax(torch.matmul(X.reshape(-1,W.shape[0]),W)+b)
随后通过训练函数更新参数:
train(net,train_iter,test_iter,cross_entropy,num_epoches,updater)
最后尝试保存并加载网络用于预测:
PATH='./net.pth' torch.save(net,PATH) saved_net=torch.load(PATH) predict(saved_net,test_iter,6)
但预测时发现训练更新后的参数W和b并未被保存和加载,询问正确的保存自定义网络及更新后参数的方法。
正确的保存与加载方法
方法一:直接保存和加载独立参数(适配函数式实现)
由于net是普通函数,W和b是独立的全局张量,并非函数的内部属性,直接保存函数本身无法带上这些训练后的参数。正确做法是单独保存参数:
保存参数
PATH = './params.pth' torch.save({'W': W, 'b': b}, PATH)
加载参数并复用网络
加载参数后,将其赋值给全局变量即可让原net函数使用:
params = torch.load(PATH) W, b = params['W'], params['b'] # 此时调用net(X)会自动使用加载后的参数 predict(net, test_iter, 6)
如果想避免全局变量的依赖,可以修改net为带参数的函数:
def net(X, W, b): return softmax(torch.matmul(X.reshape(-1,W.shape[0]),W)+b) # 调用时通过lambda包装传入参数 predict(lambda X: net(X, W, b), test_iter, 6)
方法二:改用nn.Module类实现网络(推荐方案)
PyTorch的模型保存机制专为继承自nn.Module的类设计,这类实现能自动管理参数,更规范且不易出错:
重构网络为类
import torch.nn as nn import torch.nn.functional as F class SoftmaxNet(nn.Module): def __init__(self, input_dim, output_dim): super().__init__() self.linear = nn.Linear(input_dim, output_dim) def forward(self, X): return F.softmax(self.linear(X.reshape(-1, self.linear.in_features)), dim=1)
训练类实例
# 初始化网络(替换原W的维度) net = SoftmaxNet(input_dim=W.shape[0], output_dim=W.shape[1]) # 执行训练(注意updater需适配nn.Module参数,比如使用torch.optim.SGD) train(net, train_iter, test_iter, cross_entropy, num_epoches, updater)
保存与加载模型
保存模型状态字典(推荐,灵活性更高)
PATH = './net.pth' torch.save(net.state_dict(), PATH)
加载模型
# 先初始化相同结构的网络 saved_net = SoftmaxNet(input_dim=W.shape[0], output_dim=W.shape[1]) # 加载训练后的参数 saved_net.load_state_dict(torch.load(PATH)) # 设置为评估模式(关闭 dropout 等训练相关行为) saved_net.eval() # 执行预测 predict(saved_net, test_iter, 6)
内容的提问来源于stack exchange,提问作者BuDiu
相关产品推荐
相关产品推荐

