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

LLM问答微调分词问题求助:基于Hugging Face与ChatQA数据集

问题解决:ChatQA数据集预处理时答案匹配失败及修正方案

错误原因分析

  1. 精确字符串匹配失效:数据集中的答案是对上下文内容的复述/改写(比如例子中答案是intermittent invasion of Goryeo,上下文对应内容是intermittently invaded by the Mongol Empire),直接用str.find()做精确匹配必然失败。
  2. 样本索引不匹配:开启return_overflowing_tokens=True后,tokenizer会生成比原batch更多的样本(拆分超长上下文),但后续代码仍用原batch的answers索引对应新样本,导致索引错位。
  3. 位置映射错误:原代码直接将原始字符串的字符索引作为模型的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 09:35:05