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

如何将PyTorch nn.Module转换为HuggingFace PreTrainedModel对象

如何将简单PyTorch模型转换为HuggingFace PreTrainedModel对象?

给定如下简单的PyTorch神经网络:

import torch
import torch.nn as nn

device = "cuda" if torch.cuda.is_available() else "cpu"
net = nn.Sequential(
      nn.Linear(3, 4),
      nn.Sigmoid(),
      nn.Linear(4, 1),
      nn.Sigmoid()
).to(device)

目标是将这个nn.Sequential创建的PyTorch nn.Module转换为HuggingFace PreTrainedModel,实现如下保存与加载流程:

import torch.nn as nn
from transformers.modeling_utils import PreTrainedModel


net = nn.Sequential(
      nn.Linear(3, 4),
      nn.Sigmoid(),
      nn.Linear(4, 1),
      nn.Sigmoid()
).to(device)

# 将PyTorch nn.Module转换为PreTrainedModel对象
shiny_model = do_some_magic(net, some_args, some_kwargs)

# 保存PreTrainedModel对象
shiny_model.save_pretrained("shiny-model")

# 加载预训练模型
PreTrainedModel.from_pretrained("shiny-model")

核心解决方案

要完成转换,必须构建自定义配置类和自定义PreTrainedModel子类——这是HuggingFace模型体系的核心:配置类记录模型结构参数,模型类关联配置并实现标准接口。

1. 定义自定义配置类

配置类需继承PretrainedConfig,用来记录模型的关键结构参数,确保保存后能通过配置重建模型:

from transformers import PretrainedConfig

class SimpleMLPConfig(PretrainedConfig):
    model_type = "simple_mlp"  # 必须指定模型类型,用于加载时识别
    
    def __init__(self, input_dim=3, hidden_dim=4, output_dim=1, **kwargs):
        super().__init__(**kwargs)
        self.input_dim = input_dim
        self.hidden_dim = hidden_dim
        self.output_dim = output_dim

2. 定义自定义PreTrainedModel子类

继承PreTrainedModel,关联配置类,并实现与原模型一致的结构和forward逻辑:

from transformers import PreTrainedModel

class SimpleMLP(PreTrainedModel):
    config_class = SimpleMLPConfig  # 绑定对应的配置类
    
    def __init__(self, config):
        super().__init__(config)
        # 根据配置构建与原模型完全一致的网络结构
        self.model = nn.Sequential(
            nn.Linear(config.input_dim, config.hidden_dim),
            nn.Sigmoid(),
            nn.Linear(config.hidden_dim, config.output_dim),
            nn.Sigmoid()
        )
    
    def forward(self, x):
        return self.model(x)

3. 完成模型转换

初始化自定义模型,并将原PyTorch模型的参数加载进去:

# 1. 创建匹配原模型的配置实例
config = SimpleMLPConfig(input_dim=3, hidden_dim=4, output_dim=1)
# 2. 初始化自定义PreTrainedModel
shiny_model = SimpleMLP(config)
# 3. 加载原模型的参数(strict=False兼容PreTrainedModel额外的配置参数)
shiny_model.load_state_dict(net.state_dict(), strict=False)
# 移动到目标设备
shiny_model.to(device)

4. 测试保存与加载

# 保存模型(自动生成配置文件和权重文件)
shiny_model.save_pretrained("shiny-model")

# 加载模型
loaded_model = SimpleMLP.from_pretrained("shiny-model")
loaded_model.to(device)

# 验证输出一致性
test_input = torch.randn(2, 3).to(device)
original_output = net(test_input)
loaded_output = loaded_model(test_input)
print(torch.allclose(original_output, loaded_output))  # 输出True表示一致

关键说明

  • 原生nn.Module无法直接转换为PreTrainedModel,因为后者依赖配置文件记录结构信息,必须通过自定义配置类补充
  • 配置类的参数需与模型结构一一对应,确保加载时能准确重建模型
  • 此方法完全从零开始,不依赖任何预训练模型模板,适用于任意PyTorch模型

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 08:05:21