执行摘要
- 一句话:新增KDA SM100 CuteDSL预填充内核
- 推荐动作:建议深入阅读
__init__.py中的数值修复方案(sub-chunk归一化处理per-channel gate)和_kda_workspace的缓存设计,是高性能算子实现的优秀示例。同时了解如何通过TMA和MMA编程在Blackwell上实现高效sequence-parallel内核。
功能与动机
KDA的per-channel decay gate导致fp32 exp在真实模型参数下溢出(exp参数最大可达6e4),现有Triton实现输出NaN。同时CuteDSL内核本身快速但受host overhead(每次调用重新分配约200MB张量)瓶颈,导致整体比Triton慢。需要解决数值稳定性并消除host开销。
实现拆解
- 数值修复:在
kda_blackwell/__init__.py的chunk_kda_cutedsl中,将两个会溢出的operand(kL, qg2)替换为通过sub-chunk归一化FLA内核chunk_kda_scaled_dot_kkt_fwd计算的gated矩阵,再通过恒等右operand注入不变CuteDSL MMA,确保所有exp指数≤0。
- 性能优化:引入模块级
_KDA_WS可重用工作空间,按(Hv, K, V, device, dtype, stream)键控,避免每次调用重新分配清零;eye和元数据仅在cu_seqlens对象变化时重算,无同步。
- Extend修复:在
kda_cutedsl.py的CuteDSLKDAKernel.extend中,将g和beta按q的实际token数裁剪,修复cuda-graph padding下形状不匹配导致的崩溃。
- 后端集成:在
kda_backend.py中移除临时的numerically unstable防护,使CuteDSL预填充默认启用。
- 测试与基准:新增
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逆内核;类别 source;类型 core-logic;符号 Sm100KdaChunkUWKernel, init, _make_tma_args, call): 新增SM100 KKT逆+U/W专有内核,实现per-channel门控下的矩阵求逆和U/W计算,是预填充管线的核心计算步。
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_h.py(模块 状态更新核;类别 source;类型 core-logic;符号 Sm100KdaChunkHKernel, init, _make_bf16_tma_args, _make_h_tma_args): 新增SM100 KDA chunk recurrent-state更新内核,处理per-column状态衰减,使用预缩放的kg张量。
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/kernel_o.py(模块 输出内核;类别 source;类型 core-logic;符号 Sm100KdaChunkOKernel, init, _make_bf16_tma_args, _make_h_tma_args): 新增SM100 KDA输出内核,使用预缩放qg/qg2/kg张量,无需g_cu输入。
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/__init__.py(模块 预填充管道;类别 source;类型 dependency-wiring;符号 prepare_metadata, _kda_workspace, chunk_kda_cutedsl): 编排整个CuteDSL prefill管道:prepare_metadata生成chunk索引,_kda_workspace提供可重用工作空间,chunk_kda_cutedsl实现数值修复后的prefill入口。
benchmark/bench_linear_attention/bench_kda_prefill_cutedsl.py(模块 基准测试;类别 source;类型 dependency-wiring;符号 _l2norm, kda_flops, kda_bytes, make_inputs): 提供CuteDSL vs Triton的性能和正确性基准测试,验证加速比和数值误差。
python/sglang/srt/layers/attention/linear/kernels/kda_blackwell/prologue.py(模块 融合序言;类别 source;类型 core-logic;符号 _kda_prologue_kernel, kda_prologue): Triton融合序言,计算per-chunk cumsum和五组预缩放张量,弥合per-channel门控与CuteDSL MMA的接口。
python/sglang/srt/layers/attention/linear/kernels/kda_cutedsl.py(模块 后端集成;类别 source;类型 dependency-wiring;符号 _is_blackwell, init, _ensure_extend_loaded, extend): 后端集成:CuteDSLKDAKernel现在支持extend(prefill),并修复cuda-graph padding导致的形状不匹配。
python/sglang/srt/layers/attention/linear/kda_backend.py(模块 后端适配;类别 source;类型 dependency-wiring): 移除临时数值不稳定防护,使CuteDSL prefill默认启用。
test/registered/attention/test_kda_prefill_cutedsl.py(模块 回归测试;类别 test;类型 test-coverage;符号 _l2norm, test_kda_chunk_cutedsl_correctness, test_kda_chunk_cutedsl_internal_gate_activation, test_kda_chunk_cutedsl_realistic_gate): 新增回归测试,包含真实门控测试(可检测之前未发现的数值溢出),以及varlen正确性和内部门控激活测试。
关键符号: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
新增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
编排整个CuteDSL prefill管道:prepare_metadata生成chunk索引,_kda_workspace提供可重用工作空间,chunk_kda_cutedsl实现数值修复后的prefill入口。
# File: __init__.py
# KDA SM100/Blackwell CuteDSL prefill pipeline
def 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
后端集成: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
评论区精华
风险与影响
- 风险:
- 硬件限制: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限定, 工作空间内存, 多流并发
关联脉络
参与讨论