Prhub

#31904 [KDA] Fix mixed exponent bases in Triton chunk prefill

原始 PR 作者 yuan-luo 合并时间 2026-07-22 16:03 文件变更 6 提交数 4 评论 4 代码增减 +234 / -65

执行摘要

修复 KDA chunk prefill 中 exp2/exp 混合导致的衰减错误

KDA 定义每个通道的衰减门在自然对数空间,累计门 $G_t = \sum g_s$,从位置 j 到 i 的衰减应为 $\exp(G_i-G_j)$。但部分 intra-chunk kernel 使用 exp2 直接计算,而其他路径使用自然 exp,导致混合指数基底,实际计算的衰减弱于理论值。

该 PR 值得精读,因为它清晰地展示了如何通过规范指数域来修复不匹配 bug,同时采用编译选项保证兼容性。设计决策(如 gk_scale 避免物化、编译时 use_exp2)可作为类似问题的参考。

讨论亮点

Reviewer kaixih 建议在 chunk_kda_scaled_dot_kkt_fwd 中添加 gk_scale 参数,直接加载时乘以 RCP_LN2,避免物化完整的转换张量。同时建议在测试中添加非融合路径的覆盖案例(force non-fused path),通过大量 chunk 使 _small_grid=False 来触发独立对角化和重计算 kernel。Author yuan-luo 采纳了两项建议,体现在第二个和第三个 commit 中。

实现拆解

  1. 在 KDA 的累计门生产处(kda_gate_chunk_cumsumchunk_local_cumsum)加入 RCP_LN2 = 1.4426950216293335(log2(e)),将自然对数门转换为 log2 空间。

  2. 所有 KDA chunk kernel(scaled KKT/QK 构建、独立和融合的 W/U/KG 重计算、chunk 内融合求解/重计算、状态更新、输出构建)统一改用 exp2

  3. 在共享状态 kernel chunk_gated_delta_rule_fwd_h 中添加编译选项 use_exp2,默认 False 保持 GDN 原行为,KDA 传入 True

  4. chunk_kda_scaled_dot_kkt_fwd 中添加 gk_scale 参数,避免物化转换后的门张量,直接在 kernel 加载时乘以 RCP_LN2

  5. 新增独立 PyTorch 递归参考实现,验证固定长度和变长场景下,chunk prefill 的输出和最终状态与理论值误差小于 1%。

文件 模块 状态 重要度
test/registered/attention/test_kda_kernels.py 回归测试 modified 6.96
python/sglang/kernels/ops/attention/fla/kda.py KDA 核心 modified 6.21
python/sglang/kernels/ops/attention/fla/chunk_delta_h.py 状态 kernel modified 5.56
python/sglang/kernels/ops/attention/fla/chunk_intra.py Intra kernel modified 5.25
python/sglang/srt/hardware_backend/xpu/kernels/fla/chunk_delta_h.py XPU 适配 modified 5.87

关键符号

kda_gate_chunk_cumsum chunk_local_cumsum chunk_kda_scaled_dot_kkt_fwd chunk_kda_fwd_intra chunk_gated_delta_rule_fwd_h

关键源码片段

test/registered/attention/test_kda_kernels.py test-coverage

新增 `TestKDAChunkExponentDomain` 测试类,覆盖固定长度、变长、融合 / 非融合多条代码路径,回归验证修复的正确性。

@unittest.skipIf(not torch.cuda.is_available(), "Test requires CUDA")
class TestKDAChunkExponentDomain(CustomTestCase):
    """Guard KDA prefill against mixing natural-log gates with exp2 kernels."""
​
    @staticmethod
    def _naive_recurrent(q, k, v, g, beta, initial_state, lengths):
        # Pure-PyTorch 逐 token 递归,作为无 bug 参考
        q, k, v, g, beta = (tensor.float() for tensor in (q, k, v, g, beta))
        scale = q.shape[-1] ** -0.5
        output = torch.empty_like(v)
        final_state = initial_state.float().clone()
​
        offset = 0
        for seq_idx, length in enumerate(lengths):
            state = final_state[seq_idx]
            for i in range(offset, offset + length):
                state = state * g[0, i].exp().unsqueeze(-2) # 自然 exp
                residual = v[0, i] - torch.einsum("hvk,hk->hv", state, k[0, i])
                state = state + torch.einsum(
                    "hv,hk->hvk",
                    residual * beta[0, i, :, None],
                    k[0, i],
                )
                output[0, i] = (
                    torch.einsum("hvk,hk->hv", state, q[0, i]) * scale
                )
            final_state[seq_idx] = state
            offset += length
        return output, final_state
​
    @torch.inference_mode()
    def test_chunk_prefill_matches_natural_exp_recurrence(self):
        device = get_device()
        dtype = torch.bfloat16
        num_heads, head_dim = 2, 64
​
        # 测试组合 : (lengths, use_varlen, fuse_gate)
        cases = (
            ([129], False, False), # 固定长度,pre-activated gate
            ([15, 16, 17, 63, 65], True, True), # varlen,fused raw gate
            ([2] * 129, True, False), # 强制非融合路径(_small_grid=False)
        )
        for lengths, use_varlen, fuse_gate in cases:
            with self.subTest(lengths=lengths, use_varlen=use_varlen, fuse_gate=fuse_gate):
                # 构造输入 ...(省略具体构造)
                # 运行 chunk_kda 与 naive_recurrent,断言 relative_rmse < 0.01
python/sglang/kernels/ops/attention/fla/kda.py core-logic

核心修改文件:定义 RCP_LN2,修改 scaled_dot_kkt kernel 和入口函数,统一使用 exp2,添加 gk_scale 参数。

# 位于 python/sglang/kernels/ops/attention/fla/kda.py# log2(e) 的 FP32 近似值,用于将自然对数门转换到 log2 空间
RCP_LN2 = 1.4426950216293335def chunk_kda_scaled_dot_kkt_fwd(
    q, k, gk=None, beta=None, scale=None,
    gk_scale=1.0, # <-- 新增参数,避免物化转换张量
    cu_seqlens=None, chunk_indices=None, output_dtype=None,
):
    # ... 省略形状推导
    # 将 gk_scale 传递给 Triton kernel
    grid = ...
    chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter[grid](
        ..., gk_scale=gk_scale, ...
    )# 在 kernel 内部 ( 以 _intra_sub_inter 为例 ):
# ...
b_gn = (
    tl.load(g + ...) * gk_scale # 加载时转换
)
# 所有 gate 相关的 exp 替换为 exp2
b_k = tl.load(...) * exp2(b_g - b_gn[None, :])

评论区精华

使用 gk_scale 避免物化转换张量 设计

Reviewer `kaixih` 建议在 `chunk_kda_scaled_dot_kkt_fwd` 中添加 `gk_scale` 参数,在 kernel 加载门时直接乘以 `RCP_LN2`,避免创建完整尺寸的转换副本。

结论:Author 采纳,添加 `gk_scale` 参数并传递至 Triton kernel,Blackwell 后端同样使用。 · 已解决

增加非融合路径的测试覆盖 测试

Reviewer `kaixih` 指出应添加 `([2] * 129, True, False)` 案例,使 `_small_grid=False` 以触发非融合对角化和重计算 kernel。

结论:Author 添加该案例,并在测试中说明其目的。 · 已解决

风险与影响

修改集中在 KDA chunk prefill 路径,解码路径不受影响;GDN 通过 use_exp2 编译选项隔离,默认行为不变;Blackwell 后端通过 gk_scale 同步适配。但未进行端到端模型精度基准测试,kernel 级回归测试覆盖了核心路径,可能遗漏边界条件。XPU 适配代码已同步调整但无法在 CI 中自动验证。

影响使用 KDA chunk prefill 的模型(如 DeepSeek V2/V3 等),能纠正精度错误,提升生成质量。用户无需任何配置变更。不影响显式单步解码或 GDN。预期性能中性或略有提升。

KDA 核心路径变更 缺少端到端模型精度基准 XPU 路径需人工验证

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论