本仓库只实现一个实验:NSA backward fan-in skew on H100。
研究问题是:当四种 routing workload 的有效 causal sparse-edge 数完全相同, 但 selected KV blocks 的 query fan-in 分布不同时,FLA 社区 NSA backward 的 dQ、CSR 与 fused dK/dV latency 是否显著不同?
本仓库不复制、不修改、不重写 NSA kernel:
| Work unit | 实际调用 |
|---|---|
| selection forward | FLA parallel_nsa_fwd |
| delta preprocess | FLA parallel_attn_bwd_preprocess |
| query→block graph transpose | FLA prepare_block_csr |
| dQ core | FLA parallel_nsa_bwd_kernel_dq |
| fused dK/dV core | FLA parallel_nsa_bwd_kernel_dkv |
| complete backward | FLA parallel_nsa_bwd |
| small numerical oracle | FLA naive_nsa_selection |
| staged-launcher oracle | FLA parallel_nsa_bwd |
我们只负责:
- 生成
local / uniform / sink-heavy / zipfrouting; - 保证四组具有完全相同的有效 causal edges;
- 按官方 launcher 的参数和 grid 调用内部 kernel,以便分阶段计时;
- correctness、fan-in/CSR 统计、显存口径、JSON/CSV 与 profiler 入口。
FLA 固定 commit:
0a9b9f222e86b9a895c2447767e9b4cce6c8d530
q [1,T,64,192] BF16
k [1,T,4,192] BF16
v [1,T,4,128] BF16
dO [1,T,64,128] BF16
block_indices [1,T,4,16] INT32
T 32768
block_size 64
早期 query (t) 只能选择 (\min(16,\lfloor t/64\rfloor+1)) 个 causal blocks。因此真实 edge 数不是机械的 (T H_{kv}16=2{,}097{,}152),而是
[ 4\sum_{t=0}^{32767}\min(16,\lfloor t/64\rfloor+1) =2{,}066{,}432. ]
四种 pattern 使用完全相同的 per-query valid-slot 数,因此其有效 edge 数严格 相同。
git submodule update --init --recursive
# 安装已经在 H100 PCIe 上通过 check/smoke 的固定版本
pip install -r requirements-h100.txt
pip install --no-deps -e third_party/flash-linear-attention
pip install -e .
export TRITON_F32_DEFAULT=ieee
export FLA_CACHE_RESULTS=0
export FLA_DISABLE_BACKEND_DISPATCH=1
export TRITON_CACHE_DIR="$PWD/.triton-cache"不要同时保留另一个 fla-core / flash-linear-attention wheel,以免导入错误实现。
# 1. 硬件、依赖、commit 与 shape gate
make check
# 2. 两级正确性门禁:
# T=256:local 与官方 naive oracle 比较 O/dQ/dK/dV
# T=2048:四种 routing 的 staged launcher 分别与官方 wrapper 比较 dQ/dK/dV
make smoke
# 3. T=32768,先用 local 完成一次 autotune,再以同一 config 测四种 routing
make bench若 runtime 放在默认的 /var/tmp/$USER/nsa-exp,可用包装脚本一次性设置单卡与
cache 环境:
scripts/run_h100.sh make check
scripts/run_h100.sh make smoke
scripts/run_h100.sh make bench指定输出:
PYTHONPATH=src python scripts/bench_backward.py \
--config configs/e1_h100.json \
--output results/raw/e1_h100.json单阶段 Nsight Compute:
ncu --set full \
--nvtx \
--nvtx-include "nsa_local-vs-sink-heavy_dkv_core/" \
-o profiles/e1_local_vs_sink_dkv \
python scripts/profile_stage.py \
--config configs/e1_h100.json \
--pattern local-vs-sink-heavy \
--stage dkv_core \
--benchmark-result results/raw/e1_h100.json完整 profiler 口径见 docs/NSIGHT.md。
精确的单 CTA latency 需要修改 kernel instrumentation,本实验明确不这样做。
因此 kernel tail 证据来自完整 dKV kernel latency、fan-in tail 与 Nsight
wave/occupancy/stall 指标,不能标成“最慢 CTA 时间”。
bench_backward.py 保存:
- 每个 stage 的全部 raw CUDA-event observations;
- median/p10/p90/p95;
- mean/max/p95/p99 fan-in、CV、Gini;
- CSR row query-span 与 adjacent-gap;
delta + CSR持久辅助缓冲、预分配梯度缓冲;- official wrapper 的 output-inclusive peak allocated/reserved memory;
- FLA/Torch/Triton/GPU/driver 元数据;
- Triton autotune config(运行时若可读取)。
生成 Blog 主图:
pip install -r requirements-analysis.txt
python scripts/plot_fanin_latency.py \
--input results/summaries/e1_h100.csv \
--x fanin_cv正式结果尚未产生。
Fan-in 分布统计覆盖 dKV grid 的全部 CSR rows,包括 fan-in 为 0 的 rows。首次 H100 验收后,还需根据 raw JSON 记录的 Torch/Triton 版本生成环境 lock;当前 requirements 只复用固定 FLA commit 声明的最低版本。