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

PyTorch中如何保存包含激活函数的模型架构与权重?

在PyTorch中同时保存神经网络结构与权重的方法

你遇到的问题本质是state_dict()仅保存模型的可学习参数权重,不会记录模型的结构定义(包括forward方法里的激活函数调用逻辑)。以下是几种可靠的解决方案:

方案1:直接保存整个模型实例

PyTorch支持直接保存完整的模型对象,这样会同时存储模型结构(含forward逻辑)和权重参数。

保存代码:

import torch
from torch import nn

class Test(nn.Module):
    def __init__(self):
        super(Test, self).__init__()
        self.fc1 = nn.Linear(10, 25)
        self.fc2 = nn.Linear(25, 10)
        self.relu = nn.ReLU()
        self.tanh = nn.Tanh()
    def forward(self, inputs):
        return self.tanh(self.fc2(self.relu(self.fc1(inputs))))

test = Test().float()
torch.save(test, "test_full_model.pt")  # 直接保存模型实例

加载代码:

import torch

# 无需重复定义Test类(若类定义在当前环境中可直接加载)
test_loaded = torch.load("test_full_model.pt")
# 验证输出一致性
dummy_input = torch.randn(1, 10)
print(test_loaded(dummy_input))

注意:该方式依赖保存模型时的环境,模型类的定义必须存在或可被正确导入,否则加载会报错;跨PyTorch版本可能存在兼容性问题。

方案2:使用TorchScript导出模型(推荐用于部署/跨环境)

TorchScript会将模型转换为静态图结构,同时包含权重和计算逻辑,兼容性更强,适合部署场景。

导出代码:

test = Test().float()
dummy_input = torch.randn(1, 10)  # 输入示例,用于追踪计算图
traced_model = torch.jit.trace(test, dummy_input)
traced_model.save("test_traced.pt")

加载代码:

loaded_traced_model = torch.jit.load("test_traced.pt")
dummy_input = torch.randn(1, 10)
print(loaded_traced_model(dummy_input))

这种方式不依赖原始模型类的定义,加载后直接得到可运行的模型,激活函数的调用顺序会被严格记录在静态图中,不会出现逻辑篡改问题。

方案3:手动保存模型结构配置+权重

如果需要更灵活的控制,可以手动存储模型的结构参数(如输入输出维度、激活函数顺序等),再结合state_dict()保存权重,加载时根据配置重建模型。

保存代码:

model_config = {
    "input_dim": 10,
    "hidden_dim": 25,
    "output_dim": 10,
    "activation_order": ["relu", "tanh"]  # 记录激活函数调用顺序
}

# 保存配置与权重
torch.save({
    "config": model_config,
    "state_dict": test.state_dict()
}, "test_config_and_weights.pt")

加载代码:

# 根据配置重建模型
class TestConfigurable(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.fc1 = nn.Linear(config["input_dim"], config["hidden_dim"])
        self.fc2 = nn.Linear(config["hidden_dim"], config["output_dim"])
        self.relu = nn.ReLU()
        self.tanh = nn.Tanh()
        self.activation_order = config["activation_order"]
    
    def forward(self, inputs):
        x = self.fc1(inputs)
        x = self.relu(x) if self.activation_order[0] == "relu" else self.tanh(x)
        x = self.fc2(x)
        x = self.relu(x) if self.activation_order[1] == "relu" else self.tanh(x)
        return x

# 加载配置与权重
checkpoint = torch.load("test_config_and_weights.pt")
test_config = TestConfigurable(checkpoint["config"])
test_config.load_state_dict(checkpoint["state_dict"])

该方式适合需要自定义模型结构的场景,但需自行维护配置与模型逻辑的对应关系。


内容的提问来源于stack exchange,提问作者learner

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 02:15:35