--- library_name: kernels license: apache-2.0 --- # mamba3 Mamba-3 MIMO chunked forward and recurrent decode on NVIDIA GPUs, loadable through `kernels`. The reference baselines are the `mamba3_MIMO_chunk_ref` and `mamba3_MIMO_step_ref` programs from `state-spaces/mamba`, matched to 2.2e-6 in float32. Upstream ships these paths as a pinned `tilelang==0.1.8` kernel for the forward and a CUTLASS-DSL kernel for decode, and `Mamba3(is_mimo=True)` raises in its constructor when TileLang is absent rather than falling back. This is the same two operations as one compiled extension with no TileLang, TVM-FFI, CUTLASS-DSL or quack dependency. ## Usage ```python from kernels import get_kernel m3 = get_kernel("phanerozoic/mamba3", version=1, trust_remote_code=True) # prefill: q, k are [B, S, R, Gqk, N]; v, z are [B, S, H, P] y = m3.forward(q, k, v, q_bias, k_bias, mimo_v, mimo_o, angles, adt, dt, trap, z=z, mimo_z=mimo_z, D=D, chunk_size=16) # decode: one token per call, state carried st = m3.DecodeState(B, H, headdim=P, dstate=N, rank=R, nangles=Na) y_t = st.step(q_t, k_t, v_t, q_bias, k_bias, mimo_v, mimo_o, angles_t, adt_t, dt_t, trap_t, z=z_t, mimo_z=mimo_z, D=D) ``` `version` selects the release branch; `trust_remote_code` is required by `kernels` for publishers without the trusted-publisher mark. `angles` are the raw per-step rotation rates; the cumulative sum is taken internally in float32. ## API | Symbol | Purpose | |---|---| | `forward(q, k, v, ..., chunk_size, norm_weight)` | chunked forward, `[B, S, H, P]` out, differentiable | | `DecodeState(batch, nheads, headdim, dstate, rank, nangles)` | carried decode state | | `DecodeState.step(...)` | one decode token, `[B, H, P]` out | | `cumulative_angles(angles, dt)`, `chunk_decay(adt, chunk_size)` | schedule helpers | | `dispatch_paths(mimo_rank, dstate, headdim, chunk_size, bf16)` | which kernels a geometry selects | `norm_weight` selects the fused head-wise RMSNorm pre-gate path. All three ops carry meta kernels, so the forward traces under `torch.compile(fullgraph=True)` without a graph break. ## Method The rank axis is flattened into the row index `i = c*R + r`, so a chunk tile is `(C*R, N)` against `(C*R, P)` and the forward is the SSD decomposition: `chunk_state` builds `kv[n]`, `state_passing` scans it, `chunk_scan` computes `q_inter @ states[n] + mask(Q Kᵀ) V + D v`. Keys carry `factor = γ_t + Δ_{t+1}(1 - λ_{t+1})`, the single-SSD form of the trapezoidal recurrence, and the intra diagonal block is rescaled to `γ_t`. Contractions are register-blocked 4x4 on a 16x16 thread grid only where the tensor cores are not reachable. The float32 contractions split each operand into a tf32 head and tail and run `hi*hi + hi*lo + lo*hi`; the dropped `lo*lo` term is below the fp32 rounding of the sum, and a single tf32 pass would not hold the float32 contract. Rotation is applied while staging and never materialized. Decode runs the recurrence directly, `S ← αS + β(kprevᵀ vprev) + γ(kᵀv)`, one warp per `p` row with the rank axis in registers. The trapezoidal term reaches one step back, so the previous key and value are carried state. `mimo_v`, `mimo_o` and `mimo_z` are `(H, R, P)` elementwise weights, so rank expansion is a broadcast multiply and the fold is a contraction over `r` alone. Rotation angles stay float32 at any input dtype, since they accumulate over the sequence. On the bfloat16 path the operands are staged by `ldmatrix` and contracted by `mma.sync.m16n8k16` written out directly, because nvcc's wmma never emits `ldmatrix` for bf16 and falls back to scalar loads and matrix-move fixups. The rank, the head dimension and the tile counts each warp owns are template arguments, so the accumulators index statically and stay in registers across the whole `d` loop and the epilogue, and the index arithmetic carries no integer division. Every shared tile is padded off a multiple of the 32 banks. The causal weight, the chunk decay and the `D` term all fold into the accumulator drain, `D v_r` riding the intra product on A's diagonal, so the float A tile never exists and the output tile reuses the space its operands occupied. The rank weights and the decay column are staged once per block, a rotary pair is one wide load, and the head axis varies fastest across the grid so that blocks sharing a chunk's q and k are resident together. The state scan issues its chunk loads four ahead of the carry that consumes them; the fma order is unchanged, and the writes land on the addresses just read. ## Measured RTX 6000 Ada, `H=32, P=64, N=128, R=4, C=16`. Forward against the torch reference program, which is what runs when TileLang is unavailable: | seqlen | kernel | torch | speedup | peak memory | |---|---|---|---|---| | 512 | 0.331 ms | 3.312 ms | 10.02x | 62 vs 375 MB | | 1024 | 0.586 ms | 7.860 ms | 13.42x | 131 vs 741 MB | | 2048 | 1.531 ms | 17.609 ms | 11.50x | 254 vs 1472 MB | | 4096 | 3.186 ms | 40.460 ms | 12.70x | 499 vs 2935 MB | Decode against the eager step, with the state traffic it sustains: | batch | kernel | eager | speedup | state | |---|---|---|---|---| | 1 | 0.0245 ms | 0.754 ms | 30.8x | 86 GB/s | | 4 | 0.0398 ms | 0.749 ms | 18.8x | 211 GB/s | | 16 | 0.0665 ms | 0.773 ms | 11.6x | 505 GB/s | | 64 | 0.1659 ms | 2.030 ms | 12.2x | 809 GB/s | | 256 | 0.9697 ms | 7.492 ms | 7.7x | 554 GB/s | At batch 64 the decode step is at this card's memory roofline. At batch 1 it is launch- and dependency-latency bound; capture it in a CUDA graph. Against upstream's CuTe decode kernel on an H200, `H=64, P=64, N=128, R=4`, bfloat16, gated, eight tokens a call. Upstream's kernel takes q and k already rotated, so the comparison is against its Triton rotary step and its CuTe step together, which is what one call here does; the CuTe kernel alone is the column beside it. Both sides agree with the reference program to 5.1e-3, and with each other to 5.6e-3, which is bfloat16 rounding: | batch | kernel | rotary + cute | cute alone | ratio | |---|---|---|---|---| | 1 | 0.812 ms | 2.901 ms | 0.646 ms | 3.57x | | 4 | 0.813 | 2.885 | 0.645 | 3.55x | | 16 | 1.180 | 2.894 | 0.638 | 2.45x | | 64 | 3.563 | 2.894 | 0.742 | 0.81x | | 256 | 12.738 | 3.051 | 1.973 | 0.24x | The two scale differently. This kernel is work-proportional in batch; upstream's pair is nearly flat from 1 to 256, and its rotary step is four fifths of it, so the crossover is near batch 32. Below it this kernel leads by up to 3.6x, above it upstream does, by 4x at batch 256. Against upstream's TileLang forward rather than the reference program, on an RTX 3070 Ti Laptop at sm_86 in bfloat16. Both sides run in one process from one float32 master set, each cast to its own dtype contract: | S | H | R | C | kernel | tilelang | ratio | |---|---|---|---|---|---|---| | 512 | 8 | 4 | 16 | 0.180 ms | 0.398 ms | 2.21x | | 1024 | 8 | 4 | 16 | 0.297 ms | 0.789 ms | 2.66x | | 2048 | 8 | 4 | 16 | 0.544 ms | 1.628 ms | 2.99x | | 4096 | 8 | 4 | 16 | 1.054 ms | 3.316 ms | 3.15x | | 2048 | 8 | 1 | 64 | 0.252 ms | 0.514 ms | 2.04x | | 2048 | 32 | 4 | 16 | 2.189 ms | 1.862 ms | 0.85x | On an H200 at sm_90, where upstream's float32 forward also fits its 168 KB dynamic shared-memory request and both dtypes can be compared: | S | H | R | C | kernel | tilelang | ratio | | kernel f32 | tilelang f32 | ratio | |---|---|---|---|---|---|---|---|---|---|---| | 512 | 8 | 4 | 16 | 0.114 ms | 0.204 ms | 1.79x | | 0.183 ms | 0.281 ms | 1.54x | | 1024 | 8 | 4 | 16 | 0.106 | 0.383 | 3.62x | | 0.308 | 0.554 | 1.80x | | 2048 | 8 | 4 | 16 | 0.181 | 0.754 | 4.16x | | 0.598 | 1.140 | 1.91x | | 4096 | 8 | 4 | 16 | 0.378 | 1.641 | 4.34x | | 1.236 | 2.547 | 2.06x | | 2048 | 32 | 4 | 16 | 0.712 | 1.035 | 1.45x | | 2.217 | 1.311 | 0.59x | | 2048 | 8 | 1 | 64 | 0.100 | 0.267 | 2.68x | | 0.211 | 0.365 | 1.73x | The head-count row is part-dependent. This kernel is work-proportional, one block per `(batch, head, chunk)`; upstream's unit of parallelism absorbs 32 heads, so it is underutilized below `H=32` and saturated at it. On sm_86 that saturation is enough to take the row, and on sm_90 it is not. The float32 column leads at `H=8` and trails at `H=32`, the one geometry either dtype loses on this part. Upstream's float32 forward requests 168,128 bytes of dynamic shared memory per block, and 180,992 at `C = 64`. The opt-in ceiling is 101,376 bytes on consumer and workstation Ampere and Ada. This kernel's float32 forward requests 39,104 bytes at the same geometry. The harness is `benchmarks/bench_vs_upstream.py`; `TODO.md` records the rest. ## Backward `forward` is an autograd function. The backward is two chunk-parallel passes around a reverse state scan, recomputing chunk intermediates rather than storing them: the first recovers `O` and takes the epilogue adjoint, the scan runs `dSt[n] = dSt_part[n] + decay[n] dSt[n+1]` in reverse, and the second produces `dQ`, `dKf` and `dV` before the rotation adjoint lands on `q`, `k` and the biases. Gradients cover the activations, every `(H,R,P)` rank weight, `D`, the norm weight, and the schedule itself: `dt`, `adt`, `trap` and the rotation rates. It is reproducible. The grid is `(batch, chunk)` with the heads looped inside, so every gradient indexed by `(b, s, ...)` is owned by a single block and is accumulated in head order rather than by atomics; per-head parameters go to a per-block slot reduced in block order, and the two cross-chunk terms, the shifted half of the trapezoidal factor and the scan's decay adjoint, are staged and summed in a second ordered pass. No floating-point atomic appears anywhere in the backward source, and repeated runs return the forward and all 14 gradients bitwise identical. The two chunk workspaces and the recomputed states are `[B, H, Nc, N, P]` each and `dO` is `[B, H, Nc, C*R, P]`, so peak allocation is roughly the forward's plus that. At the geometry above: | seqlen | forward | forward + backward | |---|---|---| | 512 | 55 MB | 129 MB | | 1024 | 109 MB | 258 MB | | 2048 | 218 MB | 513 MB | | 4096 | 435 MB | 1025 MB | ## Correctness Every figure below is the worst case over that program's grid, measured on an RTX 6000 Ada against `state-spaces/mamba`'s own reference programs. | program | checks | worst | |---|---|---| | `tests/test_mamba3.py` decode | ranks 1 to 8, grouped-query, both rotation fractions | 3.0e-7 | | `tests/test_mamba3.py` forward | against `mamba3_MIMO_chunk_ref` | 2.5e-6 | | `tests/test_mamba3.py` forward vs decode | partial rotation | 1.3e-6 | | `tests/test_mamba3.py` ragged | sequence length not a multiple of the chunk | 3.3e-6 | | `tests/test_mamba3.py` fused norm | against the upstream fused-norm reference | 3.6e-7 | | `tests/test_mamba3.py` bfloat16 | against the float32 kernel | 9.4e-3 | | `tests/test_mamba3.py` long sequences | `S` to 32,768, accumulated angle to 3,378 radians | 3.5e-6 | | `tests/test_bwd.py` | 15 raw gradients against autograd through the reference | 7.6e-7 | | `tests/test_autograd.py` | the wrapped function end to end, cumsum adjoints included | 6.4e-7 | | `tests/test_bwd_bf16.py` | the bfloat16 backward against the float32 one | 6.8e-3 | | `tests/test_bigchunk.py` | the streamed intra path up to `C*R = 256` | 6.1e-7 | | `ref/test_ref.py` | this repo's step reference against upstream's | 2.3e-7 | | `tests/test_dispatch.py` | all five launcher paths against the step reference | 1.5e-2 | | `tests/test_determinism.py` | forward and 14 gradients over 3 repeats | bitwise | | `tests/test_fixes.py` | meta kernels, `fullgraph=True`, the backward memory guard | pass | Upstream's chunk reference omits the cos/sin padding its step reference performs, so its pairwise path is only defined at `Na == N/2`; partial rotation is checked against decode instead. The rotation angle is a cumulative sum over the sequence and grows without bound, so it is evaluated with a range-reduced sincos rather than the approximate intrinsic, whose argument reduction would make the residual proportional to the angle. Agreement with the reference is flat in sequence length: 2.3e-6 at `S = 2048` and 3.5e-6 at `S = 32768`, where the accumulated angle reaches 3,378 radians. The residual is set by the three-term tf32 split, which accumulates over the `N/8` contraction tiles, so it is the figure to watch if `N` grows; `MAMBA3_SCAN=fma` returns the float32 scan to 1.4e-6 at a third less throughput. ## Requirements and limits - NVIDIA GPU with compute capability 8.0+; float32 or bfloat16 values, float32 schedule and rotation angles. - `N` even, `Na <= N/2`, `mimo_rank <= 8`, `nheads` divisible by the query-group count. Sequence length and chunk size are arbitrary; above `C*R = 64` the forward switches to a row-block streamed intra term so nothing scales with `(C*R)²`. `C = 64 / mimo_rank` remains the fastest setting, since the intra term is `O(S*C)` and grows faster than the state workspace shrinks. - The backward has no streamed path, so its tiles stay resident and its shared memory grows as `(C*R)²` and `C*R*P`. Both entry points check the requirement against the device's opt-in maximum and raise naming the two figures rather than failing at launch. A 99 KB device admits `C*R = 64` at `headdim = 64`; `C*R = 128`, or `headdim = 128` at `C*R = 64`, needs the forward instead. - Gradients are computed in float32 and returned in the input dtype, and the backward computes all of them whether or not each input requires grad. - `chunk_state` and `chunk_scan` both contract on the tensor cores in both dtypes, float32 through the three-pass tf32 split and bfloat16 natively. The tensor-core kernels are instantiated for `mimo_rank` in {1, 2, 4, 8} and `headdim` in {16, 32, 64, 128}, and the scans additionally want `chunk_size * mimo_rank == 64`; any other geometry takes the FMA path. `dispatch_paths` reports which kernels a geometry selects, and `MAMBA3_SCAN=fma` demotes the float32 scan to FMA for an A/B in one process. - Packaged for compute capabilities 8.0, 8.6, 8.9, 9.0, 10.0 and 12.0. The whole suite runs on 8.6 and 8.9; the others are build targets. ## References Lahoti, Li, Chen, Wang, Bick, Kolter, Dao, Gu, "Mamba-3: Improved Sequence Modeling using State Space Principles" (ICLR 2026, arXiv:2603.15569); Dao and Gu, the SSD chunked scan of Mamba-2; `state-spaces/mamba` as the reference implementation. ## License Apache-2.0.