| --- |
| library_name: kernels |
| license: apache-2.0 |
| tags: |
| - neuron |
| - nki |
| - trainium |
| - inferentia |
| - mamba3 |
| - mamba-3 |
| - state-space-model |
| - ssm |
| - linear-attention |
| - prefill |
| - decode |
| --- |
| |
| # Mamba-3 NKI Kernels for AWS Neuron |
|
|
| Full-mixer replacement for a Mamba-3 SSM block on AWS Trainium/Inferentia via PyTorch Native. Handles **prefill** (chunked forward) and **autoregressive decode** (single-token step) with NKI-accelerated SSD kernels, plus an eager fallback for correctness comparison. |
|
|
| Implements the ICLR 2026 Mamba-3 paper (Lahoti, Li, Chen, Wang, Bick, Kolter, Dao, Gu — [arXiv:2603.15569](https://arxiv.org/abs/2603.15569)): |
| - **Trapezoidal (2nd-order) discretization** via a 3-term recurrence (α, β, γ coefficients per token) |
| - **Complex-valued state via data-dependent RoPE** for state-tracking capability |
| - **MIMO rank-r matrix state update** (r=4 in the reference config) |
|
|
| **Version**: v1.0.0 |
|
|
| ## Performance (target dims: `d_model=1024, d_state=128, headdim=64, nheads=32`, trn2.3xlarge LNC=2) |
|
|
| ### Forward, batch=1 |
|
|
| | Path | seqlen=64 | seqlen=256 | seqlen=512 | seqlen=1024 | |
| |---|---:|---:|---:|---:| |
| | SISO + torch.compile + NKI | **1.72 ms** | **3.30 ms** | **4.57 ms** | **8.10 ms** | |
| | SISO eager, no NKI | 12.5 ms | 17.5 ms | 20.3 ms | 28.5 ms | |
| | MIMO + NKI (eager) | **6.35 ms** | **10.32 ms** | **15.18 ms** | **26.90 ms** | |
| | MIMO eager, no NKI | 17.7 ms | 27.0 ms | 27.2 ms | 37.3 ms | |
|
|
| **SISO with NKI + torch.compile: 4-7x faster than eager two-SSD.** MIMO with NKI: 2-2.8x faster than MIMO eager two-SSD. See [`examples/05_benchmark.py`](examples/05_benchmark.py) to reproduce. |
|
|
| ### Decode step (single-token, prefill_len=64) |
| |
| | Path | eager | with NKI decode kernel | speedup | |
| |---|---:|---:|---:| |
| | SISO | 3.59 ms | **3.22 ms** (p95 3.47 ms) | 1.12x | |
| | MIMO | 3.93 ms | **3.66 ms** (p95 3.86 ms) | 1.07x | |
| |
| Both SISO and MIMO decode use dedicated NKI kernels (state update + output matmul). MIMO decode additionally contracts the rank-R dimension in-kernel via nc_matmul. All 128 sequential decode steps pass parity vs eager at cos>0.9999. |
|
|
| ### Batched throughput (batch=32, seqlen=256, eager + NKI) |
|
|
| | Path | per-sample latency | Throughput vs batch=1 | |
| |---|---:|---:| |
| | SISO | **1.78 ms/sample** | 3.94x | |
| | MIMO | **5.76 ms/sample** | 1.79x | |
|
|
| ### Correctness |
|
|
| - **Forward parity**: SISO + MIMO forward outputs match eager reference at cos_sim > 0.999999 across all tested configs (batch=1/2, seqlen=64/256, multiple seeds). |
| - **Decode parity**: 128/128 consecutive SISO decode steps pass cos_sim > 0.999 vs CPU reference; state cache persists correctly. |
| - **Backward parity**: All 12 MIMO parameter gradients pass cos_sim > 0.999 on device via `torch.autograd`. See [`examples/04_backward_training.py`](examples/04_backward_training.py) for the training demo. |
| |
| ## Repository layout |
| |
| ``` |
| build/torch-neuron/ |
| ├── __init__.py <- public API re-exports |
| ├── metadata.json <- HF kernels library metadata |
| ├── constants.py <- DSTATE, HEADDIM, Q_SISO, C_MIMO, R_MIMO, ... |
| ├── ops.py <- eager helpers: RMSNorm, apply_rope, segsum, ssd_siso, ssd_mimo |
| ├── masks.py <- host-side mask builders (causal, triu, block_diag) with device cache |
| ├── layers.py <- NeuronMamba3Mixer, NeuronMamba3Layout, Mamba3Cache |
| └── nki_kernels/ |
| ├── __init__.py <- re-exports the NKI kernels |
| ├── mamba3_siso_ssd.py <- SISO SSD kernel (single-SSD form + diagonal correction) |
| ├── mamba3_siso_ssd_wrapper.py <- SISO wrapper (layout adapter) |
| ├── mamba3_mimo_ssd.py <- MIMO SSD kernel (rank-flatten + block-diag correction) |
| ├── mamba3_mimo_ssd_wrapper.py <- MIMO wrapper (layout adapter) |
| └── mamba3_siso_decode.py <- Decode kernel (state update + output) |
| |
| examples/ |
| ├── README.md <- how to run each sample |
| ├── _loader.py <- shared kernel loader (local or HF Hub) |
| ├── 01_smoke_test.py <- verify kernel loads + one forward call |
| ├── 02_forward_parity.py <- NKI vs eager on identical weights (SISO + MIMO) |
| ├── 03_decode_generation.py <- prefill + 128 decode steps with cache |
| ├── 04_backward_training.py <- 10-step training loop with autograd |
| └── 05_benchmark.py <- reproduce the perf tables above |
| ``` |
| |
| Users load the package via HF's `kernels` library or `KernelConfig` -- the internal file split is transparent. |
| |
| ## Requirements |
| |
| - **AWS Neuron SDK Beta 3+** (torch-neuronx 2.11.3+, NKI 0.4.0+, PyTorch 2.11+) |
| - **trn2.3xlarge or larger** (tested on trn2.3xlarge with LNC=2) |
| - **`kernels` library** (`pip install kernels>=0.15.2`) or a local clone of this repo |
| - **Environment**: `export TORCH_NEURONX_ENABLE_CONCATENATION=1` recommended (~3-8% speedup, zero cost) |
| |
| ## Direct usage |
| |
| ```python |
| import torch |
| from kernels import get_kernel |
| |
| # revision + trust_remote_code required by kernels >= 0.15 |
| mamba3 = get_kernel( |
| "jburtoft/mamba3-neuron-kernels", |
| revision="v1.0.0", |
| trust_remote_code=True, |
| ) |
| |
| # SISO mixer (mimo_rank=1) |
| mixer = mamba3.NeuronMamba3Mixer( |
| d_model=1024, |
| d_state=128, |
| headdim=64, |
| chunk_size=64, # 64 for SISO, 16 for MIMO (R=4) |
| mimo_rank=1, # 1 for SISO, 4 for MIMO |
| use_nki_ssd=True, |
| ).to("neuron") |
| |
| # Prefill: multi-token forward |
| u = torch.randn(1, 256, 1024, device="neuron") |
| y, cache = mixer(u) # y: (1, 256, 1024), cache: Mamba3Cache namedtuple |
|
|
| # Decode: single-token step with cache handoff |
| next_tok = torch.randn(1, 1, 1024, device="neuron") |
| y_step, cache = mixer.step(next_tok, cache) |
| ``` |
| |
| For optimal SISO single-request latency, wrap in `torch.compile`: |
| |
| ```python |
| mixer_compiled = torch.compile(mixer, backend="neuron") # 2-3x faster than eager |
| ``` |
| |
| **Do NOT torch.compile MIMO** — there is a known 10-18x regression on MIMO where the compiler emits excessive transpose kernels for 5D rank-dim tensors. Use eager mode for MIMO. |
| |
| ## Configuration constraints |
| |
| The NKI kernels are compiled against the reference target config. When `use_nki_ssd=True`: |
| |
| | Config | SISO | MIMO | |
| |---|---|---| |
| | `d_state` | 128 | 128 | |
| | `headdim` | 64 | 64 | |
| | `chunk_size` | 64 | 16 | |
| | `mimo_rank` | 1 | 4 | |
| |
| Other configs work if you use `use_nki_ssd=False` (eager path only, ~3x slower). |
| |
| Sequence length must be divisible by `chunk_size` (64 or 16). Batch, `nheads`, `d_model` are unconstrained. |
| |
| ## How it works |
| |
| The `NeuronMamba3Mixer.forward()` pipeline (per the ICLR 2026 paper): |
| |
| 1. **`in_proj`**: single Linear producing `[z, x, B, C, dt, lam, theta]`. |
| 2. **Phase 0**: `dt_softplus = softplus(dt + dt_bias)`, `lam_sigmoid = sigmoid(lam)`, then discretization coefficients `alpha = exp(dt·A)`, `gamma = lam·dt`, `beta = (dt - gamma)·alpha`. |
| 3. **Phase 0.5**: RMSNorm on B and C (fused variance step across the two tensors). |
| 4. **Phase 1**: angle cumsum (`cum_angles = -cumsum(dt·theta)`), then RoPE applied to B and C. |
| 5. **Phase 2-4**: the SSD core. In SISO mode this is a **single-SSD kernel with diagonal correction**; in MIMO mode it's a **rank-flattened SSD with block-diagonal correction**. Both are NKI kernels when `use_nki_ssd=True`. |
| 6. **Phase 5**: D skip connection, silu gate, MIMO rank-fold (or SISO gate), then `out_proj`. |
| |
| Decode via `step()` implements the direct recurrence per Eq. 9: `h_t = alpha_t·h_{t-1} + beta_t·prev_Bx + gamma_t·(B_t⊗x_t)`, output `y_t = C_t^T·h_t`. |
| |
| ## Kernel details |
| |
| ### `mamba3_siso_ssd_kernel` |
| |
| Chunked SISO SSD. Uses the **single-SSD form** with `scale = γ + Δ_{t+1}(1-λ_{t+1})` (mathematically equivalent to the reference two-SSD form) + diagonal correction. Chunk size 64. Reuses the Mamba-2 SSD structure with three Mamba-3-specific changes: `scale`-weighted x, diagonal correction, and separate exp_cs_last broadcast to `d_state=128` partitions. |
| |
| ### `mamba3_mimo_ssd_kernel` |
| |
| Chunked MIMO SSD via **rank flattening**. Views `(C, R, feat)` tensors as `(C·R, feat)` in SBUF so the intra-chunk matmul becomes a 64×64 GEMM (same shape as SISO). Diagonal correction is computed in-kernel by **reusing the CB matmul with a block-diagonal mask** — same PE-array cost as SISO's diagonal correction. |
| |
| ### `mamba3_siso_decode_state_kernel` |
| |
| Single-token SISO decode. Handles the state update `h_t = α·h_{t-1} + β·prev_Bx + γ·B_t⊗x_t` and the output matmul `y = C_t^T·h_t` across all `nheads=32` heads in a batched fashion. RoPE and other pre-SSD work stay eager. |
| |
| ### `mamba3_mimo_decode_kernel` |
| |
| Single-token MIMO decode. Same state-update math as SISO, plus: |
| - **Rank-R contraction for `BX = Σ_r B_rot[n,r] · x_mimo[p,r]`**: done via `nc_matmul` with R on partitions of both operands (transposed from the natural DSTATE-partition layout via `nc_transpose`). |
| - **Rank-R output for `y[p,r] = Σ_n new_state[n,p] · C_rot[n,r]`**: `nc_matmul` with DSTATE on partitions, C_rot as stationary (R free), new_state as moving (HEADDIM free). Result (R, HEADDIM) is transposed back to the mixer's (HEADDIM, R) format. |
| |
| RoPE, rank-expand (`x_mimo = x·mimo_x_proj`), rank-fold (`y·mimo_down`), and gate all stay eager -- the kernel focuses on the SSD math where nc_matmul provides a clear win. |
| |
| ## Known limitations |
| |
| 1. **MIMO `torch.compile` regression** (ticket 61): torch.compile on the MIMO mixer runs 10-18x slower than eager due to the compiler emitting many small transpose kernels for the 5D rank-dim tensors. Use eager for MIMO. |
| |
| 2. **`F.pad` autograd bug** (ticket 58): silent 2% gradient bias when `F.pad(x[:, :-1], (0,...,1,0))` flows into a multi-operand einsum on device. The mixer works around this by using `torch.cat` for the beta-term shift. Transparent to users. |
| |
| 3. **NKI kernel silent-wrong-data on reshape** (ticket 59): certain `.reshape()` patterns on strided per-head slices silently return wrong data. Works around this by using host-side rank flattening + direct 2D slicing. |
| |
| 4. **NKI kernel crash on strided partition writes** (ticket 60): the compiler crashes on writes to non-contiguous partition ranges. Works around this in the MIMO diagonal correction by using a fused matmul + block-diagonal mask. |
| |
| All four are filed as tickets with the AWS Neuron team. Workarounds are stable in v1.0.0. |
| |
| ## Environment (verified working) |
| |
| - Container: `421672808698.dkr.ecr.us-east-1.amazonaws.com/concourse-release-0461d3b:latest` (Beta 3 DLC) |
| - torch-neuronx: 2.11.3.0.1278 |
| - neuronx-cc: 2.25.1280.0 |
| - NKI: 0.4.0 |
| - Instance: trn2.3xlarge in `ap-southeast-4`, LNC=2 |
| |
| ## Attribution |
| |
| This is a Trainium port of the algorithms described in: |
| |
| > Lahoti, Li, Chen, Wang, Bick, Kolter, Dao, Gu. |
| > *Mamba-3: Improved Sequence Modeling using State Space Principles.* |
| > ICLR 2026. arXiv:2603.15569. |
| |
| Reference implementation: [`state-spaces/mamba`](https://github.com/state-spaces/mamba) (Apache 2.0). |
| |
| ## License |
| |
| Apache 2.0. See [LICENSE](LICENSE). |
| |