Skip to content

Repository files navigation

NSA-KernelExp

本仓库只实现一个实验: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 / zipf routing;
  • 保证四组具有完全相同的有效 causal edges;
  • 按官方 launcher 的参数和 grid 调用内部 kernel,以便分阶段计时;
  • correctness、fan-in/CSR 统计、显存口径、JSON/CSV 与 profiler 入口。

FLA 固定 commit:

0a9b9f222e86b9a895c2447767e9b4cce6c8d530

Frozen primary shape

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 数严格 相同。

H100 setup

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,以免导入错误实现。

Execution

# 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 时间”。

Outputs

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 声明的最低版本。

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages