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):
- 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 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. Seeexamples/04_backward_training.pyfor 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)
kernelslibrary (pip install kernels>=0.15.2) or a local clone of this repo- Environment:
export TORCH_NEURONX_ENABLE_CONCATENATION=1recommended (~3-8% speedup, zero cost)
Direct usage
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:
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):
in_proj: single Linear producing[z, x, B, C, dt, lam, theta].- Phase 0:
dt_softplus = softplus(dt + dt_bias),lam_sigmoid = sigmoid(lam), then discretization coefficientsalpha = exp(dt·A),gamma = lam·dt,beta = (dt - gamma)·alpha. - Phase 0.5: RMSNorm on B and C (fused variance step across the two tensors).
- Phase 1: angle cumsum (
cum_angles = -cumsum(dt·theta)), then RoPE applied to B and C. - 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. - 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 vianc_matmulwith R on partitions of both operands (transposed from the natural DSTATE-partition layout vianc_transpose). - Rank-R output for
y[p,r] = Σ_n new_state[n,p] · C_rot[n,r]:nc_matmulwith 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
MIMO
torch.compileregression (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.F.padautograd bug (ticket 58): silent 2% gradient bias whenF.pad(x[:, :-1], (0,...,1,0))flows into a multi-operand einsum on device. The mixer works around this by usingtorch.catfor the beta-term shift. Transparent to users.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.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 (Apache 2.0).
License
Apache 2.0. See LICENSE.
- Downloads last month
- -