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

如何格式化数据以微调HuggingFace DPR模型(PyTorch环境)

解决HuggingFace DPR微调的数据集格式与负样本机制问题

核心概念澄清

HuggingFace的DPR问题编码器+上下文编码器就是FB原始仓库里的「检索器」,二者完全对应,只是拆分了两个组件方便单独调用(比如单独生成问题/上下文embedding)。

数据集格式要求

HuggingFace的DPR Trainer对输入格式做了简化,不需要复杂的JSON结构,你只需要构造包含以下字段的数据集(用HuggingFace Dataset 对象即可,不用本地JSON):

  • question: 字符串类型,你的问题文本
  • positive_ctxs: 列表类型,每个元素是一个字典,包含text字段(对应该问题的正确上下文)
  • (可选)hard_negatives: 列表类型,每个元素是一个字典,包含text字段(手动标注的难负样本上下文)

最小可行的数据集示例(用Python构造):

from datasets import Dataset

train_data = [
    {
        "question": "PyTorch怎么定义自定义数据集?",
        "positive_ctxs": [{"text": "可以继承torch.utils.data.Dataset类,重写__len__和__getitem__方法"}]
    },
    {
        "question": "DPR的in-batch negatives是什么?",
        "positive_ctxs": [{"text": "指批次内其他样本的正上下文自动作为当前样本的负样本,用于训练时的对比学习"}]
    }
]

dataset = Dataset.from_list(train_data)

In-batch Negatives机制说明

HuggingFace的DPR Trainer默认自动启用in-batch negatives,不需要你手动给每个问题批量设置负样本:训练时,同一个batch里其他问题对应的正上下文,会被当作当前问题的负样本参与对比损失计算。这种方式效率远高于给每个问题单独配负样本,完全适配高效批处理。

如果你的数据集里有hard_negatives字段,Trainer会同时使用in-batch negatives和手动难负样本,进一步提升模型效果。

与FB原始仓库的差异

FB的原始DPR仓库要求JSON格式是因为它的预处理逻辑更复杂,支持更多自定义配置;而HuggingFace的实现做了封装,把数据处理和训练逻辑整合到了DPRTrainer里,你只需要传入标准的Dataset对象即可,不用关心底层的JSON解析。

微调代码示例

from transformers import DPRQuestionEncoder, DPRContextEncoder, DPRTrainer, DPRTrainingArguments
from datasets import Dataset

# 加载预训练模型
question_encoder = DPRQuestionEncoder.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
context_encoder = DPRContextEncoder.from_pretrained("facebook/dpr-ctx_encoder-single-nq-base")

# 构造训练数据(如上述示例)
train_data = [
    {"question": "你的问题1", "positive_ctxs": [{"text": "对应正上下文1"}]},
    {"question": "你的问题2", "positive_ctxs": [{"text": "对应正上下文2"}]}
]
dataset = Dataset.from_list(train_data)

# 设置训练参数
training_args = DPRTrainingArguments(
    output_dir="./dpr-finetuned",
    per_device_train_batch_size=8,
    num_train_epochs=3,
    logging_dir="./logs",
)

# 初始化Trainer并训练
trainer = DPRTrainer(
    question_encoder=question_encoder,
    context_encoder=context_encoder,
    args=training_args,
    train_dataset=dataset,
)

trainer.train()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 14:35:43