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

如何在 PyTorch TokenClassification 模型上添加 pytorch-crf 层并解决隐藏状态问题

问题解答

一、获取BERT隐藏状态并拼接最后四层

dslim/bert-base-NER是基于BERTForTokenClassification实现的,默认调用forward方法只返回token分类的logits。要拿到BERT各层的隐藏状态,只需在调用模型时添加output_hidden_states=True参数——此时返回的结果会包含所有层的隐藏状态序列。

修改后的核心代码如下:

# 调用BERT时指定返回隐藏状态
outputs = self.bert(input_ids, attention_mask=attention_mask, output_hidden_states=True)
# 通过.hidden_states属性直接获取所有层的隐藏状态(列表形式,从embedding层到最后一层)
last_four_hidden = outputs.hidden_states[-4:]  # 取最后4层
# 拼接最后4层的隐藏状态
sequence_output = torch.cat(last_four_hidden, dim=-1)

不用硬记索引,直接通过.hidden_states属性访问更直观,避免因返回结构变化出错。

二、更简便的BERT-CRF模型实现方式

有两种省心的实现思路:

1. 基于基础BERT+自定义CRF层

直接用未绑定分类头的基础BERT模型(如bert-base-uncased),自己添加线性映射层和CRF层,代码更灵活,也不会受预训练NER模型的分类头干扰:

import torch
import torch.nn as nn
from transformers import BertModel
from torchcrf import CRF

class BERT_CRF(nn.Module):
    def __init__(self, num_labels):
        super().__init__()
        # 加载基础BERT,指定返回隐藏状态
        self.bert = BertModel.from_pretrained('bert-base-uncased', output_hidden_states=True)
        self.dropout = nn.Dropout(0.1)
        # 拼接4层后,维度是hidden_size*4,映射到标签数
        self.classifier = nn.Linear(self.bert.config.hidden_size * 4, num_labels)
        self.crf = CRF(num_labels, batch_first=True)

    def forward(self, input_ids, attention_mask, labels=None):
        outputs = self.bert(input_ids, attention_mask=attention_mask)
        # 拼接最后4层隐藏状态
        last_four = torch.cat(outputs.hidden_states[-4:], dim=-1)
        sequence_output = self.dropout(last_four)
        logits = self.classifier(sequence_output)
        
        # 训练时返回损失,推理时返回解码后的标签序列
        if labels is not None:
            loss = -self.crf(logits, labels, mask=attention_mask.bool(), reduction='mean')
            return loss
        else:
            preds = self.crf.decode(logits, mask=attention_mask.bool())
            return preds

2. 简化版:只用最后一层隐藏状态

如果对精度要求不是极致,也可以直接用BERT最后一层的隐藏状态,省去拼接步骤,减少计算量,代码更简洁:

# 替换拼接部分的代码
sequence_output = self.dropout(outputs.hidden_states[-1])
logits = self.classifier(sequence_output)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 00:31:15