跳转至

chat_template.py — 聊天模板处理

文件路径: verl/utils/chat_template.py

文件概述

处理聊天模型的模板(Chat Template),提取 system prompt 和 generation prompt 的 token 序列。在多轮对话 SFT 训练中,需要准确识别模板中的各个部分以正确设置 loss mask。

核心函数详解

1. 提取 system prompt 和 generation prompt

def extract_system_prompt_and_generation(tokenizer):
    # 编码一条消息
    token1 = tokenizer.apply_chat_template(
        [{"role": "user", "content": ""}], add_generation_prompt=False, tokenize=True
    )
    # 编码两条消息
    token2 = tokenizer.apply_chat_template(
        [{"role": "user", "content": ""}] * 2, add_generation_prompt=False, tokenize=True
    )
    # system prompt = 一条消息的前缀 - 两条消息共有的部分
    system_prompt = token1[:-(len(token2) - len(token1))]

    # generation prompt = 带生成提示的编码 - 不带的编码
    token3 = tokenizer.apply_chat_template(
        [{"role": "user", "content": ""}], add_generation_prompt=True, tokenize=True
    )
    generate_prompt = token3[len(token1):]

    return system_prompt, generate_prompt

原理: 通过对比编码一条和两条消息的差异,推断出 system prompt(开头的固定部分)。通过对比是否加 generation prompt,推断出生成提示(如 <|im_start|>assistant\n)。

用途说明

在多轮对话 SFT 中: - system prompt 的 token 不参与 loss 计算(loss mask = 0) - generation prompt 同样不参与 loss 计算 - 只有 assistant 回复的内容参与 loss 计算

与其他模块的关系

  • 被 dataset/multiturn_sft_dataset.py 用来设置 loss mask
  • 依赖 tokenizer.py 中的 normalize_token_ids

小结

chat_template.py 通过巧妙的差分方法自动提取模板组件,让多轮对话训练能正确区分"需要学习的"和"不需要学习的"token。