为BERT模型添加CRF层实现NER时遇mask错误求助
问题:BERT+CRF医疗NER模型训练触发ValueError:“mask of the first timestep must all be on”
错误信息
ValueError Traceback (most recent call last) <ipython-input-32-99c3c401704b> in <cell line: 85>() 83 84 # Start training ---> 85 trainer.train() 7 frames /usr/local/lib/python3.10/dist-packages/torchcrf/__init__.py in _validate(self, emissions, tags, mask) 165 no_empty_seq_bf = self.batch_first and mask[:, 0].all() 166 if not no_empty_seq and not no_empty_seq_bf: ---> 167 raise ValueError('mask of the first timestep must all be on') 168 169 def _compute_score( ValueError: mask of the first timestep must all be on
相关代码
from transformers import TrainingArguments, Trainer from torchcrf import CRF import torch.nn as nn from transformers import DataCollatorForTokenClassification from transformers import AutoTokenizer, BertTokenizerFast def tokenize_and_align_labels(examples): tokenized_inputs = tokenizer(examples["tokens"], truncation=True, is_split_into_words=True) labels = [] for i, label in enumerate(examples[f"ner_tags"]): word_ids = tokenized_inputs.word_ids(batch_index=i) previous_word_idx = None label_ids = [] for word_idx in word_ids: if word_idx is None: label_ids.append(-100) elif word_idx != previous_word_idx: label_ids.append(label[word_idx]) else: label_ids.append(label[word_idx] if label_all_tokens else -100) previous_word_idx = word_idx labels.append(label_ids) tokenized_inputs["labels"] = labels return tokenized_inputs label_all_tokens = False tokenizer = BertTokenizerFast.from_pretrained('bert-base-cased') tokenized_data = my_dataset_dict.map(tokenize_and_align_labels, batched=True) data_collator = DataCollatorForTokenClassification(tokenizer=tokenizer) class BERT_CRF_Model(nn.Module): def __init__(self, bert_model, num_labels): super(BERT_CRF_Model, self).__init__() self.bert = bert_model self.dropout = nn.Dropout(0.1) self.classifier = nn.Linear(self.bert.config.hidden_size, num_labels) self.crf = CRF(num_labels, batch_first=True) def forward(self, input_ids, attention_mask, labels=None): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) sequence_output = self.dropout(outputs[0]) emissions = self.classifier(sequence_output) if labels is not None: loss = -self.crf(emissions, labels, mask=attention_mask.bool(), reduction='mean') return loss else: prediction = self.crf.decode(emissions, mask=attention_mask.bool()) return emissions class CustomTrainer(Trainer): def __init__(self, *args, crf_layer=None, **kwargs): super().__init__(*args, **kwargs) self.crf_layer = crf_layer def compute_loss(self, model, inputs, return_outputs=False): labels = inputs.pop("labels") emissions = model(**inputs) emissions = torch.stack(emissions) if isinstance(emissions, list) else emissions mask = inputs["attention_mask"].bool() if mask.size(0) == 0 or mask[:, 0].sum() == 0: raise ValueError("Le masque du premier pas de temps doit être activé") loss = -self.crf_layer(emissions, labels, mask=mask) return (loss, inputs) if return_outputs else loss from transformers import BertModel bert_model = BertModel.from_pretrained("bert-base-cased") model = BERT_CRF_Model(bert_model, num_labels=len(unique_labels)) crf_layer = CRF(num_tags=len(unique_labels)) training_args = TrainingArguments( output_dir="my_awesome_ner_model", evaluation_strategy="epoch", learning_rate=2e-5, per_device_train_batch_size=12, per_device_eval_batch_size=12, num_train_epochs=1, weight_decay=0.01, push_to_hub=True, ) trainer = CustomTrainer( model=model, args=training_args, train_dataset=tokenized_data["train"], eval_dataset=tokenized_data["val"], tokenizer=tokenizer, data_collator=data_collator, crf_layer=crf_layer ) trainer.train()
问题分析与修复方案
核心问题
错误来自torchcrf的校验逻辑:要求批量中每个样本的第一个时间步(即<CLS> token位置)的attention_mask必须为1,不能存在被mask的情况。你的代码存在以下关键问题:
- 重复实例化CRF层:模型内部已定义
self.crf,但外部又单独创建crf_layer传给自定义Trainer,导致逻辑冲突。 - 自定义Trainer的compute_loss逻辑错误:提前移除labels后,模型forward进入预测分支返回emissions,后续用外部CRF计算loss,且模型自带的CRF未被正确使用。
- 潜在数据问题:数据集可能存在空样本,导致tokenize后attention_mask异常。
修复步骤
1. 修正模型forward逻辑,统一使用内部CRF
修改BERT_CRF_Model的forward方法,正确处理-100标签并返回标准格式的输出:
class BERT_CRF_Model(nn.Module): def __init__(self, bert_model, num_labels): super(BERT_CRF_Model, self).__init__() self.bert = bert_model self.dropout = nn.Dropout(0.1) self.classifier = nn.Linear(self.bert.config.hidden_size, num_labels) self.crf = CRF(num_labels, batch_first=True) def forward(self, input_ids, attention_mask, labels=None): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) sequence_output = self.dropout(outputs[0]) emissions = self.classifier(sequence_output) if labels is not None: # 过滤掉标签为-100的无效位置 active_mask = (labels != -100) & attention_mask.bool() # 调整张量维度,只保留有效部分 active_emissions = emissions[active_mask].view(emissions.size(0), -1, emissions.size(2)) active_labels = labels[active_mask].view(labels.size(0), -1) active_mask = active_mask.view(labels.size(0), -1) loss = -self.crf(active_emissions, active_labels, mask=active_mask, reduction='mean') return {"loss": loss, "emissions": emissions} else: prediction = self.crf.decode(emissions, mask=attention_mask.bool()) return {"predictions": prediction, "emissions": emissions}
2. 删除自定义Trainer,使用原生Trainer
模型forward已返回loss,原生Trainer可直接处理,无需额外覆盖compute_loss:
# 初始化模型 bert_model = BertModel.from_pretrained("bert-base-cased") model = BERT_CRF_Model(bert_model, num_labels=len(unique_labels)) # 训练参数保持不变 training_args = TrainingArguments( output_dir="my_awesome_ner_model", evaluation_strategy="epoch", learning_rate=2e-5, per_device_train_batch_size=12, per_device_eval_batch_size=12, num_train_epochs=1, weight_decay=0.01, push_to_hub=True, ) # 使用原生Trainer trainer = Trainer( model=model, args=training_args, train_dataset=tokenized_data["train"], eval_dataset=tokenized_data["val"], tokenizer=tokenizer, data_collator=data_collator, ) trainer.train()
3. 清理数据集,移除空样本
确保数据集中没有tokens为空的样本:
# 过滤空样本 my_dataset_dict = my_dataset_dict.filter(lambda x: len(x["tokens"]) > 0)
4. 验证attention_mask正确性
随机抽取样本检查第一个时间步的mask值:
sample = tokenized_data["train"][0] print("Attention mask first value:", sample["attention_mask"][0]) # 输出应为1
内容的提问来源于stack exchange,提问作者Akram H
相关产品推荐
相关产品推荐

