Prhub

#29460 Fix SWA cache loc slicing for all attention backends

原始 PR 作者 merrymercy 合并时间 2026-06-28 11:37 文件变更 1 提交数 2 评论 2 代码增减 +9 / -0

执行摘要

修复 SWA cache loc 切片未统一应用的问题

swa_out_cache_loc 在元数据初始化时基于完整(可能填充的)out_cache_loc 计算,而分段 CUDA 图会在 radix_attention.py 中将 out_cache_loc 缩小到每层的真实 token 数,但元数据上的 swa_out_cache_loc 未被同步缩小。当 KVWriteLoc 将两者捆绑传递给 set_kv_buffer 时,形状不匹配导致错误。该修复通过自动切片 swa_loc 以匹配 loc 的长度,解决了所有 8 个后端(flashinfer, triton, flashattention, trtllm_mha, aiter, xpu, torch_native, musa)的这一问题。

值得精读。该 PR 展示了如何用一个集中的 __post_init__ 钩子解决跨后端的共性问题,是设计模式的应用范例。建议关注 KVWriteLoc 数据类的设计思想,以及如何避免在多个后端重复修复。

讨论亮点

无 review 评论。PR 由作者直接合并,说明问题明确、修复方案简洁,团队内部已达成共识。

实现拆解

  1. 定位问题根因:在 python/sglang/srt/mem_cache/memory_pool.pyKVWriteLoc 数据类中,swa_loc 可能比 loc 长,因为 swa_out_cache_loc 在元数据初始化时基于完整 out_cache_loc 计算,而分段 CUDA 图后续会缩小 out_cache_loc

  2. 集中化修复:在 KVWriteLoc__post_init__ 方法中添加逻辑:若 swa_loc 不为 None 且其第一个维度长度与 loc 不同,则截断 swa_locloc 的长度。这确保了所有注意力后端在构造 KVWriteLoc 时自动获得一致的切片。

  3. 移除重复修复:该方案替代了原先仅在 TRT-LLM MHA 后端的临时修复,避免了代码重复,且无需修改其他后端文件。

  4. 测试与配置:本次改动未包含测试文件变更,但 CI 检测到问题后合并;无配置或部署配套改动。

文件 模块 状态 重要度
python/sglang/srt/mem_cache/memory_pool.py 缓存层 modified 6.56

关键符号

__post_init__

关键源码片段

python/sglang/srt/mem_cache/memory_pool.py core-logic

核心修复文件:在 KVWriteLoc 数据类中新增 __post_init__ 方法,自动切片 swa_loc 以匹配 loc 长度,覆盖所有注意力后端。

@dataclass
class KVWriteLoc:
    """Write target(s) for ``KVCache.set_kv_buffer``.    ``loc`` is the full-pool write location; ``swa_loc`` is the pre-translated
    full->SWA location for hybrid SWA pools (``None`` otherwise). Bundling them
    lets a backend issue one ``set_kv_buffer`` call regardless of pool type.
    """
​
    loc: torch.Tensor
    swa_loc: Optional[torch.Tensor] = None
​
    def __post_init__(self):
        # swa_out_cache_loc 在 metadata 初始化时基于完整(可能填充的)
        # out_cache_loc 计算。Piecewise CUDA graphs 后续会将 out_cache_loc
        # 缩小到每层的真实 token 数,导致 swa_loc 可能比 loc 长。
        # 这里将它们切片对齐,因为两者 token 顺序一致。
        if self.swa_loc is not None and self.swa_loc.shape[0] != self.loc.shape[0]:
            self.swa_loc = self.swa_loc[: self.loc.shape[0]]

评论区精华

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

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

风险与影响

风险较低。改动仅 9 行,逻辑简单:在 __post_init__ 中对 swa_loc 进行形状检查并切片。由于切片操作仅当形状不匹配时执行,且切片方向为 [:loc.shape[0]](即截断到 loc 长度),不会导致越界或数据丢失(因为 swa_loc 原本就包含所有 token 的 SWA 位置,只是可能包含额外填充)。但需要注意:如果有后端依赖于 swa_loc 保持原始长度(例如后续有重新填充操作),则此切片可能改变行为。不过从 PR body 看,原始设计就应保持长度一致,因此回归风险小。

影响范围:所有使用 SWA 且涉及分段 CUDA 图的推理场景,尤其是多注意力后端的混合部署。影响程度:中等——修复了可能导致崩溃或错误输出的 shape mismatch bug,但只影响已启用 SWA 且触发分段图的特定路径。用户无感知,但内部正确性提升。

核心路径变更(内存池 KVWriteLoc) 缺少测试覆盖(未添加对应测试)

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论