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
相关产品推荐
相关产品推荐

