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

如何将自定义PyTorch model.pt转换为支持from_pretrained()加载的Hugging Face模型

PyTorch自定义NLP模型转Hugging Face Transformers格式教程

首先明确两个核心问题的答案:

  1. 转换为Hugging Face(以下简称HF)生态兼容的模型格式完全可以实现,转换后可以直接搭配HF提供的各类任务头部做微调、推理等操作
  2. 转换过程确实需要配套的配置文件,但不需要手动编写,通过继承HF提供的基类,调用内置方法即可自动生成所有必要文件

1. 前置准备

操作前先确认以下信息,避免后续加载权重报错:

  • 你预训练时使用的所有模型超参数:包括词汇表大小、隐层维度、Transformer层数、注意力头数、最大序列长度、中间层维度等
  • 你的model.pt保存的是模型的state_dict(仅权重,更推荐的保存方式)还是完整的模型对象
  • 你预训练时的模型结构代码,要保证和转换时写的结构完全一致

2. 继承HF基类定义自定义模型和配置

这是兼容HF生态的核心步骤,你需要分别定义继承自PretrainedConfig的配置类,和继承自PreTrainedModel的模型类:

from transformers import PreTrainedModel, PretrainedConfig
import torch
import torch.nn as nn

# 自定义配置类:存放模型所有超参数
class CustomNLPConfig(PretrainedConfig):
    # 模型唯一标识,不要和HF现有模型重名
    model_type = "custom_nlp"

    def __init__(
        self,
        vocab_size=30522,
        hidden_size=768,
        num_hidden_layers=12,
        num_attention_heads=12,
        intermediate_size=3072,
        max_position_embeddings=512,
        **kwargs
    ):
        super().__init__(**kwargs)
        # 这里加入所有你预训练时用到的自定义超参数
        self.vocab_size = vocab_size
        self.hidden_size = hidden_size
        self.num_hidden_layers = num_hidden_layers
        self.num_attention_heads = num_attention_heads
        self.intermediate_size = intermediate_size
        self.max_position_embeddings = max_position_embeddings
# 自定义模型类:结构和你预训练时的代码完全一致
class CustomNLPModel(PreTrainedModel):
    # 绑定上面定义的配置类
    config_class = CustomNLPConfig
    base_model_prefix = "custom_nlp"

    def __init__(self, config):
        super().__init__(config)
        # 以下结构完全复制你预训练时的模型定义即可
        self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size)
        self.encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(
                d_model=config.hidden_size,
                nhead=config.num_attention_heads,
                dim_feedforward=config.intermediate_size
            ),
            num_layers=config.num_hidden_layers
        )
        # 预训练时用到的专属头如果不需要通用可以删掉,后续可以直接接HF的通用任务头

    def forward(self, input_ids, attention_mask=None, **kwargs):
        # 前向传播逻辑和预训练时保持一致,返回格式尽量对齐HF标准:第一个值为last_hidden_state
        x = self.embeddings(input_ids)
        # 注意mask格式要适配你自己的模型实现,这里只是示例
        if attention_mask is not None:
            attention_mask = ~attention_mask.bool()
        last_hidden_state = self.encoder(x, src_key_padding_mask=attention_mask)
        return (last_hidden_state,)

3. 加载预训练权重

根据你model.pt的保存类型选择对应的加载方式:

情况1:保存的是state_dict(推荐)

# 初始化配置,参数和预训练时完全一致
config = CustomNLPConfig(
    vocab_size=30522,
    hidden_size=768,
    num_hidden_layers=12,
    # 其余参数和预训练时保持一致
)
# 初始化空模型
model = CustomNLPModel(config)
# 加载权重
state_dict = torch.load("model.pt", map_location="cpu")
# 如果权重key有多余前缀(比如分布式训练时加的module.),手动调整后再加载
# state_dict = {k.replace("module.", ""):v for k,v in state_dict.items()}
model.load_state_dict(state_dict)

情况2:保存的是完整模型对象

old_model = torch.load("model.pt", map_location="cpu")
state_dict = old_model.state_dict()
# 后续和上面的加载逻辑一致:初始化CustomNLPModel后加载state_dict即可

4. 保存为HF标准格式

直接调用HF模型内置的save_pretrained方法,会自动生成所有需要的配套文件:

# 保存到指定目录
model.save_pretrained("./hf_custom_nlp_model")

保存后目录下会生成两个核心文件:

  • 权重文件:pytorch_model.bin 或 model.safetensors(新版本HF默认生成更安全的safetensors格式)
  • 配置文件:config.json,包含所有模型超参数的序列化结果

5. 验证兼容性

你可以直接用HF的from_pretrained方法加载,也可以搭配通用任务头使用:

# 直接加载模型
loaded_model = CustomNLPModel.from_pretrained("./hf_custom_nlp_model")

# 注册模型到Auto类,就可以用HF的Auto API加载,搭配各类任务头
CustomNLPConfig.register_for_auto_class()
CustomNLPModel.register_for_auto_class("AutoModel")

# 示例:搭配文本分类头做微调
from transformers import AutoModelForSequenceClassification
classifier = AutoModelForSequenceClassification.from_pretrained(
    "./hf_custom_nlp_model",
    num_labels=2
)

小提示:如果你的模型有配套的自定义分词器,也可以继承PreTrainedTokenizer类,同样调用save_pretrained保存,就可以和模型配套使用。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 17:18:05