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

能否无需模型类定义保存并继续训练PyTorch模型?

无需原始模型类定义即可保存并继续训练PyTorch模型的方法

方法1:使用TorchScript脚本化模型(支持训练)

你之前认为TorchScript模型只能用于推理是个误解——标准结构的TorchScript模型完全可以继续训练,步骤如下:

保存模型

import torch
import torch.nn as nn
import torch.nn.functional as F

# 示例模型(保存时不需要后续加载环境有这个类定义)
class MyModel(nn.Module):
    def __init__(self, input_dim, hidden_dim, output_dim):
        super().__init__()
        self.fc1 = nn.Linear(input_dim, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, output_dim)
    
    def forward(self, x):
        x = F.relu(self.fc1(x))
        return self.fc2(x)

# 实例化并训练一段(可选)
model = MyModel(10, 20, 5)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
# ... 训练代码 ...

# 脚本化并保存
scripted_model = torch.jit.script(model)
scripted_model.save("trained_model.pt")

加载并继续训练

import torch
import torch.nn.functional as F

# 无需定义MyModel类
loaded_model = torch.jit.load("trained_model.pt")
loaded_model.train()  # 切换到训练模式

# 重新初始化优化器(或保存优化器状态一起加载)
optimizer = torch.optim.Adam(loaded_model.parameters(), lr=1e-3)

# 继续训练示例
criterion = nn.CrossEntropyLoss()
for x, y in your_new_dataloader:
    optimizer.zero_grad()
    output = loaded_model(x)
    loss = criterion(output, y)
    loss.backward()
    optimizer.step()

注意:如果模型包含TorchScript不支持的动态操作(比如复杂的自定义控制流、未脚本化的自定义模块),这种方法可能失效,此时需要改用其他方案。

方法2:保存模型结构配置+参数集合

如果TorchScript不符合你的需求,可以手动保存模型的结构参数(比如每层的类型、输入输出维度等)和权重、优化器状态,加载时动态重建模型:

保存模型

import torch
import torch.nn as nn

def get_module_config(module):
    """提取标准Torch模块的初始化参数"""
    config = {"type": type(module).__name__}
    if isinstance(module, nn.Linear):
        config["params"] = {
            "in_features": module.in_features,
            "out_features": module.out_features,
            "bias": module.bias is not None
        }
    elif isinstance(module, nn.Conv2d):
        config["params"] = {
            "in_channels": module.in_channels,
            "out_channels": module.out_channels,
            "kernel_size": module.kernel_size,
            "stride": module.stride,
            "padding": module.padding,
            "bias": module.bias is not None
        }
    elif isinstance(module, nn.ReLU):
        config["params"] = {"inplace": module.inplace}
    # 可扩展支持其他标准模块类型
    return config

# 示例模型
model = nn.Sequential(
    nn.Linear(10, 20),
    nn.ReLU(inplace=True),
    nn.Linear(20, 5)
)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
# ... 训练代码 ...

# 保存结构配置、权重、优化器状态
save_data = {
    "model_config": [get_module_config(m) for m in model],
    "model_state_dict": model.state_dict(),
    "optimizer_state_dict": optimizer.state_dict()
}
torch.save(save_data, "model_with_config.pt")

加载并继续训练

import torch
import torch.nn as nn

def build_model_from_config(config_list):
    """根据配置列表重建模型"""
    model = nn.Sequential()
    for idx, config in enumerate(config_list):
        module_cls = getattr(nn, config["type"])
        module = module_cls(**config["params"])
        model.add_module(f"layer_{idx}", module)
    return model

# 加载数据并重建模型
save_data = torch.load("model_with_config.pt")
model = build_model_from_config(save_data["model_config"])
model.load_state_dict(save_data["model_state_dict"])
model.train()

# 加载优化器
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
optimizer.load_state_dict(save_data["optimizer_state_dict"])

# 继续训练
criterion = nn.CrossEntropyLoss()
for x, y in your_new_dataloader:
    optimizer.zero_grad()
    output = model(x)
    loss = criterion(output, y)
    loss.backward()
    optimizer.step()

这种方法的局限性是仅支持你预先编写了配置提取逻辑的标准模块,如果模型包含自定义模块,仍需要加载环境中有该模块的类定义。

总结

  • 如果模型结构符合TorchScript的要求,优先选择TorchScript脚本化的方法,无需额外处理结构,且原生支持训练。
  • 若模型存在TorchScript不兼容的操作,可使用结构配置+参数保存的方案,需手动扩展支持的模块类型。

内容的提问来源于stack exchange,提问作者Sir Absolute 0

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 01:11:15