feat(backend): Mamba2 ops, with CPU reference implementations #85
Loading…
Reference in a new issue
No description provided.
Delete branch "feat/mamba2-cpu-ops"
Deleting a branch is permanent. Although the deleted branch may continue to exist for a short time before it actually gets removed, it CANNOT be undone in most cases. Continue?
Phase 3 of
docs/hybrid-state-cache.md: the Mamba2 compute ops, with CPU reference implementations. No architecture calls them yet; thenemotron_hmodel (phase 4) will.Ops
All four are defined against the reference
torch_forwardpath in Nemotron-H'smodeling_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 issilu(bias + Σ_k w[:,k]·in[t-kernel+1+k]), where the inputs are preceded by the sequence's previouskernel-1inputs from its conv-state slot (or zeros when fresh). Afterwards the slot holds the lastkernel-1inputs.SSMScan(dst, xbc, dt, aLog, dtBias, d, state, rb)is the selective-scan recurrence. Withdt' = softplus(dt + dt_bias), each token computesh = h·exp(-exp(A_log)·dt') + dt'·B⊗xandy = h·C + D·x. Head h reads B and C from grouph/(heads/groups), and the recurrence seeds from and writes back the sequence's SSM slot.xbcwhole and reads x, B and C at fixed column offsets. The model splitsin_projinto z, xBC and dt at load (asmistral4does), so no split op is needed.time_step_limitdefaults to(0, inf), where softplus already lies, and the config does not set it.GatedRMSNorm(dst, x, z, weight, groups, eps)computesrmsnorm(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 isdown(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 ruleBatch.StateSlotsdocuments. 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
ErrNotImplementedfor 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.Testing
0.5+3·1, 0.5+2·1+3·2, …pin which weight multiplies the current token.y = 7and then2/e + 2.5.D·x, and heads 0 and 1 untouched.go build,go vetandgo test ./...pass, and the CUDA package type-checks under-tags cuda(go vet).🤖 Generated with Claude Code
Automated review by pr-reviewer v0.52.3 | Safety Check | Mistral Small | tracking id
r-b90490-9c352fThis 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 retryto try again.