Skip to content

Latest commit

 

History

History
110 lines (79 loc) · 4.5 KB

File metadata and controls

110 lines (79 loc) · 4.5 KB

LLM Inference Kernels

English

一个 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 后端。

HGEMM(FP16,占 cuBLAS 百分比)

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%。

swizzle 不划算的场景

在很短的 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 就划算。

FlashAttention-2 Prefill(对比 SDPA-flash)

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 128

M、N、K 须为 16 的倍数,以保证每个 16×16 MMA tile 整块在界内或界外,epilogue 不必 处理半块。

许可证

MIT License。