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 コメント · GitHub で見る

AMD / HIPNVIDIA / CUDAModels & quants

本文

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

関連リンク

インストール・モデル・リリースへの站内リンク。