mT5-small问答模型训练收敛但推理输出空答案技术问询
解决mT5-small阿拉伯语QA模型推理输出空答案的问题
针对mT5-small在Arabic SQUAD数据集上训练后,训练/验证指标优异但推理输出<extra_id_0>这类空答案的问题,可从以下几个核心方向排查:
1. 数据预处理的标签映射错误
mT5做QA任务时,若采用起始/结束位置分类的方式,必须确保答案的字符位置正确映射到tokenized后的索引。常见问题:
- 未正确调用tokenizer的
char_to_token方法将答案的起始/结束字符位置转换为token索引,导致模型学习的标签是无效值(如固定为0),看似指标达标但未学到问题与答案的关联。 - 错误地将答案替换为
<extra_id_*>系列特殊token作为训练目标,模型自然会输出这类空标记。
2. 推理阶段的解码逻辑错误
训练时若为起始/结束位置分类任务,推理不能直接解码模型的logits为token,而需:
- 从模型输出的
start_logits和end_logits中取argmax得到起始、结束token索引; - 从输入的上下文token序列中截取对应区间的token,再解码为自然语言答案。
若直接解码模型输出的token序列,大概率会得到<extra_id_0>这类默认特殊token。
3. 模型输出层与任务不匹配
需确认模型输出层是否适配QA任务:
- 若做起始/结束位置预测,需从mT5的
last_hidden_state(序列输出)而非<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token输出构建两个Dense分类层,分别预测起始、结束位置的概率,输出维度应为(None, max_seq_len); - 若误用序列生成的输出层搭配分类损失,会导致模型输出混乱。
关键代码排查示例
错误预处理示例(导致输出空token)
def preprocess_fn(examples): inputs = [f"question: {q} context: {c}" for q, c in zip(examples["question"], examples["context"])] # 错误:将训练目标设为固定的<extra_id_0> targets = ["<extra_id_0>"] * len(examples) model_inputs = tokenizer(inputs, max_length=512, truncation=True) with tokenizer.as_target_tokenizer(): labels = tokenizer(targets, max_length=32, truncation=True) model_inputs["labels"] = labels["input_ids"] return model_inputs
正确预处理示例(处理起始/结束位置)
def preprocess_fn(examples): inputs = [f"question: {q} context: {c}" for q, c in zip(examples["question"], examples["context"])] model_inputs = tokenizer(inputs, max_length=512, truncation=True, padding="max_length") start_positions = [] end_positions = [] for ctx, ans in zip(examples["context"], examples["answers"]): start_char = ans["answer_start"][0] end_char = start_char + len(ans["text"][0]) # 字符位置转token索引 start_idx = tokenizer.char_to_token(0, start_char) end_idx = tokenizer.char_to_token(0, end_char - 1) # 处理无法映射的情况 start_idx = start_idx if start_idx is not None else 0 end_idx = end_idx if end_idx is not None else 0 start_positions.append(start_idx) end_positions.append(end_idx) model_inputs["start_positions"] = start_positions model_inputs["end_positions"] = end_positions return model_inputs
正确推理示例(截取答案区间)
def predict_answer(question, context): inputs = tokenizer( f"question: {question} context: {context}", return_tensors="tf", max_length=512, truncation=True, padding="max_length" ) start_logits, end_logits = model(inputs) start_idx = tf.argmax(start_logits, axis=1).numpy()[0] end_idx = tf.argmax(end_logits, axis=1).numpy()[0] # 从输入token中截取答案并解码 answer_tokens = inputs["input_ids"][0].numpy()[start_idx:end_idx+1] return tokenizer.decode(answer_tokens, skip_special_tokens=True)
内容的提问来源于stack exchange,提问作者Moustafa Banbouk
相关产品推荐
相关产品推荐

