memo(crt-a): forward_only cost investigation — Case 2 (stateless K-window)

forward_only requires the full K=seq_len window on every call and resets
Mamba2 h_s2 to zero each invocation — no cross-call state carry. Additionally,
the call is inside the stride=200 gate in harness.rs (spec §3.2 claim "already
every event" is incorrect as of HEAD). Moving to stride=1 without a forward_step
implementation would be a ~200x GPU cost increase. A0.5 must implement
forward_step with persistent h_s2, a dedicated K=1 CUDA graph, and session-reset
hook. Memo documents exact file paths, line numbers, required struct/method
changes, and open questions for the A0.5 implementer.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-20 17:43:52 +02:00
parent b508c38529
commit 2e87ed0daf

View File

@@ -0,0 +1,318 @@
# CRT Phase A — forward_only Cost Investigation
**Date:** 2026-05-20
**Branch:** `ml-alpha-phase-a`
**Task:** A0 (read-only research)
**Author:** CRT-A0 agent
---
## 1. TL;DR
**This is Case 2: stateless K-window.** Every call to `forward_only` in
`crates/ml-alpha/src/trainer/perception.rs` processes the full K-window from
scratch. The Mamba2 SSM recurrent state (`h_s2`) is zero-initialised once at
scratch allocation time and is never updated between calls — the scan kernel
treats it as a fixed zero residual, not a carry-over from the previous call.
Concretely, `h_s2` is a `[B, hidden_dim]` buffer that never receives a write
after construction.
**However:** the spec's §3.2 claim "Trunk forward pass (already every event)"
is **WRONG** as of HEAD. The current harness at `harness.rs:242` places
`forward_only` INSIDE the `if self.event_count % stride == 0` gate — it is
called every 200 events, not every event. This directly contradicts §3.2 and is
the primary finding of this investigation. Moving `forward_only` outside the
stride gate would mean calling the full K=64 window scan 200× more often,
which is a ~200× GPU cost increase for that call alone.
**Recommendation for A0.5:** A0.5 MUST implement `forward_step` — an
incremental single-step SSM inference that carries persistent `h_s2` state
across calls instead of zeroing it per call. This is the only path to ≤ 2×
wall-time on Phase A. Task CRT.A.0.5 is not a no-op.
---
## 2. Evidence
### 2.1 forward_only is inside the stride gate
File: `crates/ml-backtesting/src/harness.rs`
```rust
// Line 242 — both forward_only AND step_decision are gated:
if self.event_count % stride == 0 && self.snapshot_window.len() == self.seq_len {
let window: Vec<Mbp10RawInput> = self.snapshot_window.iter().cloned().collect();
let probs_all = self.trainer.forward_only(&window) // line 244
.context("trainer.forward_only")?;
// ...
self.sim.broadcast_alpha(&last_probs)?;
self.sim.step_decision_with_latency(raw.ts_ns, &self.sim_config)?;
}
```
The comment at line 239 says "At decision-stride boundaries: run forward inference + sim decision." Both operations are inside the same gate. The spec's §3.2 says "Trunk forward pass (already every event)" — this is incorrect.
### 2.2 forward_only signature requires the full K-window
File: `crates/ml-alpha/src/trainer/perception.rs`, line 2301:
```rust
pub fn forward_only(&mut self, snapshots: &[Mbp10RawInput]) -> Result<Vec<f32>> {
let b_sz = self.cfg.n_batch;
let k_seq = self.cfg.seq_len;
anyhow::ensure!(
snapshots.len() == b_sz * k_seq, // line 2305 — requires B*K snapshots
...
);
...
}
```
The function hard-enforces that exactly `n_batch * seq_len` snapshots are
provided. With the default `seq_len=64` and `n_batch=1`, every call requires
exactly 64 snapshots. There is no single-step variant.
### 2.3 Mamba2 h_s2 is never carried between calls
File: `crates/ml-alpha/src/mamba2_block.rs`, lines 151165:
```rust
/// `h_s2` is zero-initialised once and never written (no residual carry
/// from a prior chunk in the supervised path). The scan kernel reads it
/// as a constant addition to h_enriched_seq.
pub struct Mamba2BlockForwardScratch {
...
pub h_s2: GpuTensor, // [n_batch, hidden_dim] (zero residual)
...
}
```
And in `forward_train_seq_into` (`mamba2_block.rs:1116`), `h_s2` is passed as
the initial-state argument to the scan kernel:
```rust
.arg(scratch.h_s2.cuda_data()) // line 1116 — always zero, never updated
```
`h_s2` is allocated once in `Mamba2BlockForwardScratch::new` via
`GpuTensor::zeros` (line 191) and never written afterward. It is a constant
zero buffer. Each call to `forward_train_seq_into` starts the SSM scan from
h=0 regardless of what was computed on the previous call.
The same pattern holds for both Mamba2 layers:
- Layer 1: `self.trunk.mamba2_l1_mut().forward_train_seq_into(&self.vsn_out_d, &mut self.mamba2_fwd_scratch)` at line 2506
- Layer 2: `self.trunk.mamba2_l2_mut().forward_train_seq_into(&self.ln_a_out_d, &mut self.mamba2_l2_fwd_scratch)` at line 2531
Both scratch objects hold their own `h_s2: GpuTensor` initialised to zeros.
### 2.4 The CUDA Graph capture bakes in K=seq_len steps
File: `crates/ml-alpha/src/trainer/perception.rs`, lines 23582396:
The three-state machine (warmup → capture → replay) captures the full kernel
chain — VSN → Mamba2×2 → LN×2 → transpose → attn-pool → CfC K-loop → heads —
for a fixed `(b_sz, k_seq, total_snaps)` shape. A `forward_step` variant would
need its own graph with `(b_sz, 1, b_sz)` shape and persistent state passed via
device pointers rather than the scratch zero buffer.
### 2.5 Production seq_len
From `crates/ml-alpha/examples/alpha_train.rs`, line 62:
```rust
#[arg(long, default_value_t = 64)]
seq_len: usize,
```
Default seq_len is 64. Each `forward_only` call runs 64 Mamba2 steps for both
layers.
---
## 3. Cost Analysis
### 3.1 Per-call kernel work
Each `forward_only` call executes:
1. Host-side staging fill: `B*K * ~60` float32 writes (< 1 µs, OUTSIDE graph)
2. CUDA Graph replay or first-pass dispatch:
- `snap_feature_assemble_batched`: `B*K` rows
- VSN forward: `B*K` rows
- Mamba2 L1 `scan_fwd_seq`: K sequential steps over `[B, K, hidden_dim]`
- LN_a forward: `B*K` rows
- Mamba2 L2 `scan_fwd_seq`: K sequential steps
- LN_b forward: `B*K` rows
- Transpose: `[B, K, H]``[K, B, H]`
- Attention pool: K steps
- CfC K-loop: K steps
- Multi-horizon heads: K positions × N_HORIZONS outputs
3. Stream sync + DtoH probs: `K * B * N_HORIZONS` floats (64 × 1 × 4 = 256 f32)
The Mamba2 scan is a sequential K-step recurrence — it cannot be parallelised
over K. Cost scales linearly with K.
### 3.2 Cost at stride=200 vs stride=1
At `decision_stride=200`, `forward_only` is called once per 200 events.
If the stride gate is removed without adding `forward_step`, `forward_only`
is called once per event: **200× more calls, each still running K=64 steps**.
With CUDA Graph replay the per-call overhead is low (mostly the DtoD staging
copy + graph launch), but the actual Mamba2 scan work is O(K) per call. Calling
K=64-step inference 200× more frequently is a ~200× increase in GPU kernel
work for the encoder path.
The controller kernels (`decision_policy_*` + `seed_inflight_limits_batched`)
are indeed light (block-per-backtest, O(1) work). The spec's §3.4 cost estimate
correctly identifies these as the dominant concern — but it implicitly assumes
the trunk forward pass stays at stride=200. It is not correct that the forward
pass is "already every event."
### 3.3 Correct cost model for `forward_step`
If `forward_step` is implemented with persistent h_s2 state:
- Per-call work: 1 Mamba2 step × 2 layers = O(hidden_dim × state_dim)
- No K-window scatter/gather: only 1 snapshot staged per call
- No CfC K-loop: only the final position's head output needed
- Estimated per-call GPU time: ~1/64 of current `forward_only`
At stride=1 with `forward_step`, total encoder GPU work ≈ 200/64 ≈ 3.1× current.
Combined with controller kernels (light, 200× more calls), total wall-time
should land in the 24× range, consistent with the §3.4 budget of ≤ 2×
(achievable if the forward step is sufficiently fast relative to controller work).
---
## 4. Recommendation for Task A0.5
**A0.5 is required and non-trivial.** The following kernel changes are needed:
### 4.1 New Rust struct: `Mamba2BlockStepScratch`
Mirrors `Mamba2BlockForwardScratch` but sized for `K=1`:
- `x: GpuTensor``[B, hidden_dim]`
- `a_proj: GpuTensor``[B, state_dim]`
- `b_proj: GpuTensor``[B, state_dim]`
- `h_s2: GpuTensor`**`[B, hidden_dim]` — PERSISTENT, updated each call**
- `h_out: GpuTensor``[B, hidden_dim]` — single-step enriched output
Location: add to `crates/ml-alpha/src/mamba2_block.rs` alongside
`Mamba2BlockForwardScratch`.
### 4.2 New kernel call path: `Mamba2Block::step_into`
Mirrors `forward_train_seq_into` but:
- Input shape: `[B, 1, in_dim]` (single snapshot)
- Calls `kernel_fwd_seq` with `k_i32 = 1`
- After the scan, **copies `h_out` into `h_s2`** (carry-forward)
- Or: a dedicated `mamba2_alpha_scan_step` kernel that reads/writes `h_s2`
in-place
The Mamba2 `scan_fwd_seq` kernel at
`crates/ml-alpha/src/kernels/mamba2_alpha_kernel.cu` (referenced at
`crates/ml-alpha/build.rs:11`) will need to be called with K=1. Check whether
the kernel handles K=1 without corner-case issues (e.g., shared-memory tile
sizing assumptions). If K=1 is not safe, a dedicated `scan_fwd_step` kernel
variant is required.
### 4.3 New method: `PerceptionTrainer::forward_step`
Signature:
```rust
pub fn forward_step(&mut self, snapshot: &Mbp10RawInput) -> Result<[f32; N_HORIZONS]>
```
Behaviour:
1. Stage 1 snapshot into pinned buffers (outside CUDA Graph)
2. Run VSN on 1 row
3. Run `Mamba2Block::step_into` (L1) — updates L1 h_s2 in-place
4. Run LN_a on 1 row
5. Run `Mamba2Block::step_into` (L2) — updates L2 h_s2 in-place
6. Run LN_b on 1 row
7. Run the final CfC step + head projection for 1 position
8. Sync + DtoH: `N_HORIZONS` floats (4 f32)
9. Return probs for the current (single) position
This method should have its own CUDA Graph (capture on second call, replay
thereafter). Because K=1, the graph is small and replay overhead dominates —
confirm with profiling that graph replay is faster than eager dispatch at K=1.
### 4.4 Harness changes (in A1)
Once `forward_step` exists, `harness.rs` changes are:
```rust
// Before (stride-gated):
if self.event_count % stride == 0 && self.snapshot_window.len() == self.seq_len {
let window = self.snapshot_window.iter().cloned().collect::<Vec<_>>();
let probs_all = self.trainer.forward_only(&window)?;
...
}
// After (every event, once window bootstrapped):
if self.snapshot_window.len() == self.seq_len {
let probs = self.trainer.forward_step(&raw)?;
self.sim.broadcast_alpha(&probs)?;
self.sim.step_decision_with_latency(raw.ts_ns, &self.sim_config)?;
}
```
The sliding window (`snapshot_window: VecDeque<Mbp10RawInput>`) is still
needed for the warm-up period (until K events are seen) but thereafter serves
only as a `len() == self.seq_len` bootstrap check — the SSM state carries
forward in the L1/L2 scratch buffers.
### 4.5 Bootstrap concern
`forward_only` bootstraps Mamba2 with K=64 steps of context. `forward_step`
starts from h_s2=0. For the first `seq_len-1` events the model has no
accumulated context — the predictions will be from a cold-start state.
Mitigation options (implementer should choose):
a. Keep the first inference as `forward_only` (at `event_count == seq_len-1`),
then transfer the final `h_s2` state from that call into the step scratch,
and switch to `forward_step` thereafter. This requires a new method
`PerceptionTrainer::transfer_step_state_from_window` or similar.
b. Accept cold-start predictions for the first K events; the Wiener-α EMA
in conviction smoothing (spec §4.2) will dampen the cold-start noise.
c. Use option (b) but gate decisions on `event_count >= seq_len` — already
present in the harness via the `self.snapshot_window.len() == self.seq_len`
check.
Option (c) is the simplest: no changes to bootstrap semantics, the existing
warm-up guard stays, and the first K events remain decision-free. The h_s2
state simply accumulates naturally from zero.
---
## 5. Open Questions for A0.5 Implementer
1. **K=1 in scan_fwd_seq kernel**: Does the existing Mamba2 CUDA kernel at
`mamba2_alpha_kernel.cu` handle `K=1` safely? Shared-memory block sizing and
warp-tile assumptions may break at K=1. If so, a dedicated `scan_fwd_step`
kernel variant is required. Verify with the K=1 case before implementing the
full graph.
2. **h_s2 carry-forward mechanism**: The cleanest implementation writes h_out
over h_s2 in-place at the end of each step. But `h_s2` in the current
scratch is declared immutable (`cuda_data()` not `data_mut()`) and is passed
read-only to the scan kernel. Either add a write-back pass after the scan, or
restructure `step_into` to accept a mutable h_s2 that it writes in-place.
3. **CUDA Graph at K=1**: Graph replay overhead is typically 515 µs. At K=1,
the actual kernel work may be shorter than the replay overhead. Profile the
K=1 eager dispatch vs graph replay on the L40S to determine whether graphing
is worth it.
4. **Session resets**: The SSM state in `forward_step` accumulates across
session boundaries. The harness already has session-gap detection (via
`ts_ns` gaps). On session-gap detection, `forward_step` should reset h_s2 to
zero and require a new bootstrap period. Add a `reset_step_state()` method
and call it from the harness on session-gap events.
5. **Spec §3.2 correction**: The spec states "Trunk forward pass (already every
event)" — this is factually incorrect as of HEAD. The spec should be
corrected (or a note added) to reflect that `forward_only` is currently
stride-gated and that `forward_step` is the A0.5 deliverable. The A1
implementer should know the spec was wrong on this point.