十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

用 TileLang 编写高性能 DeepSeek MLA 内核:从 FlashAttention 到 FlashMLA 级性能的完整实战指南

用 TileLang 编写高性能 DeepSeek MLA 内核:从 FlashAttention 到 FlashMLA 级性能的完整实战指南 用 TileLang 编写高性能 DeepSeek MLA 内核从 FlashAttention 到 FlashMLA 级性能的完整实战指南【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang本文以 TileLang 官方示例examples/deepseek_mla为核心系统讲解如何在 TileLang 中编写高性能的 DeepSeek MLAMulti-Head Latent Attention内核。文章先剖析 MLA 相比传统 MHA/GQA 的性能瓶颈大 head_dim 带来的寄存器压力再逐一拆解 Layout Inference、Threadblock/Shared Memory Swizzling、Warp-Specialization、Pipeline、Split-KV 等关键优化技术并给出可直接运行、可验证正确性的完整代码路径与基准测试方法。读完本文你将掌握用约 80 行 Python 代码复现接近 FlashMLA 性能的 MLA Decode 内核的全部要点。MLA 简介为什么各家编译器都在优化它DeepSeek 的 MLA 是一种以硬件效率著称的新型注意力机制能够显著提升模型推理速度。由于它的头部维度远超传统注意力许多深度学习编译器与算子库如 Triton、FlashInfer 等都为其实现了定制内核。2025 年 2 月DeepSeek 团队在 GitHub 上开源了 FlashMLA它基于 CUTLASS 模板并融合了 FlashAttention 中的在线 Softmax 等优化技术取得了非常亮眼的性能。TileLang 仓库在 examples/deepseek_mla/README.md 中给出了用 TileLang 从零编写 MLA 内核的完整方案与性能对比其核心结论是TileLang 仅需约 80 行 Python 代码即可在大多数测试场景下达到与 FlashMLA 相当的性能并显著优于 FlashInfer 和 Triton。这正是理解 TileLang「易用性与高性能兼得」这一设计理念的最佳范例。基准测试TileLang vs FlashMLA / FlashInfer / Triton / Torch仓库在 batch size 为 64 与 128、float16 精度下对 FlashMLA、TileLang、Torch、Triton、FlashInfer 五者做了性能对比结果分别见下图。Figure 1batch size64 时的性能对比Figure 2batch size128 时的性能对比从结果看TileLang 在大多数场景下性能与 FlashMLA 相当明显超过 FlashInfer 与 Triton。需要强调的是这一性能结果来自仓库示例的实测数据具体数值与硬件环境相关其更重要的意义在于展示在保持接近手写 CUTLASS 内核性能的同时TileLang 把内核源码量压缩到了约 80 行 Python。先看核心难点MLA 的 head_dim 为什么让内核难写在展开 TileLang 实现之前先回顾传统 FlashAttention 的核心计算逻辑# acc_s: [block_M, block_N] # scores_max: [block_M] # scores_scale: [block_M] # acc_o: [block_M, dim] for i in range(loop_range): acc_s Q K[i] scores_max_prev scores_max scores_max max(acc_s, dim1) scores_scale exp(scores_max_prev - scores_max) acc_o * scores_scale acc_s exp(acc_s - scores_max) acc_o acc_s V[i] ...其中acc_s是每次迭代中Q K的结果形状为[block_M, block_N]acc_o是当前迭代的输出累加形状为[block_M, dim]。为了降低访存延迟acc_s与acc_o都必须驻留在寄存器中。与 MHA、GQA 等传统注意力算子相比优化 MLA 的最大挑战在于其 head_dim 非常大query与key的 head dim 为 576512 无 RoPE 部分 64 RoPE 部分value的 head dim 为 512。这直接引发一个严重问题acc_o过大。当线程数不足例如 128 线程时会发生寄存器溢出register spilling严重损害性能。那么如何切分矩阵乘法呢在 Hopper 架构上多数高性能内核使用wgmma.mma_async指令它将 4 个 warp128 线程组织成一个 warpgroup 进行集体 MMA 运算。但wgmma.mma_async要求 M 维最小为 64这意味着每个 warpgroup 的最小 M 维只能降到 64而单个 warpgroup 承载64 × 512的 tile 又太大依然会导致寄存器溢出。因此唯一可行的方案是沿dim维切分acc_o两个 warpgroup 分别计算acc_o的左半部分与右半部分。但这又引入一个新挑战——两个 warpgroup 都需要完整的acc_s结果作为输入。TileLang 示例给出的解法是每个 warpgroup 在Q K阶段只计算一半的acc_s然后通过共享内存从另一个 warpgroup 获取另一半。这一「先分后合」的协同计算模式正是整个 MLA 内核布局设计的出发点。完整内核实现解析example_mla_decode.py仓库中的 example_mla_decode.py 是 MLA Decode单 token、长 KV内核的完整实现它同时包含了split kernel分块计算 combine kernel合并结果两个T.prim_func由num_split参数决定启用哪一个。这是理解其余变体paged、persistent、ws的基础。JIT 入口与编译配置tilelang.jit( out_idx[4], pass_configs{tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}, ) def flashattn(batch, heads, kv_head_num, seqlen_kv, dim, pe_dim, block_N, block_H, num_split, softmax_scale): scale float(softmax_scale * 1.44269504) # log2(e) dtype T.float16 accum_dtype T.float32 kv_group_num heads // kv_head_num VALID_BLOCK_H min(block_H, kv_group_num) assert kv_head_num 1, kv_head_num must be 1几个值得注意的细节out_idx[4]声明第 4 个参数Output为输出张量便于 TileLang 做内存管理与自动 shape 推断。TL_ENABLE_FAST_MATH对应 tilelang/transform/pass_config.py 中的tl.enable_fast_math配置项开启后编译器可用更激进的数学近似例如把 exp 换算成 exp2。这里scale softmax_scale * log2(e)把 Softmax 中的exp全部换算成exp2配合 fast math 进一步提升效率。VALID_BLOCK_H min(block_H, kv_group_num)处理 MQA 场景下block_H超过实际 KV 头组数的边界情况。Split Kernel沿 KV 维切分 在线 Softmax当num_split 1时启用main_split。它使用三维 Grid(batch, heads // min(block_H, kv_group_num), num_split)第三维即 KV 切分维度with T.Kernel(batch, heads // min(block_H, kv_group_num), num_split, threads256) as (bid, hid, bz): Q_shared T.alloc_shared([block_H, dim], dtype) S_shared T.alloc_shared([block_H, block_N], dtype) Q_pe_shared T.alloc_shared([block_H, pe_dim], dtype) KV_shared T.alloc_shared([block_N, dim], dtype) K_pe_shared T.alloc_shared([block_N, pe_dim], dtype) acc_s T.alloc_fragment([block_H, block_N], accum_dtype) acc_o T.alloc_fragment([block_H, dim], accum_dtype) ... T.use_swizzle(10) loop_range T.ceildiv((seqlen_kv // num_split), block_N) for k in T.Pipelined(loop_range, num_stages2): kv_start (seqlen_kv // num_split) * bz k * block_N T.copy(KV[bid, kv_start:kv_end, cur_kv_head, :], KV_shared) T.copy(K_pe[bid, kv_start:kv_end, cur_kv_head, :], K_pe_shared) T.clear(acc_s) T.gemm(Q_shared, KV_shared, acc_s, transpose_BTrue, policyT.GemmWarpPolicy.FullCol) T.gemm(Q_pe_shared, K_pe_shared, acc_s, transpose_BTrue, policyT.GemmWarpPolicy.FullCol) ...注意 MLA 的打分计算被拆成两个 GEMM无 RoPE 部分Q KV与 RoPE 部分Q_pe K_pe都累加进同一个acc_s等效于[Q | Q_pe] [KV | K_pe]^T但避免拼接开销。随后的在线 Softmax 使用scores_max/scores_max_prev/scores_scale/scores_sum/logsum一组 fragment 实现标准的 FlashAttention 数值稳定流程T.copy(scores_max, scores_max_prev) T.fill(scores_max, -T.infinity(accum_dtype)) T.reduce_max(acc_s, scores_max, dim1, clearFalse) for i in T.Parallel(block_H): scores_max[i] T.max(scores_max[i], scores_max_prev[i]) for i in T.Parallel(block_H): scores_scale[i] T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale) for i, j in T.Parallel(block_H, block_N): acc_s[i, j] T.exp2(acc_s[i, j] * scale - scores_max[i] * scale) T.reduce_sum(acc_s, scores_sum, dim1)acc_s经 softmax 后先T.copy到S_shared再读回acc_s_castfloat32→float16为下一步acc_s V做准备T.copy(acc_s, S_shared) T.copy(S_shared, acc_s_cast) for i in T.Parallel(block_H): logsum[i] logsum[i] * scores_scale[i] scores_sum[i] for i, j in T.Parallel(block_H, dim): acc_o[i, j] * scores_scale[i] T.gemm(acc_s_cast, KV_shared, acc_o, policyT.GemmWarpPolicy.FullCol)循环结束后对acc_o做归一化并把每段的 logsum 写入全局缓冲区glse、部分输出写入Output_partialfor i, j in T.Parallel(block_H, dim): acc_o[i, j] / logsum[i] for i in T.Parallel(block_H): logsum[i] T.log2(logsum[i]) scores_max[i] * scale T.copy(logsum, glse[bid, hid * VALID_BLOCK_H : (hid 1) * VALID_BLOCK_H, bz]) T.copy(acc_o, O_shared) T.copy(O_shared, Output_partial[bid, hid * VALID_BLOCK_H : (hid 1) * VALID_BLOCK_H, bz, :])Combine Kernel合并 Split-KV 结果main_split内嵌的第二个T.Kernel(heads, batch, threads128)负责把num_split段的部分输出按各自 logsum 重新加权合并与 FlashDecoding 的 combine 阶段同理lse_max_local -T.infinity(accum_dtype) for k in T.serial(num_split): lse_max_local T.max(lse_max_local, glse[bz, hid, k]) for k in T.Pipelined(num_split, num_stages1): lse_local_split glse[bz, hid, k] lse_logsum_local T.exp2(lse_local_split - lse_max_local) lse_logsum_local T.log2(lse_logsum_local) lse_max_local for k in T.serial(num_split): ... scale_local T.exp2(lse_local_split - lse_logsum_local) o_accum_local[i] po_local[i] * scale_local它先求出各 split 段 logsum 的最大值lse_max_local再以该最大值为锚点做在线 logsum 累加最终按exp2(lse_split - lse_logsum)加权求和各段部分输出保证数值稳定。参考实现与正确性验证仓库中的 torch_refs.py 与example_mla_decode.py内置的ref_program给出了纯 PyTorch 参考实现将q/q_pe与kv/k_pe沿最后一维拼接成完整 query/key用einsum计算 scores经F.softmax后与 value 做加权求和。验证方式则是通过kernel.get_profiler(tensor_supply_typetilelang.TensorSupplyType.Randn)生成随机输入调用profiler.assert_allclose(ref_program, rtol1e-4, atol1e-4)校验正确性再用profiler.do_bench(warmup500)统计延迟并换算 TFLOPS。变体内核针对不同的部署场景同目录下还提供了多个变体核心计算逻辑与上述版本一致example_mla_decode_paged.py面向 PagedAttention 的分页 KV cache 版本通过block_table做间接寻址并对cache_seqlens做掩码T.if_then_else(...) - -infexample_mla_decode_persistent.py持久化内核版本Grid 大小固定为 SM 数量在单个 kernel 内以T.serial(waves)循环遍历所有 tile并用T.sync_grid()同步 split 与 combine 两个阶段example_mla_decode_ws.py手动 Warp-Specialization 版本显式使用T.alloc_barrier、T.ptx_cp_async、T.wgmma_gemm、T.set_max_nreg等底层原语experimental/example_mla_decode_kv_fp8.py实验性的 KV 缓存 FP8 量化版本。六大关键优化技术逐项拆解下面回到 README 的核心内容逐项讲解 TileLang 是如何用少量 Python 注解把这些复杂优化落地的并给出仓库源码佐证。1. Layout Inference让编译器替你推导 buffer 布局图 3 与图 4 展示了 MLA 前端的 TileLang 脚本与它对应的执行计划。其中T.gemm表示矩阵乘法transpose_BTrue表示对 B 矩阵转置policyFullCol指定每个 warpgroup 计算一列即沿结果矩阵的垂直方向切分T.copy表示 buffer 间的拷贝。Figure 3Q K 中的 buffer 形状Figure 4acc_s V 中的 buffer 形状从 TileLang 前端代码到执行计划的映射由Layout Inference完成。它是 TileLang 的核心优化技术之一基于 Tile-Operator如T.gemm、T.copy自动推导所需的 buffer 形状与最优布局再生成对应代码。下面以 MLA 中的 buffer 形状推导为例具体说明计算Q K时根据T.gemm上的policyFullCol注解TileLang 推断每个 warpgroup 的acc_s_0形状应为[blockM, blockN / 2]由于随后是policyFullCol的acc_s V它要求每个 warpgroup 拥有完整的acc_s结果于是 TileLang 推断此时acc_s的形状应为[blockM, blockN]继续向前传播T.copy(S_shared, acc_s)中的S_shared与acc_s都应为[blockM, blockN]。可以看到「每个 warpgroup 先算一半 acc_s、再经共享内存补齐另一半」这一人工设计在 TileLang 中只需一个policyFullCol注解即可驱动整个布局链路的自动推导。GemmWarpPolicy定义在 tilelang/tileop/base.py共三种取值Square各维均衡切分、FullRow所有 warp 分到行方向、FullCol所有 warp 分到列方向。值得指出这套调度方案与 FlashMLA 的实现策略不同FlashMLA 把Q K分配给单个 warpgroup而acc_o的切分方式与本文一致。即便如此TileLang 的调度仍能达到与 FlashMLA 相当的性能。2. Threadblock Swizzling一行代码提升 L2 命中率Threadblock swizzling 是 GPU 内核优化中常见的性能技巧。GPU 的 L2 cache 由多个 SM 共享threadblock swizzling 通过重映射 threadblock 的调度顺序来优化数据访问模式从而提升 L2 命中率。传统调度通常按 grid 的自然顺序执行 threadblock这会导致相邻 threadblock 之间的数据访问不连续缓存数据利用率低下swizzle 技术则采用数学映射如对角或交错映射调整执行顺序使连续调度的 threadblock 访问相邻或重叠的数据区域。在 TileLang 中一行 Python 即可开启T.use_swizzle(panel_size: int, order: str row)其中panel_size指定 swizzle 的 threadblock 组宽度order指定 swizzle 模式。从 tilelang/language/annotations.py 的实现看order实际支持row对应rasterization2DRow、column对应rasterization2DColumn、mlx对应rasterization2DMLX三种取值且可通过enableFalse关闭。MLA 内核中实际使用的是T.use_swizzle(10)即 panel_size 为 10 的行优先 swizzle。3. Shared Memory Swizzling消除 bank conflictCUDA 中共享内存被划分为多个 memory bank每个 bank 每时钟周期可并行服务一次线程请求。当多个线程同时访问映射到同一 bank 的不同地址时会发生 bank conflict迫使这些访问串行化降低性能。常用的对策是 shared memory swizzling通过重映射数据在共享内存中的存储方式例如在地址计算中引入 XOR 或其它位运算把原本落在同一 bank 的地址分散到不同 bank使连续线程的访存分布更均匀。这对矩阵乘法、卷积等高吞吐计算尤为重要。TileLang 同样支持共享内存 swizzling同样只需一行T.annotate_layout({ S_shared: TileLang.layout.make_swizzled_layout(S_shared), })T.annotate_layout允许用户为 buffer 指定任意布局为了方便TileLang 提供了make_swizzled_layout原语自动生成 swizzle 布局定义见 tilelang/layout/swizzle.py支持k_major、allow_pad参数。此外该模块还针对不同硬件指令提供了make_volta_swizzled_layout、make_wgmma_swizzled_layout、make_tcgen05mma_swizzled_layout、make_full_bank_swizzled_layout等专用变体。4. Warp-Specialization生产-消费模型自动编排Hopper 架构上常用的性能优化手段是 warp specialization典型做法是指定一个 warpgroup 作为 producer用 TMATensor Memory Accelerator负责数据搬运其余 warpgroups 作为 consumer 负责计算。但这种编程模式很复杂开发者需要手动管理 producer/consumer 的执行逻辑包括通过mbarrier对象做同步。在 TileLang 中用户完全无需关心这些实现细节前端脚本会自动被转换成 warp-specialized 形式所有 producer/consumer 同步均由 TileLang 自动处理。也就是说example_mla_decode.py中普通的T.copyT.gemm写法在 Hopper 后端编译时即可被自动映射为「TMA 搬运 WGMMA 计算」的流水形态。如果确实需要手动控制可以参考 example_mla_decode_ws.py它通过T.alloc_barrier、T.ptx_cp_async、T.cp_async_barrier_noinc、T.wgmma_gemm、T.wait_wgmma等显式原语完整刻画了这套生产-消费流水。5. Pipeline多级流水重叠访存与计算Pipeline 通过重叠访存与计算来提升内存密集型算子的性能。在 TileLang 中用T.Pipelined注解实现T.Pipelined(range: int, stage: int)这里range指定流水线的循环范围stage指定流水级数。多级流水让计算与内存访问重叠执行对内存密集型算子能显著提升性能但级数越高共享内存占用越大需要根据具体场景权衡。MLA 内核中实际使用的是T.Pipelined(loop_range, num_stages2)即 2 级流水。从 tilelang/language/loop.py 的 API 定义看num_stages表示 producer 与 consumer 之间最多使用的 buffer 份数当num_stages0时表示关闭流水。该 API 还支持order、stage、sync、group等手动调度参数便于高级用户精细控制流水结构。6. Split-KV小 batch 下榨干 SM 并行度仓库还实现了类似 FlashDecoding 的Split-KV优化当 batch size 较小时并行度不足会导致 SM 资源无法被充分利用。此时可以把kv_ctx维度切分到多个 SM 上并行计算再合并结果。实现上分为split kernel 与 combine kernel两个阶段用户通过num_split参数控制切分大小详见上文example_mla_decode.py中main_split/ combine 两段代码。当num_split 1时走 splitcombine 路径否则退化为main_no_split单 kernel 路径。此外paged 版本的example_mla_decode_paged.py更进一步把 KV 切分与页表寻址block_table结合并为不同 batch 的变长cache_seqlens计算各自的分块数blocks_per_split与remaining_blocks实现真正面向生产推理场景的变长 Split-KV。如何运行、验证与基准测试运行 MLA Decode 内核直接运行 example_mla_decode.py 即可完成正确性校验与性能统计python examples/deepseek_mla/example_mla_decode.py --batch 132 --heads 128 --kv_heads 1 --kv_ctx 8192 --dim 512 --pe_dim 64脚本支持--batch、--heads、--kv_heads、--kv_ctx、--dim、--pe_dim等命令行参数默认配置为 batch132、128 个 Q head、单 KV head、KV 上下文 8192、head dim 512 RoPE dim 64即 MLA 标准配置 576 512 64。运行后会打印延迟ms与 TFLOPS。paged 变体同理python examples/deepseek_mla/example_mla_decode_paged.py --batch 128 --h_q 128 --h_kv 1 --cache_seqlen 8192 --d 576 --dv 512复现 README 中的多框架基准对比benchmark_mla.py 在统一的 paged KV 布局block_tableblocked_k下封装了 Torch、FlashMLA、FlashInfer、Triton、TileLang 五种实现FUNC_TABLE并内置正确性互检torch.testing.assert_close与带宽/TFLOPS 统计。用法# 对比两个实现 python examples/deepseek_mla/benchmark_mla.py --baseline torch --target tilelang --compare # 只测某一个实现 python examples/deepseek_mla/benchmark_mla.py --target flash_mla --one # 跑全部实现 python examples/deepseek_mla/benchmark_mla.py --all默认 shape 配置为 batch128、s_q1、128 个 Q head、单 KV head、d576、dv512seqlen 覆盖 1024~32768。结果会写入{benchmark_type}_perf.csv。注意 FlashMLA / FlashInfer 属于可选依赖需要按各自官方方式安装后才能参与对比。测试与回归仓库为 MLA 示例提供了完整的测试与回归入口test_example_mla_decode.pypytest 用例标注了requires_cuda与requires_cuda_compute_version_ge(9, 0)即要求 Hopper 及以上的 CUDA 计算能力并注明 CuTeDSL 后端暂不支持alloc_global时会跳过regression_example_mla_decode.py性能回归脚本调用run_regression_perf用 CUPTI 后端计时便于 CI 中监控性能波动。向 AMD MI300X 移植的要点仓库还提供了面向 AMD MI300X 的移植说明 examples/deepseek_mla/amd/README.md 及配套基准脚本benchmark_mla_decode_amd_tilelang.py、benchmark_mla_decode_amd_triton.py、benchmark_mla_decode_amd_aiter.py其要点可归纳为指令集差异MI300X 不需要显式 TMA 与 warp specializationHopper 上由编译器自动完成的处理在 AMD 上同样无感知源码层面几乎无差别共享内存约束MI300X 只有 64KB 共享内存Hopper 为 228KB需要减少软件流水级数并把 Q 矩阵从共享内存改为寄存器缓存T.alloc_shared→T.alloc_fragmentTile 尺寸灵活性MI300X 无 WGMMA 指令block_m 不必是 64 的倍数tile 选择更自由bank conflict 规则差异AMD 与 NVIDIA 的共享内存 bank 规则不同swizzle 策略也不同——这一差异同样由 TileLang 自动处理代码层面无可见区别。AMD 版本报告的性能结论是在大多数测试场景下 TileLang 与手写汇编内核 aiter-asm 基本持平0.73x~1.21x相比 Triton 最高可快约 6.5 倍而实现仅约 70 行 Python。总结以 MLA 为例TileLang 展示了「高表达力 自动优化」的完整路径开发者只需要用T.gemm配合GemmWarpPolicy.FullCol、T.copy、T.Parallel、T.Pipelined等高层原语描述计算意图TileLang 的 Layout Inference 会自动推导 buffer 形状与 warp 切分T.use_swizzle、T.annotate_layout一行注解即可获得线程块级与共享内存级 swizzleWarp-Specialization 与 TMA 的复杂同步被编译器自动编排再叠加 Split-KV 处理小 batch 场景。最终约 80 行 Python 即可获得与 FlashMLA 相当、且明显优于 FlashInfer 与 Triton 的性能——这正是「易用性与高性能兼得」这一 DSL 设计目标的直观体现。如果你正在为 DeepSeek 系模型编写推理内核建议从 example_mla_decode.py 入手理解基础流程再按部署形态paged / persistent / fp8选择对应变体并借助benchmark_mla.py与回归测试在目标硬件上做量化验证。【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表