File size: 10,886 Bytes
e95c403
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d1ebb85
 
 
 
 
 
e95c403
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2a17ff8
 
 
 
 
 
e95c403
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d1ebb85
 
 
 
 
 
 
 
e95c403
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
---
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).