跳转至

multiturn_sft_dataset.py — 多轮对话 SFT 数据集

文件路径: verl/utils/dataset/multiturn_sft_dataset.py

文件概述

支持多轮对话的 Supervised Fine-Tuning (SFT) 数据集。每个样本包含多轮 user/assistant 对话,只有 assistant 的回复参与 loss 计算。

核心类

MultiTurnSFTDataset

class MultiTurnSFTDataset(Dataset):
    def __init__(self, parquet_files, tokenizer, config, processor=None, max_samples=-1):
        # 加载 Parquet 文件
        # 提取消息列表
        # 获取 system prompt 和 generation prompt

数据处理流程

def __getitem__(self, item):
    # 1. 构建消息(替换 <image>/<video> 占位符)
    messages = self._build_messages(row_dict)

    # 2. 逐条消息 tokenize
    for i, message in enumerate(messages):
        _input_ids, _loss_mask, _attention_mask, _inputs = self._process_single_message(...)
        # assistant 消息: loss_mask = 1(需要学习)
        # user/system 消息: loss_mask = 0(不学习)

    # 3. 拼接所有消息
    input_ids = torch.cat(input_ids, dim=0)
    loss_mask = torch.cat(loss_mask, dim=0)

    # 4. 处理 padding/truncation
    if self.pad_mode == DatasetPadMode.RIGHT:
        # 右 padding 到 max_length
    elif self.pad_mode == DatasetPadMode.NO_PADDING:
        # 不 padding,截断即可

Loss Mask 机制

消息序列:  [system] [user msg1] [assistant reply1] [user msg2] [assistant reply2]
loss_mask: [0 0 0]  [0 0 0 0]   [0 1 1 1 1 1]      [0 0 0 0]  [0 1 1 1 1 1]
                                  ↑ generation_prompt 部分不计 loss

Sanity Check

def sanity_check(self, input_ids, messages, tools, enable_thinking):
    """验证逐条 tokenize 拼接的结果与整体 tokenize 的结果一致"""
    inputs = processor.apply_chat_template(messages, tokenize=True, ...)
    if not torch.equal(input_ids, inputs["input_ids"].squeeze(0)):
        if self.ignore_input_ids_mismatch:
            logger.warning_once(error_message)
        else:
            raise AssertionError(error_message)

核心功能列表

功能 说明
多轮对话 支持任意多轮 user/assistant 对话
Loss mask 只有 assistant 回复参与训练
多模态 支持图像和视频输入
工具调用 支持 tools 参数
Thinking mode 支持 enable_thinking 模式
Sanity check 验证 tokenization 一致性

与其他模块的关系

  • 依赖 chat_template.py 提取 system/generation prompt
  • 依赖 vision_utils.py 处理图像/视频
  • 依赖 dataset_utils.py 的 DatasetPadMode
  • 依赖 py_functional.py 处理嵌套数据

小结

multiturn_sft_dataset.py 是一个功能完善的多轮对话 SFT 数据集,通过精确的 loss mask 控制确保只学习 assistant 的回复内容,同时支持多模态和工具调用等高级特性。