Pull requests / #800
#800 qsa_prompt_attn: Volta (sm_70) m8n8k4 kernel - keep the hi and lo halves in separate accumulator chains
closed · @eelgaev · 0 コメント · GitHub で見る
BenchmarksSetup & installMulti-GPUNVIDIA / CUDAModels & quantsDocumentation
本文
Small follow-up to #600. Thanks @fks for the kernel, it's a great speedup on V100. ### What's going on The m8n8k4 kernel splits q (for q·k) and p (for p·v) into an FP16 hi half and a lo half, so the products come out close to FP32 precision. Right now both halves go into the same accumulator: ```cpp mma884(tg[rb], ah.x, ah.y, b0, b1); // hi mma884(tg[rb], al.x, al.y, b0, b1); // lo, same accumulator ``` Volta's tensor cores truncate the running sum as they accumulate. They don't round it like a plain FP32 add would. By the time the lo products arrive, the accumulator already holds the hi products, which are orders of magnitude larger, so most of the lo contribution gets cut off. The split ends up buying much less than it should. That's why the kernel sits around 5e-6 against FP64 while the FP32 kernel it replaces is around 2e-6. It's probably also part of why the 99k document in the #600 discussion came out a bit above the controls. ### The fix The lo products get their own accumulator (`tgl` for the scores, `tmpl` for p·v), and the two are added in plain FP32 at the end. Nothing else changes: same tiling, same shared memory, same launch. Register counts are identical before and after (166 / 226 / 163 for modes 0 / 1 / 3), with no spills. ### Numbers `qsa_prompt_attn_parity`, one V100-SXM2-16GB, 2,048 queries, error vs the FP64 reference as the tool reports it: | context | FP32 kernel | before (one chain) | after (two chains) | kernel time | |---|---:|---:|---:|---:| | 8,192 | 2.14e-6 | 5.88e-6 | **3.02e-6** | 11.94 → 12.20 ms | | 32,768 | 2.37e-6 | 5.08e-6 | **2.33e-6** | 12.20 → 12.38 ms | | 131,072 | 1.81e-6 | 5.35e-6 | **2.75e-6** | 12.34 → 12.51 ms | The fix brings the kernel to the FP32 kernel's accuracy for about 1.5-2% more kernel time. Attention is a small slice of a prompt, so I'd expect that to vanish end to end. I didn't A/B prompt speed for this change on its own, though. On a real model it moves in the same direction. Teacher-forced with `STRATA_LOGPOS` (600 positions after a 32K and a 2K prompt, KL vs an FP32-path reference, UD-Q4_K_XL, int8 KV, 4x V100), KL went 0.0240 → 0.0220 on the 32K text and 0.0168 → 0.0154 on the 2K one. That setup also runs some of our fork's own opt-in kernels, so take it as a direction rather than a clean A/B. The parity table above is the clean comparison. ### How it was tested This was tested on our AC922 fork ([eelgaev/Strata-AC922](https://github.com/eelgaev/Strata-AC922), branch `ac922`: an IBM AC922 with 2x POWER9 and 4x V100-SXM2 16 GB). main doesn't build on ppc64le yet, so here's what I did: - compiled this patch against main (the kernel object builds clean, no new warnings); - ran the parity tool from the fork, where this kernel is the same code as main's apart from this change. The before/after rows above come from that tree with the two accumulator lines toggled; - the fork's launch configuration has used the fixed kernel since then: full 4-GPU benchmark runs, prompts up to 252K tokens, decode text checks, and the teacher-forced runs above. To reproduce: ```sh cmake -S . -B build -DSTRATA_EXPERIMENTAL_SM60=ON -DSTRATA_PARITY_PROMPT_ATTN=ON ... cmake --build build --target qsa_prompt_attn_parity CUDA_VISIBLE_DEVICES=0 ./build/qsa_prompt_attn_parity 32768 2048 5 ``` V100 only. sm_60 never takes this kernel (no m8n8k4 there). I also updated the error figure in `docs/NVIDIA_V100.md`.
関連リンク
インストール・モデル・リリースへの站内リンク。