head_dim>128 (256/512) GQA Flash-Decode NKI Kernel for AWS Trainium
Version 1.0.1. A flash-attention online-softmax decode (single-token,
S_q == 1) NKI kernel for GQA attention with head_dim > 128 (e.g. the
Gemma family: 256 for sliding-window layers, 512 for global layers). It handles
head dims that exceed the NeuronCore 128-partition cap by D-tiling, and packs
the GQA query heads that share one KV head onto the matmul output partition.
This kernel is a portable, hardware-proven-correct decode kernel for
head_dim > 128 GQA on Trainium. Note (updated Task 020, SDK 2.32): the stock
attention_block_tkg head_dim>128 decode kernel is numerically correct in
isolation on trn2 (cos ≥ 0.99999 across many reproducers incl. QK-norm, real
RoPE, per-head block table, 2-step autoregressive). Its known failure is
specific: inside a full multi-layer model graph, the stock d256 GQA-fold
decode (update_cache=False + W_out=None + d_head>128 + GQA per-head 3D
table) produces wrong output, with severity scaling by per-rank KV heads — an
in-graph nkilib/compiler defect (reported to AWS), not a math error. This kernel
is a self-contained alternative that sidesteps that in-graph path, so it is
the recommended choice for d>128 GQA decode until the stock path is fixed
upstream. It is an option, not the only-correct option — the stock path works
for configurations that don't hit the full-model GQA-fold defect.
Value proposition — read this first
This kernel's value is CORRECTNESS + LONG-CONTEXT ENABLEMENT, not raw speed.
- Correctness: per-head cosine ~1.0 vs a CPU-FP32 oracle; hardware-validated end-to-end on real Gemma models (table below).
- Long-context enablement: the decode graph compiles and runs at
max_model_len= 2048 / 4096 on a single LNC=2 core, exactly where the eager decode fallback fails to even compile (a hard mask-reshape error atmax_model_len > 1024). This is capability the eager path cannot provide at all. - NOT a speedup. The single-core dispatch (
wrapped[1],n_prgs=1: one program owns all GQA heads and re-streams the whole KV cache on one core each step) is memory-bandwidth-bound. At short context it runs at roughly 0.14–0.8× eager decode tok/s (5–7× slower). Do not enable it for short-context (≤ 1024) serving where the eager path both compiles and is faster. See the perf caveat below.
What ships
A single functional entry point plus its eligibility gate and torch reference, re-exported at package level:
| Symbol | Purpose |
|---|---|
attention_decode(q, k, v, attn_mask, softmax_scale, num_kv_groups, force_torch=False) |
3-tier dispatch: NKI kernel when eligible and on Neuron, else the SDPA torch reference. |
_can_use_nki_kernel(q, k, num_kv_groups) |
Static-shape + device eligibility gate; returns False on CPU or on ineligible shapes. |
_torch_attention_decode(...) |
The PyTorch SDPA reference (fallback + correctness oracle). |
Supported shapes (the eligibility gate)
_can_use_nki_kernel returns True only when all hold:
- tensors on a Neuron device (or NKI sim) — a Neuron NKI runtime is present;
head_dim % 128 == 0(partition-dim D-tiling);S_ctx % 128 == 0(K-tile width);S_decode == 1(single-token decode; spec-decodeS_q > 1is unsupported);- GQA head match:
num_q_heads == num_kv_heads * num_kv_groups.
Otherwise (or with force_torch=True) the call transparently uses the SDPA
torch reference.
Input contract
q: [B, num_q_heads, S_decode=1, head_dim]
k: [B, num_kv_heads, S_ctx, head_dim] # S_ctx % 128 == 0
v: [B, num_kv_heads, S_ctx, head_dim]
attn_mask: [B, 1, S_decode=1, S_ctx] # additive: 0 keep, -inf mask
softmax_scale: float # 1.0 for Gemma
num_kv_groups: int # num_q_heads // num_kv_heads
-> out: [B, num_q_heads, S_decode=1, head_dim]
The kernel internally pre-transposes K to [B, Hkv, D, S_ctx] (K-stationary
MM1) and runs the flash online-softmax accumulation over S_ctx / 128 K-tiles.
Hardware validation on record
| Model / config | Instance / SDK | head_dim | Result |
|---|---|---|---|
| Gemma4-31B-it, TP=32, on-device greedy | trn2.48xl, vLLM beta5 (SDK 2.30) | 256 (SWA) / 512 (global) | 7/7 checks incl. three ~27K-context needle tests (needle@5/50/95%); per-head cosine ~1.0 vs CPU-FP32 |
| Gemma3 8Q/4KV/d256, greedy parity vs CPU-FP32 (12 multilingual prompts) | trn2.3xl, DLAMI 20260818 (SDK 2.32) | 256 | TP=1: 11/12 exact (argmax 12/12), TP=4: 10/12 (> eager TP=4 8/12); compiles+runs decode at max_model_len 2048/4096 where eager can't compile (>1024) |
| NKI-sim parity (TP1/TP4, global + sliding, s_prior 512–2048) | CPU simulator | 256/512 | cos ≥ 0.999997 |
The Gemma4-31B run also drove a coherent long-context serving stack (France→Paris, 2+2=4, and BANANA-7731 needle retrieval at ~27K tokens).
Usage via get_kernel
import torch
import libtorch_neuronx_lite # REQUIRED before get_kernel on the vLLM 0.24 venv:
# registers torch.neuron so the kernels backend
# detector selects `neuron` (not `cuda`).
from kernels import get_kernel
# version=1 resolves the `v1` BRANCH; revision="v1.0.1" pins the exact tag;
# revision="main" tracks latest. There is no zero-arg default — always pass one.
adk = get_kernel(
"jburtoft/attention-decode-d256-neuron-kernels",
version=1,
trust_remote_code=True,
)
# GQA decode, head_dim=256, 8 query heads / 4 KV heads (num_kv_groups=2), S_ctx=512:
import torch
out = adk.attention_decode(
q, # [B, 8, 1, 256] on Neuron
k, # [B, 4, 512, 256]
v, # [B, 4, 512, 256]
attn_mask, # [B, 1, 1, 512] additive (0 keep, -inf mask)
softmax_scale=1.0,
num_kv_groups=2,
) # -> [B, 8, 1, 256]
kernels backend detection requires a torch.neuron-registered torch build
(a PyTorch-Native / vLLM-Neuron env). On the vLLM 0.24 venv (SDK 2.32) torch is a
CUDA-tagged build (2.11.0+cu130) and hasattr(torch, "neuron") is False until
import libtorch_neuronx_lite (or torch_xla) runs — so you must import it
before get_kernel, otherwise the detector picks the cuda variant and refuses
neuron. On a plain CPU box the package still imports (the NKI path is simply
unavailable and attention_decode uses the torch reference).
Perf caveat (measured)
Steady-state decode throughput, NKI vs eager, at the two contexts where both compile (Gemma3 8Q/4KV/d256, TP=1, LNC=2, SDK 2.32):
| max_model_len | batch | eager tok/s | NKI tok/s | NKI / eager |
|---|---|---|---|---|
| 512 | 1 | 15.69 | 2.16 | 0.14× |
| 512 | 8 | 99.0 | 18.66 | 0.19× |
| 1024 | 1 | 15.67 | 2.48 | 0.16× |
| 1024 | 8 | 90.03 | 18.76 | 0.21× |
Single-token decode is memory-bandwidth-bound; the single-core (n_prgs=1)
dispatch re-streams the entire KV cache on one core every step with no
cross-core parallelism. Use this kernel for long context (> 1024), not for
short-context throughput.
Long-context enablement (why it exists)
| max_model_len | eager decode | NKI decode |
|---|---|---|
| 1024 | compiles + runs | compiles + runs |
| 2048 | FAILS to compile (mask reshape) | compiles + runs |
| 4096 | FAILS | compiles + runs |
| 8192 | FAILS | compiles; needs TP=4 / LNC=1 to fit HBM at load (capacity, not graph, limit) |
Future performance follow-on
The single-core wrapped[1] dispatch is the throughput bottleneck. A
multi-program (n_prgs > 1) dispatch that shards the GQA heads / KV across
cores — and/or a windowed-KV gather for the sliding-window layers (only the
last ~window KV columns are needed) — would cut the per-step KV re-stream and
materially close the gap. This is a documented follow-on, not shipped in v1.0.0.
Retirement path
Once AWS fixes the stock attention_block_tkg head_dim>128 in-graph GQA-fold
decode defect (reported; the stock kernel is already correct in isolation on
trn2, so the fix is in the full-model compilation/scheduling path), the stock
NF.attention_decode becomes the preferred path and this kernel can be retired.
Until then this kernel is the recommended portable option for d>128 GQA decode on
Trainium — an alternative that avoids the in-graph defect, not the only correct
kernel.
Requirements
- AWS Trainium (trn2), Neuron SDK 2.30 (beta5) or 2.32 (0.24 vLLM venv).
- A
torch.neuron-registered torch build forget_kernelbackend detection. - NKI (bundled in the Neuron venv). Imports cleanly on beta5 venv, 0.24 venv, or a plain CPU box (venv-agnostic import shim).
License
Apache-2.0. This is an inference runtime kernel package (not a model). No customer-specific code or data.
- Downloads last month
- -