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 at max_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-decode S_q > 1 is 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 for get_kernel backend 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
-
kernel
neuron
nki
trainium
attention
decode
flash-attention
gqa
head-dim-256
head-dim-512
gemma
long-context
bf16
apache-2.0