使用 TileLang 实现 FlashAttention:多级内存抽象与在线 Softmax 的完整实战指南

使用 TileLang 实现 FlashAttention:多级内存抽象与在线 Softmax 的完整实战指南 使用 TileLang 实现 FlashAttention多级内存抽象与在线 Softmax 的完整实战指南【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang本篇技术指南以 examples/flash_attention/README.md 为骨架围绕 TileLangtile-lang如何以极简方式表达 FlashAttention 这一复杂融合算子展开先剖析其在不同内存层级定义 Buffer的核心思想再逐行拆解官方 README 中的内核骨架并结合仓库内完整可运行的 example_mha_fwd_bhsd.py 与 example_mha_fwd_varlen.py 讲解网格设计、流水线循环、因果掩码、在线 Softmax、Warp 切分策略GemmWarpPolicy与自动调优实战。读完后你将掌握如何用 TileLang 的T.alloc_shared/T.alloc_fragment规划共享内存与寄存器如何用T.Pipelined重叠访存与计算以及如何对照 PyTorch 参考实现完成正确性校验与性能评测。FlashAttention 为什么值得用 TileLang 表达标准的 Attention 计算流程是S Q·Kᵀ缩放→Softmax(S)→O P·V。朴素实现会把完整的S矩阵写回全局内存显存占用随序列长度平方增长而 FlashAttention 通过分块tiling 在线 Softmaxonline softmax让每个 block 只保存一个block_N × block_M的分数块避免了对完整注意力矩阵的物化同时把显存带宽消耗从 O(N²) 降为 O(N)。在传统 CUDA 中要把逐块读 K/V → 计算分数 → 在线更新 max/sum → 累积输出这一融合逻辑手写出来需要同时管理 shared memory 布局、线程同步、warp 级 GEMM 切分与流水线。而 TileLang 提供的核心抽象正是 README 开头强调的那句话Using tile-lang, we can define buffers at different memory layers. For instance,Q_shared,K_shared, andV_sharedcan be defined in shared memory, whileacc_sandacc_ocan be placed in registers. This flexibility allows us to represent a complex fusion pattern like FlashAttention in a simple way.也就是说显式、分层的 Buffer 声明是这套表达力的根基编译期即可明确每个数据块所处的存储层级后续的存储重写、流水线调度与访存优化都由编译器接管。内存分层抽象shared 与 fragmentTileLang 通过T.alloc_shared与T.alloc_fragment两类分配原语把存储层级直接暴露给开发者。其实现位于 tilelang/language/allocate.pyT.alloc_shared(shape, dtype, scopeshared.dyn)分配共享内存缓冲区用于线程间通信全局内存到片上数据的搬运中转对应 CUDA 的__shared__T.alloc_fragment(shape, dtype, scopelocal.fragment)分配寄存器片段缓冲区用于 GEMM 累加器与逐线程私有数据是性能关键的片上存储同模块还提供alloc_local、alloc_var、alloc_global等分别对应线程私有局部内存、单元素标量与全局工作区。在 FlashAttention 内核中二者的分工非常清晰代码来自 README 骨架# Allocate shared memory for Q, K, V to reduce global memory accesses Q_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) # Allocate buffers on register acc_s T.alloc_fragment([block_M, block_N], accum_dtype) acc_s_cast T.alloc_fragment([block_M, block_N], dtype) acc_o T.alloc_fragment([block_M, dim], accum_dtype) scores_max T.alloc_fragment([block_M], accum_dtype) scores_max_prev T.alloc_fragment([block_M], accum_dtype) scores_scale T.alloc_fragment([block_M], accum_dtype) scores_sum T.alloc_fragment([block_M], accum_dtype) logsum T.alloc_fragment([block_M], accum_dtype)值得注意的工程细节是acc_s注意力分数与acc_o输出累加都使用accum_dtype如T.float32而非输入dtype如T.float16因为 GEMM 累加必须保持足够精度acc_s_cast则用于在第二次 GEMM 前把分数从 fp32 转换回 fp16从而满足矩阵乘法对精度的要求同时利用 fp16 Tensor Core。网格与缓冲区布局一次读懂 Kernel 启动READMME 骨架中的内核通过T.Kernel启动一个三维网格把整个 Attention 问题自然地切分到 GPU 上with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch, threadsthread_num) as (bx, by, bz):bx序列维度的 block 索引每个 block 处理block_M个 query tokenbyhead 维度索引bzbatch 维度索引threadsthread_num每个 block 的线程数。从源码看T.Kernel定义于 tilelang/language/kernel.py它接收 13 个表示gridDim.(x|y|z)的 extent以及表示blockDim的threads参数并以 target-neutral 形式发射为thread_binding循环再由各后端流水线如MaterializeKernelLaunch物化因此同一份内核代码可以面向 CUDA、ROCm、CPU 等多个目标编译。在真实实现 example_mha_fwd_bhsd.py 中网格被组织为T.Kernel(T.ceildiv(seq_q, block_M), heads, batch, threadsthreads)缓冲区还额外增加了一个O_shared T.alloc_shared([block_M, dim], dtype)用于在写回全局内存前暂存输出以提升写回的合并效率T.copy(acc_o, O_shared) T.copy(O_shared, Output[bz, by, bx * block_M : (bx 1) * block_M, :])逐行拆解主循环分块 GEMM 在线 Softmax整个 FlashAttention 的核心在主循环中。README 骨架使用T.Pipelined包裹循环体让编译器自动做多级流水线重叠拷贝与计算并依据是否因果掩码动态计算循环上界loop_range ( T.ceildiv((bx 1) * block_M, block_N) if is_causal else T.ceildiv(seq_len, block_N) ) # Pipeline the loop to overlap copies/gemm stages for k in T.Pipelined(loop_range, num_stagesnum_stages):T.Pipelined定义于 tilelang/language/loop.py其num_stages参数表示生产者和消费者之间最多可用的缓冲区级数0 表示不启用流水线并支持通过order/stage/sync/group做手动精细调度。每次迭代完成四件事1. 加载 K 块并构造分数T.copy(K[bz, k * block_N : (k 1) * block_N, by, :], K_shared) if 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)因果模式下凡是query 位置 key 位置的分数都被置为-infSoftmax 后权重为 0非因果模式直接T.clear(acc_s)清零。这里T.if_then_else与T.Parallel提供了类似 CUDA 中 per-thread 条件写的表达能力。2. Q·Kᵀ 分数 GEMMT.gemm(Q_shared, K_shared, acc_s, transpose_BTrue, policyT.GemmWarpPolicy.FullRow)T.gemm是 TileLang 的同步 GEMM 接口tilelang/language/gemm_op.py。transpose_BTrue表示对K_shared做转置从而以 GEMM 形式计算 Q·KᵀpolicyT.GemmWarpPolicy.FullRow指定 warp 切分策略见下文专节。从注释可知该 GEMM 结果直接保留在寄存器片段acc_s中。3. 在线 Softmaxrescaling 与归一化统计这是 FlashAttention 数值正确性的关键也是 README 骨架注释最密集的部分for i, j in T.Parallel(block_M, block_N): acc_s[i, j] * scale # Save old scores_max, then reset scores_max T.copy(scores_max, scores_max_prev) T.fill(scores_max, -T.infinity(accum_dtype)) # Compute the maximum value per row on dimension 1 (block_N) T.reduce_max(acc_s, scores_max, dim1, clearFalse) for i in T.Parallel(block_M): scores_max[i] T.max(scores_max[i], scores_max_prev[i]) # Compute the factor by which we need to rescale previous partial sums for i in T.Parallel(block_M): scores_scale[i] T.exp2(scores_max_prev[i] - scores_max[i]) # Rescale the partial output accumulation to keep exponents consistent for i, j in T.Parallel(block_M, dim): acc_o[i, j] * scores_scale[i] # Exponentiate (scores - max) for the new block for i, j in T.Parallel(block_M, block_N): acc_s[i, j] T.exp2(acc_s[i, j] - scores_max[i]) # Make a cast of acc_s to fp16 for the next GEMM T.copy(acc_s, acc_s_cast)每一步的数值动机都写在注释里scores_max_prev保存上一块的每行最大值T.reduce_max(..., clearFalse)只归约不初始化输出再与旧最大值取T.max得到当前见过的最大的行最大值scores_scale是旧块与新最大值的指数差用于把上一块的累积和acc_o与logsum缩放到同一指数基准acc_s重新指数化得到归一化前的权重T.copy(acc_s, acc_s_cast)将分数从 fp32 片段拷贝为 fp16 片段供第二次 GEMM 使用。T.reduce_max/T.reduce_sum的实现位于 tilelang/language/reduce_op.py其clear参数控制归约前是否把输出初始化为-infsum 则为 0dim1表示沿block_N维度归约。4. P·V 输出 GEMM 与 logsum 更新T.gemm(acc_s_cast, V_shared, acc_o, policyT.GemmWarpPolicy.FullRow) 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]第二次 GEMM 计算Softmax(Q·Kᵀ)·V并累加到acc_ologsum以与acc_o完全相同的 rescaling 规则递推更新保证最后一步归一化的正确性。5. 收尾除以 logsum 并写回# Final step: divide each partial output by logsum (completing the softmax) for i, j in T.Parallel(block_M, dim): acc_o[i, j] / logsum[i] # Write back the final output block from acc_o to the Output buffer T.copy(acc_o, Output[bz, bx * block_M : (bx 1) * block_M, by, :])至此一个 block 的完整 FlashAttention 前向流程结束分块 GEMM → 在线最大/求和统计 → 指数缩放 → 输出累加 → 最终归一化全程没有物化完整注意力矩阵。GemmWarpPolicyWarp 切分策略的源码级解析README 骨架中两次使用policyT.GemmWarpPolicy.FullRow其语义在 tilelang/tileop/base.py 中有明确定义Square 0在 M/N 两个维度上尽量均衡地分配 warp按矩阵长宽比求最平衡的切分FullRow 1把所有 warp 分配给 M行维度m_warp num_warps, n_warp 1FullCol 2把所有 warp 分配给 N列维度。compute_warp_partitiontilelang/tileop/base.py展示了切分背后的约束每个 warp 在 M 方向至少需要 16 行M % (m_warp * 16)在 N 方向至少需要 8 列N % (n_warp * 8)。若FullRow下 M 无法被m_warp*16整除它会自动把部分 warp 分到 N 维Square策略则遍历所有满足m * n num_warps的(m, n)组合选择每 warp 负载比最接近理想比例的方案。在 FlashAttention 场景中分数矩阵acc_s的维度是[block_M, block_N]如 64×64而block_M通常与block_N同量级采用FullRow可以让每个 warp 负责一整行配合acc_o的逐行 rescaleacc_o[i, j] * scores_scale[i]与logsum[i]的逐行更新warp 间的数据依赖被最小化这也是前向内核在多数配置下的推荐选择。从骨架到可运行内核example_mha_fwd_bhsd.py 实战README 的骨架代码经过tilelang.jit装饰后即可变成可编译、可调优、可评测的完整内核。example_mha_fwd_bhsd.py 给出了标准封装方式autotune(configsget_configs(), warmup10, rep10) tilelang.jit( out_idx[3], pass_configs{ tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, }, ) def flashattn(batch, heads, seq_q, seq_kv, dim, is_causal, block_M64, block_N64, num_stages1, threads128): scale (1.0 / dim) ** 0.5 * 1.44269504 # log2(e) ...要点拆解tilelang.jit(out_idx[3])声明第 3 个参数Output为输出张量JIT 编译后直接返回它pass_configs{tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}开启快速数学优化scale (1.0 / dim) ** 0.5 * 1.44269504这是 README 骨架中先乘 scale 再T.exp2做法的由来——将exp(x·ln2)与exp2(x·log2e)等价变换后编译器可以把乘法和指数运算融合为ffma指令源码注释中明确提到这一优化动机函数闭包内的T.prim_func def main(...)才是真正的内核定义外层 Python 函数充当配置工厂config factoryblock_M、block_N、num_stages、threads都是可调参数。该文件中与 README 骨架的差异在于布局为[batch, heads, seq, dim]BHSD并通过past_len seq_kv - seq_q支持seq_kv seq_q的场景因果循环上界变为T.min(T.ceildiv(seq_kv, block_N), T.ceildiv((bx 1) * block_M past_len, block_N))同时用q_idx bx * block_M i past_len对齐因果掩码。正确性参考实现与性能评测同文件提供了 PyTorch 参考实现ref_programexample_mha_fwd_bhsd.py用torch.einsum(bhqd,bhkd-bhqk, Q, K)计算分数、F.softmax做归一化、再与 V 做einsum因果掩码用torch.tril生成。主流程main函数的核心验证代码为kernel flashattn(batch, heads, seq_q, seq_kv, dim, is_causal, block_M64, block_N64, num_stages1, threads128) profiler kernel.get_profiler() profiler.assert_allclose(ref_program_processed, rtol0.01, atol0.01) print(All checks pass.) latency profiler.do_bench(ref_program_processed, warmup500)assert_allclose在给定的rtol0.01, atol0.01容差内对比 TileLang 内核与 PyTorch 参考输出profiler.do_bench分别对参考实现与 TileLang 内核计时并通过total_flops / latency * 1e-9计算 TFlopsFLOPs 按2.0 * batch * heads * seq_q * seq_kv * dim估算两个矩阵乘共两倍因果模式再乘 0.5命令行入口支持--batch --heads --seq_q --seq_kv --dim --is_causal --tune等参数。自动调优autotune(configs...)装饰器配合--tune参数可以自动遍历配置空间并挑选最优配置。get_configs默认提供block_M[128], block_N[128], num_stages[2], threads[256]调优模式下调用kernel.latency、kernel.config、kernel.ref_latency即可获得最优延迟、最优配置与参考延迟。不同 GPU 世代对配置有不同偏好例如 example_mha_fwd_varlen.py 的注释明确建议 Hopper 架构使用(128, 128, 2or3, 256)。变长序列varlen扩展cu_seqlens 与边界处理真实推理场景中一个 batch 内各序列长度往往不同为此仓库提供了 example_mha_fwd_varlen.py。它把 padded 张量压成 unpadded 形式并额外传入cu_seqlens_q、cu_seqlens_kint32 前缀和数组与max_seqlen_qexample_mha_fwd_varlen.pyQ_unpad: T.Tensor(q_shape, dtype), K_unpad: T.Tensor(k_shape, dtype), V_unpad: T.Tensor(v_shape, dtype), cu_seqlens_q: T.Tensor([batch_size 1], T.int32), cu_seqlens_k: T.Tensor([batch_size 1], T.int32), max_seqlen_q: T.int32, Output_unpad: T.Tensor(o_shape, dtype),内核内部通过q_start_idx cu_seqlens_q[batch_idx]、q_end_idx cu_seqlens_q[batch_idx 1]等从前缀和中还原每个 batch 的起止位置并对越界OOB位置做显式处理拷贝时直接切出可能越界的区间分数初始化时把 OOB 位置置为-1e9而非-inf以避免后续统计受 NaN 干扰输出写回前再按bx * block_M i q_current_seqlen过滤。其参考对比直接使用flash_attn.flash_attn_varlen_func并通过torch.testing.assert_close(out, fla_out, rtol1e-2, atol1e-2)校验。测试与回归体系该示例不是孤立的 demo而是被纳入仓库的测试与回归体系test_example_flash_attention.py 通过tilelang.testing.requires_cuda装饰器批量跑example_mha_fwd_bhsd、example_mha_fwd_bshd、example_mha_fwd_varlen、example_gqa_bwd、example_mha_bwd_bhsd等前向/反向/变长用例其中example_gqa_bwd_tma_reduce_varlen还要求 CUDA compute capability 恰好为 9.0即 Hopper 的 sm_90 与 TMA reduce 特性regression_example_flash_attention.py 则调用各示例的run_regression_perf基于tilelang.testing.process_func并以do_bench(backendcupti)采集性能用于性能回归监控。同目录还包含example_gqa_fwd_bshd.py、example_gqa_bwd.py、example_mha_bwd_bhsd.py、example_mha_bwd_bshd.py、example_gqa_bwd_tma_reduce_varlen.py等 GQA/反向/变长变体以及varlen_utils.py提供generate_random_padding_mask、generate_qkv等数据生成工具和bert_padding.py。若要在本机验证可按仓库说明安装 TileLang 与 CUDA 依赖后直接运行python example_mha_fwd_bhsd.py默认 batch1、heads1、seq256、dim64观察 All checks pass. 与延迟/TFlops 输出或加--is_causal --tune体验自动调优。小结通过examples/flash_attention/README.md这份不足百行的骨架可以看到 TileLang 表达 FlashAttention 的完整范式分层 Buffer 声明T.alloc_shared/T.alloc_fragment把共享内存与寄存器两个层级显式化编译器据此完成存储重写与访存优化T.Kernel三维网格自然映射 batch/head/sequence 三个维度T.Pipelined以极低成本获得多级流水线重叠访存与 GEMMT.gemmGemmWarpPolicy一行完成带 warp 切分策略的矩阵乘法T.ParallelT.reduce_max/reduce_sumT.exp2组合实现数值稳定的在线 Softmax。配合 example_mha_fwd_bhsd.py 的 JIT 封装、自动调优与基准评测以及 varlen 变体与回归测试这套示例既是一份可运行的性能内核也是学习 TileLang 语言能力的进阶教材——从表达一个复杂融合算子到在生产中持续验证其正确性与性能链条完整、开箱即用。【免费下载链接】tilelangDomain-specific language designed to streamline the development of high-performance GPU/CPU/Accelerators kernels项目地址: https://gitcode.com/GitHub_Trending/ti/tilelang创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考