Prhub

#27488 [KDA] Add CuteDSL Prefill Kernel on SM100

原始 PR 作者 yuan-luo 合并时间 2026-06-10 21:25 文件变更 9 提交数 2 评论 6 代码增减 +3045 / -5

执行摘要

新增 KDA SM100 CuteDSL 预填充内核

KDA的per-channel decay gate导致fp32 exp在真实模型参数下溢出(exp参数最大可达6e4),现有Triton实现输出NaN。同时CuteDSL内核本身快速但受host overhead(每次调用重新分配约200MB张量)瓶颈,导致整体比Triton慢。需要解决数值稳定性并消除host开销。

建议深入阅读__init__.py中的数值修复方案(sub-chunk归一化处理per-channel gate)和_kda_workspace的缓存设计,是高性能算子实现的优秀示例。同时了解如何通过TMA和MMA编程在Blackwell上实现高效sequence-parallel内核。

讨论亮点
  • gemini-code-assist[bot] (高优先级):指出全局工作空间_KDA_WS在多个CUDA流并发时存在数据竞争,建议将流ID加入键。 → 提交者已采纳。
  • gemini-code-assist[bot] (中优先级):建议使用int32创建tok避免int64转换,消除不必要的.long()开销。 → 提交者已采纳并微小调整。
  • BBuf:要求在kernel_h.py头注释中添加仓库链接。 → 提交者已修正注释。

实现拆解

  1. 数值修复:在kda_blackwell/__init__.pychunk_kda_cutedsl中,将两个会溢出的operand(kL, qg2)替换为通过sub-chunk归一化FLA内核chunk_kda_scaled_dot_kkt_fwd计算的gated矩阵,再通过恒等右operand注入不变CuteDSL MMA,确保所有exp指数≤0。
  2. 性能优化:引入模块级_KDA_WS可重用工作空间,按(Hv, K, V, device, dtype, stream)键控,避免每次调用重新分配清零;eye和元数据仅在cu_seqlens对象变化时重算,无同步。
  3. Extend修复:在kda_cutedsl.pyCuteDSLKDAKernel.extend中,将gbetaq的实际token数裁剪,修复cuda-graph padding下形状不匹配导致的崩溃。
  4. 后端集成:在kda_backend.py中移除临时的numerically unstable防护,使CuteDSL预填充默认启用。
  5. 测试与基准:新增test_kda_prefill_cutedsl.py(含真实门控测试)和bench_kda_prefill_cutedsl.py,验证正确性和性能。
文件 模块 状态 重要度
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_kkt_inv_uw.py KKT 逆内核 added 9.18
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_h.py 状态更新核 added 9.18
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_o.py 输出内核 added 9.18
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py 预填充管道 added 8.81
benchmark/bench_linear_attention/bench_kda_prefill_cutedsl.py 基准测试 added 8.78
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/prologue.py 融合序言 added 8.13
python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py 后端集成 modified 8.47
python/sglang/srt/layers/attention/linear/kda_backend.py 后端适配 modified 6.59
test/registered/attention/test_kda_prefill_cutedsl.py 回归测试 added 7.88

关键符号

chunk_kda_cutedsl prepare_metadata _kda_workspace kda_prologue Sm100KdaChunkUWKernel.__call__ Sm100KdaChunkHKernel.__call__ Sm100KdaChunkOKernel.__call__ CuteDSLKDAKernel.extend CuteDSLKDAKernel._ensure_extend_loaded

关键源码片段

python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_kkt_inv_uw.py core-logic

新增 SM100 KKT 逆 +U/W 专有内核,实现 per-channel 门控下的矩阵求逆和 U/W 计算,是预填充管线的核心计算步。

# File: kernel_kkt_inv_uw.py
# SPDX-License-Identifier: Apache-2.0
# KDA SM100 KKT-inverse + U/W kernel (gate folded outside into pre-scaled keys)class Sm100KdaChunkUWKernel:
    def __call__(
        self,
        KL, # k * exp(g_cu - g_cu_last) [T, Hv, K]
        KR, # k * exp(g_cu_last - g_cu) [T, Hv, K]
        KG, # k * exp(g_cu) [T, Hv, K]
        V, U, W, beta, cu_seqlens, chunk_indices, total_chunks, num_sms, stream
    ):
        tma_g2s = cpasync.CopyBulkTensorTileG2SOp()
        tma_s2g = cpasync.CopyBulkTensorTileS2GOp()
        KL_args = self._make_tma_args(KL, self.K_dim, self.num_stages, tma_g2s)
        KR_args = self._make_tma_args(KR, self.K_dim, self.num_stages, tma_g2s)
        KG_args = self._make_tma_args(KG, self.K_dim, self.num_stages, tma_g2s)
        V_args = self._make_tma_args(V, self.V_dim, self.num_stages, tma_g2s)
        U_args = self._make_tma_args(U, self.V_dim, 1, tma_s2g)
        W_args = self._make_tma_args(W, self.K_dim, 1, tma_s2g)
        grid = (num_sms // self.Hv, self.Hv, 1)
        block = (self.num_warps * 32, 1, 1)
        self.kernel(
            KL_args, KR_args, KG_args, V_args, U_args, W_args,
            beta, cu_seqlens, chunk_indices, total_chunks,
        ).launch(grid=grid, block=block, stream=stream)
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py dependency-wiring

编排整个 CuteDSL prefill 管道:prepare_metadata 生成 chunk 索引,_kda_workspace 提供可重用工作空间,chunk_kda_cutedsl 实现数值修复后的 prefill 入口。

# File: __init__.py
# KDA SM100/Blackwell CuteDSL prefill pipelinedef prepare_metadata(cu_seqlens, chunk_size=64):
    cs = cu_seqlens.to(torch.int64)
    seqlens = cs[1:] - cs[:-1]
    nchunks = (seqlens + chunk_size - 1) // chunk_size
    chunk_offsets = torch.zeros(nchunks.numel()+1, dtype=torch.int32, device=dev)
    chunk_offsets[1:] = nchunks.cumsum(0).to(torch.int32)
    total = int(chunk_offsets[-1].item())
    seq_id = torch.repeat_interleave(torch.arange(n, device=dev), nchunks)
    local = torch.arange(total, device=dev) - chunk_offsets[seq_id].to(torch.int64)
    chunk_indices = torch.stack([seq_id.to(torch.int32), local.to(torch.int32)], dim=1)
    return chunk_indices, chunk_offsets, torch.tensor([total], device=dev), total_KDA_WS = {}
def _kda_workspace(q, T, Hv, K, V, cu_seqlens):
    stream = torch.cuda.current_stream(device=q.device).cuda_stream
    key = (Hv, K, V, q.device, q.dtype, stream)
    ws = _KDA_WS.get(key)
    if ws is None or ws['cu'] is not cu_seqlens:
        ci, co, tcs, total = prepare_metadata(cu_seqlens)
        # ... 分配或重用 scratch
    return ws
python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py dependency-wiring

后端集成:CuteDSLKDAKernel 现在支持 extend(prefill),并修复 cuda-graph padding 导致的形状不匹配。

# File: kda_cutedsl.py (CuteDSLKDAKernel class)
def extend(self, q, k, v, g, beta, *, ssm_states, cache_indices, query_start_loc, A_log=None, dt_bias=None, lower_bound=None, **kwargs):
    head_k_dim = k.shape[-1]
    self._ensure_extend_loaded(head_k_dim)
    # L2 normalize and convert to bf16
    q_n = self._l2norm_fn(q[0].contiguous()).to(torch.bfloat16)
    k_n = self._l2norm_fn(k[0].contiguous()).to(torch.bfloat16)
    v_in = v[0].contiguous().to(torch.bfloat16)
    # Trim g and beta to real token count (fix for cuda-graph padding)
    num_tokens = q_n.shape[0]
    g_in = g[0][:num_tokens]
    beta_in = beta[0][:num_tokens]
    # Call the cutedsl prefill pipeline
    o, ht = self._extend_fn(q_n, k_n, v_in, g_in, beta_in, ...)
    return o[None] # restore batch dim

评论区精华

全局工作空间多流并发 正确性

gemini-code-assist[bot] 指出全局字典 _KDA_WS 按 (Hv,K,V,dev,dtype) 键控,在多个 CUDA 流上并发执行 KDA forward 时互相覆盖导致数据损坏。建议包含流 ID。

结论:yuan-luo 采纳建议,在 commit 中将当前 CUDA 流 ID 加入键。 · 已解决

int32 tok 避免 int64 转换 性能

gemini-code-assist[bot] 建议用 int32 创建 tok 张量,避免 cu_seqlens int32 到 int64 的两次转换,减少开销。

结论:yuan-luo 应用了建议并稍作调整。 · 已解决

kernel_h.py 缺少仓库链接 documentation

BBuf 要求在 kernel_h.py 头注释中添加仓库链接。

结论:yuan-luo 修订了注释。 · 已解决

风险与影响

  • 硬件限制:CuteDSL内核仅SM100+(Blackwell)支持,其他GPU自动回退到Triton,可能因未启用新路径而无影响。
  • 工作空间内存:_KDA_WS按配置缓存各种张量,可能在动态shape场景下保留大量内存,但通过键匹配可及时重用,未释放时不增常驻内存。
  • 多流并发:初始实现键未包含流,已被修复,经查无竞争。
  • 数值精度:与token-by-token基准相比,o_err为4.88e-4,全有限,与无修复时的NaN相比已正确。
  • 用户影响:使用Kimi-Linear模型在Blackwell GPU上的推理用户将获得正确结果和显著加速(1.08x-1.52x)。其他模型或GPU无影响(回退)。
  • 系统影响:增加约4.5k行代码,主要在内核和测试目录。构建需要CuteDSL和cuTENSOR支持,已在ci覆盖。
  • 团队影响:为后续KDA decode内核或更多linear attention的Blackwell优化提供了参考模式。
Blackwell 限定 工作空间内存 多流并发

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论