Prhub

#6120 [sglang] feat: Patch sglang to support on-policy distillation teacher

原始 PR 作者 mingruimingrui 合并时间 2026-04-23 16:51 文件变更 2 提交数 2 评论 4 代码增减 +78 / -1

执行摘要

SGLang 引擎支持在线策略蒸馏教师推理

在线策略蒸馏目前仅支持vLLM作为教师推理引擎,为提供更多引擎选择,需要增加对SGLang的支持。PR body明确说明动机:“On-policy distillation currently only supports vLLM as teacher inference engine. This MR adds support for SGlang.”

建议精读此PR,尤其关注_extract_prompt_logprobs_sglang函数的设计——它清晰地展示了如何将SGLang的logprob输出格式对齐vLLM接口,是引擎适配的典范。值得关注的设计决策:采用硬断言而非防御性填充,在测试充分的前提下可保证数据质量,但生产环境中风险较高。建议后续添加单元测试覆盖异常情况。

讨论亮点

Reviewer: gemini-code-assist[bot] 指出_extract_prompt_logprobs_sglanginput_top_logprobs可能短于input_token_logprobs导致IndexError,并建议将对len(ids) == num_prompt_logprobs的硬断言改为填充或截断,以避免因SGLang返回数量不一致时worker崩溃。

Author: mingruimingrui 回应边缘情况处理正确,如果input_token_logprobs为空,内层循环会被跳过。但对第二个关于断言严格性的评论未直接回复。

讨论集中在对logprob数据健壮性的处理上:reviewer建议防御性编程(填充/截断),而作者维持了硬断言方案,仅在主线逻辑中确保一致性。该讨论未完全解决,但PR仍被合并。

实现拆解

  1. 配置验证扩展verl/workers/config/distillation.py):在DistillationTeacherModelConfig._validate_topk_logprobsmatch语句中新增case "sglang"分支,直接pass。原因是SGLang的top_logprobs_num是逐请求参数,无需引擎级启动上限对齐(不像vLLM的max_logprobs)。

  2. 新增logprob提取函数verl/workers/rollout/sglang_rollout/async_sglang_server.py):定义_extract_prompt_logprobs_sglang函数,将SGLang返回的input_token_logprobsinput_top_logprobs重塑为与vLLMextract_prompt_logprobs相同的输出合约——填充result_dict中的prompt_idsprompt_logprobs列表,形状为[sequence_length, max(num_prompt_logprobs, 1)]

# verl/workers/rollout/sglang_rollout/async_sglang_server.py
def _extract_prompt_logprobs_sglang(
    meta_info: dict,
    num_prompt_logprobs: int,
    sequence_length: int,
    result_dict: dict[str, list],
) -> None:
    # 获取SGLang的input_token_logprobs列表,每个元素为(logprob, token_id, _)
    input_token_logprobs = meta_info.get("input_token_logprobs") or []
    # 如果需要top-k,则获取input_top_logprobs
    if num_prompt_logprobs > 0:
        input_top_logprobs = meta_info.get("input_top_logprobs") or []
    prompt_ids_ls: list[list[int]] = []
    prompt_logprobs_ls: list[list[float]] = []
    # 跳过位置0(logprob=None,无预测上下文),与vLLM约定一致
    for position in range(1, len(input_token_logprobs)):
        if num_prompt_logprobs == 0:
            logprob, token_id, _ = input_token_logprobs[position]
            prompt_ids_ls.append([int(token_id)])
            prompt_logprobs_ls.append([float(logprob)])
        else:
            top_entries = input_top_logprobs[position]
            # SGLang按排名返回最佳优先,保持此顺序以匹配vLLM提取器的rank-1槽位
            ids = [int(tok_id) for _, tok_id, _ in top_entries]
            logprobs = [float(logprob) for logprob, _, _ in top_entries]
            assert len(ids) == num_prompt_logprobs, (
                f"SGLang returned {len(ids)} top logprobs at position {position}, "
                f"expected {num_prompt_logprobs}."
            )
            prompt_ids_ls.append(ids)
            prompt_logprobs_ls.append(logprobs)
    # 追加一行填充数据,使总长度等于sequence_length,匹配vLLM约定
    pad_width = max(num_prompt_logprobs, 1)
    prompt_ids_ls.append([0] * pad_width)
    prompt_logprobs_ls.append([0.0] * pad_width)
    # 最终长度断言,确保与消费者期望一致
    assert len(prompt_ids_ls) == sequence_length, (
        f"SGLang prompt_logprobs length ({len(prompt_ids_ls)}) does not match "
        f"sequence length ({sequence_length}); check logprob_start_len=0 invariant."
    )
    result_dict["prompt_ids"] = prompt_ids_ls
    result_dict["prompt_logprobs"] = prompt_logprobs_ls
  1. 集成到SGLang生成请求verl/workers/rollout/sglang_rollout/async_sglang_server.py):在SGLangHttpServer.generate方法中,当sampling_params包含prompt_logprobs时,将其转换为SGLang的return_logprob + logprob_start_len=0 + top_logprobs_num参数,调用生成后在meta_info中提取input_token_logprobs,并调用_extract_prompt_logprobs_sglang填充结果字典。
# 在async_sglang_server.py的SGLangHttpServer.generate方法中(约第410行)
# 解析prompt_logprobs参数(来自蒸馏教师配置)
prompt_logprobs = sampling_params.pop("prompt_logprobs", None)
if prompt_logprobs is not None:
    # 设置SGLang的logprob参数
    sampling_params["return_logprob"] = True
    sampling_params["logprob_start_len"] = 0
    sampling_params["top_logprobs_num"] = prompt_logprobs
    # ... 调用生成后 ...
    # 在获取meta_info后,如果启用了prompt_logprobs,则提取并填充result_dict
    if prompt_logprobs is not None:
        _extract_prompt_logprobs_sglang(
            meta_info, prompt_logprobs, sequence_length, result_dict
        )
  1. 测试配套:PR说明中作者提到回测了examples/on_policy_distillation_trainer/run_qwen_gsm8k.shrun_qwen_gsm8k_megatron.sh,损失/奖励/验证分数与主分支大体一致,并附上了FSDP和Megatron的实验曲线对比图。但本次提交中未包含新的单元测试或CI测试文件。
文件 模块 状态 重要度
verl/workers/rollout/sglang_rollout/async_sglang_server.py 推理引擎 modified 7.55
verl/workers/config/distillation.py 配置层 modified 5.19

关键符号

_extract_prompt_logprobs_sglang

关键源码片段

verl/workers/rollout/sglang_rollout/async_sglang_server.py core-logic

核心变更文件,新增 `_extract_prompt_logprobs_sglang` 函数和 `generate` 方法中的 logprob 提取逻辑,是实现 SGLang 教师支持的主要入口。

# verl/workers/rollout/sglang_rollout/async_sglang_server.py ( 新增函数 )def _extract_prompt_logprobs_sglang(
    meta_info: dict,
    num_prompt_logprobs: int,
    sequence_length: int,
    result_dict: dict[str, list],
) -> None:
    # 从 meta_info 中获取 SGLang 返回的 input_token_logprobs 列表
    input_token_logprobs = meta_info.get("input_token_logprobs") or []
    # 如果请求了 top-k logprobs,则获取 input_top_logprobs
    if num_prompt_logprobs > 0:
        input_top_logprobs = meta_info.get("input_top_logprobs") or []
    prompt_ids_ls: list[list[int]] = []
    prompt_logprobs_ls: list[list[float]] = []
    # 跳过位置 0(logprob=None,无预测上下文),与 vLLM 约定一致
    for position in range(1, len(input_token_logprobs)):
        if num_prompt_logprobs == 0:
            # 仅取采样 token 的 logprob 和 token_id
            logprob, token_id, _ = input_token_logprobs[position]
            prompt_ids_ls.append([int(token_id)])
            prompt_logprobs_ls.append([float(logprob)])
        else:
            top_entries = input_top_logprobs[position]
            # SGLang 返回排名最佳优先,保持此顺序以匹配 vLLM 提取器的 rank-1 槽位
            ids = [int(tok_id) for _, tok_id, _ in top_entries]
            logprobs = [float(logprob) for logprob, _, _ in top_entries]
            # 硬断言确保数量一致(review 中建议改为填充,但作者保持断言)
            assert len(ids) == num_prompt_logprobs, (
                f"SGLang returned {len(ids)} top logprobs at position {position}, "
                f"expected {num_prompt_logprobs}."
            )
            prompt_ids_ls.append(ids)
            prompt_logprobs_ls.append(logprobs)
    # 追加一行填充数据,使总长度等于 sequence_length,匹配 vLLM 约定
    pad_width = max(num_prompt_logprobs, 1)
    prompt_ids_ls.append([0] * pad_width)
    prompt_logprobs_ls.append([0.0] * pad_width)
    # 最终长度断言,确保与消费者期望一致
    assert len(prompt_ids_ls) == sequence_length, (
        f"SGLang prompt_logprobs length ({len(prompt_ids_ls)}) does not match "
        f"sequence length ({sequence_length}); check logprob_start_len=0 invariant."
    )
    # 将结果填充到 result_dict 中,供蒸馏教师管理器消费
    result_dict["prompt_ids"] = prompt_ids_ls
    result_dict["prompt_logprobs"] = prompt_logprobs_ls

评论区精华

logprob 数据健壮性:空列表和 IndexError 风险 正确性

gemini-code-assist[bot] 指出如果 input_token_logprobs 为空,循环被跳过且最终断言会失败;如果 input_top_logprobs 短于 input_token_logprobs,会 IndexError。

结论:作者回复边缘情况处理正确(空列表时循环跳过),但未对 IndexError 场景给出方案。PR 合并时未修改代码,风险保留。 · unresolved

硬断言 vs 填充 / 截断的权衡 设计

gemini-code-assist[bot] 建议将 len(ids) == num_prompt_logprobs 的硬断言改为填充或截断,避免因 SGLang 返回数量不一致时 worker 崩溃。

结论:作者未直接回应,代码维持硬断言。可能认为在正确配置下不应出现不一致,但生产环境存在风险。 · unresolved

风险与影响

  1. 生产环境健壮性风险_extract_prompt_logprobs_sglang中的assert语句如果SGLang返回的top-k数量与num_prompt_logprobs不一致,或最终列表长度与sequence_length不匹配,将导致进程崩溃。虽在测试中未触发,但在不同模型或极端情况下可能暴露。
    - 涉及文件:verl/workers/rollout/sglang_rollout/async_sglang_server.py
  2. 回退风险:该修改影响SGLang rollouter的generate接口,当未启用蒸馏(prompt_logprobs为None)时,行为应与之前完全一致,回归风险较低。
  3. 缺失测试覆盖:目前无针对_extract_prompt_logprobs_sglang的单元测试,也未集成到CI中,边角情况可能不会被及时发现。
  • 用户影响:使用者只需在教师配置中将distillation.teacher_models.teacher_model.inference.name设为sglang即可使用SGLang作为蒸馏推理引擎,无需其他改动。
  • 系统影响:改动仅作用于蒸馏训练场景下的SGLang推理路径,不影响常规rollout或非蒸馏训练。
  • 团队影响:提供了独立于vLLM的教师引擎选项,有助于降低对单一推理引擎的依赖,并可能在多教师蒸馏场景中利用SGLang的性能优势。
缺乏测试覆盖 断言可能在生产环境崩溃

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论