Prhub

#31624 [Refactor] Move output logprob processing into the logprob_processor layer

原始 PR 作者 hnyls2002 合并时间 2026-07-18 13:33 文件变更 2 提交数 5 评论 6 代码增减 +226 / -163

执行摘要

将输出 logprob 处理逻辑从 Sampler 移至 logprob_processor 层

PR #20071 引入了 logprob_processor 层处理输入 logprob,本 PR 完成输出侧的迁移,实现 logprob 处理的完整分离,提高代码组织性和可维护性。

建议仔细审核 OutputLogprobProcessorattach_logprobs_to_output 与原有 _attach_logprobs_to_output 的一致性,以及 preprocess_fn 注入逻辑。适合对 logprob 处理感兴趣的开发人员阅读,了解重构模式。

讨论亮点

本 PR 无 review 评论或讨论线程。

实现拆解

  1. get_token_ids_logprobs_batch_optimized 函数从 sampler.py 移动到 logprob_processor.py,并保留其优化性能的向量化实现。
  2. logprob_processor.py 中添加 OutputLogprobsResult 数据类,提供类似 InputLogprobsResult 的结构。
  3. 提取 OutputLogprobProcessor 类,包含 write_to 方法(将结果写入 LogitsProcessorOutput)和 attach_logprobs_to_output 方法,替代 Sampler 中原有的 _attach_logprobs_to_output 方法。
  4. 在 Sampler 的 __init__ 中创建 OutputLogprobProcessor 实例,并在 forward 中调用其方法。同时将 _preprocess_logits 作为 preprocess_fn 注入,以便在处理器中使用。
  5. 保留 Sampler.compute_logprobs_only 作为外层代理,保持外部调用兼容性。未增加测试用例,但现有测试应覆盖功能。
文件 模块 状态 重要度
python/sglang/srt/layers/logprob_processor.py logprob 层 modified 8.87
python/sglang/srt/layers/sampler.py 采样器 modified 8.17

关键符号

get_token_ids_logprobs_batch_optimized OutputLogprobsResult write_to OutputLogprobProcessor attach_logprobs_to_output compute_logprobs_only

关键源码片段

python/sglang/srt/layers/logprob_processor.py core-logic

主要变更文件,新增 OutputLogprobProcessor 类、OutputLogprobsResult 数据类和 get_token_ids_logprobs_batch_optimized 函数,完成 output logprob 的集中处理。

def get_token_ids_logprobs_batch_optimized(
    logprobs: torch.Tensor,
    token_ids_logprobs: List[List[int]],
) -> Tuple[List, List]:
    """向量化批量处理 token ID logprobs 提取,使用单次 GPU gather 替代逐项 gather,适合大批量。"""
    batch_size = len(token_ids_logprobs)
    device = logprobs.device
​
    # 计算每个请求的长度,将 None 视为空列表
    token_lengths = torch.tensor(
        [len(token_ids or []) for token_ids in token_ids_logprobs], device=device
    )
    total_tokens = int(token_lengths.sum().item())
​
    if total_tokens == 0:
        return [logprobs.new_empty(0) for _ in token_ids_logprobs], [
            [] for _ in token_ids_logprobs
        ]
​
    # 构建扁平索引:row_indices 和 col_indices
    row_indices = torch.repeat_interleave(
        torch.arange(batch_size, device=device), token_lengths
    )
    col_indices = torch.tensor(
        [
            token_id
            for token_ids in token_ids_logprobs
            for token_id in (token_ids or [])
        ],
        device=device,
        dtype=torch.long,
    )
​
    # 单次 gather 操作
    gathered_logprobs = logprobs[row_indices, col_indices]
​
    # 根据每个请求的长度 split 回去
    values = torch.split(gathered_logprobs, token_lengths.tolist())
    indices = torch.split(col_indices, token_lengths.tolist())
    return list(values), [split.tolist() for split in indices]

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险较低,因为 PR 声称字节级一致(除两处声明替换)。但潜在风险包括:对 sampling_info.device 的替换为 batch_next_token_ids.device 可能在不同设备场景下有差异;attach_logprobs_to_output 方法现在在 OutputLogprobProcessor 中,需要确保所有调用路径正确;由于涉及 Sampler 核心逻辑,若 preprocess_fn 注入有误可能影响 logprob 计算。回归风险可控,现有 CI 测试已通过。

影响开发者:logprob 处理逻辑更集中,便于后续维护和扩展。对用户:功能无变化,API 保持兼容。对系统:性能理论上无变化,但未来优化 logprob 处理时 logprob_processor 成为唯一入口。

核心路径变更 外部 API 兼容需注意

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论