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 的回复内容,同时支持多模态和工具调用等高级特性。