简介:命名实体识别(NER)是自然语言处理的基础任务,旨在从非结构化文本中识别人名、地名、机构名等关键实体。传统方法依赖 BiLSTM-CRF 建模序列标签约束,但语义表征能力有限;预训练模型 BERT 虽能提供强大的上下文语义向量,却难以独立保证标签间的合法转移。将 BERT 作为语义特征抽取器,接入 BiLSTM-CRF 进行序列标注,既引入迁移学习带来的通用语言理解,又通过维特比解码实现全局最优标签序列。该方案在金融、医疗、法律等领域的小样本中文语料中表现稳健,特别适合实体边界复杂、标注质量不均的实际工程场景。本文基于 PyTorch 完整呈现这一组合模型的实现细节,包括数据对齐、CRF 手写与分层学习率调参。 做中文命名实体识别(NER)的时候,很多朋友会在“直接用BERT微调”和“传统BiLSTM-CRF”之间反复纠结。这个项目标题给的答案很明确:把BERT作为预训练特征抽取器,后面接上BiLSTM-CRF做序列标注,用PyTorch完整实现一套可落地的中文NER代码。它的核心价值在于解决了两个痛点——单用BERT时标签之间缺乏约束,单用BiLSTM-CRF时语义特征又不够强。这套组合在学术界和工业界都经过了大量验证,尤其适合小样本、标注数据质量一般、以及对实体边界要求比较严格的中文场景。无论你是刚入门NLP的在校生,还是在公司里需要快速搭一个NER服务的工程师,这篇文章里都有能直接用上的内容。

下面我按实际做项目的顺序,从方案选型、数据预处理、代码实现到训练调参,把整套流程完整拆开讲一遍。

1. 方案选型:为什么是BERT + BiLSTM-CRF

1.1 三个组件各自解决什么问题

先说清楚这套架构为什么合理。BERT是一个预训练语言模型,它的强项是上下文语义表示——输入一个句子,它能输出每个字(或者每个词)带有全局语义信息的向量。在中文场景下,BERT用的是字粒度,每个汉字对应一个向量,这个向量不是简单查表得到的静态词向量,而是根据整个句子动态计算出来的,所以同一个“深”字在"深水"和"深夜"里的向量完全不同。

BiLSTM接在BERT后面,作用是进一步建模序列依赖。别看BERT本身已经很强了,但它的Transformer结构对位置信息的建模方式和LSTM不同,BiLSTM天生就更适合捕捉序列中的局部依赖和长距离转移特征。它在每个时间步上会同时利用前向和后向的信息,输出一个拼接后的隐状态向量。

CRF层是这个组合里的“纪律委员”。序列标注任务里有个经典问题——预测出的标签序列可能是非法的。比如B-PER后面紧跟一个I-ORG,这种边界错乱的序列在纯BERT和BiLSTM输出后直接取argmax的时候经常出现,因为每个token是独立计算分类概率的。CRF通过一个转移矩阵显式学习标签之间的约束关系,比如"B后面只能跟I或O,不能跟另一个B",这种全局最优解码能力是前两者不具备的。

1.2 对比其他方案,这套组合赢在哪里

我在实际项目里对比过下面几种方案,结果可以参考下:

方案 语义特征 标签约束 训练成本 适用场景
纯BiLSTM-CRF(随机初始化) 训练数据充足、领域单一
纯BERT微调 + Softmax 数据量大、实体边界规整
BERT + Softmax + 规则后处理 标签约束规则容易手工定义
BERT + BiLSTM-CRF 较高 标注数据有限、实体类型多、边界复杂

从表格里能看出来,BERT + BiLSTM-CRF在语义特征和标签约束上都拿了高分。代价是训练参数量大、推理速度略慢,但在绝大多数业务场景下,这点算力开销换来的精度提升完全是划算的。

我做一个金融领域的命名实体识别项目时,用纯BERT微调F1在0.82左右,换成BERT + BiLSTM-CRF后直接到了0.87,提升主要集中在长实体和嵌套实体上。原因其实很好理解:金融文本里“中国工商银行北京分行”这种长实体,中间每个字单独看语义都是中性的,但CRF通过标签转移约束把整个边界串起来了。

1.3 标题里的“预训练”到底指什么

这里必须澄清一个容易混淆的概念。项目标题说的“预训练”,指的是加载BERT在大规模中文语料上训练好的权重,把它作为初始化参数用到我们的NER任务上,并不是让你在BERT的基础上再去做Masked Language Model预训练。后者成本极高,需要海量无标注语料和专业GPU集群,不是单个项目的范畴。实际要做的,是把这个已经具备通用语言理解能力的BERT“搬到”当前任务的模型结构里,然后让BERT的参数随着任务进行微调(finetune)。这种范式也叫迁移学习——利用在其他任务上学到的通用知识,加速当前任务的学习收敛,这在标注数据有限的场景下尤其重要。

BiLSTM-CRF部分的参数则是从头开始训练的,因为不同的NER任务标注体系不同,这部分没有通用的预训练权重可以加载。这点在设计优化器时有直接影响,后面我会专门讲分层学习率的设置。

2. 环境准备与中文数据预处理

2.1 PyTorch环境与依赖安装

这套代码最核心的几个依赖是:PyTorch、Hugging Face Transformers、pytorch-crf(或者自己写的CRF层)。安装本身不复杂,但不少新手在环境上栽跟头,最常见的坑是版本不匹配。

# 建议在conda环境里操作
conda create -n ner python=3.9 -y
conda activate ner

# CPU版本的PyTorch(先跑通逻辑再用GPU)
pip install torch==2.0.1
# GPU版本请根据你的CUDA版本去PyTorch官网选对应的安装命令,这里不展开

# Transformers库,用来加载BERT预训练模型
pip install transformers==4.30.2

# CRF层,也可以自己实现(后面我会给代码)
pip install pytorch-crf

这里有一个来自实战的提醒:PyTorch 2.x的版本对于 torch.load 默认参数有调整,比如 weights_only 的默认值在某些版本里变成了True,如果预训练模型是用旧版本保存的,加载时可能报错或者行为不一致。最稳妥的做法是固定版本环境,不要盲目追新。我在一个项目里就因为升级了PyTorch导致BERT权重加载行为变化,排查了半天才发现是版本问题。

2.2 中文数据的标签体系与粒度选择

中文NER的标注体系常用的是BIO(Begin, Inside, Outside)或者BIESO(加End, Single)。我推荐用BIO,简单直接,配合CRF足够。

如果你标注的是人名、地名、机构名这三类,标签集合就是:

  • B-PER, I-PER(人名)
  • B-LOC, I-LOC(地名)
  • B-ORG, I-ORG(机构名)
  • O(非实体)

数据格式就是每行一个字/词 + 空格 + 标签,句子之间空一行。中文BERT是基于字的,所以这里按字切分,一个汉字对应一个标签。很多新手在做中文NER时总想着先分词再标注,这是个误区——基于字的模型天生不受分词错误影响,BERT的字向量已经能表达词级语义了。

样例数据看起来像这样:

中 O
国 B-LOC
工 I-LOC
商 I-LOC
银 I-LOC
行 I-LOC
发 O
布 O
了 O
人 B-PER
工 I-PER
智 I-PER
能 I-PER
报 O
告 O

2.3 数据加载与标签对齐:最容易出bug的环节

用BERT做中文NER,数据预处理最核心的一步是让BERT的token和标签一一对应。这一部分出错率最高,值得详细说。

BERT的分词器(Tokenizer)处理中文时,默认每个字是一个token,但遇到英文、数字、特殊符号时,可能被切成多个subword。比如"iPhone14"可能被切成"i"、"phone"、"14"三个token。这就导致一个问题:原始的一个"字"对应的标签,要复制给切出的多个subword。

from transformers import BertTokenizer

tokenizer = BertTokenizer.from_pretrained("bert-base-chinese")
labels = ["O", "B-LOC", "I-LOC", "O"]

# 模拟原始句子分字后对应标签
chars = ["中", "国", "银", "行"]
encoded = tokenizer(chars, is_split_into_words=True)
word_ids = encoded.word_ids()
print(encoded.tokens())    # ['[CLS]', '中', '国', '银', '行', '[SEP]']
print(word_ids)            # [None, 0, 1, 2, 3, None]

# 根据word_ids把标签对齐到token级别
aligned_labels = []
previous_word_idx = None
for word_idx in word_ids:
    if word_idx is None:
        aligned_labels.append(-100)  # 特殊token的标签,-100在loss计算时被忽略
    elif word_idx != previous_word_idx:
        aligned_labels.append(label2id[labels[word_idx]])
    else:
        aligned_labels.append(label2id[labels[word_idx]])  # subword沿用第一个标签
    previous_word_idx = word_idx

这里的 -100 是个约定俗成的技巧,PyTorch的CrossEntropyLoss默认会忽略值为-100的target位置,这样[CLS]和[SEP]就不会参与损失计算。另一个关键点是 is_split_into_words=True ,这个参数告诉tokenizer输入已经按词切分了,返回的 word_ids 能帮我们把每个token映射回原始字。

2.4 Padding与注意力掩码的细节

模型要求batch内序列长度一致,所以需要padding。但BERT在计算自注意力时必须知道哪些位置是真实的、哪些是padding,这就是attention_mask的作用。CRF在解码时也要排除padding位置,否则会把padding的标签当成有效标签参与计算。

完整的数据批处理逻辑我建议自己写一个Dataset类,而不是依赖transformers自带的,因为NER任务需要同时处理input_ids、attention_mask、labels三个序列,而且labels和input_ids的长度必须严格一致。

import torch
from torch.utils.data import Dataset

class NERDataset(Dataset):
    def __init__(self, texts, labels, tokenizer, max_len=128, label2id=None):
        self.texts = texts
        self.labels = labels
        self.tokenizer = tokenizer
        self.max_len = max_len
        self.label2id = label2id

    def __len__(self):
        return len(self.texts)

    def __getitem__(self, idx):
        text_chars = list(self.texts[idx])
        labels = self.labels[idx]
        encoded = self.tokenizer(
            text_chars,
            is_split_into_words=True,
            max_length=self.max_len,
            truncation=True,
            padding='max_length',
            return_tensors='pt'
        )
        word_ids = encoded.word_ids()

        label_ids = [-100] * self.max_len  # 如果padding到max_len,先全部填-100
        previous_word_idx = None
        i = 0
        for word_idx in word_ids:
            if word_idx is None:
                label_ids[i] = -100
            elif word_idx >= len(labels):
                label_ids[i] = -100  # 截断超出部分
            elif word_idx != previous_word_idx:
                label_ids[i] = self.label2id[labels[word_idx]]
            else:
                label_ids[i] = self.label2id[labels[word_idx]]
            previous_word_idx = word_idx
            i += 1

        return {
            'input_ids': encoded['input_ids'].squeeze(0),
            'attention_mask': encoded['attention_mask'].squeeze(0),
            'labels': torch.tensor(label_ids, dtype=torch.long)
        }

这里有个处理细节容易被忽略:截断之后, word_idx 可能超出原来labels的长度,需要加一个边界判断,否则会报IndexError。如果你截断的位置恰好落在一个实体的中间,那被保留的部分也要有标签——好在 word_ids 能帮我们做映射,原始标签在截断范围内时直接取对应标签即可。

3. 模型结构与核心代码实现

3.1 从BERT输出到BiLSTM输入

模型定义是整个项目的核心。在PyTorch里,我们通过继承 nn.Module 来搭建这个组合模型。BERT部分用transformers加载预训练权重,BiLSTM和CRF部分则使用PyTorch的基础模块。

import torch
import torch.nn as nn
from transformers import BertModel, BertPreTrainedModel

class BertBiLSTMCRF(nn.Module):
    def __init__(self, bert_pretrained_path, num_labels, lstm_hidden=256, dropout=0.1):
        super().__init__()
        self.num_labels = num_labels
        self.bert = BertModel.from_pretrained(bert_pretrained_path)
        self.dropout = nn.Dropout(dropout)
        self.bilstm = nn.LSTM(
            input_size=768,          # BERT base 输出维度
            hidden_size=lstm_hidden,
            num_layers=1,
            batch_first=True,
            bidirectional=True
        )
        self.fc = nn.Linear(lstm_hidden * 2, num_labels)  # 双向LSTM拼接后是2倍hidden
        self.crf = CRF(num_labels, batch_first=True)

为什么BERT的输出维度是768?因为 bert-base-chinese 属于BERT base系列,hidden_size固定为768。 lstm_hidden 我习惯设成256,这样双向拼接后是512,经过一个全连接层后映射到标签数量,参数量比较均衡。

3.2 前向传播与CRF解码逻辑

前向传播的逻辑并不复杂:输入先过BERT得到序列特征,再经过dropout防止过拟合,接着送到BiLSTM建模序列依赖,然后通过全连接层获得每个token在每个标签上的得分(logits),最后把这个得分矩阵交给CRF计算损失或解码。

    def forward(self, input_ids, attention_mask, labels=None):
        bert_outputs = self.bert(
            input_ids=input_ids,
            attention_mask=attention_mask
        )
        sequence_output = bert_outputs.last_hidden_state  # [batch, seq_len, 768]
        sequence_output = self.dropout(sequence_output)
        lstm_output, _ = self.bilstm(sequence_output)     # [batch, seq_len, 512]
        lstm_output = self.dropout(lstm_output)
        logits = self.fc(lstm_output)                     # [batch, seq_len, num_labels]

        if labels is not None:
            # 训练阶段:计算CRF损失(负对数似然)
            loss = -self.crf(logits, labels, mask=attention_mask.bool(), reduction='mean')
            return loss
        else:
            # 推理阶段:维特比解码
            decoded = self.crf.decode(logits, mask=attention_mask.bool())
            return decoded

推理时的 decoded 是一个list,每个元素是对应句子解码出的标签索引序列。这里有个要点:CRF返回的解码序列长度是真实的序列长度(不包含padding),所以在后处理时需要记录每个句子的原始长度,否则把padding位置当成有效标签会把你彻底搞晕。

注意我这里的 CRF 类不是 pytorch-crf 库,而是自己实现的版本。原因后面问题排查部分会详细说,但先说结论:自己实现CRF虽然代码多一点,但你能完全掌控每一处细节,排查问题时心里有底。

3.3 手写CRF层:损失计算与维特比解码

CRF层的核心是转移矩阵。在序列标注里,模型不仅要预测每个token的标签,还要考虑标签之间的转移概率。比如"B-PER后面跟I-PER"是合理的,但"I-PER后面跟B-LOC"基本不可能。转移矩阵就是用来存储这种标签间转移的得分。

计算损失时,CRF用的方法是“真实路径得分 / 所有可能路径得分之和”,然后取负对数。这里的难点在分母——所有可能路径是指数级的,不能暴力枚举,要用动态规划。

import torch
import torch.nn as nn

class CRF(nn.Module):
    def __init__(self, num_tags, batch_first=True):
        super().__init__()
        self.num_tags = num_tags
        self.batch_first = batch_first
        # 转移矩阵:start_transitions[i]表示从start转移到标签i的得分
        # transitions[i][j]表示从标签i转移到标签j的得分
        self.start_transitions = nn.Parameter(torch.empty(num_tags))
        self.end_transitions = nn.Parameter(torch.empty(num_tags))
        self.transitions = nn.Parameter(torch.empty(num_tags, num_tags))
        nn.init.uniform_(self.start_transitions, -0.1, 0.1)
        nn.init.uniform_(self.end_transitions, -0.1, 0.1)
        nn.init.uniform_(self.transitions, -0.1, 0.1)

    def _compute_score(self, emissions, tags, mask):
        # 计算一条真实路径的总得分 = 发射得分 + 转移得分
        seq_len = emissions.shape[1]
        score = self.start_transitions[tags[:, 0]]
        for i in range(seq_len - 1):
            current_tag = tags[:, i]
            next_tag = tags[:, i + 1]
            score += emissions[torch.arange(emissions.shape[0]), i, current_tag]
            score += self.transitions[current_tag, next_tag] * mask[:, i + 1]
        last_tag = tags.gather(1, mask.sum(dim=1, keepdim=True) - 1).squeeze(1)
        score += emissions[torch.arange(emissions.shape[0]), mask.sum(dim=1) - 1, last_tag]
        score += self.end_transitions[last_tag]
        return score

    def _compute_normalizer(self, emissions, mask):
        # 用前向算法计算所有可能路径的得分和(log-sum-exp)
        seq_len = emissions.shape[1]
        score = self.start_transitions.unsqueeze(0) + emissions[:, 0, :]
        for i in range(1, seq_len):
            broadcast_score = score.unsqueeze(2)  # [batch, num_tags, 1]
            broadcast_emissions = emissions[:, i, :].unsqueeze(1)  # [batch, 1, num_tags]
            next_score = broadcast_score + self.transitions.unsqueeze(0) + broadcast_emissions
            next_score = torch.logsumexp(next_score, dim=1)
            score = torch.where(mask[:, i].unsqueeze(1).bool(), next_score, score)
        score = score + self.end_transitions.unsqueeze(0)
        return torch.logsumexp(score, dim=1)

    def forward(self, emissions, tags, mask=None, reduction='mean'):
        if mask is None:
            mask = torch.ones(emissions.shape[:2], dtype=torch.long, device=emissions.device)
        if self.batch_first:
            emissions = emissions.transpose(0, 1)
            tags = tags.transpose(0, 1)
            mask = mask.transpose(0, 1)
        numerator = self._compute_score(emissions, tags, mask)
        denominator = self._compute_normalizer(emissions, mask)
        llh = numerator - denominator
        if reduction == 'none':
            return llh
        elif reduction == 'sum':
            return llh.sum()
        else:
            return llh.mean()

    def decode(self, emissions, mask=None):
        # 维特比解码:找出得分最高的路径
        if mask is None:
            mask = torch.ones(emissions.shape[:2], dtype=torch.long, device=emissions.device)
        if self.batch_first:
            emissions = emissions.transpose(0, 1)
            mask = mask.transpose(0, 1)
        seq_len = emissions.shape[0]
        batch_size = emissions.shape[1]
        score = self.start_transitions.unsqueeze(0) + emissions[0]
        history = []
        for i in range(1, seq_len):
            broadcast_score = score.unsqueeze(2)
            broadcast_emissions = emissions[i].unsqueeze(1)
            next_score = broadcast_score + self.transitions.unsqueeze(0) + broadcast_emissions
            next_score, indices = next_score.max(dim=1)
            score = torch.where(mask[i].unsqueeze(1).bool(), next_score, score)
            history.append(indices)
        score = score + self.end_transitions.unsqueeze(0)
        best_tag = torch.argmax(score, dim=1)
        best_paths = []
        for b in range(batch_size):
            length = mask[:, b].sum().item()
            best_path = [best_tag[b].item()]
            for h in reversed(history[:length - 1]):
                best_tag[b] = h[b][best_tag[b]]
                best_path.append(best_tag[b].item())
            best_path.reverse()
            best_paths.append(best_path)
        return best_paths

这里要特别提一下 torch.where 的作用。在padding位置上,前向算法的分数不应该继续累加,而是保持上一步的分数不变。如果不加这个判断,padding位置会污染整条路径的得分计算。

理解了CRF的实现,你就能明白为什么BERT + BiLSTM-CRF比纯BERT + Softmax好——Softmax对每个位置独立分类,完全不考虑相邻标签的关系;CRF则用转移矩阵编码了这种关系,并且在解码时使用维特比算法寻找全局最优路径。

3.4 分层学习率:让预训练权重和新层各得其所

BERT部分已经有了很好的语义理解能力,如果直接让它跟BiLSTM-CRF用同一个学习率,尤其是用一个较大的学习率,极其容易灾难性遗忘——把BERT在大规模语料学到的通用知识给冲掉。而BiLSTM和CRF是随机初始化的,学习率太小时收敛极慢。

我的做法是给BERT层设置相对较小的学习率(比如2e-5),给下游新建的层设置相对较大的学习率(比如1e-3),这个比例大约1:50。下面是用PyTorch实现这个策略的方式。

from transformers import AdamW, get_linear_schedule_with_warmup

bert_params = set(model.bert.parameters())
crf_params = set(model.crf.parameters())
lstm_fc_params = set(model.bilstm.parameters()) | set(model.fc.parameters())

optimizer_grouped_parameters = [
    {'params': bert_params, 'lr': 2e-5},
    {'params': lstm_fc_params, 'lr': 1e-3},
    {'params': crf_params, 'lr': 1e-3},
]

optimizer = AdamW(optimizer_grouped_parameters, weight_decay=0.01)
total_steps = len(train_dataloader) * num_epochs
scheduler = get_linear_schedule_with_warmup(
    optimizer,
    num_warmup_steps=int(total_steps * 0.1),
    num_training_steps=total_steps
)

关于优化器这里有经验变化:transformers早就把AdamW集成到了自己的 optimization 模块里,并且PyTorch 2.0之后原生的 torch.optim.AdamW 已经足够稳定,不需要额外引入transformers的版本。我现在的习惯是直接用 torch.optim.AdamW ,少一个依赖少一份坑。

4. 训练流程设计与参数调优

4.1 训练循环:梯度裁剪是必需品

到了训练环节,你可能觉得就是标准的三板斧——forward、backward、step。但在BERT + BiLSTM-CRF这种结构里,梯度裁剪是必须加的。BiLSTM在长序列上容易出现梯度爆炸(特别是序列长度超过200的时候,及时加上了dropout),CRF的反向传播也比较剧烈。我用了一个通用的 clip_grad_norm_ 函数来限制梯度范数。

def train_epoch(model, dataloader, optimizer, scheduler, device):
    model.train()
    total_loss = 0
    for batch in dataloader:
        input_ids = batch['input_ids'].to(device)
        attention_mask = batch['attention_mask'].to(device)
        labels = batch['labels'].to(device)

        optimizer.zero_grad()
        loss = model(input_ids, attention_mask, labels)
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)
        optimizer.step()
        scheduler.step()
        total_loss += loss.item()
    return total_loss / len(dataloader)

梯度范数我通常设在 1.0 5.0 之间。如果设得太小,模型训练不充分;设得太大,梯度裁剪形同虚设。我一般先设为5.0,如果观察到loss曲线不稳定再调小。

4.2 验证指标:不要只看token准确率

很多新手训练NER模型时只看整体的准确率(Accuracy),但这其实是个严重的误区。因为一个NER数据集里90%以上的token都是O标签,模型哪怕把所有非实体都预测成O,准确率也能到90%以上,而真正的实体一个都没预测出来。

正确的评估方式是按实体级别计算精确率(Precision)、召回率(Recall)、F1值。也就是说,只有预测的实体边界和类型跟真实标注完全一致才算正确。比如真实实体是“中国工商银行”(ORG),模型预测“中国”是ORG后就没往下接,这个实体就算预测错误。

计算实体级F1的代码逻辑比较复杂,但核心思想是:把每个句子解码出的标签序列转成实体列表,然后和真实实体列表做对比,统计出精确匹配的数量。

def extract_entities(tags, id2label):
    # 把标签序列转成实体列表 [(entity_type, start, end), ...]
    entities = []
    i = 0
    while i < len(tags):
        if tags[i] in id2label and id2label[tags[i]].startswith('B-'):
            entity_type = id2label[tags[i]][2:]
            start = i
            i += 1
            while i < len(tags):
                tag = id2label.get(tags[i], 'O')
                if tag == 'I-' + entity_type:
                    i += 1
                else:
                    break
            entities.append((entity_type, start, i - 1))
        else:
            i += 1
    return entities

id2label 是标签ID到标签字符串的逆映射表。实体级F1的计算公式和平常一样:F1 = 2 * P * R / (P + R),其中P是预测出的实体中有多少是正确匹配的,R是标准实体中有多少被正确预测出来了。

4.3 超参数设置参考

不同数据集的最优超参数不同,但有一套经过验证的起点值可以参考。以BERT base + BiLSTM(256) + CRF为例:

超参数 建议值 说明
max_seq_len 128 中文句子大多在100字以内,太长显存吃不消
batch_size 16 单卡12G显存左右可跑,不够就减半
BERT学习率 2e-5 微调BERT的通用值
BiLSTM/FC学习率 1e-3 新层从零训练,需要大一点
CRF学习率 1e-3 CRF参数少,收敛快
warmup比例 0.1 前10%步数线性预热
训练轮数 5 观察验证集F1,早停

batch size这个参数很有意思。BERT base模型在12G显存上,batch_size=32就已经比较紧张了。如果显存不够,优先降低batch_size而不是max_seq_len,因为序列长度影响的是每个样本的相对复杂度,batch_size影响的是训练稳定性。batch_size降到8甚至4都能训练,只是收敛略慢。

5. 常见问题与排查技巧实录

5.1 显存不足(OOM)

BERT模型占用的显存大头其实是在保存中间激活值用于反向传播,而不是参数本身。如果OOM了,按这个顺序排查:

  1. 调小batch_size,从16调到8甚至4,看是否还报错。
  2. 调小max_seq_len,尽量避免过长的序列。
  3. 使用梯度累积,模拟大的batch_size。
accumulation_steps = 4
optimizer.zero_grad()
for i, batch in enumerate(train_dataloader):
    loss = model(batch) / accumulation_steps  # 损失也要除以累积步数
    loss.backward()
    if (i + 1) % accumulation_steps == 0:
        optimizer.step()
        scheduler.step()
        optimizer.zero_grad()

注意一个细节:梯度累积时,loss要除以累积步数,这样等效于用一个更大的batch size做一次更新,梯度范数是多个小batch的平均。

还有一个容易忽略的因素——PyTorch的 ctx 缓存机制。如果你在同一个进程里反复创建和销毁大模型,显存碎片化会很严重。在代码里显式加一句 torch.cuda.empty_cache() 只能清理未使用的缓存块,本质上解决不了显存碎片化的问题,更有效的办法是让batch size稳定,不让模型动态分配太多中间变量。

5.2 训练loss不下降或下降极慢的情况

loss不下降通常有这几种原因:

第一,学习率设置不对。如果你把BERT层和BiLSTM-CRF层设成同一个学习率,并且这个率又偏大(比如1e-3),BERT可能直接学崩了。解决办法是分层学习率,BERT层2e-5起步。

第二,CRF的mask没传对。如果你在训练时没有给CRF传attention_mask,CRF会把padding位置也当成有效标签来计算转移概率,这会导致loss严重异常。检查一下你的CRF调用是否传了mask参数,并且mask的类型是 torch.bool

第三,标签对齐错误。常见的问题是padding后label_ids用的是0而不是-100,导致[CLS]位置也参与了loss计算。这种问题不会让loss完全不降,但会让模型学到错误的模式。排查方法很简单:打印一个batch的input_ids和labels,肉眼看看[CLS]位置的标签是不是-100。

5.3 预测结果中实体边界漂移

BERT + BiLSTM-CRF在预测时,如果实体边界经常“多一个字”或者“少一个字”,往往是两类问题。

一类是训练数据本身边界标注不统一。比如"中国工商银行"有时候标成"工商银行",这种标注噪声会让CRF学到错误转移模式。如果你的数据允许,建议统一标注规范,尤其是那种由多个子部分组成的机构名。

另一类是CRF转移矩阵学到的偏好过于顽固。比如训练数据里"银行"后面的词经常是O,CRF就倾向于在任何"银行"后面收尾实体。这时候可以调整CRF的初始转移矩阵—— start_transitions end_transitions 可以设置偏置,但这属于比较后期的手段,一般我建议先检查数据。

5.4 BERT权重加载失败或警告其他问题

BertModel.from_pretrained("bert-base-chinese") 加载权重时偶尔会遇到一些警告,比如某些层的权重缺失或者shape不匹配。这通常是因为你在 BertModel.from_pretrained 里传了 config 参数并修改了某些配置(比如改了hidden_size),而预训练权重是按默认配置保存的。解决方案很简单:尽量保持BERT部分的配置跟预训练时一致,不要随意改动hidden_size和num_hidden_layers。

另一个常见问题是网络原因无法下载预训练权重。如果 from_pretrained 卡住或超时,可以手动下载权重文件到本地目录,然后 from_pretrained("/path/to/local/bert") 。哈工大的 bert-base-chinese 权重质量很高,也不存在网络问题,推荐使用。

还有一个坑值得提醒:PyTorch 2.6版本之后 torch.load 的默认参数有过变化, weights_only 参数默认值在某些版本里改为True,如果加载老版本保存的模型,可能出现 WeightsUnpickler 相关的报错。遇到时,要么升级到含新逻辑的transformers版本,要么在 torch.load 里显式指定 weights_only=False 。但务必注意, weights_only=False 会带来反序列化风险,仅在你确定权重文件来源可信的情况下使用。

5.5 加了BiLSTM-CRF反而比纯BERT差

这种返祖现象确实存在,而且不少。排查顺序:

  1. 先看是不是CRF的学习率过大。CRF的转移矩阵是全局参数,受序列长度影响极大,学习率跟BiLSTM设成一样可能本身就是错的。我建议CRF层用比BiLSTM稍小的学习率,比如5e-4。

  2. 再看训练数据量。CRF是有向图模型,对标签共现频率非常敏感。如果训练数据只有几千条,标签转移事件很少,CRF学到的转移矩阵会带有很强的噪声。这时候可以考虑用 pytorch-crf 里默认的正则化,或者减少训练轮数防止过拟合。

  3. 最后检查loss是否真的在下降。有时候加了CRF之后,loss下降了但F1也下降了,这很可能是CRF过度拟合了训练集里的标签转移模式,在验证集上泛化变差。此时应该提高dropout比例(从0.1调到0.3),或者给CRF转移矩阵加L2正则。

5.6 推理速度太慢的优化经验

BERT + BiLSTM-CRF的推理速度确实不算快。CPU上单条短文本(128字内)大约需要20-30ms,GPU上大概2-5ms。如果你要处理高并发场景,有几个经验:

  1. 使用ONNX Runtime或者TensorRT对BERT部分做加速。BiLSTM和CRF比较轻量,瓶颈在BERT的Transformer计算上。将BERT导出为ONNX后,CPU推理性能提升较明显,因为Pytorch的CPU算子对Transformer的支持不够极致,ONNX Runtime针对这类算子有额外优化。

  2. 减少max_seq_len。如果你知道业务中96%的句子都在50字以内,就不要把max_seq_len设成512,无谓的padding计算浪费显存和算力。

  3. 如果你追求极致的推理性能,可以尝试蒸馏BERT——用大模型作为teacher,训练一个小模型(比如tiny-bert)作为student,这个方案在实际业务中特别香,精度损失通常控制在2%以内,速度能提升5-8倍。

做了几个项目之后,我最大的感受是:BERT + BiLSTM-CRF这套结构,最大的威力不在于单个模型有多强,而在于它把“语义理解”和“序列约束”两个优势合理结合了。在实际调试中,我很少遇到模型本身的问题,更多的问题出现在数据对齐、超参设置和CRF实现细节上。最后再分享一个受用很久的小技巧:训练前先打印一个batch的预测结果,肉眼过一遍实体边界,很多时候比盯着训练曲线管用得多。千万不要刚拿到一个NER任务就冲上去调参,20%的时间花在数据质量检查上,能省下你80%的调参时间。

本文还有配套的精品资源,点击获取 menu-r.4af5f7ec.gif

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐