Prhub

#31312 Fix LongCat n-gram token-table crashes on padded batches

原始 PR 作者 whn09 合并时间 2026-07-22 20:24 文件变更 2 提交数 5 评论 2 代码增减 +22 / -12

执行摘要

修复 LongCat n-gram 在填充批次上的崩溃

在 cuda-graph 或 EP eager 填充模式下,seq_lensnext_token_idsreq_pool_indices 被填充到 graph batch size,而 batch_size 是真实请求数,导致逐请求张量索引不一致,进而触发 RuntimeError: The expanded size of the tensor (N) must match the existing size (M)cudaErrorIllegalAddress 崩溃。

该 PR 值得合并,但建议后续补充单元测试覆盖填充批次场景,以及采纳 Gemini 建议添加 _ng_bs 守卫到 update_ngram_token_table_after_sampling 中以增强鲁棒性。

讨论亮点

Gemini Code Assist 机器人建议在 update_ngram_token_table_after_sampling 中增加与 NgramEmbedding.forward 类似的 _ng_bs 守卫,以防止 batch_size 超过输出张量分配大小导致越界。该建议未在后续讨论中被接受或实施,但 PR 作者未进一步回应,审阅者 BBuf 已批准该 PR。

实现拆解

  1. 修复 update_ngram_token_table_after_samplingngram_embedding_manager.py:将 seq_lensnext_token_idsreq_pool_indices 及输出张量切片至 batch_size,确保只有真实请求更新 token table,填充行不会污染表格。
  2. 修复 NgramEmbedding.forwardlayers/n_gram_embedding.py:引入 _ng_bs 守卫变量,取 min(batch_size, req_lens.shape[0], column_starts.shape[0]),驱动 cumsum 和 n-gram kernel 的请求循环,防止 kernel 越界读取 column_starts/req_lens
  3. 测试验证:在 meituan-longcat/LongCat-2.0-FP8、tp16/ep16 over EFA 上验证,并发数 1-64 下修复后不再崩溃且输出正确。
文件 模块 状态 重要度
python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py 嵌入管理 modified 6.04
python/sglang/srt/layers/n_gram_embedding.py 嵌入层 modified 5.98

关键符号

update_ngram_token_table_after_sampling NgramEmbedding.forward

关键源码片段

python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py data-contract

修复了更新 token table 时未对 padded 张量切片的问题,是导致 RuntimeError 的直接原因。

# 文件 : python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py
# 函数 : update_ngram_token_table_after_samplingdef update_ngram_token_table_after_sampling(
    *,
    ngram_embedding_info,
    next_token_ids: torch.Tensor,
    req_pool_indices: torch.Tensor,
    seq_lens: torch.Tensor,
    batch_size: int,
) -> bool:
    """使用采样后的 token 更新 ngram token table。"""
    skip_token_table_update = ngram_embedding_info.skip_token_table_update
    if skip_token_table_update is not None:
        # 跳过未完成预填充的请求
        indices = (~skip_token_table_update).nonzero(as_tuple=True)[0]
        if indices.numel() == 0:
            return False
        update_token_table(
            ne_token_table=ngram_embedding_info.token_table,
            tokens=next_token_ids[indices].to(torch.int32),
            row_indices=req_pool_indices[indices],
            column_starts=seq_lens[indices].to(torch.int32),
            req_lens=torch.ones(indices.numel(), dtype=torch.int32, device=next_token_ids.device),
            ignore_tokens=None,
        )
        return True
​
    # 修复 : seq_lens / next_token_ids / req_pool_indices 可能被填充到 cuda-graph
    # 的 batch size, 而 batch_size 是真实请求数。切片到 batch_size 可防止
    # 填充行污染 token table ( 同时确保形状匹配 )。
    ngram_embedding_info.out_column_starts[:batch_size] = seq_lens[:batch_size]
    ngram_embedding_info.out_req_lens[:batch_size] = 1
    update_token_table(
        ne_token_table=ngram_embedding_info.token_table,
        tokens=next_token_ids[:batch_size].to(torch.int32),
        row_indices=req_pool_indices[:batch_size],
        column_starts=ngram_embedding_info.out_column_starts[:batch_size],
        req_lens=ngram_embedding_info.out_req_lens[:batch_size],
        ignore_tokens=None,
    )
    return True
python/sglang/srt/layers/n_gram_embedding.py core-logic

修复了 n-gram kernel 读取越界的问题,通过引入 _ng_bs 守卫确保索引不超出数组范围。

# 文件 : python/sglang/srt/layers/n_gram_embedding.py
# 类 : NgramEmbeddingdef forward(self, input_ids: torch.Tensor, forward_batch: ForwardBatch):
    if (
        forward_batch.forward_mode.is_extend()
        or forward_batch.forward_mode.is_decode()
    ):
        ngram_embedding_info = forward_batch.ngram_embedding_info
        # 修复 : ngram_info 数组可能比 forward_batch.batch_size 短
        # ( 混合 / 重叠批次 )。基于数组长度驱动请求循环,防止 kernel 越界
        # 读取 column_starts / req_lens (cudaErrorIllegalAddress)。
        _ng_bs = min(
            forward_batch.batch_size,
            ngram_embedding_info.req_lens.shape[0],
            ngram_embedding_info.column_starts.shape[0],
        )
        torch.cumsum(
            ngram_embedding_info.req_lens[:_ng_bs],
            dim=0,
            dtype=torch.int32,
            out=self.exclusive_req_len_sums[1 : 1 + _ng_bs],
        )
        compute_n_gram_ids(
            ne_n=self.over_embedding_n,
            ne_k=self.over_embedding_k,
            ne_weights=self.oe_weights,
            ne_mods=self.oe_mods,
            tokens=input_ids.to(torch.int32),
            exclusive_ne_embedder_size_sums=self.exclusive_oe_embedder_size_sums,
            exclusive_req_len_sums=self.exclusive_req_len_sums[: _ng_bs + 1],
            ne_token_table=ngram_embedding_info.token_table,
            row_indices=forward_batch.req_pool_indices[:_ng_bs],
            column_starts=ngram_embedding_info.column_starts[:_ng_bs],
            n_gram_ids=self.oe_n_gram_ids[: len(input_ids)],
            eos_token_id=self.eos_token_id,
        )
    # ... 后续代码不变

评论区精华

在 update_ngram_token_table_after_sampling 中添加 _ng_bs 守卫 正确性

Gemini Code Assist 机器人建议在 `update_ngram_token_table_after_sampling` 中也引入类似 `NgramEmbedding.forward` 中的 `_ng_bs` 守卫,以防止 `batch_size` 超过输出张量分配大小导致越界。

结论:未采纳:PR 作者未回应,但审阅者 BBuf 已批准 PR。 · unresolved

风险与影响

最低风险:修复范围明确(仅 LongCat n-gram 路径),改动量小(+22/-12 行),且已在真实模型上验证。但未添加单元测试,回归依赖于现有集成测试。

影响范围较小:仅影响使用 LongCat-2.0 n-gram embedding 并且在 cuda-graph 或 EP eager 模式下填充批次的用户。修复后这些用户将避免 decode 崩溃,输出正确。

缺少测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论