如何在保存PyTorch模型时整合所使用的激活函数?
在PyTorch中保存模型时同时保留激活函数的几种方法
方法一:将激活函数定义为模型的子模块
把激活函数作为模型类的成员变量(子模块),PyTorch会自动将其纳入模型结构展示,同时在保存/加载时关联起来。修改代码如下:
import torch import torch.nn as nn import copy class NNM(nn.Module): def __init__(self, num_features, num_hidden): super(NNM, self).__init__() self.fc1 = nn.Linear(num_features, num_hidden) self.fc2 = nn.Linear(num_hidden, 1) # 将激活函数定义为模型子模块 self.act_func = nn.Sigmoid() self.saved_parameters = [] def forward(self, x): x = self.act_func(self.fc1(x)) return self.fc2(x) def save_parameters(self): self.saved_parameters.append(copy.deepcopy(self.state_dict())) model = NNM(28, 100) print(model)
此时打印模型结构会输出:
NNM( (fc1): Linear(in_features=28, out_features=100, bias=True) (fc2): Linear(in_features=100, out_features=1, bias=True) (act_func): Sigmoid() )
如果需要明确激活函数的作用位置,可以把模块名改成self.fc1_post_act = nn.Sigmoid(),这样名称更直观。
方法二:用nn.Sequential组合层与激活函数
如果想清晰展示激活函数和对应线性层的绑定关系,可将两者打包成nn.Sequential模块:
class NNM(nn.Module): def __init__(self, num_features, num_hidden): super(NNM, self).__init__() # 组合fc1和sigmoid为一个子模块 self.fc1_with_act = nn.Sequential( nn.Linear(num_features, num_hidden), nn.Sigmoid() ) self.fc2 = nn.Linear(num_hidden, 1) self.saved_parameters = [] def forward(self, x): x = self.fc1_with_act(x) return self.fc2(x) def save_parameters(self): self.saved_parameters.append(copy.deepcopy(self.state_dict())) model = NNM(28, 100) print(model)
打印结果会是:
NNM( (fc1_with_act): Sequential( (0): Linear(in_features=28, out_features=100, bias=True) (1): Sigmoid() ) (fc2): Linear(in_features=100, out_features=1, bias=True) )
这种方式能直接看到激活函数与对应线性层的从属关系,结构更清晰。
方法三:保存参数时额外记录激活函数信息
如果不需要修改模型结构展示,仅需在保存参数时记录激活函数的相关信息,可扩展save_parameters方法,把激活函数的类型、作用位置和参数字典一起保存:
class NNM(nn.Module): def __init__(self, num_features, num_hidden): super(NNM, self).__init__() self.fc1 = nn.Linear(num_features, num_hidden) self.fc2 = nn.Linear(num_hidden, 1) # 记录激活函数的关键信息 self.act_info = { 'type': 'Sigmoid', 'applied_after': 'fc1', 'applied_before': 'fc2' } self.saved_parameters = [] def forward(self, x): x = torch.sigmoid(self.fc1(x)) return self.fc2(x) def save_parameters(self): # 保存参数字典+激活函数信息 saved_data = { 'state_dict': copy.deepcopy(self.state_dict()), 'act_info': self.act_info } self.saved_parameters.append(saved_data) model = NNM(28, 100) model.save_parameters() # 查看保存的激活函数信息 print(model.saved_parameters[0]['act_info'])
后续加载参数时,就能从保存的字典中获取激活函数信息,确保推理时使用正确的激活逻辑。
内容的提问来源于stack exchange,提问作者pyaj
相关产品推荐
相关产品推荐

