flash_attn.cute.interface
flash_attn/cute/interface.py 1 文件 | FA4 的 Python 入口,把 flash_attn_func / flash_attn_varlen_func 派发到具体 kernel 入口:flash_attn_func, flash_attn_varlen_func 依赖:flash_attn.cute.flash_fwd, flash_attn.cute.flash_bwd, flash_attn.cute.tile_scheduler, flash_attn.cute.prepare_scheduler AGENTS.md 与 README 都把它作为 from flash_attn.cute import flash_attn_func 的落点 |
flash_attn.cute.flash_fwd
flash_attn/cute/flash_fwd.py 1 文件 + 多架构变体 | FA4 forward 主 kernel(SM90/SM100/SM120 多架构派发) 入口:flash_fwd_sm90, flash_fwd_sm100, flash_fwd_sm120(文件名即入口,具体类名需验证) 依赖:flash_attn.cute.pipeline, flash_attn.cute.copy_utils, flash_attn.cute.mma_sm100_desc, flash_attn.cute.mask, flash_attn.cute.softmax, flash_attn.cute.tile_scheduler commit 热点集中在 SM100 hd256、ping-pong、MLA 等 forward 变体 |
flash_attn.cute.flash_bwd
flash_attn/cute/flash_bwd.py 1 文件 + 多架构变体 | FA4 backward 主 kernel(dQ/dK/dV 计算) 入口:flash_bwd_sm90, flash_bwd_sm100, flash_bwd_sm120, flash_bwd_preprocess, flash_bwd_postprocess 依赖:flash_attn.cute.flash_fwd(共享 tiling 逻辑), flash_attn.cute.pipeline, flash_attn.cute.tile_scheduler 最近提交在 hd256 backward、deterministic、seqused_q/k 支持上频繁改动 |
flash_attn.cute.flash_fwd_mla
flash_attn/cute/flash_fwd_mla_sm100.py 3 文件 | FA4 MLA(Multi-head Latent Attention)forward,服务 DeepSeek 类模型 入口:flash_fwd_mla_sm100, flash_fwd_mla_1cta_sm100, flash_fwd_mla_1cta_kb64_sm100 依赖:flash_attn.cute.seqlen_info, flash_attn.cute.block_info, flash_attn.cute.pack_gqa, flash_attn.cute.tile_scheduler 最近 commit 把 dense / sparse top-k / paged / 1CTA-2CTA 都合进 MLA dispatch |
flash_attn.cute.flash_bwd_mla
flash_attn/cute/flash_bwd_mla_sm100.py 4 文件 | FA4 MLA backward(dK/dV/dQ 拆分 kernel) 入口:flash_bwd_mla_sm100, flash_bwd_mla_dk_sm100, flash_bwd_mla_dq_dqv_sm100, flash_bwd_mla_dq_dqv_sm100_h64 依赖:flash_attn.cute.flash_bwd, flash_attn.cute.tile_scheduler MLA backward 按 head 数拆分 1CTA/2CTA,是近期热点 |
flash_attn.cute.tile_scheduler
flash_attn/cute/tile_scheduler.py 1 文件 | SM90/SM100/SM120 的 tile 调度器:static / dynamic-persistent / CLC 入口:SchedulerState, ClcSchedulerState, DynamicPersistentSchedulerState, SingleTileVarlenScheduler, Sm100FmhaClcDynamicTileScheduler, compute_sm100_fmha_varlen_grid 依赖:cutlass.pipeline, cutlass.utils.ClcDynamicPersistentTileScheduler, quack.cute_dsl_utils 封装了 CLC 硬件调度器与 dynamic-persistent 软件调度器,是 SM100 性能关键 |
flash_attn.cute.prepare_scheduler
flash_attn/cute/prepare_scheduler.py 1 文件 | 在 host 侧把 (batch, seqlen, heads) 转成 scheduler 元数据(cu_blocks、cu_seqlens、split 信息) 入口:FlashPrepareScheduler, SchedulerMetadataTensorsTorch 依赖:flash_attn.cute.seqlen_info, flash_attn.cute.block_info, flash_attn.cute.cu_blocks_kernel 对应 FA3 的 hopper/flash_prepare_scheduler.cu,是 FA4 的纯 CuTeDSL 重写 |
flash_attn.cute.hardware_pipeline
flash_attn/cute/pipeline.py 1 文件 | 封装 Hopper/Blackwell 的 async pipeline(TMA/UMMA/CpAsync/NamedBarrier) 入口:PipelineStateSimple, make_pipeline_state, PipelineTmaAsync, PipelineUmmaAsync, PipelineClcFetchAsync 依赖:cutlass.pipeline, cutlass._mlir.dialects.nvvm 把 CUTLASS pipeline 包装成 Python-DSL 友好的形式,是 kernel 层的主要依赖 |
flash_attn.cute.fmha_sm100_hd256
flash_attn/cute/sm100_hd256_2cta_fmha_backward_*.py 3 文件 | SM100 headdim=256 的 2CTA fused MHA backward(DK/DV 与 DQ 拆 kernel) 入口:BlackwellFusedMultiHeadAttentionBackwardDKDVKernel, BlackwellFusedMultiHeadAttentionBackwardDQKernel 依赖:flash_attn.cute.tile_scheduler, flash_attn.cute.mask, flash_attn.cute.copy_utils, cutlass.utils.blackwell_helpers 最近提交热点:deterministic、stack 回收、length 推导 |
flash_attn.cute.block_sparse
flash_attn/cute/block_sparsity.py 3 文件 | block-sparse attention 的稀疏度计算与 top-K gather 入口:compute_block_sparsity, topk_gather_kv(同目录 compute_block_sparsity.py、topk_gather_kv.py) 依赖:flash_attn.cute.block_info, flash_attn.cute.seqlen_info 配合 benchmarks/benchmark_sparse_mla_fwd.py 与 AI/SPARSE_MLA_*.md 使用 |
flash_attn.cute.cute_dsl_build
flash_attn/cute/cute_dsl_ptxas.py 2 文件(含 cute_dsl_utils.py) | 把 CuTeDSL @jit 函数编译到 PTX/CUBIN,管理持久化 cache 入口:需验证(应为 compile / ptxas 封装) 依赖:nvidia-cutlass-dsl, apache-tvm-ffi, torch AGENTS.md 提到的 FLASH_ATTENTION_CUTE_DSL_CACHE_ENABLED 与 two-pass 测试都依赖它 |
flash_attn.cute.testing
flash_attn/cute/testing.py 1 文件 | FA4 测试基础设施:FakeTensor 模式、env 装饰器、kernel 变体遍历 入口:maybe_fake_tensor_mode, is_fake_mode 依赖:flash_attn.cute.cute_dsl_build, torch._subclasses.fake_tensor tests/cute/*.py 都依赖它实现 two-pass 编译/执行分离 |
csrc.flash_attn_ck
csrc/flash_attn_ck/ 约 10 文件 | ROCm CK 后端:FA2 的 AMD 实现(mha_fwd/mha_bwd/varlen) 入口:mha_fwd, mha_bwd, mha_varlen_fwd, mha_varlen_bwd(flash_api.cpp) 依赖:composable_kernel, ATen, HIP 最近修复了 num_splits heuristic;README 明确这是 ROCm 默认后端 |
hopper.flash_attn_3
hopper/ 507 文件 | FA3(Hopper-only,C++/CUDA,FP16/BF16/FP8) 入口:flash_fwd_kernel_sm90, flash_bwd_kernel_sm90, flash_api.cpp, flash_api_stable.cpp 依赖:CUTLASS SM90 TMA/WGMMA, CUDA >= 12.3 README 标 beta;仍被维护但活跃开发已切到 FA4 |
csrc.flash_attn
csrc/flash_attn/ 约 197 文件(含 src/) | FA2 原始 CUDA 实现(Ampere/Ada/Hopper) 入口:flash_fwd_kernel, flash_bwd_kernel, flash_api.cpp, generate_kernels.py 依赖:CUDA, CUTLASS 2.x 老一代,仍被 PyTorch 生态大量依赖,但提交极少 |
flash_attn.cute.bench
flash_attn/cute/benchmark.py 2 文件(benchmark.py、bench_utils.py) | FA4 端到端 benchmark 入口 入口:benchmark_flash_attention(目录级,具体函数需验证) 依赖:flash_attn.cute.interface, flash_attn.cute.bench_utils 与 benchmarks/bench_sm90.py、benchmark_sparse_mla_fwd.py 互补 |