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

BERT模型训练出现NaN损失,无法完成训练且无输出问题求助

损失NaN问题修复方案

核心问题

你的代码出现NaN损失是因为输入编码与标签位置完全不匹配:

  • 你把问题和答案拼成元组传入编码函数,得到的是问题+答案的拼接序列,但计算start/end位置时,用的是单独编码答案的token索引,这两个索引和拼接序列中的实际位置毫无关系,导致损失计算时出现异常值,最终变为NaN。
  • 同时数据集仅2条样本,训练稳定性极差,进一步加剧了异常。

修复步骤

1. 修正输入编码与标签计算

抽取式QA模型(如BERTForQuestionAnswering)要求答案必须存在于输入的上下文文本中,且标签位置要对应拼接序列中的实际索引:

# Preprocess your dataset
questions = ["What is the capital of France", "Who invented the telephone"]
answers = ["Paris", "Alexander Graham Bell"]

# 构造包含答案的上下文(问题+答案,确保答案在输入序列内)
contexts = [f"{q} {a}" for q, a in zip(questions, answers)]

# 按QA任务标准格式编码:(问题, 上下文)
encoded_inputs = tokenizer.batch_encode_plus(
    list(zip(questions, contexts)),
    padding=True,
    truncation=True,
    max_length=256,
    return_tensors='pt'
)

# 计算正确的start/end位置
start_positions = []
end_positions = []
for q, a, ctx in zip(questions, answers, contexts):
    # 获取完整编码序列(含特殊token)
    full_tokens = tokenizer(q, ctx, add_special_tokens=True)['input_ids']
    # 获取答案的token序列(不含特殊token)
    answer_tokens = tokenizer.encode(a, add_special_tokens=False)
    # 在完整序列中匹配答案的位置
    match_len = len(answer_tokens)
    for i in range(len(full_tokens) - match_len + 1):
        if full_tokens[i:i+match_len] == answer_tokens:
            start_positions.append(i)
            end_positions.append(i + match_len - 1)
            break

start_positions = torch.tensor(start_positions)
end_positions = torch.tensor(end_positions)

2. 提升训练稳定性

  • 扩充数据集:仅2条样本无法让模型学到有效规律,至少收集数百条同类型问答样本。
  • 调整学习率:尝试将学习率降至5e-7,或添加学习率调度器:
    from torch.optim.lr_scheduler import StepLR
    optimizer = torch.optim.AdamW(model.parameters(), lr=5e-7)
    scheduler = StepLR(optimizer, step_size=2, gamma=0.5)
    # 训练循环中每次step后执行
    scheduler.step()
    
  • 梯度裁剪:防止梯度爆炸引发NaN:
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)  # 添加梯度裁剪
    optimizer.step()
    

3. 修正测试逻辑

测试时需保持和训练一致的输入格式(问题+包含答案的上下文):

# Test the model
test_question = "What is the capital of France?"
test_context = f"{test_question} Paris"

encoded_test = tokenizer(test_question, test_context, padding=True, truncation=True, return_tensors='pt')
test_input_ids = encoded_test['input_ids']
test_attention_mask = encoded_test['attention_mask']

model.eval()
with torch.no_grad():
    outputs = model(test_input_ids, attention_mask=test_attention_mask)
    predicted_start = torch.argmax(outputs.start_logits)
    predicted_end = torch.argmax(outputs.end_logits)
    predicted_answer = tokenizer.decode(test_input_ids[0][predicted_start:predicted_end+1])

print(f"Predicted Answer: {predicted_answer}")

4. 模型选型建议

如果目标是打造聊天机器人,抽取式QA模型并不适合,建议改用生成式模型(如T5、GPT系列、ChatGLM等),这类模型可直接生成自由文本答案,无需依赖答案出现在输入序列中。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 15:07:55