LLM问答微调分词问题求助:基于Hugging Face与ChatQA数据集
问题解决:ChatQA数据集预处理时答案匹配失败及修正方案
错误原因分析
- 精确字符串匹配失效:数据集中的答案是对上下文内容的复述/改写(比如例子中答案是
intermittent invasion of Goryeo,上下文对应内容是intermittently invaded by the Mongol Empire),直接用str.find()做精确匹配必然失败。 - 样本索引不匹配:开启
return_overflowing_tokens=True后,tokenizer会生成比原batch更多的样本(拆分超长上下文),但后续代码仍用原batch的answers索引对应新样本,导致索引错位。 - 位置映射错误:原代码直接将原始字符串的字符索引作为模型的
start_positions/end_positions,但模型需要的是token化后的token索引,而非原始字符位置。
解决步骤
1. 先确认数据集answers字段的完整结构
先打印单个样本的answers字段,确认是否包含位置标注:
print(dataset['train'][0]['answers'])
如果answers中包含start/end等位置信息,直接用这些标注好的位置,避免字符串匹配。
2. 修改预处理函数,解决核心问题
- 关闭
return_overflowing_tokens=True(如果不需要拆分超长样本),或添加逻辑处理样本映射关系; - 用offset_mapping将字符位置转换为token索引;
- 处理答案与上下文不精确匹配的情况,可选择跳过这类样本,或用模糊匹配定位大致位置。
修改后的预处理函数示例
import torch from transformers import RobertaForQuestionAnswering, RobertaTokenizerFast, Trainer, TrainingArguments, DefaultDataCollator from peft import LoraConfig from datasets import load_dataset, DatasetDict from transformers import pipeline dataset = load_dataset("nvidia/ChatQA-Training-Data", "drop") pretrained_model_name = "deepset/roberta-base-squad2" tokenizer = RobertaTokenizerFast.from_pretrained(pretrained_model_name) model = RobertaForQuestionAnswering.from_pretrained(pretrained_model_name) def preprocess_function(examples): questions = [msg[0]['content'] for msg in examples['messages']] contexts = [] for doc in examples['document']: if isinstance(doc, list): contexts.append(doc[0] if len(doc) > 0 else "") else: contexts.append(doc) # 提取答案文本 answers_list = [ans[0] for ans in examples['answers']] inputs = tokenizer( questions, contexts, max_length=512, truncation="only_second", # 先关闭overflowing tokens,避免样本数不匹配 return_overflowing_tokens=False, return_offset_mapping=True, padding="max_length" ) offset_mapping = inputs.pop("offset_mapping") start_positions = [] end_positions = [] for i in range(len(questions)): answer = answers_list[i] context = contexts[i] offsets = offset_mapping[i] if not isinstance(context, str): start_positions.append(0) end_positions.append(0) continue # 尝试小写精确匹配 answer_lower = answer.lower() context_lower = context.lower() start_idx = context_lower.find(answer_lower) # 精确匹配失败则尝试关键词匹配 if start_idx == -1: answer_tokens = answer_lower.split() for token in answer_tokens: start_idx = context_lower.find(token) if start_idx != -1: end_idx = start_idx + len(token) break else: # 完全找不到,标记为无效样本 start_positions.append(0) end_positions.append(0) continue else: end_idx = start_idx + len(answer) # 将字符位置转换为token索引 start_token_idx = None end_token_idx = None for idx, (start_char, end_char) in enumerate(offsets): if start_char <= start_idx and end_char >= start_idx: start_token_idx = idx if start_char <= end_idx and end_char >= end_idx: end_token_idx = idx break # 标记无效样本 if start_token_idx is None or end_token_idx is None: start_positions.append(0) end_positions.append(0) else: start_positions.append(start_token_idx) end_positions.append(end_token_idx) inputs["start_positions"] = start_positions inputs["end_positions"] = end_positions return inputs
3. 过滤无效样本(可选)
预处理后,过滤掉start_positions和end_positions都为0的无效样本:
train_dataset = dataset["train"].map(preprocess_function, batched=True, remove_columns=dataset["train"].column_names) train_dataset = train_dataset.filter(lambda x: x["start_positions"] != 0 or x["end_positions"] != 0)
额外提示
- 若要处理超长上下文,重新开启
return_overflowing_tokens=True时,必须用inputs["overflow_to_sample_mapping"]来映射新样本对应的原样本索引,确保答案对应正确; - 更可靠的方式是使用数据集自带的答案位置标注,如果ChatQA数据集的answers字段包含位置信息,直接使用会比字符串匹配更准确;
- 对于复杂的改写型答案,可以考虑用语义匹配模型来定位上下文对应的片段。
内容的提问来源于stack exchange,提问作者LeGOATJames23
相关产品推荐
相关产品推荐

