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

如何基于PyTorch保存加载含CRF层的自定义Hugging Face模型及config.json

自定义BERT+CRF模型的保存与加载问题

问题描述

我在TokenClassification模型顶部添加了简单的自定义pytorch-crf层,让模型更鲁棒。成功训练后,保存模型时文件夹内没有config.json文件,该如何为自定义模型保存config.json?另外,加载训练后的模型时,最后的CRF层消失了,这是怎么回事?

训练代码

from torchcrf import CRF

model_checkpoint = "dslim/bert-base-NER"
tokenizer = BertTokenizer.from_pretrained(model_checkpoint,add_prefix_space=True)
bert_model = BertForTokenClassification.from_pretrained(
                        model_checkpoint,id2label=id2label,label2id=label2id)
bert_model.config.output_hidden_states=True


class BERT_CRF(nn.Module):
    
    def __init__(self, bert_model, num_labels):
        super(BERT_CRF, self).__init__()
        self.bert = bert_model
        self.dropout = nn.Dropout(0.25)
        
        self.classifier = nn.Linear(768, num_labels)

        self.crf = CRF(num_labels, batch_first = True)
    
    def forward(self, input_ids, attention_mask,  labels=None, token_type_ids=None):
        outputs = self.bert(input_ids, attention_mask=attention_mask)
        
        sequence_output = torch.stack((outputs[1][-1], outputs[1][-2], outputs[1][-3], outputs[1][-4])).mean(dim=0)
        sequence_output = self.dropout(sequence_output)
        
        emission = self.classifier(sequence_output) # [32,256,17]
        labels=labels.reshape(attention_mask.size()[0],attention_mask.size()[1])
        
        if labels is not None:    
            loss = -self.crf(log_soft(emission, 2), labels, mask=attention_mask.type(torch.uint8), reduction='mean')
            prediction = self.crf.decode(emission, mask=attention_mask.type(torch.uint8))
            return [loss, prediction]
                
        else:         
            prediction = self.crf.decode(emission, mask=attention_mask.type(torch.uint8))
            return prediction


model = BERT_CRF(bert_model, num_labels=len(label2id))
model.to(device)

args = TrainingArguments(
    "spanbert_crf_ner2",
    # evaluation_strategy="epoch",
    save_strategy="epoch",
    learning_rate=2e-5,
    num_train_epochs=1,
    weight_decay=0.01,
    per_device_train_batch_size=8,
    # per_device_eval_batch_size=32
    fp16=True
    # bf16=True #Ampere GPU
)

trainer = Trainer(
    model=model,
    args=args,
    train_dataset=train_data,
    # eval_dataset=train_data,
    # data_collator=data_collator,
    # compute_metrics=compute_metrics,
    tokenizer=tokenizer)

trainer.train()
trainer.save_model("model_spanbert_ner")

保存的模型信息

Saving model checkpoint to spanbert_crf_ner2/checkpoint-62500
Trainer.model is not a `PreTrainedModel`, only saving its state dict.
tokenizer config file saved in spanbert_crf_ner2/checkpoint-62500/tokenizer_config.json
Special tokens file saved in spanbert_crf_ner2/checkpoint-62500/special_tokens_map.json


Training completed. Do not forget to share your model on huggingface.co/models =)


100%|██████████| 62500/62500 [15:30:27<00:00,  1.12it/s]
Saving model checkpoint to model_spanbert_ner
Trainer.model is not a `PreTrainedModel`, only saving its state dict.
{'train_runtime': 55837.6817, 'train_samples_per_second': 17.909, 'train_steps_per_second': 1.119, 'train_loss': 1.8942613859863282, 'epoch': 2.0}
tokenizer config file saved in model_spanbert_ner/tokenizer_config.json
Special tokens file saved in model_spanbert_ner/special_tokens_map.json

训练后模型的最后几层

(11): BertLayer(
            (attention): BertAttention(
              (self): BertSelfAttention(
                (query): Linear(in_features=768, out_features=768, bias=True)
                (key): Linear(in_features=768, out_features=768, bias=True)
                (value): Linear(in_features=768, out_features=768, bias=True)
                (dropout): Dropout(p=0.1, inplace=False)
              )
              (output): BertSelfOutput(
                (dense): Linear(in_features=768, out_features=768, bias=True)
                (LayerNorm): LayerNorm((768,), eps=1e-12, elementwise_affine=True)
                (dropout): Dropout(p=0.1, inplace=False)
              )
            )
            (intermediate): BertIntermediate(
              (dense): Linear(in_features=768, out_features=3072, bias=True)
              (intermediate_act_fn): GELUActivation()
            )
            (output): BertOutput(
              (dense): Linear(in_features=3072, out_features=768, bias=True)
              (LayerNorm): LayerNorm((768,), eps=1e-12, elementwise_affine=True)
              (dropout): Dropout(p=0.1, inplace=False)
            )
          )
        )
      )
    )
    (dropout): Dropout(p=0.1, inplace=False)
    (classifier): Linear(in_features=768, out_features=21, bias=True)
  )
  (dropout): Dropout(p=0.25, inplace=False)
  (classifier): Linear(in_features=768, out_features=21, bias=True)
  (crf): CRF(num_tags=21)
)

加载模型后的结构

model = AutoModelForTokenClassification.from_pretrained("model_spanbert_ner",ignore_mismatched_sizes=True)
tokenizer = AutoTokenizer.from_pretrained("model_spanbert_ner",model_max_length=512)



(11): BertLayer(
          (attention): BertAttention(
            (self): BertSelfAttention(
              (query): Linear(in_features=768, out_features=768, bias=True)
              (key): Linear(in_features=768, out_features=768, bias=True)
              (value): Linear(in_features=768, out_features=768, bias=True)
              (dropout): Dropout(p=0.1, inplace=False)
            )
            (output): BertSelfOutput(
              (dense): Linear(in_features=768, out_features=768, bias=True)
              (LayerNorm): LayerNorm((768,), eps=1e-12, elementwise_affine=True)
              (dropout): Dropout(p=0.1, inplace=False)
            )
          )
          (intermediate): BertIntermediate(
            (dense): Linear(in_features=768, out_features=3072, bias=True)
            (intermediate_act_fn): GELUActivation()
          )
          (output): BertOutput(
            (dense): Linear(in_features=3072, out_features=768, bias=True)
            (LayerNorm): LayerNorm((768,), eps=1e-12, elementwise_affine=True)
            (dropout): Dropout(p=0.1, inplace=False)
          )
        )
      )
    )
  )
  (dropout): Dropout(p=0.1, inplace=False)
  (classifier): Linear(in_features=768, out_features=21, bias=True)

问题解答

1. 为什么没有生成config.json?怎么解决?

原因

你的BERT_CRF类继承的是PyTorch原生的nn.Module,而非Hugging Face的PreTrainedModel。Trainer组件只会对PreTrainedModel自动处理配置文件的保存,普通nn.Module仅会保存模型状态字典,不会生成config.json。

解决方法

有两种可行方案:

方案一:让自定义模型继承PreTrainedModel

这种方式更贴合Hugging Face的生态,后续保存、加载都会更顺畅:

from transformers import BertConfig, PreTrainedModel, BertModel
import torch.nn as nn
from torchcrf import CRF
import torch.nn.functional as F

# 1. 定义自定义Config类,扩展BertConfig
class BertCRFConfig(BertConfig):
    def __init__(self, use_crf=True, dropout_rate=0.25, **kwargs):
        super().__init__(**kwargs)
        self.use_crf = use_crf
        self.dropout_rate = dropout_rate

# 2. 自定义模型继承PreTrainedModel
class BERT_CRF(PreTrainedModel):
    config_class = BertCRFConfig  # 指定对应的Config类

    def __init__(self, config):
        super().__init__(config)
        self.bert = BertModel(config)
        self.dropout = nn.Dropout(config.dropout_rate)
        self.classifier = nn.Linear(config.hidden_size, config.num_labels)
        self.crf = CRF(config.num_labels, batch_first=True)

    def forward(self, input_ids, attention_mask, labels=None, token_type_ids=None):
        outputs = self.bert(input_ids, attention_mask=attention_mask, output_hidden_states=True)
        # 取最后四层隐藏层均值
        sequence_output = torch.stack(outputs.hidden_states[-4:]).mean(dim=0)
        sequence_output = self.dropout(sequence_output)
        emission = self.classifier(sequence_output)
        
        if labels is not None:
            loss = -self.crf(F.log_softmax(emission, 2), labels, mask=attention_mask.type(torch.uint8), reduction='mean')
            prediction = self.crf.decode(emission, mask=attention_mask.type(torch.uint8))
            return {"loss": loss, "predictions": prediction}
        else:
            prediction = self.crf.decode(emission, mask=attention_mask.type(torch.uint8))
            return {"predictions": prediction}

# 初始化模型的方式改为基于Config
config = BertCRFConfig.from_pretrained(model_checkpoint, num_labels=len(label2id), id2label=id2label, label2id=label2id)
model = BERT_CRF(config)

这样训练后调用trainer.save_model(),就会自动生成config.json。

方案二:手动保存原BERT模型的Config

如果不想修改模型继承关系,可以手动将原BERT模型的配置保存到目标文件夹:

# 在调用trainer.save_model()之后执行
bert_model.config.save_pretrained("model_spanbert_ner")

2. 加载模型时CRF层消失的原因及解决方法

原因

你使用AutoModelForTokenClassification.from_pretrained()加载模型,这个方法只会初始化标准的TokenClassification模型结构(即BERT+线性分类器),它完全不知道你自定义的CRF层存在,所以加载后的模型自然没有CRF层。

解决方法

必须基于你自定义的BERT_CRF类来加载状态字典,步骤如下:

import torch
from torchcrf import CRF
from transformers import BertForTokenClassification, BertTokenizer
import torch.nn as nn
import torch.nn.functional as F

# 1. 必须重新定义和训练时完全一致的BERT_CRF类
class BERT_CRF(nn.Module):
    def __init__(self, bert_model, num_labels):
        super(BERT_CRF, self).__init__()
        self.bert = bert_model
        self.dropout = nn.Dropout(0.25)
        self.classifier = nn.Linear(768, num_labels)
        self.crf = CRF(num_labels, batch_first=True)
    
    def forward(self, input_ids, attention_mask, labels=None, token_type_ids=None):
        outputs = self.bert(input_ids, attention_mask=attention_mask)
        sequence_output = torch.stack((outputs[1][-1], outputs[1][-2], outputs[1][-3], outputs[1][-4])).mean(dim=0)
        sequence_output = self.dropout(sequence_output)
        emission = self.classifier(sequence_output)
        labels = labels.reshape(attention_mask.size()[0], attention_mask.size()[1]) if labels is not None else None
        
        if labels is not None:    
            loss = -self.crf(F.log_softmax(emission, 2), labels, mask=attention_mask.type(torch.uint8), reduction='mean')
            prediction = self.crf.decode(emission, mask=attention_mask.type(torch.uint8))
            return [loss, prediction]
        else:         
            prediction = self.crf.decode(emission, mask=attention_mask.type(torch.uint8))
            return prediction

# 2. 加载Tokenizer和原BERT模型
tokenizer = BertTokenizer.from_pretrained("model_spanbert_ner", model_max_length=512)
bert_model = BertForTokenClassification.from_pretrained("dslim/bert-base-NER", id2label=id2label, label2id=label2id)
bert_model.config.output_hidden_states=True

# 3. 初始化自定义模型
num_labels = len(label2id)
model = BERT_CRF(bert_model, num_labels=num_labels)

# 4. 加载保存的状态字典
model.load_state_dict(torch.load("model_spanbert_ner/pytorch_model.bin"))
model.to(device)

这样加载后的模型就会包含CRF层,且参数是训练好的。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 08:01:22