Instructions to use jburtoft/kda-neuron-kernels with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Kernels
How to use jburtoft/kda-neuron-kernels with Kernels:
# !pip install kernels from kernels import get_kernel kernel = get_kernel("jburtoft/kda-neuron-kernels") - Notebooks
- Google Colab
- Kaggle
File size: 19,781 Bytes
dab7980 c8ae402 dab7980 bf6d2ce dab7980 c8ae402 dab7980 6a4016a c8ae402 6a4016a c8ae402 6a4016a dab7980 bf6d2ce dab7980 bf6d2ce dab7980 c8ae402 6a4016a c8ae402 6a4016a c8ae402 6a4016a c8ae402 6a4016a c8ae402 dab7980 bf6d2ce dab7980 bf6d2ce dab7980 c8ae402 dab7980 6a4016a c8ae402 dab7980 bf6d2ce dab7980 bf6d2ce dab7980 bf6d2ce dab7980 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 | ---
license: apache-2.0
library_name: kernels
tags:
- kernel
- neuron
- trainium
- kda
- linear-attention
- fla-core
- training
- backward
---
# kda-neuron-kernels
**Neuron NKI kernels for KDA (Kernel-based Decomposed Attention) linear attention.**
Model-agnostic implementation of the KDA algorithm described in the [flash-linear-attention (fla-core) library](https://github.com/fla-org/flash-linear-attention). Compatible with any HuggingFace Transformers model whose attention layer follows the KDA algorithm.
Compatible with AWS Trainium (trn2 tested; trn1 / trn3 not yet verified). Runs under PyTorch Native (Beta 3+ / Beta 4).
## What this package provides
**Inference forward kernels** (v1.0/v1.1) — three raw NKI entry points:
- **`kda_recurrent_fwd(q, k, v, g, beta)`** — TKG (decode) per-token recurrence. One (batch, head) invocation processes S tokens sequentially.
- **`kda_recurrent_fwd_state(q, k, v, g, beta)`** — same as above, also returns the final recurrent state for CTE→TKG hand-off.
- **`kda_chunk_step(q, k, v, beta_broadcast, g_cumsum, g_last, state_in)`** — CTE (prefill) per-chunk step. Processes one 128-token chunk given the state from the previous chunk.
**Training / backward** (v1.2/v1.3) — differentiable `torch.autograd.Function` wrappers plus the underlying backward kernels:
- **`kda_recurrent(q, k, v, g, beta, initial_state=None)`** → `(output, final_state)`, differentiable. `loss.backward()` routes through the NKI recurrent backward.
- **`kda_chunked(q, k, v, g, beta, initial_state=None)`** → `(output, final_state)`, differentiable. Loops chunks in Python around the single-chunk kernels.
- **`kda_chunked_fused(q, k, v, g, beta, initial_state=None)`** (v1.3) → same numerics as `kda_chunked`, but processes all chunks in **one NKI launch** per direction (state carried in SBUF). **2–5× faster** on hardware for multi-chunk sequences (speedup grows with sequence length). Recommended for training.
- Raw kernels also exported: `kda_recurrent_fwd_v2`, `kda_chunk_step_v2`, `kda_recurrent_bwd`, `kda_chunk_bwd`, and the fused `kda_fused_chunked_fwd`, `kda_fused_chunked_bwd`.
Backward gradients verified against fla-core autograd on both the NKI simulator (trn2-pinned) and real trn2 hardware — see the Parity and Performance sections.
**Not yet provided**: a full `NeuronKDA(nn.Module)` drop-in replacement for HF Transformers' `Kda` layer (gated on the upstream transformers KDA integration finalizing).
For a full-layer example, see the sibling kernel package: [`jburtoft/qwen35-deltanet-neuron-kernels`](https://huggingface.co/jburtoft/qwen35-deltanet-neuron-kernels) which does the equivalent for Qwen3.5's Gated DeltaNet.
## Installation & Environment
Requires:
- **PyTorch Native Beta 3+** (torch-neuronx 2.11.3+, PyTorch 2.11+)
- **NKI ≥ 0.4.0** (tested on NKI 0.5.0 / SDK 2.31 / Beta 4)
- **`transformers` with `KernelConfig` support** (main branch commit ≥ `4b0a02931b`, i.e. `5.10.0.dev0` or later)
- **`kernels==0.15.2`**
Consult [Neuron documentation](https://awsdocs-neuron.readthedocs-hosted.com/) for the current Beta setup.
## Usage
### Direct kernel calls (v1.0 supported path)
```python
import torch
import torch.nn.functional as F
from kda_neuron_kernels.build.torch_neuron import (
kda_recurrent_fwd, kda_recurrent_fwd_state, kda_chunk_step,
)
# Example: prefill one chunk of 128 tokens for a single (batch, head) slice
B, S, H, Dk = 1, 128, 64, 128
q_raw = torch.randn(B, S, H, Dk)
k_raw = torch.randn(B, S, H, Dk)
v = torch.randn(B, S, H, Dk) # v_dim == k_dim == 128
g = -torch.rand(B, S, H, Dk) * 0.01 # per-dim log-decay, negative
beta = torch.rand(B, S, H) # per-head scalar
# Wrapper preprocessing:
# - L2-norm q, k and scale q by 1/sqrt(Dk) (fla-core convention)
q = F.normalize(q_raw, p=2, dim=-1) * (Dk ** -0.5)
k = F.normalize(k_raw, p=2, dim=-1)
# Dispatch chunked kernel per (b, h)
state = torch.zeros(Dk, Dk, dtype=torch.float32).to("neuron")
for b in range(B):
for h in range(H):
q_c = q[b, :, h].contiguous().to("neuron") # (128, 128)
k_c = k[b, :, h].contiguous().to("neuron")
v_c = v[b, :, h].contiguous().to("neuron")
g_c = g[b, :, h].contiguous()
gc = torch.cumsum(g_c, dim=0).to("neuron")
gl = gc[-1:, :].expand(128, Dk).contiguous()
beta_c = beta[b, :, h]
beta_bc = beta_c.unsqueeze(-1).expand(128, Dk).contiguous().to("neuron")
chunk_out, state = kda_chunk_step(q_c, k_c, v_c, beta_bc, gc, gl, state)
# chunk_out: (128, 128) float32 per-token output
# state: (128, 128) carries to next chunk
```
For a fully-worked example (with all wrapper preprocessing done in PyTorch),
see `tests/example_usage.py` in this repo.
### Via `KernelConfig` (v1.1 planned)
v1.1 will support the following API (currently blocked on `NeuronKDA` full-layer
wrapper implementation, planned once the upstream transformers KDA integration
finalizes):
```python
from transformers import AutoModelForCausalLM, KernelConfig
kernel_config = KernelConfig({
"Kda": "jburtoft/kda-neuron-kernels:NeuronKDA",
})
model = AutoModelForCausalLM.from_pretrained(
"<your-kda-model>", # any HF model whose attention layer is KDA
dtype=torch.bfloat16,
kernel_config=kernel_config,
device_map="neuron",
trust_remote_code=True,
)
```
### Training (v1.2 — differentiable)
```python
import torch, torch.nn.functional as F
from kda_neuron_kernels.build.torch_neuron import kda_chunked_fused, kda_recurrent
S, D = 256, 128 # S must be divisible by 128 for the chunked paths
q = F.normalize(torch.randn(S, D), p=2, dim=-1) * (D ** -0.5)
k = F.normalize(torch.randn(S, D), p=2, dim=-1)
v = torch.randn(S, D) * 0.3
g = -torch.rand(S, D) * 0.01 # per-K raw log-decay
beta = (torch.rand(S, 1) - 0.5 + 1.0).expand(S, D).contiguous() # per-token scalar bcast
for t in (q, k, v, g, beta):
t.requires_grad_(True)
# CPU -> simulator; or .to("neuron") for trn2 hardware.
# kda_chunked_fused: single NKI launch for all chunks (fastest). kda_chunked: Python chunk loop.
out, final_state = kda_chunked_fused(q, k, v, g, beta, initial_state=None)
loss = out.sum()
loss.backward() # gradients flow through the fused NKI backward
# q.grad, k.grad, v.grad, g.grad, beta.grad now populated
# Decode-style recurrent training is also available (zero initial_state only):
# out, fs = kda_recurrent(q, k, v, g, beta)
```
Per (batch, head): loop `B*H` in the caller. `kda_recurrent` currently requires
zero `initial_state`; use `kda_chunked` for state carry-over across sequence packs.
## Wrapper Contract (Read This)
**All three kernels take raw q, k already L2-normed by the caller**, with `q` additionally scaled by `1/sqrt(dk)` (fla-core convention). The kernels compute all decay-related pre-scaling (`exp(gc_mean)`, `exp(-gc_mean)`, `exp(gc)`, `exp(g_last-gc)`, `exp(g_last)`) internally from `gc = cumsum(g)`.
For `kda_chunk_step`:
- `q`, `k`: RAW L2-normed q, k (q scaled by 1/sqrt(dk)) — shape `(128, 128)`
- `v`: raw value tensor — shape `(128, 128)`
- `beta_broadcast`: per-token scalar beta broadcast to `(128, 128)`
- `g_cumsum`: per-dim cumsum(g) within the chunk — shape `(128, 128)`
- `g_last`: `g_cumsum[-1:, :]` broadcast to `(128, 128)`
- `state_in`: recurrent state from previous chunk — shape `(128, 128)`
- Returns: `(chunk_out (128, 128), state_out (128, 128))`
For `kda_recurrent_fwd`:
- `query`, `key`: `(S, 128)` RAW L2-normed
- `value`, `g_in`, `beta_in`: `(S, 128)` per-dim / broadcast form (see `nki_kda.py` docstring)
- Returns: `output (S, 128)` per-token output
## Hard Constraints
- **`head_k_dim == head_v_dim == 128`** — matches NeuronCore SBUF partition width. Not currently portable to other head dims.
- **`chunk_size == 128`** for `kda_chunk_step`.
- **float32 inputs** (kernel internally handles precision).
- **trn2 tested**. trn1 / trn3 not verified in v1.0.
## Parity
Measured against fla-core `naive_recurrent_kda` and `naive_chunk_kda` PyTorch references on random inputs at a typical KDA operating regime (`g_scale=0.01`, seq_len=128, single (b, h)):
| Kernel | cos_sim vs fla reference | max_abs_diff |
|--------|-------------------------|--------------|
| `kda_recurrent_fwd` (S=128) | **1.00000** | 3.4e-8 |
| `kda_chunk_step` (C=128) | **0.99988** | 1.22e-3 |
Precision floor sources:
- Recurrent: after the v1.1 K-vs-V axis fix (see Fix history), the recurrent kernel matches fla to `max_abs_diff ~= 3e-8` (fp32 rounding floor) across S = 16, 128, 256, 512. The v1.0 number (0.99977) was dominated by the state-decay axis bug, not by BF16 accumulation.
- Chunked: scalar-mean approximation in the intra-chunk attention (`exp(gc_mean_i - gc_mean_j)` instead of exact per-dim `exp(gc_i - gc_j)`). A deliberate O(BT^2) vs O(BT^2 * K) tradeoff; the ceiling (~0.99988) is inherent to the algorithm choice and unaffected by the v1.1 fix.
Note: an earlier version of the chunked kernel (predating v1.0) had a latent accuracy bug (cos_sim ≈ 0.78 on random inputs at typical KDA g-scale). This bug was fixed in v1.0. A second axis-convention bug (state decayed per-V instead of per-K) affecting both kernels was fixed in v1.1. See "Fix history" below.
## Performance
Measured on trn2.3xlarge, LNC=2, single logical core, PyTorch Native Beta 4, SDK 2.31, single (batch, head) invocation:
### Chunked (prefill)
| Metric | Value |
|--------|-------|
| Wall-clock per chunk (C=128) | **87 μs** |
| Per-token effective | 0.68 μs |
| Achieved TFLOPS | 1.89 |
| MFU (BF16 peak 158 TFLOPS/LNC=2) | **1.19%** |
| MFU (FP32 peak 40 TFLOPS/LNC=2) | 4.71% |
| Achieved HBM GB/s | 6.77 |
| MBU (empirical peak 218 GB/s/LNC=2) | 3.11% |
### Recurrent (decode)
| Metric | Value |
|--------|-------|
| Wall-clock per call (S=128) | **838 μs** |
| Per-token (amortized) | 6.55 μs |
| Achieved TFLOPS | 0.023 |
| MFU (BF16 peak 158 TFLOPS/LNC=2) | **0.01%** |
| MFU (FP32 peak 40 TFLOPS/LNC=2) | 0.06% |
| Achieved HBM GB/s | 0.47 |
| MBU (empirical peak 218 GB/s/LNC=2) | 0.22% |
**The recurrent kernel is overhead-dominated at S=128** (per-token wall ≫ per-token useful work). For real decode throughput, batch multiple tokens per invocation:
| S per invocation | Per-token wall-clock |
|------------------|----------------------|
| 1 | 70.4 μs (all overhead) |
| 8 | 9.7 μs |
| 128 | 6.55 μs |
| 512 | 6.38 μs |
Increasing `S` from 1 to 8 gives a **7.3× per-token improvement** with no kernel changes. This is the dominant lever for decode throughput.
## Comparison to `torch.compile(backend="neuron")` on the same reference
The direct "is NKI worth it?" question, measured on the same trn2.3xlarge with the same fla-core `naive_*` PyTorch reference compiled through the Neuron XLA backend:
### Recurrent (decode) -- NKI vs `torch.compile(naive_recurrent_kda)`
| S | NKI (μs) | torch.compile (μs) | NKI vs torch.compile |
|---|----------|--------------------|-----------------------|
| 1 | 70 | **66** | 0.94× (torch marginally faster in the overhead-dominated regime) |
| 8 | **77** | 97 | **1.26×** |
| 32 | **231** | 272 | **1.17×** |
| 128 | **848** | 939 | **1.11×** |
At S ≥ 8, NKI is 11-26% faster per invocation. Both converge to ~6.5-7.3 μs per token as S grows. The NKI advantage is largest in the S=8-32 range where per-call overhead is amortized but the sequential recurrence stays short.
### Chunked (prefill) -- NKI vs `torch.compile(naive_chunk_kda)`
| C | NKI (μs) | torch.compile (μs) | NKI vs torch.compile |
|---|----------|--------------------|-----------------------|
| 128 | **87** | 1660 | **19.08×** |
On prefill, NKI is **19× faster** than the same algorithm compiled through the Neuron XLA backend. The NKI kernel packs the 24 128×128 matmuls of the Neumann series plus all elementwise ops into a single NEFF with all intermediates staying in SBUF; the XLA-compiled path does many HBM round-trips.
### Compilation time (first-run)
| Path | S=1 | S=8 | S=32 | S=128 recurrent | C=128 chunked |
|------|-----|-----|------|-----------------|---------------|
| NKI | ~8s | ~8s | ~9s | ~9s | ~2s |
| torch.compile | 1.9s | 3.4s | 10.6s | 45.3s | **284s (4.7 min)** |
torch.compile is competitive on tiny recurrent-workload first-runs (S ≤ 8) but takes 5x longer at S=128 and **142× longer on chunked** because the entire Neumann-series graph has to be traced through Python loops and lowered by XLA. NKI's `@nki.jit` compiles the whole algorithm as one function.
**Bottom line**: NKI is a modest win for short recurrent workloads (1.1-1.3× per-invocation) and a step-change win for chunked prefill (19×) on both wall-clock and compile time. Full report at [`torch_compile_comparison.md`](../perf/torch_compile_comparison.md) in the kda-kernel project working tree.
## Peak references
All MFU / MBU denominators cited above are per LNC=2 core on trn2 (NeuronCore-v3), from AWS documentation and empirical measurements:
- **PE peak**: 158 BF16 TFLOPS (spec) / 40 FP32 TFLOPS (spec). Trainium2 architecture guide.
- **HBM peak**: 750 GB/s theoretical (3 TB/s device / 8 physical cores × 2 per LNC=2), 218 GB/s empirical achievable (per project measurements on similar workloads).
Peak numbers are per LNC=2 core on trn2 (NeuronCore-v3), from the AWS Trainium2 architecture guide and empirical measurements.
## Fix history
### v1.3 (2026-08-06) — Fused multi-chunk backward
Added `kda_chunked_fused` + the raw fused kernels (`kda_fused_chunked_fwd`,
`kda_fused_chunked_bwd`): all chunks processed in **one NKI launch** per direction,
recurrent state / `dSn` carried in SBUF across the internal chunk loop. Same numerics
as `kda_chunked` (verified cos_sim ≈ 1.0 end-to-end through `loss.backward()`), but
eliminates the per-chunk launch overhead of the Python chunk loop.
**Performance** (fwd+bwd wall-clock, trn2, single (b,h)):
| S | Chunks | Python-loop | Fused | Speedup |
|---|--------|-------------|-------|---------|
| 256 | 2 | 695 μs | 336 μs | 2.07× |
| 512 | 4 | 1320 μs | 338 μs | 3.90× |
| 1024 | 8 | 2588 μs | 529 μs | 4.89× |
The fused fwd+bwd is nearly flat from S=256→512 (launch overhead paid once, not
NC times); the speedup grows with sequence length. Use `kda_chunked_fused` for training.
### v1.2 (2026-08-06) — Training / backward support
Added NKI backward kernels + `torch.autograd.Function` wrappers (`kda_recurrent`,
`kda_chunked`), making KDA trainable on Trainium. Backward math verified against
fla-core autograd on the simulator (trn2-pinned) and real trn2 hardware, end-to-end
through `loss.backward()`.
**Parity** (backward, vs fla-core naive autograd):
- Recurrent: all 5 gradients (dq, dk, dv, dg, dbeta) cos_sim ≥ 0.9998 — exact algorithm.
- Chunked: dq/dk/dv/dbeta cos_sim ≥ 0.9998; dg differs from fla (chunked forward uses
the scalar-mean intra-chunk-attention approximation, so our dg is the exact gradient
of *our* forward — self-consistent cos_sim=1.0 — but differs from fla's exact per-dim dg).
**Performance** (backward, trn2, single (b,h), from `perf/backward_perf.md` in the source project):
- recurrent bwd S=128: 2031 μs (2.4× the forward); chunked bwd C=128: 174 μs (2.0× the forward).
- vs `torch.compile(backend="neuron")` fwd+bwd: **9.0× faster** (recurrent) and **86.6× faster**
(chunked); NKI compiles in seconds vs 200-585 s for torch.compile's autograd graph.
**Two NKI techniques** used in the backward: reverse-cumsum via a single matmul
(`U @ dgc` with upper-tri ones), and per-free-dim reductions via `tensor_reduce(axis=(1,))`.
Deferred: a **fused multi-chunk backward** (one NKI launch for all chunks) to remove the
Python chunk-loop's per-chunk launch overhead at long sequences.
### v1.1 (2026-08-05) — State-decay axis fix (T-KDA-02)
**Accuracy fix (T-KDA-02)**: Both kernels decayed the recurrent state per-V column (`state[k, v] *= exp(g[v])`) instead of per-K row (`state[k, v] *= exp(g[k])`) as fla-core's canonical KDA does. The bug was invisible with uniform g (per-dim identical values) but diverged on per-dim varying g at typical KDA scale.
- **Recurrent** (`kda_recurrent_fwd`): parity vs fla went from cos_sim 0.99977 → **1.00000** (max_abs_diff 2.65e-3 → 3.4e-8). The fix also **simplifies** the kernel — the per-V decay required a transpose-scale-transpose dance (2 `nc_transpose` + 1 `tensor_scalar` per token); per-K decay is a single `tensor_scalar` on the partition axis, so v1.1 is also marginally faster.
- **Chunked** (`kda_chunk_step`): the per-K state decay is fixed (matters for state carry-over across chunks). Forward-output parity is unchanged at 0.99988 because that number is dominated by the separate scalar-mean intra-chunk-attention approximation.
All 6 recurrent backward gradients (dq, dk, dv, dg, dbeta, dinitial_state) now match fla `naive_recurrent_kda` autograd at cos_sim = 1.0. (Backward kernels themselves ship separately; see the kda-backward work.)
### v1.0 (2026-08-05) — Initial public release
**Accuracy fix (T-KDA-01)**: The pre-release version of the chunked kernel had a wrong-sign / extra-factor pattern in FOUR internal steps (`k_beta` for QK/A construction, `k_beta * exp_gc` for `k_cumdecay`, `q_c * exp_gc` for `attn_inter`, `k_c * exp_gl_minus_gc` for `k_state_decay`) that produced cos_sim ≈ 0.78 vs fla-core reference at typical KDA g-scale. **This is fixed in v1.0**: the wrapper contract now passes raw q, k (previously required wrapper to pre-multiply by `exp(±gc_mean)`), and the kernel computes all decay flavors internally, correctly distinguishing between the four different scaled versions of `k` needed.
**Prefill performance optimization**: 9× `nc_matmul(stationary=X, moving=eye)` transpose-via-matmul calls replaced with `nisa.nc_transpose(dst, data=X)`. This gives the compiler an explicit transpose hint and delivers a **−6.7% wall-clock** improvement (93.3 → 87.1 μs per chunk) with zero parity risk.
### Not addressed in v1.1 (deferred to v2.0)
- **Chunked dg parity vs fla**: because the chunked forward uses a scalar-mean intra-chunk-attention approximation, the g-gradient differs from fla's exact per-dim form (dq/dk/dv/dbeta match at 0.99986). Fixing requires exact per-dim intra-chunk attention (O(BT^2·K) instead of O(BT^2)).
- **Deeper R2 fix**: operand-order refactoring to compute `QK.T` directly and skip the transpose pairs (estimated additional ~5-10% wall-clock).
- **Split-Neumann across LNC=2**: dividing the Neumann series across two physical cores (up to 2× wall-clock, requires cross-core state management).
- **Full-layer wrapper**: `NeuronKDA(nn.Module)` drop-in replacement for HF Transformers' `Kda` layer (once upstream integration finalizes).
- **Backward / training kernels**: NKI backward kernels for training (chunked + recurrent) are in development; the math is verified but the NKI ports are not yet in this package.
## References
- **Algorithm**: [flash-linear-attention (fla-core)](https://github.com/fla-org/flash-linear-attention) — KDA is defined in `fla/ops/kda/`.
- **Sibling kernel package**: [`jburtoft/qwen35-deltanet-neuron-kernels`](https://huggingface.co/jburtoft/qwen35-deltanet-neuron-kernels) — same publication pattern, DeltaNet variant.
## License
Apache-2.0. This kernel package is an inference-runtime component, not a fine-tuned model. The fla-core algorithm reference is MIT-licensed and compatible.
## Contributing
This package's development happens in the internal `kda-kernel` project. External contributions welcome via PR to this HuggingFace Hub repo. For issues affecting the underlying kernels, please file on the Neuron team's internal ticketing system (not on public GitHub) so we can route them correctly.
|