kda-neuron-kernels / README.md
Jim Burtoft
v1.6: kernel-repo distribution fix
eacd70d
|
Raw
History Blame
12.2 kB
metadata
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. 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 library on a Trainium machine:

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

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

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

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.