使用SQuAD数据集微调LaBSE做QA任务时遇Loss返回错误求助
解决LaBSE微调SQuAD问答任务时的Loss返回错误
问题原因
你加载的BertModel是基础预训练模型,仅输出last_hidden_state和pooler_output,没有针对问答任务的预测头,也不会自动利用预处理好的start_positions和end_positions计算损失。而Trainer默认期望模型返回包含loss的字典,因此触发该错误。
解决方案
方法一:使用内置问答模型类(最简方案)
直接替换模型加载代码,使用BertForQuestionAnswering——这个类内置了问答任务的预测头,会自动处理start_positions和end_positions并计算损失:
# 替换原来的BertModel导入和初始化 from transformers import BertForQuestionAnswering model = BertForQuestionAnswering.from_pretrained(model_checkpoint)
方法二:自定义模型类(灵活扩展)
如果需要保留自定义空间,可以继承BertModel,手动添加问答预测头和损失计算逻辑:
import torch import torch.nn as nn from transformers import BertModel class BertForQA(nn.Module): def __init__(self, model_checkpoint): super().__init__() self.bert = BertModel.from_pretrained(model_checkpoint) # 新增预测层,输出start和end位置的logits self.qa_outputs = nn.Linear(self.bert.config.hidden_size, 2) def forward(self, input_ids, token_type_ids=None, attention_mask=None, start_positions=None, end_positions=None): # 获取Bert的隐藏层输出 outputs = self.bert(input_ids=input_ids, token_type_ids=token_type_ids, attention_mask=attention_mask) sequence_output = outputs.last_hidden_state # 计算start/end位置的logits logits = self.qa_outputs(sequence_output) start_logits, end_logits = logits.split(1, dim=-1) start_logits = start_logits.squeeze(-1) end_logits = end_logits.squeeze(-1) # 训练阶段计算损失 loss = None if start_positions is not None and end_positions is not None: # 处理超出序列长度的标签,避免报错 ignored_index = start_logits.size(1) start_positions.clamp_(0, ignored_index) end_positions.clamp_(0, ignored_index) # 使用交叉熵损失计算start和end位置的损失 loss_fct = nn.CrossEntropyLoss(ignore_index=ignored_index) start_loss = loss_fct(start_logits, start_positions) end_loss = loss_fct(end_logits, end_positions) loss = (start_loss + end_loss) / 2 # 返回包含loss的字典,符合Trainer的要求 return {"loss": loss, "start_logits": start_logits, "end_logits": end_logits} # 初始化自定义模型 model = BertForQA(model_checkpoint)
注意事项
- 你的数据集预处理代码已经正确生成了
start_positions和end_positions,无需修改。 - 方法一适合快速实现任务,方法二适合需要自定义模型结构或损失函数的场景。
内容的提问来源于stack exchange,提问作者Mateusz Pasierbek
相关产品推荐
相关产品推荐

