
基于 TileLang 的 HISA 分层稀疏注意力 Prefill 索引器实现剖析【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang本文档基于 examples/dsa_hisa 目录下的完整实现深入讲解如何用 TileLang 为 DeepSeek 系稀疏注意力Sparse Attention的 indexer索引器构建两阶段分层预填充prefill流水线先对 K 做粗粒度块级 mean-pooling 评分选出候选块再对候选块内的原始 token 做细粒度 MQA 评分选出最终 top-k。读完本文你将掌握 HISA 四颗 TileLang kernelfp8 块均值池化、fp8 块级 MQA、logits 就地清洗、fp8 块稀疏 MQA的接口契约、Tile 级 GEMM 实现细节、端到端编排与正确性/性能测试方法可直接在 examples/dsa_hisa 中运行验证。什么是 HISAHISAHIerarchical Sparse Attention是针对 DeepSeek 稀疏注意力的一种即插即用plug-and-play的 indexer 替换方案。与直接在整条 token 序列上做扁平扫描不同HISA 将搜索路径重写为两阶段分层过程参见 READMEStage 1 — 粗粒度块级选择把 N 个 K token 按k_block_size分组成 pool block对每个块做 mean-pooling然后让每个 query 对所有 pool block 打分选出每 query 的block_topk个块Stage 2 — 细粒度 token 级打分对每个 query只在其选中的块内的原始 token 上做全分辨率 MQAMulti-Query Attention再选出每 query 的topk_tokens个 token。这种先粗筛块、再细筛 token的层级结构把一次性的大范围 top-k 搜索拆解成两个小规模问题在保持召回质量的同时显著缩小候选规模是 HISA 的核心思想。该实现以k_block_size128作为 DeepSeek-V3.2 场景下的默认 pool block 尺寸见 hisa.py 的test_hisa默认参数也可按需调整。目录文件总览文件步骤作用fp8_block_mean_pooling.py1.1把原始 K 做 fp8 块均值池化含逐块 f32 scalepool_mqa_fp8.py1.2fp8×fp8 计算Q · pooled_K得分 → 每个 (query, pool block) 一个 logitclean_and_maintain_logits.py1.3对 stage-1 logits 就地掩码越界置-inf首/末有效块置infblock_sparse_mqa_fp8.py2.1在block_topk个选中块的原始 token 上做 fp8×fp8 细粒度打分hisa.py—端到端编排四颗 kernel 两次torch.topk 索引换算后处理tilelang_utils.py—公共工具cu_seqlens→ 逐 queryks/ke换算、fp8 量化辅助、相似度断言等每个 kernel 文件都自带一个test_*测试入口其流程统一为(a) 运行 kernel 与 torch 参考实现(b) 用torch.testing.assert_close断言一致(c) 打印 kernel 延迟。hisa.py中的test_hisa则运行完整流水线、校验输出索引掩码不变量并打印端到端延迟。Kernel 1.1fp8 块均值池化fp8_block_mean_pooling.py函数与接口函数fp8_native_block_mean_poolingTileLang kernel封装fp8_native_block_mean_pooling_interface(k, k_scale, k_block_size)blocked_k, blocked_k_scale fp8_native_block_mean_pooling_interface( k, # [N, D] fp8 k_scale, # [N] f32 — per-token scale from indexer_k_quant_and_cache k_block_size, ) # blocked_k: [num_blocks, D] fp8 # blocked_k_scale: [num_blocks] f32k_scale是indexer_k_quant_and_cache产出的逐 token 反量化 scaleREADME 中明确说明其来源。N 个 K token 被分成ceildiv(N, k_block_size)个 pool block。语义对第b个 pool block大小为kb k_block_size逐 token 反量化k_f[i] k_fp8[i] * k_scale[i]在 f32 上对块内求均值mean sum_i k_f[i] / kb不完整尾块按实际有效个数做除数用逐块 scale 重新量化为 fp8block_scale max(max_abs(mean) / 448, 1e-10)写出blocked_k[b] fp8(mean / block_scale)与blocked_k_scale[b] block_scale。448 是 fp8 e4m3 的最大可表示值FP8_MAX_INV 1.0 / 448.0因此max_abs(mean) / 448恰好是让均值的主元素贴近 fp8 满量程的缩放因子最大化量化精度1e-10下界用于防御全零块除零。该 kernel 是序列无关的它对连续k_block_size个 token 做扁平池化不感知序列边界源码注释已注明此点因此适合在索引前先对全局 K 做一次性降采样。TileLang 实现要点从 fp8_block_mean_pooling.py 的源码可见网格为T.Kernel(num_blocks, threadsthreads)即一个 block 对应一个 pool block块内按block_N × dim的 fragment 分片加载使用T.reduce_sum(..., dim0)做 f32 累加再乘1/cur_pooling_block_size得到均值用T.reduce_absmax求逐维最大绝对值得到逐块 scale对外统一由fp8_native_block_mean_pooling_interface以dimk.shape[1]、pooling_block_sizek_block_size调用无需关心内部block_N、num_stages等 tile 参数编译期启用tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True该配置在 tilelang/cuda/backend.py 中控制 CUDA 后端的 fast-math 编译开关换取算力换精度的取舍。正确性与性能验证test_fp8_block_mean_pooling用ref_fp8_block_mean_pooling做参考逐块反量化后按实际有效个数求均值对 fp8 重新量化结果以rtol5e-2, atol5e-3断言fp8 重量化相对误差约 1/256叠加在 bf16 精度之上随后用tilelang.profiler.do_bench测延迟并按bytes_moved N*D N*4 num_blocks*D num_blocks*4折算有效带宽 GB/s。测试内置配置覆盖N 16384/65536/131072、num_seqs 1/4/8/16多组规模。Kernel 1.2块级 fp8 MQApool_mqa_fp8.py函数与接口函数pool_mqa_attn_return_logits_fp8封装pool_mqa_attn_return_logits_fp8_interface(q_fp8, blocked_kv_fp8, blocked_kv_scale, weights_f32, cu_seqlen_blocked_ks, cu_seqlen_blocked_ke, block_N256)block_k_score pool_mqa_attn_return_logits_fp8_interface( q_fp8, # [M, H, D] fp8 blocked_kv_fp8, # [Nb, D] fp8 (来自 1.1) blocked_kv_scale, # [Nb] f32 (来自 1.1) weights_f32, # [M, H] f32 cu_seqlen_blocked_ks, # [M] int32 — 每个 query 在 pool-block 坐标系下的起始 cu_seqlen_blocked_ke, # [M] int32 — 每个 query 在 pool-block 坐标系下的结束 ) # block_k_score: [M, Nb] f32这里cu_seqlen_blocked_ks/ke是块坐标由 hisa.py 中cu_seqlen_ks // k_block_size起点向下取整与(cu_seqlen_ke k_block_size - 1) // k_block_size终点向上取整换算得到保证每个 query 可见的 K 范围在块坐标下被完整覆盖。语义对每个 querym与每个位于[cu_seqlen_blocked_ks[m], cu_seqlen_blocked_ke[m])内的 pool blocknblock_k_score[m, n] sum_h ReLU(q[m, h] · blocked_k[n]) * blocked_k_scale[n] * weights[m, h]注意得分公式包含三段先做 fp8 反量化乘逐块 scale、再 ReLU、最后按 query 逐 head 的weights加权归约。参考实现ref_pool_mqa_fp8用torch.einsum(mhd,nd-mnh, ...)先在 f32 下还原完整[M, Nb, H]打分网格再归约语义完全一致。TileLang 实现要点这是整条流水线中最值得细读的一颗 kernelQ 的分块block_Q默认为128 // headsblock_Q 0时按 heads 推导Q 被切块为block_Q × heads的 fp8 shared memory tileK 的可见区间联合kernel 先计算当前 Q tile 内所有 query 的cu_k_s_min min(cu_seqlen_blocked_ks)与cu_k_e_max max(cu_seqlen_blocked_ke)均对seq_len_blocked_kv截断随后只对该联合可见区间[cu_k_s_min, cu_k_e_max)做T.Pipelined流水化 GEMM 扫描block_N256、num_stages3fp8×fp8→f32 GEMMT.gemm(index_k_shared, index_q_shared, s, transpose_BTrue, clear_accumTrue, policyT.GemmWarpPolicy.FullCol)逐块 scale 在 GEMM 之后以 elementwise 方式乘入写入语义重要kernel 写出的是当前 tile 的联合可见 K 范围——联合范围内但不在某个 query 自身范围内的条目仍携带原始点积值未掩码由下一步 1.3 统一清洗联合范围之外的条目保持接口层torch.zeros零初始化值。这就是封装函数中zero-inits logits so positions the kernel doesnt touch are 0注释的含义。正确性与性能验证由于裸 kernel 输出在联合范围内未掩码test_pool_mqa_fp8在比较前对 kernel 输出与 torch 参考都套用clean_and_maintain_logits掩码这正是流水线紧随其后的动作并精确断言inf/-inf掩码模式完全一致、有限值在rtol5e-2, atol5e-2下对齐。性能侧按total_flops 2 * M * H * Nb * D折算 fp8 TFLOPS。测试还要求N_blocked % block_N 0例如k_block_size128时需M % 32768 0选择 M 时需注意这一约束。Kernel 1.3logits 就地清洗与边界维持clean_and_maintain_logits.py函数与接口函数clean_and_maintain_logits_封装clean_and_maintain_logits_interface(logits, cu_seqlen_ks, cu_seqlen_ke)clean_and_maintain_logits_interface( logits, # [M, Nb] f32 — stage-1 输出就地修改 cu_seqlen_ks, # [M] int32 — 每行起始含 cu_seqlen_ke, # [M] int32 — 每行结束不含 )语义对每一行m[cu_seqlen_ks[m], cu_seqlen_ke[m])之外的位置 → 置-inf使torch.topk忽略它们位置cu_seqlen_ks[m]与cu_seqlen_ke[m] - 1→ 置inf强制保留边界块它们在随后的 top-block 选择中必然被选中——这是 HISA 的经典技巧用于保住 sink起始块与 local局部块位置。从 clean_and_maintain_logits.py 源码看kernel 逐行一个 block用T.Pipelined(ceildiv(seq_len_kv, block_K))以block_K4096分片扫描整行仅对命中idx cu_k_s、idx cu_k_e - 1、idx cu_k_s、idx cu_k_e的位置写±inf。参考实现ref_clean_and_maintain_logits是纯 torch 版本masked_fill置-inf再对out[m, cu_ks]与out[m, (cu_ke-1).clamp(min0)]置inf。验证test_clean_and_maintain_logits断言inf/-inf掩码模式与参考完全一致且有限值位rtol0, atol0逐元素相等该 kernel 只写±inf其余位置原样保留。带宽按约2*M*N*4字节折算——大部分位置是空转实际只有掩码边界在写。Kernel 2.1块稀疏 fp8 MQAblock_sparse_mqa_fp8.py函数与接口函数fp8_native_block_sparse_mqa_attn_return_logits封装fp8_native_block_sparse_mqa_attn_return_logits_interface(q, k, k_scale, topk_block_index, kv_block_size, weights, cu_seqlen_ks, cu_seqlen_ke)block_sparse_logits fp8_native_block_sparse_mqa_attn_return_logits_interface( q, # [M, H, D] fp8 k, # [N, D] fp8 k_scale, # [N] f32 topk_block_index, # [M, block_topk] int64 — 来自 stage-1 得分的 torch.topk kv_block_size, # k_block_size weights, # [M, H] f32 cu_seqlen_ks, # [M] int32 — 每 query 的绝对 K 起始原始 token 坐标 cu_seqlen_ke, # [M] int32 — 每 query 的绝对 K 结束 ) # block_sparse_logits: [M, block_topk * kv_block_size] f32注意此处cu_seqlen_ks/ke回到绝对 token 坐标。kernel 工厂会按kv_block_size与block_N的关系自动分派两种变体README 说明一般情形kv_block_size block_N对每个选中块内部再做kv_block_size / block_N个子块的流水化T.Pipelined内层循环小池化尺寸kv_block_size block_N单趟完成无流水线。语义对每个 querym、每个选中块t ∈ [0, block_topk)blk topk_block_index[m, t]、每个块内偏移i ∈ [0, kv_block_size)k_abs blk * kv_block_size i if k_abs ∉ [cu_seqlen_ks[m], cu_seqlen_ke[m]) 或 k_abs N: block_sparse_logits[m, t * kv_block_size i] -inf else: block_sparse_logits[m, t * kv_block_size i] sum_h ReLU(q[m, h] · k[k_abs]) * k_scale[k_abs] * weights[m, h]与 stage 1 不同越界掩码由本 kernel 直接写入源码中k_i cu_k_s_min or k_i cu_k_e_max时写-T.infinity无需像 stage 1 那样额外做一次掩码 pass。参考实现ref_fp8_block_sparse_mqa通过torch.einsum(mhd,mtid-mtih, ...)完成同样的逐 (query, block, offset, head) 计算后masked_fill置-inf语义完全对齐。TileLang 实现要点从源码看该 kernel 每 query 一个 blockT.Kernel(seq_len, ...)外层T.serial(topk)串行遍历每个选中块内层按block_N min(block_N, kv_block_size // 2)的子块做流水化 GEMMscale 放在共享内存而非 fragment 中源码注释指出在串行 topk 循环下 shared 略快于 fragment。GEMM 使用T.GemmWarpPolicy.FullRow与 stage 1 的FullCol形成对照——不同的 warp 分配策略适配不同的 tile 几何。验证与约束test_fp8_block_sparse_mqa用随机torch.randperm生成块索引部分块会落在 query 自身序列之外恰好锻炼 kernel 内置掩码路径断言掩码模式精确一致、有限值在rtol1e-1, atol2e-1下对齐。参考路径会物化[M, topk, B, D]fp32 的 gathered_k约 M GB 量级因此测试配置刻意保持 M 适中以免 OOM源码注释已提示。端到端编排hisa_indexerhisa.py函数与接口函数hisa_indexertopk_indices hisa_indexer( q, # [M, H, D] fp8 k, # [N, D] fp8 k_scale, # [N] f32 weights, # [M, H] f32 cu_seqlen_ks, # [M] int32 — 每 query 的 K 起始 cu_seqlen_ke, # [M] int32 — 每 query 的 K 结束 *, k_block_size, # pool block 大小DeepSeek-V3.2 中 128 block_topk, # 每 query 保留的 top pool block 数 topk_tokens, # 交给稀疏注意力的最终 top-k 大小 ) # topk_indices: [M, topk_tokens] int32 — 每行是 query 在其 [cu_ks, cu_ke) # 窗口内的 top-k K 偏移越界槽位为 -1完整流水线(1.1) fp8_native_block_mean_pooling K, k_scale → blocked_k, blocked_k_scale (1.2) pool_mqa_attn_return_logits_fp8 Q × blocked_k → block_k_score[M, Nb] (1.3) clean_and_maintain_logits block_k_score 就地掩码-inf/inf (1.4) torch.topk(block_k_score.bfloat16(), kblock_topk, sortedFalse) → topk_block_indices[M, block_topk] int64 (2.1) fp8_native_block_sparse_mqa_… Q × K[选中块] → block_sparse_logits [M, block_topk * k_block_size] (2.2) torch.topk(block_sparse_logits, ktopk_tokens) → relevant_topk_indices[M, topk_tokens] int64 (2.3) (Python) gather topk_block_indices arith 减 cu_seqlen_ks 掩码 → topk_indices[M, topk_tokens] int32源码级细节对应 hisa.py块坐标换算cu_seqlen_blocked_ks cu_seqlen_ks // k_block_size、cu_seqlen_blocked_ke (cu_seqlen_ke k_block_size - 1) // k_block_size向池化块坐标系取整Stage 1.4 的 bf16 优化对block_k_score转 bf16 再做torch.topk(kblock_topk, sortedFalse)。源码注释明确说明 bf16 非排序比 f32 快约 40%且下游 sparse MQA 不依赖顺序同时block_topk_eff min(block_topk, block_k_score.shape[-1])防御块数不足Stage 2.2torch.topk(ktopk_tokens_eff)同样用min(topk_tokens, 候选总数)防御Stage 3 索引换算最易出错的一环slot 编号满足slot block_id_in_topk × k_block_size offset_in_block因此absolute_topk_block_indices torch.gather(topk_block_indices, dim-1, indexrelevant_topk_indices // k_block_size)反查真实块 idtopk_indices absolute_topk_block_indices * k_block_size (relevant_topk_indices % k_block_size)得到绝对 K token 位置topk_indices - cu_seqlen_ks[:, None]转为 query 自身 K 窗口内的相对偏移0表示该 query 的 K 起点与 vLLM indexer 的输出缓冲格式对齐用mask_lo topk_indices 0与mask_hi topk_indices - (cu_ke - cu_ks)[:, None] 0做窗口合法性校验不合法槽位masked_fill(..., -1)。最终输出[M, topk_tokens]int32每行为该 query 的 top-k K 位置相对自身 K 窗口的偏移越界槽位为-1。test_hisa端到端冒烟与性能test_hisa将num_seqs条等长因果序列打包进扁平的[M, H, D]Q 与[NM, D]K借助prepare_ks_ke_from_cu_seqlens见 tilelang_utils.py把cu_seqlens展开为逐 token 的ks/ke使每个 query 只能看到自己序列的前缀。校验点包括输出形状(M, topk_tokens)且为 int32索引掩码不变量所有非-1偏移必须落在[0, cu_ke[m] - cu_ks[m])内valid (valid in_range)有效槽位计数每 query 有效槽位数必须等于min(窗口长度, min(topk_tokens, block_topk * k_block_size))并打印per-query valid count match比例性能用do_bench(fn, warmup20, rep50)测端到端延迟tilelang.profiler.do_bench默认带 L2 缓存管理与 warmup/rep 控制实现在 tilelang/profiler/bench.py。__main__内置六组配置覆盖M ∈ {1024, 4096, 8192}、block_topk ∈ {16, 32, 64}、topk_tokens ∈ {256, 1024, 2048}、num_seqs ∈ {1, 4, 8}每轮结束调用torch.cuda.empty_cache()释放显存。运行与验证方式目录下每个.py均可独立运行其测试与基准cd examples/dsa_hisa python fp8_block_mean_pooling.py # 1.1 正确性 带宽 python pool_mqa_fp8.py # 1.2 正确性 fp8 TFLOPS python clean_and_maintain_logits.py # 1.3 正确性 带宽 python block_sparse_mqa_fp8.py # 2.1 正确性 fp8 TFLOPS python hisa.py # 端到端流水线 延迟运行前提CUDA 环境与已安装的 tilelang 运行时kernel 通过tilelang.jit在首次调用时编译T.Tensor[[seq_len, ...]]中的T.const标注表示序列长度这类符号维度可动态推导。各 kernel 测试均要求torch.manual_seed(0)保证可复现do_bench提供 warmup/rep 参数以稳定测时。小结HISA 的 TileLang 实现演示了如何把分层稀疏索引这一算法流程高效映射为四颗各司其职的 GPU kernel块均值池化压缩 K 规模、块级 fp8 MQA 粗筛候选块、logits 清洗维持边界语义、块稀疏 fp8 MQA 细筛最终 token最后由 hisa.py 以两次torch.topk与一段精密的索引换算收尾。每颗 kernel 都配有语义等价的 torch 参考实现与严格的掩码模式断言形成可独立验证、可增量替换的模块化设计k_block_size、block_topk、topk_tokens三个参数则完整暴露了 HISA 在召回质量与计算量之间的权衡旋钮可直接适配到 DeepSeek-V3.2 一类稀疏注意力场景的 indexer 中。【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考