feat(backend): Mamba2 ops, with CPU reference implementations #85

Merged
rcsheets merged 1 commit from feat/mamba2-cpu-ops into main 2026-09-27 12:17:13 +00:00
Owner

Phase 3 of docs/hybrid-state-cache.md: the Mamba2 compute ops, with CPU reference implementations. No architecture calls them yet; the nemotron_h model (phase 4) will.

Ops

All four are defined against the reference torch_forward path in Nemotron-H's modeling_nemotron_h.py:

  • CausalConv1D(dst, x, weight, bias, state, rb) is the depthwise causal conv over each sequence's run of tokens, fused with SiLU. Output row t is silu(bias + Σ_k w[:,k]·in[t-kernel+1+k]), where the inputs are preceded by the sequence's previous kernel-1 inputs from its conv-state slot (or zeros when fresh). Afterwards the slot holds the last kernel-1 inputs.
  • SSMScan(dst, xbc, dt, aLog, dtBias, d, state, rb) is the selective-scan recurrence. With dt' = softplus(dt + dt_bias), each token computes h = h·exp(-exp(A_log)·dt') + dt'·B⊗x and y = h·C + D·x. Head h reads B and C from group h/(heads/groups), and the recurrence seeds from and writes back the sequence's SSM slot.
    • It takes the conv output xbc whole and reads x, B and C at fixed column offsets. The model splits in_proj into z, xBC and dt at load (as mistral4 does), so no split op is needed.
    • There is no dt clamp: the reference's time_step_limit defaults to (0, inf), where softplus already lies, and the config does not set it.
  • GatedRMSNorm(dst, x, z, weight, groups, eps) computes rmsnorm(x·silu(z))·weight, with the RMS taken per group (MambaRMSNormGated, norm_before_gate=False).
  • ReLU2(dst, x) is the experts' activation (mlp_hidden_act: relu2). Their MLP is down(relu2(up(x))), with no gate.

The two stateful ops take a backend.RecurrentBatch, which gives each sequence's run length, slot, and whether it is fresh. A fresh sequence starts from zero and never reads its slot, so a previous owner's leftovers are never seen; this is the rule Batch.StateSlots documents. The CPU implementations reject a slot out of range, a slot shared by two sequences in one batch, and run lengths that don't cover the rows. They never write out of bounds or let two sequences share state.

CPU reference

On the CPU the recurrence runs one token at a time, which is its definition. A chunked scan is a GPU throughput technique, validated against this loop rather than the other way round.

The state is rounded to F32 after every token, which is the precision it is stored in between steps anyway. As a result, a run split across steps is bit-identical to the same run in one step, and the tests check state carry for exact equality rather than within a tolerance.

Other backends

  • CUDA returns ErrNotImplemented for all four ops until phase 5 writes the kernels. A hybrid model served on CUDA fails its first forward with a 501 and never runs without its recurrent layers.
  • The sizing backend stubs the ops like every other compute op.

Testing

  • Each op against values worked out by hand:
    • The conv's pre-activations 0.5+3·1, 0.5+2·1+3·2, … pin which weight multiplies the current token.
    • The scan uses a dt chosen so softplus gives exactly 1, giving y = 7 and then 2/e + 2.5.
  • Split versus whole, with exact equality, for both stateful ops, including a one-token step shorter than the conv's history.
  • Resuming a scan from a saved mid-sequence state reproduces the rest exactly.
  • Two sequences batched on swapped slots, one fresh and one resuming, match their solo runs.
  • A fresh sequence ignores a dirty slot.
  • The SSM head-to-group mapping: zeroing group 1's B and C leaves heads 2 and 3 with only D·x, and heads 0 and 1 untouched.
  • Malformed batches are rejected.
  • go build, go vet and go test ./... pass, and the CUDA package type-checks under -tags cuda (go vet).

🤖 Generated with Claude Code

Phase 3 of `docs/hybrid-state-cache.md`: the Mamba2 compute ops, with CPU reference implementations. No architecture calls them yet; the `nemotron_h` model (phase 4) will. ## Ops All four are defined against the reference `torch_forward` path in Nemotron-H's `modeling_nemotron_h.py`: - **`CausalConv1D(dst, x, weight, bias, state, rb)`** is the depthwise causal conv over each sequence's run of tokens, fused with SiLU. Output row t is `silu(bias + Σ_k w[:,k]·in[t-kernel+1+k])`, where the inputs are preceded by the sequence's previous `kernel-1` inputs from its conv-state slot (or zeros when fresh). Afterwards the slot holds the last `kernel-1` inputs. - **`SSMScan(dst, xbc, dt, aLog, dtBias, d, state, rb)`** is the selective-scan recurrence. With `dt' = softplus(dt + dt_bias)`, each token computes `h = h·exp(-exp(A_log)·dt') + dt'·B⊗x` and `y = h·C + D·x`. Head h reads B and C from group `h/(heads/groups)`, and the recurrence seeds from and writes back the sequence's SSM slot. - It takes the conv output `xbc` whole and reads x, B and C at fixed column offsets. The model splits `in_proj` into z, xBC and dt at load (as `mistral4` does), so no split op is needed. - There is no dt clamp: the reference's `time_step_limit` defaults to `(0, inf)`, where softplus already lies, and the config does not set it. - **`GatedRMSNorm(dst, x, z, weight, groups, eps)`** computes `rmsnorm(x·silu(z))·weight`, with the RMS taken per group (`MambaRMSNormGated`, `norm_before_gate=False`). - **`ReLU2(dst, x)`** is the experts' activation (`mlp_hidden_act: relu2`). Their MLP is `down(relu2(up(x)))`, with no gate. The two stateful ops take a `backend.RecurrentBatch`, which gives each sequence's run length, slot, and whether it is fresh. A fresh sequence starts from zero and never reads its slot, so a previous owner's leftovers are never seen; this is the rule `Batch.StateSlots` documents. The CPU implementations reject a slot out of range, a slot shared by two sequences in one batch, and run lengths that don't cover the rows. They never write out of bounds or let two sequences share state. ## CPU reference On the CPU the recurrence runs one token at a time, which is its definition. A chunked scan is a GPU throughput technique, validated against this loop rather than the other way round. **The state is rounded to F32 after every token**, which is the precision it is stored in between steps anyway. As a result, a run split across steps is bit-identical to the same run in one step, and the tests check state carry for exact equality rather than within a tolerance. ## Other backends - **CUDA** returns `ErrNotImplemented` for all four ops until phase 5 writes the kernels. A hybrid model served on CUDA fails its first forward with a 501 and never runs without its recurrent layers. - **The sizing backend** stubs the ops like every other compute op. ## Testing - Each op against values worked out by hand: - The conv's pre-activations `0.5+3·1, 0.5+2·1+3·2, …` pin which weight multiplies the current token. - The scan uses a dt chosen so softplus gives exactly 1, giving `y = 7` and then `2/e + 2.5`. - Split versus whole, with exact equality, for both stateful ops, including a one-token step shorter than the conv's history. - Resuming a scan from a saved mid-sequence state reproduces the rest exactly. - Two sequences batched on swapped slots, one fresh and one resuming, match their solo runs. - A fresh sequence ignores a dirty slot. - The SSM head-to-group mapping: zeroing group 1's B and C leaves heads 2 and 3 with only `D·x`, and heads 0 and 1 untouched. - Malformed batches are rejected. - `go build`, `go vet` and `go test ./...` pass, and the CUDA package type-checks under `-tags cuda` (`go vet`). 🤖 Generated with [Claude Code](https://claude.com/claude-code)
feat(backend): Mamba2 ops, with CPU reference implementations
All checks were successful
ci / test_and_build (pull_request) Successful in 24s
a60d86ae01
Phase 3 of docs/hybrid-state-cache.md. Four new backend.Ops, defined
against the reference torch_forward path in Nemotron-H's modeling code:

- CausalConv1D: the depthwise causal conv over each sequence's run of
  tokens, fused with SiLU. It seeds from the sequence's conv-state slot
  (its last kernel-1 inputs) and writes the new ones back.
- SSMScan: the selective-scan recurrence, h = h*exp(-exp(A_log)*dt') +
  dt'*B*x and y = C.h + D*x, where dt' = softplus(dt + dt_bias). Head h
  reads B and C group h/(heads/groups). It seeds from and writes back
  the sequence's SSM slot. It takes the conv output whole and reads x, B
  and C at fixed column offsets, so the model needs no split op.
- GatedRMSNorm: rmsnorm(x*silu(z))*weight, with the RMS taken per group.
- ReLU2: the experts' activation.

The stateful two take a backend.RecurrentBatch (per sequence: run
length, slot, fresh). A fresh sequence starts from zero and never reads
its slot, so a previous owner's leftovers are never seen. The CPU
references reject a slot out of range, a slot shared by two sequences
in one batch, and run lengths that don't cover the rows, rather than
writing out of bounds or letting two sequences share state.

On the CPU the recurrence runs one token at a time, which is its
definition. A chunked scan is a GPU throughput technique, and will be
validated against this loop. The state is rounded to F32 after every
token, the precision it is stored in between steps anyway, so a run
split across steps is bit-identical to the same run in one step. The
tests rely on that: state carry is checked for exact equality, not
within a tolerance.

The CUDA backend returns ErrNotImplemented for the four ops until
phase 5 writes the kernels, so a hybrid model served on CUDA fails with
a 501 and never runs without its recurrent layers. The sizing backend
stubs them like every other compute op. No architecture calls them yet.

Tests: each op against values worked out by hand; split-versus-whole
equality for both stateful ops (including a one-token step shorter than
the conv's history); resuming a scan from a saved state; two sequences
batched on swapped slots, one fresh and one resuming, matching their
solo runs; a fresh sequence ignoring a dirty slot; the SSM
head-to-group mapping; and rejection of malformed batches.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Collaborator

Automated review by pr-reviewer v0.52.3 | Safety Check | Mistral Small | tracking id r-b90490-9c352f
This is an AI-generated review and may contain mistakes.

Status: ❌ Failed


This review couldn't be completed: that model isn't loaded on the inference service right now, and no alternate model was able to review it either. Consider splitting this PR into smaller changes. Tracking id r-b90490-9c352f.

Comment @pr-reviewer-bot retry to try again.

<!-- pr-reviewer:review --> *Automated review by [pr-reviewer](https://git.brooktrails.org/brooktrails/pr-reviewer) v0.52.3 | Safety Check | Mistral Small | tracking id `r-b90490-9c352f`* *This is an AI-generated review and may contain mistakes.* **Status:** ❌ Failed --- This review couldn't be completed: that model isn't loaded on the inference service right now, and no alternate model was able to review it either. Consider splitting this PR into smaller changes. Tracking id `r-b90490-9c352f`. Comment `@pr-reviewer-bot retry` to try again.
rcsheets deleted branch feat/mamba2-cpu-ops 2026-09-27 12:17:13 +00:00
Sign in to join this conversation.
No reviewers
No labels
No milestone
No project
No assignees
2 participants
Notifications
Due date
The due date is invalid or out of range. Please use the format "yyyy-mm-dd".

No due date set.

Dependencies

No dependencies set

Reference
brooktrails/gllm!85
No description provided.