一个 Triton FlashAttention-2 Prefill 算子与一个 CUDA WMMA FP16 GEMM,从 naive 基线 起,经五个逐步优化的版本迭代而成,在 RTX 4090 上完成实测。
- HGEMM 达 cuBLAS 的 91.7%(8192³,4096³ 为 84.1%),纯手写 Tensor Core 算子。
- 从 SIMT FMA 换到 Tensor Core,吞吐 3.0× 跃升(8192³:37.2 → 111.5 TFLOPS)。
- Triton FlashAttention-2 Prefill 达 PyTorch SDPA-flash 的 0.81–0.89×,峰值 125 TFLOPS,中间显存 O(N)。
- GEMM 五个递进版本,每版只隔离一项优化,逐版独立对 cuBLAS 校验。
实测环境:RTX 4090 D(Ada, sm_89)/ CUDA 12.8(nvcc V12.8.93)/ PyTorch 2.8.0+cu128 / Triton 3.4.0。GEMM 基线为 cuBLAS(C++ driver 内闭环计时); FA2 基线为 PyTorch SDPA,限定 flash 后端。
| M=N=K | v1 | v2 | v3 | v4 | cuBLAS (TFLOPS) |
|---|---|---|---|---|---|
| 1024 | 25.4% | 66.7% | 81.0% | 81.0% | 61.7 |
| 2048 | 22.0% | 59.1% | 70.3% | 70.3% | 117.3 |
| 4096 | 26.0% | 69.2% | 82.8% | 84.1% | 140.8 |
| 8192 | 25.9% | 77.3% | 90.1% | 91.7% | 144.3 |
v4 占 cuBLAS 的比例随规模增大而升高,8192³ 达 91.7%。
在很短的 reduction 维度上(M=N=4096, K=64),v4 反而低于 v3:
| M=N=4096, K=64 | v2 | v3 | v4 |
|---|---|---|---|
| 占 cuBLAS | 80.6% | 86.2% | 53.0% |
K=64 在 BK=32 下只有 2 个 K-tile。swizzle 的 ldmatrix 索引计算与 cp.async
流水线预热都是每个 K-tile 的固定开销,短 K 上无法被足够的主循环迭代摊薄,净收益
转负。这是随形状变化的取舍,不是缺陷——该形状仍 --verify 通过,K 大了 swizzle
就划算。
batch=4,heads=32/8(GQA),head_dim=128,causal,fp16:
| seqlen | Triton (ms) | SDPA (ms) | TFLOPS | speedup |
|---|---|---|---|---|
| 512 | 0.127 | 0.114 | 67.6 | 0.895× |
| 1024 | 0.376 | 0.303 | 91.4 | 0.807× |
| 2048 | 1.260 | 1.059 | 109.1 | 0.841× |
| 4096 | 4.599 | 3.995 | 119.5 | 0.869× |
| 8192 | 17.544 | 15.570 | 125.3 | 0.887× |
在 N≥1024、D=128、causal 上稳定达到 SDPA-flash 的 0.81–0.89×,吞吐随序列增长
升至 125 TFLOPS。关键实现点:online softmax 使中间显存保持 O(N)(不物化 N×N 分数
矩阵);1/l 归一化延迟到循环外一次做(相对 FA1 的关键区别);causal 走两段循环
——先是对角线下方的满速无 mask 段,再只处理与对角线相交的 tile;GQA 从 k.shape
推断 Hkv,不额外增加 API 参数。
以上数据均在 --verify / pytest 通过的前提下测得:GEMM 逐版本对 cuBLAS 按相对
Frobenius 范数校验,FA2 对 SDPA 用双基准容差。
每版只隔离一项优化,性能差异才能归因到具体改动。百分比取自上面的 4096³ 行。
| 版本 | 优化手段 | @4096³ 占 cuBLAS | 相对上一版 | Takeaway |
|---|---|---|---|---|
| v0 | 每线程一个输出元素 | —(慢数个量级,超 2048³ 跳过) | — | 正确性锚点,不碰 Tensor Core |
| v1 | 共享内存 tiling + 8×8 寄存器 tile | 26.0% | — | SIMT FMA 的天花板 |
| v2 | WMMA 三级 tiling(CTA 128×128×32 / warp 64×32 / MMA 16×16×16) | 69.2% | 2.7× | 全项目最大一级台阶 |
| v3 | cp.async 双缓冲 |
82.8% | +13.6 pts | 用异步拷贝掩盖 global→shared 延迟 |
| v4 | XOR swizzle(手写 ldmatrix + mma.sync) |
84.1% | +1.3 pts | 方阵收益小;共享内存 37→32 KB |
需要 Linux、Python 3.10–3.12、CUDA Toolkit 12.x、PyTorch ≥ 2.4、Triton ≥ 3.0、
CMake ≥ 3.24,以及计算能力 8.0 及以上的 NVIDIA GPU——kernel 用了 cp.async,
这是 Ampere(sm_80)引入的特性。
# 安装
python -m pip install -e ".[dev]"
# 构建 GEMM 可执行文件(默认目标 sm_89;sm_80 也可)
cmake -S . -B build && cmake --build build -j
# 正确性
build/csrc/gemm/hgemm_bench --kernel v4 --m 4096 --n 4096 --k 4096 --verify
pytest tests/ -v
# 性能(CSV 落在 benchmarks/results/,已 gitignore)
benchmarks/gemm/run_bench.sh
python benchmarks/prefill/bench_prefill.py
# profiling
profiling/ncu_gemm.sh v4 4096
profiling/ncu_prefill.sh 2048 128M、N、K 须为 16 的倍数,以保证每个 16×16 MMA tile 整块在界内或界外,epilogue 不必 处理半块。
MIT License。