File size: 12,666 Bytes
f145768 07e0355 f145768 6709b3f f145768 6709b3f 24e7c33 15b4561 eacd70d 15b4561 eacd70d 24e7c33 15b4561 eacd70d f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f 07e0355 6709b3f 07e0355 6709b3f 07e0355 6709b3f 07e0355 6709b3f 07e0355 6709b3f 07e0355 6709b3f 07e0355 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f f145768 6709b3f | 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 | ---
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. Runs on AWS Trainium (trn2) under PyTorch Native.
This is a `kernel`-type repository (build variant `torch-neuron`, backend `neuron`).
Load it with the [`kernels`](https://github.com/huggingface/kernels) library on a
Trainium machine:
```python
from kernels import get_kernel
# Recommended: pin the stable major version (currently the latest).
k = get_kernel("jburtoft/kda-neuron-kernels", version=2, trust_remote_code=True)
# Or always track the latest published kernels:
# k = get_kernel("jburtoft/kda-neuron-kernels", revision="main", trust_remote_code=True)
# k.kda_chunk_step_exact(...), k.kda_chunk_step_exact_bwd(...), k.kda_recurrent_fwd(...), etc.
```
> `get_kernel` requires either `version=` or `revision=` — there is no zero-arg
> default (this is a `kernels`-library requirement, not a repo setting). `version=2`
> is the current latest major version; `revision="main"` always tracks the newest
> published tree. Both resolve to the same content today.
**Runtime requirement (important):** load this from a **PyTorch Native
(torch-neuronx)** environment where `torch` is a CPU/Neuron build and `torch.neuron`
is registered (e.g. the DLAMI venvs `aws_neuronx_venv_pytorch_2_9_nxd_inference`
or a PyTorch-Native Beta venv). The `kernels` library selects the build variant
from the active torch backend; in a **CUDA** torch build (some DLAMI base venvs ship
`torch ...+cuXXX`) it will detect backend `cuda` and refuse the `neuron` variant
with *"backend (neuron) does not match system backend (cuda)"*. If you hit that,
switch to a Neuron/PyTorch-Native venv (verify with
`python -c "from kernels.backends import _backend; print(_backend().name)"` → should
print `neuron`).
## What this package provides
### Inference (forward)
- **`kda_recurrent_fwd(q, k, v, g, beta)`** — decode / token-generation per-token
recurrence. One (batch, head) invocation processes `S` tokens sequentially.
- **`kda_recurrent_fwd_state(q, k, v, g, beta)`** — same, and also returns the final
recurrent state for prefill→decode hand-off.
- **`kda_chunk_step(q, k, v, beta, g_cumsum, g_last, state_in)`** — prefill per-chunk
step. Processes one 128-token chunk given the state from the previous chunk. Uses a
scalar-mean decay approximation for the intra-chunk term — **see the warning below**.
- **`kda_chunk_step_exact(q, k, v, beta, g, state_in)`** — numerically exact
per-channel prefill (sub-chunk + WY reformulation). Use this when the model's gate
decay is non-trivial (see warning below). ~1.05× the latency of `kda_chunk_step`.
- **`kda_chunk_step_exact_multihead(q, k, v, beta, g, state_in)`** — head-interleaved
exact prefill over `NV` heads (`[NV, C, dk]` shapes).
- **`kda_decode_batch(q, k, v, g, beta, state_in)`** — batched multi-(request, head)
decode; advances all `B*nv` items one token in a single call. Shapes `[B, nv, dk]`.
- **`kda_chunk_step_exact_bwd(q, k, v, beta, g, state_in, d_output, dS_final)`** —
exact chunked BACKWARD (gradient of `kda_chunk_step_exact`). Returns
`dq, dk, dv, dg, dbeta, dstate_in`, matching fla-core autograd at **cos_sim 1.0
for all five gradients in every gate regime** (g = 0.01…2.0). Pair this with
`kda_chunk_step_exact` for training. The older approximate chunked backward has
a broken `dg` (cos_sim ≈ 0.09 in ALL regimes) and NaNs at large decay.
### Training (differentiable, `loss.backward()`-ready)
- **`kda_recurrent(q, k, v, g, beta, initial_state=None)`** → `(output, final_state)`.
Differentiable; routes through the NKI recurrent backward. Requires zero
`initial_state`.
- **`kda_chunked(q, k, v, g, beta, initial_state=None)`** → `(output, final_state)`.
Differentiable; loops chunks in Python around the single-chunk kernels. Supports
state carry-over across chunks.
- **`kda_chunked_fused(q, k, v, g, beta, initial_state=None)`** → `(output, final_state)`.
Same numerics as `kda_chunked`, but processes all chunks in one NKI launch per
direction. Faster on multi-chunk sequences. **Recommended for training.**
- Raw kernels also exported: `kda_recurrent_fwd_v2`, `kda_chunk_step_v2`,
`kda_recurrent_bwd`, `kda_chunk_bwd`, `kda_fused_chunked_fwd`, `kda_fused_chunked_bwd`.
## Requirements
- **Hardware**: AWS Trainium (tested on trn2.3xlarge).
- **SDK / runtime**: PyTorch Native (`device="neuron"`), torch-neuronx 2.11+, PyTorch 2.11+.
- **NKI** ≥ 0.4.0.
- **`kernels`** ≥ 0.15.2 (to load via `get_kernel`).
- **`transformers`** with `KernelConfig` support, if using the `KernelConfig` path.
## Usage
### Inference — direct kernel calls
```python
import torch
import torch.nn.functional as F
from kda_neuron_kernels.build.torch_neuron import kda_chunk_step_exact
# Prefill one 128-token chunk for a single (batch, head) slice.
S, Dk = 128, 128
q_raw = torch.randn(S, Dk)
k_raw = torch.randn(S, Dk)
v = torch.randn(S, Dk)
g = -torch.rand(S, Dk) * 0.01 # per-channel log-decay (negative)
beta = torch.rand(S) # per-token scalar
# Caller preprocessing (fla-core convention): L2-norm q, k and scale q by 1/sqrt(Dk).
q = F.normalize(q_raw, p=2, dim=-1) * (Dk ** -0.5)
k = F.normalize(k_raw, p=2, dim=-1)
beta_bc = beta.unsqueeze(-1).expand(S, Dk).contiguous()
state = torch.zeros(Dk, Dk, dtype=torch.float32).to("neuron")
chunk_out, state = kda_chunk_step_exact(
q.to("neuron"), k.to("neuron"), v.to("neuron"),
beta_bc.to("neuron"), g.to("neuron"), state,
)
# chunk_out: (128, 128) per-token output; state: (128, 128) carries to the next chunk.
```
See `tests/example_usage.py` for a fully-worked example.
### Training
```python
import torch, torch.nn.functional as F
from kda_neuron_kernels.build.torch_neuron import kda_chunked_fused
S, D = 256, 128 # S must be divisible by 128
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-channel 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)
out, final_state = kda_chunked_fused(q, k, v, g, beta, initial_state=None) # .to("neuron") for hardware
loss = out.sum()
loss.backward() # gradients flow through the NKI backward
# q.grad, k.grad, v.grad, g.grad, beta.grad now populated
```
Kernels operate per (batch, head); loop `B*H` in the caller.
## Input contract
Callers pass raw `q`, `k` **already L2-normed**, with `q` additionally scaled by
`1/sqrt(dk)` (fla-core convention). The kernels compute all decay-related scaling
internally from `g`.
For `kda_chunk_step` (and `_v2`):
- `q`, `k`: L2-normed q (scaled by `1/sqrt(dk)`), L2-normed k — shape `(128, 128)`
- `v`: value tensor — `(128, 128)`
- `beta`: per-token scalar, broadcast to `(128, 128)`
- `g_cumsum`: per-channel `cumsum(g)` within the chunk — `(128, 128)`
- `g_last`: `g_cumsum[-1:, :]` broadcast to `(128, 128)`
- `state_in`: recurrent state from the previous chunk — `(128, 128)`
- Returns `(chunk_out, state_out)`, each `(128, 128)`.
For `kda_recurrent_fwd`:
- `q`, `k`: `(S, 128)` L2-normed (q scaled)
- `v`, `g`, `beta`: `(S, 128)` (`beta` per-token scalar, broadcast across the dim)
- Returns `output (S, 128)`.
## Constraints
- **`head_k_dim == head_v_dim == 128`** (matches the NeuronCore SBUF partition width).
Other head dims are not supported.
- **`chunk_size == 128`** for the chunked kernels; `S` must be divisible by 128.
- **float32 inputs.**
- `kda_recurrent` (training wrapper) requires zero `initial_state`; use `kda_chunked`
for state carry-over across sequence packs.
## ⚠️ Functional warning — chunked gate-decay approximation
`kda_chunk_step` (and its training wrappers `kda_chunked` / `kda_chunked_fused`) use a
**scalar-mean approximation** for the intra-chunk attention decay
(`exp(mean_c(gc)_i - mean_c(gc)_j)` instead of the exact per-channel
`exp(gc_i - gc_j)`). This is a compute/accuracy tradeoff.
**The approximation is only accurate for small gate decay.** Measured single-chunk
cosine similarity vs the fla-core reference:
| gate scale `g` | cos_sim (`kda_chunk_step`) |
|----------------|----------------------------|
| ~0.01 (small) | ~0.99 |
| ~0.3 | ~0.49 |
| ~2.0 | ~0.22 |
**If your model has non-trivial gate decay, use `kda_chunk_step_exact`** (or
`kda_chunk_step_exact_multihead`), which is numerically exact (cos_sim ≥ 0.9999999
across all gate regimes) at ~1.05× the latency. The recurrent kernels
(`kda_recurrent_fwd`, `kda_recurrent`) are exact in all regimes.
Because `kda_chunked` / `kda_chunked_fused` differentiate the approximate forward,
their `dg` gradient is the exact gradient of the *approximate* forward — self-consistent
for training with these kernels, but not equal to the exact-per-channel `dg` unless the
approximation is accurate (i.e. small gate decay).
## Parity
Against the fla-core `naive_recurrent_kda` / `naive_chunk_kda` PyTorch references
(random inputs, `g_scale=0.01`, seq_len=128, single (batch, head)):
| Kernel | cos_sim vs fla | max_abs_diff |
|--------|----------------|--------------|
| `kda_recurrent_fwd` (S=128) | 1.00000 | 3.4e-8 |
| `kda_chunk_step` (C=128, small gate) | 0.99988 | 1.2e-3 |
| `kda_chunk_step_exact` (C=128, all gate regimes) | ≥ 0.9999999 | ~1e-6 |
Backward gradients vs fla-core autograd, across gate regimes (g = 0.01 / 0.3 / 1.0 / 2.0):
| Backward kernel | dq | dk | dv | dg | dbeta |
|-----------------|----|----|----|----|-------|
| `kda_chunk_step_exact_bwd` (all regimes) | 1.0000 | 1.0000 | 1.0000 | **1.0000** | 1.0000 |
| approximate `kda_chunk_bwd` @ g=0.01 | 0.9999 | 0.9999 | 0.9999 | **0.085** | 0.9999 |
| approximate `kda_chunk_bwd` @ g=0.3 | 0.990 | 0.990 | 0.991 | **0.121** | 0.993 |
| approximate `kda_chunk_bwd` @ g=2.0 | NaN | NaN | NaN | **NaN** | NaN |
The exact backward fixes the approximate kernel's uncorrelated `dg` (which is wrong
in **every** regime, not just at high decay) and its NaN at large gate decay.
Training backward gradients (`kda_recurrent`, `kda_chunked`) verified end-to-end
through `loss.backward()` against fla-core autograd: recurrent all five gradients
cos_sim ≥ 0.9998; chunked `dq/dk/dv/dbeta` ≥ 0.9998 (with the `dg` caveat above —
use `kda_chunk_step_exact_bwd` to fix it).
## Performance
Measured on trn2.3xlarge, LNC=2, single logical core, single (batch, head) invocation.
### Prefill (chunked)
| 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% |
| MBU (empirical 218 GB/s/LNC=2) | 3.11% |
### Decode (recurrent)
| Metric | Value |
|--------|-------|
| Wall-clock per call (S=128) | 527 μs |
| Per-token (amortized) | 4.1 μs |
The recurrent kernel is overhead-dominated at small sequence lengths. For real decode
throughput, batch tokens (or requests via `kda_decode_batch`): per-token wall-clock
drops from ~70 μs at S=1 to ~6 μs at S=128, and `kda_decode_batch` amortizes launch
overhead across a whole serving batch.
### Training (fused vs Python chunk loop, fwd+bwd)
| S | Chunks | `kda_chunked` (loop) | `kda_chunked_fused` |
|---|--------|----------------------|---------------------|
| 256 | 2 | 695 μs | 336 μs |
| 512 | 4 | 1320 μs | 338 μs |
| 1024 | 8 | 2588 μs | 529 μs |
The fused path pays per-launch overhead once instead of per chunk, so its advantage
grows with sequence length. Prefer `kda_chunked_fused` for training.
MFU/MBU denominators are per LNC=2 core on trn2 (NeuronCore-v3), from the AWS
Trainium2 architecture guide and empirical measurements.
## Not provided
- A full `nn.Module` drop-in replacement for a HuggingFace `Kda` layer (planned).
## References
- **Algorithm**: [flash-linear-attention (fla-core)](https://github.com/fla-org/flash-linear-attention)
— KDA is defined in `fla/ops/kda/`.
## License
Apache-2.0. This is an inference/training runtime kernel package, not a model. The
fla-core algorithm reference is MIT-licensed and compatible.
|