kda-neuron-kernels / README.md
Jim Burtoft
v1.6: kernel-repo distribution fix
eacd70d
|
Raw
History Blame
12.2 kB
---
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
k = get_kernel("jburtoft/kda-neuron-kernels", version=2, trust_remote_code=True)
# k.kda_chunk_step_exact(...), k.kda_chunk_step_exact_bwd(...), k.kda_recurrent_fwd(...), etc.
```
**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.