
TileLang 块稀疏 Flash-Attention 内核实现与实战指南【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang块稀疏注意力Block-Sparse Attention是当前长上下文推理与训练中降低注意力计算量的主流手段之一与其在 token 粒度做稀疏不如把序列划分为固定大小的块只对重要的块执行 QK 与 PV 计算。本文基于 examples/blocksparse_attention 目录下的完整实现系统讲解 TileLang 版块稀疏 Flash-Attention 内核的设计思路、掩码构建方法、因果 softmax 的数值技巧、GQA 变长解码与 Paged KV Cache 等进阶形态并给出运行、验证与性能回归的完整操作方式。读完本文你将能够理解并复现这套内核并把它接入自己的稀疏注意力场景。1. 背景块稀疏注意力解决什么问题现代注意力机制的 FLOPs 与 KV Cache 访存量随序列长度平方增长而真实场景中大部分注意力分数集中在少数关键 token 上。块稀疏注意力将序列按固定块大小如 64切分先通过某种启发式如 top-k 或阈值挑选出每个查询块需要关注的 KV 块集合再把完整注意力矩阵约减为对这些块的稀疏访问从而在保持精度的同时显著降低计算与带宽开销。根据 examples/blocksparse_attention/README.md 的说明本目录中的 TileLang 内核已被用于两项相关工作Rectified Sparse AttentionRSA一种训练阶段即引入的稀疏注意力方案通过在训练中学习并裁剪不重要的注意力块来加速训练与推理。SeerAttention-R基于 SeerAttention 的稀疏注意力变体利用学习到的稀疏模式指导块级注意力计算。本目录的代码正是这两个方案中块稀疏注意力内核的通用实现基础覆盖了prefill等长序列、GQA decode变长 分页两类核心场景。2. 仓库结构总览本主题相关的全部代码集中在两个目录中文件作用example_tilelang_block_sparse_attn.pyTileLang 块稀疏 Flash-Attention 主内核prefill / 等长序列example_tilelang_sparse_gqa_decode_varlen_indice.pyGQA 变长解码内核以 block_indices 表示稀疏块-1填充example_tilelang_sparse_gqa_decode_varlen_mask.pyGQA 变长解码内核以 block_mask 布尔掩码表示稀疏块example_tilelang_sparse_gqa_decode_paged.py基于 Paged KV Cacheblock_table 逻辑块到物理块映射的稀疏解码block_sparse_attn_triton.pyTriton 对照实现等长 prefillexample_triton_sparse_gqa_decode_varlen_indice.pyTriton 对照实现GQA 变长解码split-merge 两阶段heuristic.pyKV 块分裂数num_splits启发式算法test_example_blocksparse_attention.py正确性测试入口regression_example_blocksparse_attention.py性能回归入口benchmark/blocksparse_attention/benchmark_configs.py性能基准配置3. 块稀疏掩码的构建从重要性分数到块级掩码稀疏模式由[batch, heads, downsample_len, downsample_len]的块级掩码描述其中downsample_len ceil(seq_len / BLOCK)即序列被切分后的块数。掩码中的每个布尔值表示对应 KV 块是否需要参与注意力计算。目录中的两个工具函数在两个 TileLang 示例与 block_sparse_attn_triton.py 中均有同名实现负责把重要性分数转换为这种掩码。3.1 top-k 方式def get_sparse_attn_mask_from_topk(x, topk, use_dense_for_last_blockFalse): bsz, num_head, downsample_len, _ x.shape # N_CTX downsample_len * BLOCK sparse_index torch.topk(x, topk, dim-1).indices dense_mask torch.full([bsz, num_head, downsample_len, downsample_len], False, dtypetorch.bool, devicex.device) dense_mask.scatter_(-1, sparse_index, True) if use_dense_for_last_block: dense_mask[:, :, -2:, :] True dense_mask.tril_() return dense_mask对每个查询块的重要性分数行取topk个最高分的块索引通过scatter_写入Trueuse_dense_for_last_blockTrue时强制最后两行最近的 KV 块全为True——解码场景中最近的 token 通常至关重要保证不丢失最后调用tril_()施加因果约束保证只保留下三角块。3.2 阈值方式def get_sparse_attn_mask_from_threshold(x, threshold, use_dense_for_last_blockFalse): dense_mask x threshold if use_dense_for_last_block: dense_mask[:, :, -2:, :] True dense_mask.tril_() return dense_mask与 top-k 的唯一区别是筛选规则分数超过threshold的块全部保留。两种方式都可以叠加use_dense_for_last_block与tril_()的因果处理可根据稀疏率控制需求自由选择。4. 核心内核blocksparse_flashattn实现解析example_tilelang_block_sparse_attn.py 中的blocksparse_flashattn是整套实现的基础内核面向 prefillseq_q seq_kv场景数据结构为[batch, heads, seq_len, dim]。4.1 内核参数与 JIT 装饰器tilelang.jit( out_idx[4], pass_configs{ tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, }, ) def blocksparse_flashattn(batch, heads, seq_len, dim, downsample_len, is_causal): block_M 64 block_N 64 num_stages 1 threads 128 scale (1.0 / dim) ** 0.5 * 1.44269504 # log2(e)out_idx[4]声明第 4 个张量参数Output为输出JIT 编译后内核调用直接返回输出张量TL_ENABLE_FAST_MATH开启快速数学优化允许使用更快的近似指令block_M block_N 64Q 与 K/V 的块大小均为 64与掩码下采样块大小一致num_stages 1流水线级数此内核未开启多级软件流水线threads 128每个 block 使用 128 个线程scale (1.0 / dim) ** 0.5 * 1.44269504缩放因子乘以log2(e)配合下文exp2技巧使用。4.2 张量声明与片上存储分配shape [batch, heads, seq_len, dim] block_mask_shape [batch, heads, downsample_len, downsample_len] dtype T.float16 accum_dtype T.float32 block_mask_dtype T.bool主函数接受Q, K, V, BlockSparseMask, Output五个张量。内核体内先分配共享内存shared与寄存器片段fragmentQ_shared T.alloc_shared([block_M, dim], dtype) K_shared T.alloc_shared([block_N, dim], dtype) V_shared T.alloc_shared([block_N, dim], dtype) O_shared T.alloc_shared([block_M, dim], dtype) acc_s T.alloc_fragment([block_M, block_N], accum_dtype) # QK 分数 acc_s_cast T.alloc_fragment([block_M, block_N], dtype) # softmax 概率转回 fp16 供 gemm acc_o T.alloc_fragment([block_M, dim], accum_dtype) # 输出累加器 scores_max / scores_max_prev / scores_scale / scores_sum / logsum: T.alloc_fragment([block_M], accum_dtype) block_mask T.alloc_fragment([downsample_len], block_mask_dtype) # 当前查询块的整行掩码fragment是 TileLang 中的寄存器级数据抽象acc_o以 fp32 累积避免数值误差累积。4.3 网格划分与掩码加载with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch, threadsthreads) as (bx, by, bz):网格为(seq_len / block_M, heads, batch)每个线程块负责一个(batch, head, 查询块)组合。接着T.copy(Q[bz, by, bx * block_M : (bx 1) * block_M, :], Q_shared) T.fill(acc_o, 0) T.fill(logsum, 0) T.fill(scores_max, -T.infinity(accum_dtype)) T.copy(BlockSparseMask[bz, by, bx, :], block_mask)将当前查询块载入共享内存初始化 online softmax 状态并把该查询块对应的整行稀疏掩码读入寄存器。4.4 因果循环范围与块级稀疏跳转loop_range ( T.min(T.ceildiv(seq_len, block_N), T.ceildiv((bx 1) * block_M, block_N)) if is_causal else T.ceildiv(seq_len, block_N) ) for k in T.Pipelined(loop_range, num_stagesnum_stages): if block_mask[k] ! 0: T.copy(K[bz, by, k * block_N : (k 1) * block_N, :], K_shared) ...因果模式下循环上界取min(总块数, 当前查询块允许访问的块数)直接裁剪掉未来块稀疏性的核心if block_mask[k] ! 0是块级分支——只有掩码为 1 的 KV 块才被加载并参与计算。非稀疏块完全跳过 K/V 的拷贝与后续 GEMM这是相对稠密 Flash-Attention 的全部收益来源注意该分支是设备端运行时的动态分支掩码来自显存因此块稀疏只按被选中块的占比线性降低计算量。4.5 因果掩码写入与 QK^T GEMMif is_causal: for i, j in T.Parallel(block_M, block_N): acc_s[i, j] T.if_then_else(bx * block_M i k * block_N j, 0, -T.infinity(acc_s.dtype)) else: T.clear(acc_s) T.gemm(Q_shared, K_shared, acc_s, transpose_BTrue, policyT.GemmWarpPolicy.FullRow)因果块内把对角线以下置 0、以上置-inf随后 GEMM 累加到acc_sGEMM 采用累加语义-inf会一直保持为-inf非因果块直接clearT.gemm(..., transpose_BTrue)计算Q K^TT.GemmWarpPolicy.FullRow指定 warp 级负载分配策略保证行方向block_M 方向的归约数据由同一 warp 持有。4.6 Online Softmaxexp2 与 Check_inf 技巧T.copy(scores_max, scores_max_prev) T.fill(scores_max, -T.infinity(accum_dtype)) T.reduce_max(acc_s, scores_max, dim1, clearFalse)这是标准的 online softmax 流程先保存上一轮最大值再求本轮块内行最大值随后更新。for i in T.Parallel(block_M): scores_scale[i] T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale) for i, j in T.Parallel(block_M, block_N): # Instead of computing exp(x - max), we compute exp2(x * log_2(e) - max * log_2(e)) # This allows the compiler to use the ffma instruction instead of fadd and fmul separately. acc_s[i, j] T.exp2(acc_s[i, j] * scale - scores_max[i] * scale) T.reduce_sum(acc_s, scores_sum, dim1) for i in T.Parallel(block_M): logsum[i] logsum[i] * scores_scale[i] scores_sum[i] T.copy(acc_s, acc_s_cast) for i, j in T.Parallel(block_M, dim): acc_o[i, j] * scores_scale[i]实现中注释提到两点值得注意的数值优化用exp2替代exp由于scale已预先乘上log2(e)exp(x·scale_orig)被改写为exp2(x·scale)使得编译器可将fadd fmul融合为一条ffmafused multiply-add指令Check_inf被注释掉的代码展示了 FlashAttention-3 中Check_inf的思想——当某个查询块首轮最大值仍为-inf说明它没有命中任何选中块时需要把它钳制为 0。本实现未显式执行这一步注释在源码 example_tilelang_block_sparse_attn.py 中依赖exp2(-inf · scale) exp2(-inf) 0的数学性质在概率计算中自然处理。随后acc_s转回 fp16acc_s_cast与V_shared执行第二次 GEMMT.copy(V[bz, by, k * block_N : (k 1) * block_N, :], V_shared) T.gemm(acc_s_cast, V_shared, acc_o, policyT.GemmWarpPolicy.FullRow)4.7 归一化与写回for i, j in T.Parallel(block_M, dim): acc_o[i, j] / logsum[i] T.copy(acc_o, O_shared) T.copy(O_shared, Output[bz, by, bx * block_M : (bx 1) * block_M, :])循环结束后用累积的logsum对输出做最终归一化经共享内存写回全局内存。5. 掩码与参考实现对拍示例末尾的test_topk_sparse_attention()给出了端到端验证流程BATCH, N_HEADS, SEQ_LEN, D_HEAD 1, 1, 256, 64 TOPK 2 # 每行保留的块数 BLOCK 64 # 下采样块大小构造随机 Q/K/Vfp16与随机重要性分数x_ds并把x_ds[:, :, :, 0] 100强制第一列被选中保证每行至少有一个有效块随后生成掩码并调用内核kernel blocksparse_flashattn(BATCH, N_HEADS, SEQ_LEN, D_HEAD, downsample_len, is_causalTrue) tilelang_output kernel(q, k, v, block_mask)参考实现把块掩码通过torch.kron展开为完整注意力掩码叠加因果后与标准einsum softmax结果对比容差为atol1e-2, rtol1e-2full_mask torch.kron(block_mask.float(), torch.ones(BLOCK, BLOCK, devicecuda)) full_mask full_mask[..., :SEQ_LEN, :SEQ_LEN].bool() full_mask full_mask torch.tril(torch.ones_like(full_mask)) ... torch.testing.assert_close(tilelang_output, ref_output, atol1e-2, rtol1e-2)6. GQA 变长稀疏解码内核prefill 内核要求等长序列且掩码为稠密的[B, H, n, n]矩阵。而在自回归推理的 decode 阶段每个 batch 的 query 只有一行形状[batch, heads, dim]KV 来自长度各不相同的 KV Cache因此目录提供了专门的GQAGrouped Query Attention变长解码内核支持 KV 头数小于查询头数的 GQA 结构并针对不同的稀疏表示提供两个版本example_tilelang_sparse_gqa_decode_varlen_indice.py用block_indices: [batch, heads_kv, max_selected_blocks]记录被选中块的下标未选中位置用-1填充example_tilelang_sparse_gqa_decode_varlen_mask.py用block_mask: [batch, heads_kv, num_blocks]的布尔掩码标记有效块。两个版本的内核逻辑几乎一致区别仅在于循环内的分支判断block_indices[..., startk] 0与block_mask[..., startk]以及 K/V 拷贝时的取址方式。6.1 动态形状与 split-merge 架构解码内核使用T.dynamic声明运行期确定的形状num_split T.dynamic(num_split) max_cache_seqlen T.dynamic(max_cache_seqlen) max_selected_blocks T.dynamic(max_selected_blocks)内核整体采用split-merge 两阶段架构类似 split-K 注意力Split 阶段with T.Kernel(batch, heads // valid_block_H, num_split, threadsthreads)——把每个(batch, KV 头组)的 KV 块序列切成num_split段并行处理每段产生部分输出Output_partial与对应 log-sum-expglseMerge 阶段with T.Kernel(heads, batch, threads128)——按标准 split-K 合并公式先求各 split 的 LSE 最大值lse_max_local再以exp2(lse - lse_max)为权重对部分输出加权求和。其中valid_block_H min(block_H, kv_group_num)block_H 64表示一次处理 64 个查询头cur_kv_head hid // (kv_group_num // valid_block_H)实现 GQA 中多个查询头共享同一个 KV 头的映射。6.2 块分裂与 cache_seqlens 边界处理每个 split 的循环范围与起始位置按如下方式划分blocks_per_split T.floordiv(num_blocks, num_split) remaining_blocks T.floormod(num_blocks, num_split) loop_range blocks_per_split T.if_then_else(sid remaining_blocks, 1, 0) start blocks_per_split * sid T.min(sid, remaining_blocks)即前remaining_blocks个 split 各多处理一个块保证负载均匀。循环内i_s block_indices[bid, cur_kv_head, start k] if i_s 0: has_valid_block True T.copy(K[bid, i_s * block_N : (i_s 1) * block_N, cur_kv_head, :], K_shared) ... if k 0: # assume block_indices is sorted in reverse order, otherwise, remove this if condition for i, j in T.Parallel(block_H, block_N): acc_s[i, j] T.if_then_else(i_s * block_N j cache_seqlens[bid], -T.infinity(accum_dtype), acc_s[i, j])索引为-1的填充槽直接跳过由于block_indices按降序排列最近块在前只有首个有效块k 0需要做cache_seqlens边界检查——其余块必然完全落在有效长度内。这一优化在源码注释中明确说明见 example_tilelang_sparse_gqa_decode_varlen_indice.py每个 split 的 LSE 写入glse、部分输出写入Output_partial供 merge 阶段读取。6.3 num_splits 启发式算法分裂数由 heuristic.py 中的num_splits_heuristic决定目标是在尽量接近一次 wave 填满 GPU 的前提下取最小分裂数def num_splits_heuristic(total_mblocks, num_SMs, num_n_blocks, num_m_blocks, size_one_kv_head, is_causal_or_local, max_splits):判定流程total_mblocks 0.8 * num_SMs总块数足以接近填满 SMs优先取 1 个 split除非单个 KV 头超过假定的 L2 大小50 MB且num_m_blocks 2 * num_SMs且非因果/局部注意力此时按ceil(size_one_kv_head / size_l2)取 split 数上限max_splitsnum_n_blocks 4KV 块太少不分裂直接返回 1否则在[1, min(max_splits, num_SMs, num_n_blocks)]内计算每个候选分裂数的 wave 效率eff n_waves / ceil(n_waves)其中n_waves total_mblocks * num_splits / num_SMs返回第一个达到最大效率 85% 的最小分裂数。该函数被所有解码示例TileLang 与 Triton 版本统一复用例如在SparseFlashAttn模块中通过torch.cuda.get_device_properties获取实际 SM 数后调用num_split num_splits_heuristic( total_mblocks, num_sm, num_n_blocks, num_m_blocks, size_one_kv_head, is_causal_or_localTrue, max_splits128 )其中size_one_kv_head max_selected_blocks * block_size * (dim dim_v) * 2估算单 KV 头的字节数用于 L2 驻留判断。7. Paged KV Cache 稀疏解码example_tilelang_sparse_gqa_decode_paged.py 在前述 indice 版本的基础上增加了PagedAttention 式分页兼容 vLLM 风格的 KV Cache 布局shape_k [num_pages, page_block_size, heads_kv, dim] shape_v [num_pages, page_block_size, heads_kv, dim_v] shape_block_table [batch, max_num_blocks_per_seq]K/V 不再连续存储而是以页为单位存放在num_pages个物理页中block_table[batch, 逻辑块号] - 物理页号完成映射。内核开头有约束断言assert block_N page_block_size and page_block_size % block_N 0 block_ratio page_block_size // block_Nblock_ratio表示一个物理页包含多少个注意力块。循环内通过两级索引定位数据logical_block_idx block_indices[bid, cur_kv_head, start k] if logical_block_idx 0: block_table_idx T.floordiv(logical_block_idx, block_ratio) block_tile_idx T.floormod(logical_block_idx, block_ratio) physical_block_idx block_table[bid, block_table_idx] T.copy(K[physical_block_idx, block_tile_idx * block_N : (block_tile_idx 1) * block_N, cur_kv_head, :], K_shared)即逻辑块号 →(页号, 页内块号)→ 查block_table得到物理页号 → 以block_tile_idx * block_N为起点拷贝。数值流程online softmax、边界检查、split-merge与 6 节完全一致因此本实现可以无缝替换推理框架中的稠密 PagedAttention。8. Triton 对照实现为验证正确性与评估相对性能目录同时提供了两套 Triton 对照实现block_sparse_attn_triton.py等长 prefill 场景_fwd_kernel中同样按k_block_col_idx逐块检查掩码mask_val命中才执行tl.dot(q, k)与后续 online softmax支持PAST_LENqlen klen 的增量解码场景见test_topk_sparse_attention_qlt_kl并通过LAST_K_BLOCK标记在最后一个块追加因果掩码。其backward抛出不支持错误明确只做前向example_triton_sparse_gqa_decode_varlen_indice.pyGQA 变长解码使用triton.autotune在num_warps ∈ {1,2,4}与num_stages ∈ {1,2,3,4,7}共 15 组配置间自动调优同样采用 split-merge 两阶段_split_kernel_merge_kernel并复用同一个num_splits_heuristic。Triton 版与 TileLang 版共享相同的输入约定与参考实现便于在同一问题规模下做同一算法、两种 DSL的横评。9. 运行、测试与基准9.1 单测test_example_blocksparse_attention.py 汇总了全部正确性测试其中两个 GQA 解码用例通过tilelang.testing.requires_package(flash_attn)标记仅在安装了flash_attn用于 FA2 的flash_attn_with_kvcache参考实现时执行cd examples/blocksparse_attention python test_example_blocksparse_attention.py直接运行单个示例python example_tilelang_block_sparse_attn.py # prefill topk 稀疏注意力正确性 python example_tilelang_sparse_gqa_decode_varlen_indice.py # 变长解码 indice python example_tilelang_sparse_gqa_decode_varlen_mask.py # 变长解码 mask9.2 命令行参数三个解码示例均通过argparse暴露以下可调参数默认值以 indice/mask 版本为例参数默认值说明--batch8batch 大小--heads32查询头数--heads_kv8KV 头数GQA--max_cache_seqlen8192KV Cache 最大序列长度--dim128查询/键维度--dim_v128值维度--sparse_ratio0.8稀疏率1 - 稀疏率决定保留块数--block_size32注意力块大小block_Npaged 版本额外提供--block_N 64、--page_block_size 256、--num_pages 1024。示例运行时会同时输出 TileLang 稀疏内核与 FA2 稠密flash_attn_with_kvcache的耗时对比dense time/sparse time并在sparse_ratio0.0时以 FA2 输出做端到端对拍assert allclose其余稀疏率则以纯 PyTorch 参考实现校验。9.3 性能基准benchmark/blocksparse_attention 目录提供独立的基准脚本与配置benchmark_configs.py以[BATCH, N_HEADS, SEQ_LEN, D_HEAD, TOPK, BLOCK]形式配置测试规模默认配置为[4, 2, 256, 64, 2, 64]benchmark_library_dense_fmha.py、benchmark_torch_block_sparse_fmha.py、benchmark_triton_block_sparse_fmha.py、benchmark_tilelang_block_sparse_fmha.py分别对库内稠密实现、PyTorch 稀疏参考、Triton 稀疏实现与 TileLang 稀疏实现做统一基准。9.4 性能回归regression_example_blocksparse_attention.py 通过tilelang.testing.process_func与tilelang.profiler.do_bench(..., backendcupti)测量纯内核耗时见各示例中的run_regression_perf用于在 CI 中监控性能退化。回归用例以batch1, max_cache_seqlen2048的小规模运行保证测试时长可控。10. 小结与扩展本目录围绕块稀疏 Flash-Attention提供了从掩码构建到高性能内核的完整闭环掩码层top-k / 阈值两种块稀疏模式生成支持因果与最后块全稠密选项内核层等长 prefill 内核blocksparse_flashattn、GQA 变长解码内核indice/mask 双表示、Paged KV Cache 内核统一采用块级分支跳过 online softmaxexp2/ffma 优化 split-merge 架构工程层PyTorch 参考实现对拍、FA2 稠密基线对比、Triton 同算法对照、CUPTI 基准与 CI 回归脚本一应俱全。该实现已服务于 Rectified Sparse Attention 与 SeerAttention-R 等研究工作。若需将此内核用于自定义场景可沿三条路径扩展修改掩码生成逻辑以接入学习型稀疏如 examples/seer_attention/block_sparse_attn_tilelang.py 中的 SeerAttention 变体其支持seq_q ! seq_kv且掩码以int8存储调整block_M/block_N与num_stages适配不同硬件与块粒度或基于 paged 版本对接推理框架的 KV Cache 管理器。【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考