Pull requests / #1367
#1367 prompt attention: STRATA_PROMPT_ATTN_IMMA=1 - the int8-KV kernel on INT8 tensor cores (opt-in, closer to FP32)
open · @sergqwer · 0 comentarios · En GitHub
AMD / HIPNVIDIA / CUDAModels & quants
Descripción
The int8-KV prompt attention (`prompt_attn_i8_kernel`, v2) drifts away from the FP32 kernel on short prompts. On IQ2_XS the first-token KL to the FP32 kernel is 0.025 after a 2K prompt, and the 32 greedy tokens change. This PR adds an opt-in kernel (`STRATA_PROMPT_ATTN_IMMA=1`) that puts both products on INT8 tensor cores with an operand split exact to 24 bits. It is closer to FP64 than v2 at every shape tested and 1.35-1.53x faster per launch. Unset (the default), nothing changes. ### What it does `prompt_attn_i8v3_kernel` keeps v2's block, its cp.async pipeline and its online softmax. - **Operands.** QK^T and PV run as `mma m16n8k32` on the int8 K and V codes as stored, with no int8 -> fp16 conversion. The other operand is split into three exact 8-bit parts: - q is scaled per (head, 64-dim group) and p·v_scale per (row, chunk), both to 127; - then X = rint(x · 65536) = h · 65536 + l · 256 + m, with h and l signed bytes and m an unsigned one (a `u8 x s8` MMA). That is 24 bits, finer than v2's FP16 hi + lo split (~22 bits). Each part's int32 sum is exact, and the three parts are combined in FP32. - **Occupancy.** One K stage and one V stage per warp (K(c+1) is gathered during softmax(c) and PV(c)), 12-row partial sums, and `__launch_bounds__(128, 3)`: ~25.5 KB of shared memory and 168 registers, so 3 blocks per SM instead of v2's 2. The kernel waits on latency, not on its MMAs: with one block per SM it ran 1.54x slower than with two. - **Scope.** sm_80+ only: the device code is behind the file's `STRATA_PA_SM80`, and Turing / Volta / HIP keep their kernels. The file compiles for 75/80/86/89/120. A test hook, `qsa_prompt_attn_set_imma()`, is added. `qsa_prompt_attn_parity` now checks IMMA against FP64 and v2 and times both. ### Measured RTX 5090, Ryzen 9 9950X3D, IQ2_XS (ISTA GSQ-RCO), `--kv int8`, `--expert-cache 12000`. **Kernel accuracy against FP64** (`qsa_prompt_attn_parity`, max relative error, 0 failures): | context / queries | IMMA | v2 | FP32 kernel | |---|---|---|---| | 32K / 2048 | **2.24e-6** | 3.17e-6 | 2.51e-6 | | 1500 / 1500 | **2.15e-6** | 3.30e-6 | 1.96e-6 | | 2100 / 256 | **2.64e-6** | 2.78e-6 | 1.90e-6 | **End to end, first-token KL to the FP32 kernel** (`STRATA_PROMPT_ATTN_OLD=1`); the last column says whether the 32 greedy tokens match the FP32 kernel's: | prompt | v2 | IMMA | tokens = FP32 (v2 / IMMA) | |---|---|---|---| | 600 | 0.0010 | 0.0016 | yes / yes | | 2K | 0.0248 | **0.0026** | no / **yes** | | 8K | 0.0009 | 0.0028 | no / **yes** | | 32K | 0.0340 | **0.0196** | yes / no | - Median KL: 0.0027 for IMMA against 0.013 for v2. - The top token is the same everywhere. - The greedy tokens match the FP32 kernel's on 3 of 4 prompts with IMMA, and on 2 of 4 with v2. On the same model in NVFP4 (our fork) the gap at 2K was wider: KL 0.10 for v2 against 0.007 for IMMA. **Speed:** - Per launch (`qsa_prompt_attn_parity`): 32K / 2048 queries 2.09 -> 1.55 ms (1.35x); 1500 / 1500 0.61 -> 0.40 ms (1.53x); 2100 / 256 0.28 -> 0.26 ms (1.09x). - In the engine, 32K prompt in 8K chunks (nsys, 48 launches): 358.1 -> 311.6 ms (1.15x). - Prompt time, 3 interleaved pairs, which the machine's noise dominates: | | v2 | IMMA | |---|---|---| | 2K | 701 / 764 / 853 ms | 713 / 832 / 806 ms | | 32K | 5091 / 5063 / 5419 ms | 5191 / 5208 / 4915 ms | The kernel saves ~47 ms of a ~5 s 32K prompt, which is below that noise. **Identity.** With the switch unset, the first-token logits and the 32 greedy tokens equal v0.1.40.2 on IQ2_XS after 2K and after 32K (identical bytes, repeated). 🤖 Generated with [Claude Code](https://claude.com/claude-code) https://claude.ai/code/session_01VZy1yKaDDiA8a7svdwaHio
En el sitio
Enlaces a install, modelos, releases.