--- 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.