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

如何在保存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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 05:55:14