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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 20:48:14