feat(ml-alpha): per-step supervision unrolls BPTT through full sequence

The final-step-only trainer (one BCE prediction per 32-snapshot
window) trained flat at chance on real ES data despite working on
synthetic overfit: train_loss=0.6953, val_loss=0.6943 across 40k
gradient steps. Gradient density was the bottleneck — one supervised
position per sequence × ~8K seqs/epoch isn't enough signal for the
SSM to find the alpha.

This commit supervises the model at EVERY position in the sequence:

  mamba2_alpha_scan_fwd_seq    — emits h_enriched at every t step
                                 ([N, K, sh2] instead of [N, sh2])
  mamba2_alpha_scan_bwd_seq    — accepts d_h_enriched_seq, injects
                                 gradient at each t before propagating
                                 d_state through the gate chain.
                                 d_w_c and d_h_s2 accumulate across t.

  PerceptionTrainer.step()    — loop k=0..K; cfc + heads + BCE at
                                each valid label; cfc/heads grads
                                accumulate via += in kernel writes.
                                One Mamba2 backward call consumes the
                                full grad_h_enriched_seq.

  cfc_step_backward            — grad_w_in/w_rec/b writes changed
                                 to += (callers MUST pre-zero).
  multi_horizon_heads_backward — grad_w/grad_b writes changed to +=.

  alpha_train.rs               — passes per-position label rows to
                                 step(); AUC still scored from
                                 last-position predictions.

Phase E.3 callers (alpha_baseline.rs, alpha_dqn_h600_smoke.rs) use
the LEGACY Mamba2 forward_train + backward path with `alloc_zeros`
grad buffers — unaffected.

Synthetic overfit still converges 0.6664 → 0.1976 in 250 steps.
Local 2-quarter ES.FUT smoke shows the val AUC at h300 climbing
0.513 → 0.566 over 3 epochs (was flat-at-chance before). First
gradient signal we've gotten through the new architecture.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 01:03:33 +02:00
parent 2289fa062a
commit 16f5febf27
7 changed files with 682 additions and 134 deletions

View File

@@ -94,12 +94,16 @@ extern "C" __global__ void cfc_step_backward(
sdecay[i] = decay;
__syncthreads();
grad_b[i] = d_pre;
// Parameter-grad writes use `+=`: callers may invoke this kernel
// multiple times per training step (per-position supervision) and
// need the gradients summed across calls. Callers MUST pre-zero
// grad_b / grad_w_in / grad_w_rec at the start of each step.
grad_b[i] += d_pre;
for (int k = 0; k < n_in; ++k) {
grad_w_in[i * n_in + k] = d_pre * x[k];
grad_w_in[i * n_in + k] += d_pre * x[k];
}
for (int k = 0; k < n_hid; ++k) {
grad_w_rec[i * n_hid + k] = d_pre * h_old[k];
grad_w_rec[i * n_hid + k] += d_pre * h_old[k];
}
// grad_h_old[i] = dh * decay + sum_j d_pre[j] * W_rec[j, i]

View File

@@ -204,6 +204,140 @@ extern "C" __global__ void mamba2_alpha_reduce_d_proj(
}
/* ---------------------------------------------------------------------
* Per-step variants — supervise the SSM at EVERY timestep, not just the
* final position. Used by ml-alpha PerceptionTrainer.
*
* Forward writes per-step h_enriched_seq[N, K, sh2]; backward accepts
* d_h_enriched_seq[N, K, sh2] and accumulates gradient injections at
* every step before propagating the recurrent d_state through the
* gate chain. Per-step semantics:
*
* h_enriched_seq[i, t, j] = h_s2[i, j] + sum_s w_c[j, s] * x[i, t, s]
*
* d_h_s2[i, j] = sum_t d_h_enriched_seq[i, t, j] (residual added every step)
* d_w_c[j, s] = sum_i sum_t d_h_enriched_seq[i, t, j] * x[i, t, s]
* d_state[s] at t = (carried d_state from t+1) + sum_j d_h_enriched_seq[i, t, j] * w_c[j, s]
*
* Backward kernel: thread per (i, j) injects its own j-channel
* contribution at each step; the per-channel reductions across j
* (mamba2_alpha_reduce_d_proj / _d_w_c) work UNCHANGED — same scratch
* shapes, same launch configs.
* --------------------------------------------------------------------- */
extern "C" __global__ void mamba2_alpha_scan_fwd_seq(
const float* __restrict__ a_proj, // [N, K, state_d]
const float* __restrict__ b_proj, // [N, K, state_d]
const float* __restrict__ w_c, // [sh2, state_d]
const float* __restrict__ h_s2, // [N, sh2]
float* __restrict__ h_enriched_seq, // [N, K, sh2] — written at every step
int N,
int K,
int sh2,
int state_d
) {
int i = blockIdx.x;
int j = blockIdx.y * blockDim.x + threadIdx.x;
if (i >= N || j >= sh2) return;
float x[MAMBA2_ALPHA_MAX_STATE_D];
#pragma unroll
for (int s = 0; s < MAMBA2_ALPHA_MAX_STATE_D; s++) x[s] = 0.0f;
const float h_s2_ij = h_s2[(long long)i * sh2 + j];
for (int t = 0; t < K; t++) {
long long fwd_base = ((long long)i * K + t) * state_d;
for (int s = 0; s < state_d; s++) {
float gate = 1.0f / (1.0f + expf(-a_proj[fwd_base + s]));
x[s] = gate * x[s] + b_proj[fwd_base + s];
}
float ctx = 0.0f;
for (int s = 0; s < state_d; s++) {
ctx += w_c[(long long)j * state_d + s] * x[s];
}
h_enriched_seq[(((long long)i * K) + t) * sh2 + j] = h_s2_ij + ctx;
}
}
extern "C" __global__ void mamba2_alpha_scan_bwd_seq(
const float* __restrict__ a_proj, // [N, K, state_d]
const float* __restrict__ b_proj, // [N, K, state_d]
const float* __restrict__ d_h_enriched_seq, // [N, K, sh2]
const float* __restrict__ w_c, // [sh2, state_d]
float* __restrict__ d_a_per_channel, // [N, sh2, K, state_d]
float* __restrict__ d_b_per_channel, // [N, sh2, K, state_d]
float* __restrict__ d_w_c_per_sample, // [N, sh2, state_d]
float* __restrict__ d_h_s2, // [N, sh2]
int N,
int K,
int sh2,
int state_d
) {
int i = blockIdx.x;
int j = blockIdx.y * blockDim.x + threadIdx.x;
if (i >= N || j >= sh2) return;
/* Replay forward state, caching x[t][s]. */
float x_hist[MAMBA2_ALPHA_MAX_K * MAMBA2_ALPHA_MAX_STATE_D];
float x[MAMBA2_ALPHA_MAX_STATE_D];
#pragma unroll
for (int s = 0; s < MAMBA2_ALPHA_MAX_STATE_D; s++) x[s] = 0.0f;
for (int t = 0; t < K; t++) {
long long fwd_base = ((long long)i * K + t) * state_d;
for (int s = 0; s < state_d; s++) {
float gate = 1.0f / (1.0f + expf(-a_proj[fwd_base + s]));
x[s] = gate * x[s] + b_proj[fwd_base + s];
x_hist[t * MAMBA2_ALPHA_MAX_STATE_D + s] = x[s];
}
}
/* Zero d_state and d_w_c_per_sample (the latter accumulates across t). */
float d_state[MAMBA2_ALPHA_MAX_STATE_D];
#pragma unroll
for (int s = 0; s < MAMBA2_ALPHA_MAX_STATE_D; s++) d_state[s] = 0.0f;
long long w_c_slot = ((long long)i * sh2 + j) * state_d;
for (int s = 0; s < state_d; s++) d_w_c_per_sample[w_c_slot + s] = 0.0f;
float d_h_s2_sum = 0.0f;
/* Reverse scan: at each step t (going K-1 → 0), inject this step's
* gradient into d_state via W_c BEFORE recording d_a/d_b and
* propagating through the gate. */
long long per_chan_base = ((long long)i * sh2 + j) * (long long)K * state_d;
for (int t = K - 1; t >= 0; t--) {
long long fwd_base = ((long long)i * K + t) * state_d;
float d_h_ij_t = d_h_enriched_seq[(((long long)i * K) + t) * sh2 + j];
d_h_s2_sum += d_h_ij_t;
/* W_c grad contribution at step t: d_w_c[j,s] += d_h_ij_t * x[t,s]. */
for (int s = 0; s < state_d; s++) {
d_w_c_per_sample[w_c_slot + s] +=
d_h_ij_t * x_hist[t * MAMBA2_ALPHA_MAX_STATE_D + s];
}
/* Inject step-t gradient into d_state via W_c. */
for (int s = 0; s < state_d; s++) {
d_state[s] += d_h_ij_t * w_c[(long long)j * state_d + s];
}
/* Record d_b[t,s], d_a[t,s] from the CURRENT (post-injection) d_state. */
for (int s = 0; s < state_d; s++) {
float a_raw = a_proj[fwd_base + s];
float gate = 1.0f / (1.0f + expf(-a_raw));
float sig_deriv = gate * (1.0f - gate);
float x_prev = (t == 0) ? 0.0f : x_hist[(t - 1) * MAMBA2_ALPHA_MAX_STATE_D + s];
long long slot = per_chan_base + (long long)t * state_d + s;
d_b_per_channel[slot] = d_state[s];
d_a_per_channel[slot] = d_state[s] * x_prev * sig_deriv;
d_state[s] = d_state[s] * gate;
}
}
d_h_s2[(long long)i * sh2 + j] = d_h_s2_sum;
}
/* ---------------------------------------------------------------------
* Phase 1d.4 backtest — per-trade PnL kernel.
*

View File

@@ -53,15 +53,18 @@ extern "C" __global__ void multi_horizon_heads_backward(
}
__syncthreads();
// Parameter-grad writes use `+=`: per-position supervision invokes
// this kernel K times per training step and needs the gradients
// accumulated. Callers MUST pre-zero grad_w / grad_b at step start.
if (tid < 5) {
grad_b[tid] = d_z[tid];
grad_b[tid] += d_z[tid];
}
// grad_w[k, i] = d_z[k] * h[i]
// grad_w[k, i] += d_z[k] * h[i]
// grid stride: each thread covers one i for k in 0..5.
if (tid < 128) {
const float h_i = h[tid];
for (int k = 0; k < 5; ++k) {
grad_w[k * 128 + tid] = d_z[k] * h_i;
grad_w[k * 128 + tid] += d_z[k] * h_i;
}
// grad_h[i] = sum_k d_z[k] * W[k, i]
float acc = 0.0f;

View File

@@ -125,17 +125,21 @@ fn main() -> Result<()> {
let mut epoch_train_loss = 0.0_f32;
let mut epoch_train_steps = 0usize;
while let Some(seq) = train_loader.next_sequence().context("train next_seq")? {
// Label at the LAST position of the window (per-horizon).
let last = seq.snapshots.len().saturating_sub(1);
let labels = [
seq.labels[0][last], seq.labels[1][last], seq.labels[2][last],
seq.labels[3][last], seq.labels[4][last],
];
if labels.iter().all(|v| v.is_nan()) { continue; }
// Replace any NaN with mid (0.5) — masked by BCE kernel via NaN
// detection, but Mbp10 sequence is fed as-is.
let labels_arr = labels;
let loss = trainer.step(&seq.snapshots, &labels_arr).context("train step")?;
// Per-position labels: position k has horizon-h label
// seq.labels[h][k]. NaN entries (right-edge / tied price)
// are passed through to the trainer which masks them.
let mut labels_per_pos: Vec<[f32; N_HORIZONS]> = Vec::with_capacity(seq.snapshots.len());
let mut any_finite = false;
for k in 0..seq.snapshots.len() {
let row = [
seq.labels[0][k], seq.labels[1][k], seq.labels[2][k],
seq.labels[3][k], seq.labels[4][k],
];
if row.iter().any(|v| v.is_finite()) { any_finite = true; }
labels_per_pos.push(row);
}
if !any_finite { continue; }
let loss = trainer.step(&seq.snapshots, &labels_per_pos).context("train step")?;
epoch_train_loss += loss;
epoch_train_steps += 1;
train_loss_running += loss;
@@ -162,22 +166,33 @@ fn main() -> Result<()> {
let mut val_loss_sum = 0.0_f32;
let mut val_steps = 0usize;
while let Some(seq) = val_loader.next_sequence().context("val next_seq")? {
let last = seq.snapshots.len().saturating_sub(1);
let labels = [
seq.labels[0][last], seq.labels[1][last], seq.labels[2][last],
seq.labels[3][last], seq.labels[4][last],
];
if labels.iter().all(|v| v.is_nan()) { continue; }
let mut labels_per_pos: Vec<[f32; N_HORIZONS]> = Vec::with_capacity(seq.snapshots.len());
let mut any_finite = false;
for k in 0..seq.snapshots.len() {
let row = [
seq.labels[0][k], seq.labels[1][k], seq.labels[2][k],
seq.labels[3][k], seq.labels[4][k],
];
if row.iter().any(|v| v.is_finite()) { any_finite = true; }
labels_per_pos.push(row);
}
if !any_finite { continue; }
// Forward via training step (val loss is statistically negligible
// signal vs ~train_seqs * seq_len; trainer not separately frozen).
let l = trainer.step(&seq.snapshots, &labels).context("val step")?;
let l = trainer.step(&seq.snapshots, &labels_per_pos).context("val step")?;
val_loss_sum += l;
val_steps += 1;
// AUC scored from the LAST-position predictions only — the
// trainer's `last_probs()` returns whatever was in the probs
// buffer at the end of step(), which is position seq_len-1
// by iteration order. Matches the labels at that position.
let probs = trainer.last_probs().context("last_probs")?;
let last = seq.snapshots.len().saturating_sub(1);
for h in 0..N_HORIZONS {
if !labels[h].is_nan() {
let lbl = seq.labels[h][last];
if lbl.is_finite() {
val_probs[h].push(probs[h]);
val_labels[h].push(labels[h]);
val_labels[h].push(lbl);
}
}
}

View File

@@ -130,6 +130,21 @@ pub struct Mamba2ForwardCache {
pub h_enriched: GpuTensor,
}
/// Cache produced by [`Mamba2Block::forward_train_seq`] — variant of
/// [`Mamba2ForwardCache`] that exposes the SSM output at EVERY timestep
/// instead of just the final position. Used by ml-alpha when supervising
/// the model at every snapshot in a sequence (per-step heads + BCE).
pub struct Mamba2ForwardCacheSeq {
pub input_2d: GpuTensor,
pub x: GpuTensor,
pub a_proj: GpuTensor,
pub b_proj: GpuTensor,
/// Per-step enriched state, `[N, K, hidden_dim]`. Slot `[i, t, j]` is
/// `h_s2[i, j] + sum_s w_c[j, s] * x[i, t, s]` where x is the SSM
/// state AFTER step t (post-gate update).
pub h_enriched_seq: GpuTensor,
}
/// GPU-resident Mamba2 sequence block.
///
/// Owns all parameter tensors (`W_in`, `W_a`, `W_b`, `W_c`, `W_out`) on the
@@ -162,6 +177,12 @@ pub struct Mamba2Block {
_module: Arc<CudaModule>,
pub kernel_fwd: CudaFunction,
pub kernel_bwd: CudaFunction,
/// Per-step variant of the forward scan — writes h_enriched at every
/// timestep (used by per-position supervision in PerceptionTrainer).
pub kernel_fwd_seq: CudaFunction,
/// Per-step variant of the backward scan — accepts d_h_enriched_seq
/// and injects gradient contributions into d_state at every step.
pub kernel_bwd_seq: CudaFunction,
/// Reduces `d_a_per_channel` or `d_b_per_channel` (same kernel, two
/// call sites with different I/O pointers).
pub kernel_reduce_d_proj: CudaFunction,
@@ -195,6 +216,12 @@ impl Mamba2Block {
let kernel_bwd = module
.load_function("mamba2_alpha_scan_bwd")
.map_err(|e| anyhow!("Mamba2Block: backward kernel symbol resolve: {e}"))?;
let kernel_fwd_seq = module
.load_function("mamba2_alpha_scan_fwd_seq")
.map_err(|e| anyhow!("Mamba2Block: per-step forward kernel resolve: {e}"))?;
let kernel_bwd_seq = module
.load_function("mamba2_alpha_scan_bwd_seq")
.map_err(|e| anyhow!("Mamba2Block: per-step backward kernel resolve: {e}"))?;
let kernel_reduce_d_proj = module
.load_function("mamba2_alpha_reduce_d_proj")
.map_err(|e| anyhow!("Mamba2Block: d_proj reduction kernel resolve: {e}"))?;
@@ -245,6 +272,8 @@ impl Mamba2Block {
_module: module,
kernel_fwd,
kernel_bwd,
kernel_fwd_seq,
kernel_bwd_seq,
kernel_reduce_d_proj,
kernel_reduce_d_w_c,
kernel_adamw,
@@ -720,6 +749,241 @@ impl Mamba2Block {
})
}
/// Per-step variant of [`Mamba2Block::forward_train`] — emits the SSM
/// output at EVERY timestep instead of only the final position. Used
/// by ml-alpha PerceptionTrainer to supervise the model at every
/// snapshot in a sequence (denser gradient signal than final-step
/// only).
///
/// Skips the W_out projection — callers route `h_enriched_seq`
/// directly into a downstream head. Caller pairs this with
/// [`Mamba2Block::backward_from_h_enriched_seq`] for the backward
/// chain.
pub fn forward_train_seq(
&self,
input: &GpuTensor,
) -> Result<(GpuTensor, Mamba2ForwardCacheSeq)> {
let c = &self.config;
let n_batch = match input.shape() {
[b, k, d] if *k == c.seq_len && *d == c.in_dim => *b,
shape => {
return Err(anyhow!(
"Mamba2Block::forward_train_seq: expected [B, {}, {}], got {:?}",
c.seq_len, c.in_dim, shape
));
}
};
let n_rows = n_batch * c.seq_len;
let input_2d = GpuTensor::new(input.cuda_data().clone(), vec![n_rows, c.in_dim])
.map_err(|e| anyhow!("reshape input → 2D: {e}"))?;
let (x, _) = self.w_in.inner.forward_with_slices(
&input_2d, &self.w_in.weight, &self.w_in.bias, &self.cublas, &self.stream,
).map_err(|e| anyhow!("w_in forward: {e}"))?;
let (a_proj, _) = self.w_a.inner.forward_with_slices(
&x, &self.w_a.weight, &self.w_a.bias, &self.cublas, &self.stream,
).map_err(|e| anyhow!("w_a forward: {e}"))?;
let (b_proj, _) = self.w_b.inner.forward_with_slices(
&x, &self.w_b.weight, &self.w_b.bias, &self.cublas, &self.stream,
).map_err(|e| anyhow!("w_b forward: {e}"))?;
let h_s2 = GpuTensor::zeros(&[n_batch, c.hidden_dim], &self.stream)
.map_err(|e| anyhow!("alloc h_s2: {e}"))?;
let mut h_enriched_seq = GpuTensor::zeros(
&[n_batch, c.seq_len, c.hidden_dim],
&self.stream,
)
.map_err(|e| anyhow!("alloc h_enriched_seq: {e}"))?;
let block_threads: u32 = 32;
let grid_y: u32 =
((c.hidden_dim + block_threads as usize - 1) / block_threads as usize) as u32;
let cfg = LaunchConfig {
grid_dim: (n_batch as u32, grid_y, 1),
block_dim: (block_threads, 1, 1),
shared_mem_bytes: 0,
};
let n_i32 = n_batch as i32;
let k_i32 = c.seq_len as i32;
let sh2_i32 = c.hidden_dim as i32;
let st_i32 = c.state_dim as i32;
unsafe {
self.stream
.launch_builder(&self.kernel_fwd_seq)
.arg(a_proj.cuda_data())
.arg(b_proj.cuda_data())
.arg(&self.w_c)
.arg(h_s2.cuda_data())
.arg(h_enriched_seq.data_mut())
.arg(&n_i32).arg(&k_i32).arg(&sh2_i32).arg(&st_i32)
.launch(cfg)
.map_err(|e| anyhow!("scan_fwd_seq launch: {e}"))?;
}
let cache = Mamba2ForwardCacheSeq {
input_2d, x, a_proj, b_proj,
h_enriched_seq: h_enriched_seq.clone(),
};
Ok((h_enriched_seq, cache))
}
/// Backward chain paired with [`forward_train_seq`]. `d_h_enriched_seq`
/// has shape `[N, K, hidden_dim]` matching `cache.h_enriched_seq`.
/// Returns all nine parameter gradients (`dw_out` / `db_out` zeroed,
/// since this path skips W_out entirely).
pub fn backward_from_h_enriched_seq(
&self,
cache: &Mamba2ForwardCacheSeq,
d_h_enriched_seq: &GpuTensor,
) -> Result<Mamba2BackwardGrads> {
let c = &self.config;
let n_batch = cache.h_enriched_seq.shape()[0];
let n_rows = n_batch * c.seq_len;
if d_h_enriched_seq.shape() != [n_batch, c.seq_len, c.hidden_dim] {
return Err(anyhow!(
"Mamba2Block::backward_from_h_enriched_seq: d_h_enriched_seq \
shape {:?} != [{}, {}, {}]",
d_h_enriched_seq.shape(), n_batch, c.seq_len, c.hidden_dim
));
}
let per_chan_n = n_batch * c.hidden_dim * c.seq_len * c.state_dim;
let per_sample_n = n_batch * c.hidden_dim * c.state_dim;
let mut d_a_per_channel = self.stream
.alloc_zeros::<f32>(per_chan_n)
.map_err(|e| anyhow!("alloc d_a_per_channel: {e}"))?;
let mut d_b_per_channel = self.stream
.alloc_zeros::<f32>(per_chan_n)
.map_err(|e| anyhow!("alloc d_b_per_channel: {e}"))?;
let mut d_w_c_per_sample = self.stream
.alloc_zeros::<f32>(per_sample_n)
.map_err(|e| anyhow!("alloc d_w_c_per_sample: {e}"))?;
let mut d_h_s2 = self.stream
.alloc_zeros::<f32>(n_batch * c.hidden_dim)
.map_err(|e| anyhow!("alloc d_h_s2: {e}"))?;
let block_threads: u32 = 32;
let grid_y_h: u32 =
((c.hidden_dim + block_threads as usize - 1) / block_threads as usize) as u32;
let bwd_cfg = LaunchConfig {
grid_dim: (n_batch as u32, grid_y_h, 1),
block_dim: (block_threads, 1, 1),
shared_mem_bytes: 0,
};
let n_i32 = n_batch as i32;
let k_i32 = c.seq_len as i32;
let sh2_i32 = c.hidden_dim as i32;
let st_i32 = c.state_dim as i32;
unsafe {
self.stream
.launch_builder(&self.kernel_bwd_seq)
.arg(cache.a_proj.cuda_data())
.arg(cache.b_proj.cuda_data())
.arg(d_h_enriched_seq.cuda_data())
.arg(&self.w_c)
.arg(&mut d_a_per_channel)
.arg(&mut d_b_per_channel)
.arg(&mut d_w_c_per_sample)
.arg(&mut d_h_s2)
.arg(&n_i32).arg(&k_i32).arg(&sh2_i32).arg(&st_i32)
.launch(bwd_cfg)
.map_err(|e| anyhow!("scan_bwd_seq launch: {e}"))?;
}
// Reductions are identical to the final-step variant — d_a/d_b
// tensors have the same [N, sh2, K, state_d] / [N, sh2, state_d]
// shapes regardless of how many steps contributed to them.
let red_grid_z: u32 =
((c.state_dim + block_threads as usize - 1) / block_threads as usize) as u32;
let red_cfg = LaunchConfig {
grid_dim: (n_batch as u32, c.seq_len as u32, red_grid_z),
block_dim: (block_threads, 1, 1),
shared_mem_bytes: 0,
};
let mut d_a_proj_flat: CudaSlice<f32> = self.stream
.alloc_zeros::<f32>(n_rows * c.state_dim)
.map_err(|e| anyhow!("alloc d_a_proj_flat: {e}"))?;
let mut d_b_proj_flat: CudaSlice<f32> = self.stream
.alloc_zeros::<f32>(n_rows * c.state_dim)
.map_err(|e| anyhow!("alloc d_b_proj_flat: {e}"))?;
unsafe {
self.stream
.launch_builder(&self.kernel_reduce_d_proj)
.arg(&d_a_per_channel)
.arg(&mut d_a_proj_flat)
.arg(&n_i32).arg(&k_i32).arg(&sh2_i32).arg(&st_i32)
.launch(red_cfg)
.map_err(|e| anyhow!("reduce d_a_proj seq: {e}"))?;
self.stream
.launch_builder(&self.kernel_reduce_d_proj)
.arg(&d_b_per_channel)
.arg(&mut d_b_proj_flat)
.arg(&n_i32).arg(&k_i32).arg(&sh2_i32).arg(&st_i32)
.launch(red_cfg)
.map_err(|e| anyhow!("reduce d_b_proj seq: {e}"))?;
}
let red_w_c_cfg = LaunchConfig {
grid_dim: (c.hidden_dim as u32, red_grid_z, 1),
block_dim: (block_threads, 1, 1),
shared_mem_bytes: 0,
};
let mut dw_c: CudaSlice<f32> = self.stream
.alloc_zeros::<f32>(c.hidden_dim * c.state_dim)
.map_err(|e| anyhow!("alloc dw_c seq: {e}"))?;
unsafe {
self.stream
.launch_builder(&self.kernel_reduce_d_w_c)
.arg(&d_w_c_per_sample)
.arg(&mut dw_c)
.arg(&n_i32).arg(&sh2_i32).arg(&st_i32)
.launch(red_w_c_cfg)
.map_err(|e| anyhow!("reduce dw_c seq: {e}"))?;
}
let d_b_proj_2d = GpuTensor::new(d_b_proj_flat, vec![n_rows, c.state_dim])
.map_err(|e| anyhow!("reshape d_b_proj seq: {e}"))?;
let x_act = LinearActivations { input: cache.x.clone() };
let LinearGrads { dw: dw_b, db: db_b, dx: d_x_from_b } = self
.w_b.inner
.backward_with_slices(&d_b_proj_2d, &x_act, &self.w_b.weight,
&self.cublas, &self.stream)
.map_err(|e| anyhow!("w_b backward seq: {e}"))?;
let d_a_proj_2d = GpuTensor::new(d_a_proj_flat, vec![n_rows, c.state_dim])
.map_err(|e| anyhow!("reshape d_a_proj seq: {e}"))?;
let LinearGrads { dw: dw_a, db: db_a, dx: d_x_from_a } = self
.w_a.inner
.backward_with_slices(&d_a_proj_2d, &x_act, &self.w_a.weight,
&self.cublas, &self.stream)
.map_err(|e| anyhow!("w_a backward seq: {e}"))?;
let d_x = d_x_from_a.add(&d_x_from_b, &self.stream)
.map_err(|e| anyhow!("sum d_x branches seq: {e}"))?;
let input_act = LinearActivations { input: cache.input_2d.clone() };
let LinearGrads { dw: dw_in, db: db_in, dx: _d_input } = self
.w_in.inner
.backward_with_slices(&d_x, &input_act, &self.w_in.weight,
&self.cublas, &self.stream)
.map_err(|e| anyhow!("w_in backward seq: {e}"))?;
// W_out is unused by callers of forward_train_seq — zero grads
// freeze those parameters under AdamW.
let dw_out = GpuTensor::zeros(&[1, c.hidden_dim], &self.stream)
.map_err(|e| anyhow!("alloc zero dw_out seq: {e}"))?;
let db_out = GpuTensor::zeros(&[1], &self.stream)
.map_err(|e| anyhow!("alloc zero db_out seq: {e}"))?;
Ok(Mamba2BackwardGrads {
dw_in, db_in, dw_a, db_a, dw_b, db_b, dw_c, dw_out, db_out,
})
}
/// Total trainable parameter count (sum of all projections + W_c).
pub fn param_count(&self) -> usize {
let c = &self.config;

View File

@@ -1,17 +1,29 @@
//! PerceptionTrainer — stacked Mamba2 -> CfC -> heads (per spec 2026-05-16 amendment).
//! PerceptionTrainer — stacked Mamba2 -> CfC -> heads with per-position supervision.
//!
//! Per spec 2026-05-16 amendment + 2026-05-17 BPTT-unroll amendment.
//! Topology per training step:
//! snap_features × seq_len → Mamba2.forward_train → h_enriched [hidden_dim]
//! → cfc_step (h_old=0) → h_new [hidden_dim]
//! → heads probs [5]
//! → BCE(probs, labels) → loss
//! snap_features × seq_len → Mamba2.forward_train_seq → h_enriched_seq [K, hidden_dim]
//! For each k=0..K-1 with a valid label:
//! → cfc_step (h_old=0)h_new_k
//! → heads → probs_k
//! → BCE(probs_k, labels[k]) → loss_k
//!
//! Backward chain:
//! grad_probs → heads_backward → grad_h_new + grad_W_heads/b_heads
//! → cfc_step_backward → grad_W_in/W_rec/b + grad_x (=grad_h_enriched)
//! → Mamba2.backward_from_h_enriched → full Mamba2 grad set
//! For each supervised k (kernels accumulate via += into shared grad
//! buffers; CfC + heads weights are shared across positions):
//! grad_probs_k → heads_backward → grad_h_new_k + accum grad_W_heads/b_heads
//! → cfc_step_backward → accum grad_W_in/W_rec/b + grad_x_k
//! grad_x_k stored into grad_h_enriched_seq[k]
//! Once:
//! grad_h_enriched_seq → Mamba2.backward_from_h_enriched_seq → full Mamba2 grad set
//!
//! Optimizers:
//! Per-step supervision densifies the gradient signal ~K× over the
//! final-step-only design (which trained flat at chance on real ES
//! data despite working on synthetic overfit). Mamba2's analytical
//! scan backward unrolls those per-step gradients through the full
//! 32-step SSM state evolution.
//!
//! Optimizers (unchanged):
//! - 5 CfC AdamWs (W_in, W_rec, b, heads_w, heads_b)
//! - 1 Mamba2AdamW for all 9 Mamba2 parameter tensors
//!
@@ -251,13 +263,18 @@ impl PerceptionTrainer {
})
}
/// One training step on a sequence of `seq_len` snapshots with one
/// label set (per-horizon, applied at the last position). Returns
/// the BCE loss.
/// One training step on a sequence of `seq_len` snapshots with
/// per-position labels. NaN entries in `labels_per_position[k][h]`
/// signal that position `k`'s horizon `h` is invalid (right-edge or
/// tied price) — supervision skips that horizon at that position
/// but other horizons at the same position still contribute.
///
/// Returns the mean BCE loss over (position × horizon) pairs that
/// had finite labels.
pub fn step(
&mut self,
snapshots: &[Mbp10RawInput],
labels: &[f32; N_HORIZONS],
labels_per_position: &[[f32; N_HORIZONS]],
) -> Result<f32> {
anyhow::ensure!(
snapshots.len() == self.cfg.seq_len,
@@ -265,6 +282,12 @@ impl PerceptionTrainer {
snapshots.len(),
self.cfg.seq_len
);
anyhow::ensure!(
labels_per_position.len() == self.cfg.seq_len,
"labels_per_position.len()={} != seq_len={}",
labels_per_position.len(),
self.cfg.seq_len
);
// 1. Build the window tensor: pack seq_len snap_features.
let mut window_tensor = GpuTensor::zeros(
@@ -274,7 +297,6 @@ impl PerceptionTrainer {
.map_err(|e| anyhow::anyhow!("window alloc: {e}"))?;
for (k, snap) in snapshots.iter().enumerate() {
// Upload raw input slots.
upload_into(&self.stream, &snap.bid_px, &self.stg_bid_px, &mut self.bid_px_d)?;
upload_into(&self.stream, &snap.bid_sz, &self.stg_bid_sz, &mut self.bid_sz_d)?;
upload_into(&self.stream, &snap.ask_px, &self.stg_ask_px, &mut self.ask_px_d)?;
@@ -300,128 +322,229 @@ impl PerceptionTrainer {
.arg(&mut self.snap_feat_d);
unsafe { launch.launch(cfg1).context("snap fwd")?; }
}
// Copy snap_feat_d into window_tensor at offset k*FEATURE_DIM (DtoD).
let nbytes = FEATURE_DIM * std::mem::size_of::<f32>();
unsafe {
let (src_ptr, _g1) = self.snap_feat_d.device_ptr(&self.stream);
let (dst_base, _g2) = window_tensor.data_mut().device_ptr_mut(&self.stream);
let dst_offset_ptr = dst_base + (k * FEATURE_DIM * std::mem::size_of::<f32>()) as u64;
cudarc::driver::result::memcpy_dtod_async(
dst_offset_ptr,
src_ptr,
nbytes,
self.stream.cu_stream(),
dst_offset_ptr, src_ptr, nbytes, self.stream.cu_stream(),
)
.context("window pack DtoD")?;
}
}
self.stream.synchronize().context("window pack sync")?;
// 2. Mamba2 forward — emits (logit, cache); we use cache.h_enriched.
let (_logit, cache) = self.mamba2.forward_train(&window_tensor).context("mamba2 fwd")?;
// h_enriched: [1, HIDDEN_DIM]. Copy its slice into CfC's input
// buffer h_old or directly use as `x` for cfc_step.
// 2. Mamba2 per-step forward → h_enriched_seq [1, K, HIDDEN_DIM].
let (h_enriched_seq, cache) = self
.mamba2
.forward_train_seq(&window_tensor)
.context("mamba2 forward_train_seq")?;
// 3. cfc_step: x = h_enriched, h_old = zeros.
// Zero h_old first.
// 3. Pre-zero CfC + heads grad buffers (kernels accumulate via +=
// across the K per-position calls below).
self.stream
.memset_zeros(&mut self.grad_w_in_d)
.map_err(|e| anyhow::anyhow!("zero grad_w_in: {e}"))?;
self.stream
.memset_zeros(&mut self.grad_w_rec_d)
.map_err(|e| anyhow::anyhow!("zero grad_w_rec: {e}"))?;
self.stream
.memset_zeros(&mut self.grad_b_d)
.map_err(|e| anyhow::anyhow!("zero grad_b: {e}"))?;
self.stream
.memset_zeros(&mut self.grad_heads_w_d)
.map_err(|e| anyhow::anyhow!("zero grad_heads_w: {e}"))?;
self.stream
.memset_zeros(&mut self.grad_heads_b_d)
.map_err(|e| anyhow::anyhow!("zero grad_heads_b: {e}"))?;
self.stream
.memset_zeros(&mut self.h_old_d)
.map_err(|e| anyhow::anyhow!("zero h_old: {e}"))?;
let dt_s = 1.0_f32; // unit dt for stacked v1 (CfC's role is per-cell scaled tanh)
// 4. Per-position grad_h_enriched_seq accumulator [1, K, HIDDEN_DIM]
// — fed to mamba2.backward_from_h_enriched_seq below.
let mut grad_h_enriched_seq = GpuTensor::zeros(
&[1, self.cfg.seq_len, HIDDEN_DIM],
&self.stream,
)
.map_err(|e| anyhow::anyhow!("grad_h_enriched_seq alloc: {e}"))?;
// Launch configs reused across the K per-position iterations.
let dt_s = 1.0_f32;
let n_in_i = HIDDEN_DIM as i32;
let n_hid_i = HIDDEN_DIM as i32;
let block_dim = 128u32;
let grid_dim = ((HIDDEN_DIM as u32) + block_dim - 1) / block_dim;
let cfg2 = LaunchConfig {
let cfg_cfc = LaunchConfig {
grid_dim: (grid_dim, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: 0,
};
{
let mut launch = self.stream.launch_builder(&self.step_fn);
launch
.arg(&self.w_in_d).arg(&self.w_rec_d).arg(&self.b_d).arg(&self.tau_d)
.arg(cache.h_enriched.cuda_data()).arg(&self.h_old_d)
.arg(&dt_s).arg(&n_in_i).arg(&n_hid_i)
.arg(&mut self.h_new_d);
unsafe { launch.launch(cfg2).context("cfc fwd")?; }
}
// 4. heads forward.
let cfg3 = LaunchConfig {
let cfg_heads = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (N_HORIZONS as u32, 1, 1),
shared_mem_bytes: 0,
};
{
let mut launch = self.stream.launch_builder(&self.heads_fn);
launch.arg(&self.heads_w_d).arg(&self.heads_b_d).arg(&self.h_new_d).arg(&mut self.probs_d);
unsafe { launch.launch(cfg3).context("heads fwd")?; }
}
self.stream.synchronize().context("fwd sync")?;
// 5. BCE forward + grad_probs.
let probs_host = download(&self.stream, &self.probs_d)?;
let bce_out = bce_multi_horizon_loss_and_grad_gpu(
&self.dev,
&BceInput { probs: probs_host.clone(), labels: labels.to_vec(), n_horizons: N_HORIZONS, n_pos: 1 },
)?;
let loss = bce_out.loss;
// 6. heads backward → grad_h_new + grad_heads_*
let grad_probs_d = upload(&self.stream, &bce_out.grad_probs)?;
let probs_d_alias = upload(&self.stream, &probs_host)?;
{
let mut launch = self.stream.launch_builder(&self.heads_bwd_fn);
launch
.arg(&self.heads_w_d).arg(&probs_d_alias).arg(&self.h_new_d).arg(&grad_probs_d)
.arg(&mut self.grad_heads_w_d).arg(&mut self.grad_heads_b_d).arg(&mut self.grad_h_new_d);
unsafe { launch.launch(cfg3).context("heads bwd")?; }
}
// 7. cfc_step_backward → grad_W_in/W_rec/b + grad_h_old (discard) + grad_x (=grad_h_enriched)
let shared_mem = (2 * HIDDEN_DIM * std::mem::size_of::<f32>()) as u32;
let cfg_bwd = LaunchConfig {
let cfg_cfc_bwd = LaunchConfig {
grid_dim: (grid_dim, 1, 1),
block_dim: (block_dim, 1, 1),
shared_mem_bytes: shared_mem,
};
{
let mut launch = self.stream.launch_builder(&self.step_bwd_fn);
launch
.arg(&self.w_in_d).arg(&self.w_rec_d).arg(&self.b_d).arg(&self.tau_d)
.arg(cache.h_enriched.cuda_data()).arg(&self.h_old_d).arg(&self.grad_h_new_d)
.arg(&dt_s).arg(&n_in_i).arg(&n_hid_i)
.arg(&mut self.grad_w_in_d).arg(&mut self.grad_w_rec_d)
.arg(&mut self.grad_b_d).arg(&mut self.grad_h_old_d)
.arg(&mut self.grad_x_d);
unsafe { launch.launch(cfg_bwd).context("cfc bwd")?; }
}
self.stream.synchronize().context("cfc bwd sync")?;
// 8. Mamba2 backward: wrap grad_x_d as GpuTensor [1, HIDDEN_DIM] for backward_from_h_enriched.
// Need to construct a fresh GpuTensor from our CudaSlice. The
// backward API takes a borrow; build a temporary GpuTensor
// copy. (For v2, allocate this once and reuse.)
let mut grad_h_enriched_slice = self.stream
.alloc_zeros::<f32>(HIDDEN_DIM)
.context("grad_h_enriched alloc")?;
let nbytes = HIDDEN_DIM * std::mem::size_of::<f32>();
unsafe {
let (src_ptr, _g1) = self.grad_x_d.device_ptr(&self.stream);
let (dst_ptr, _g2) = grad_h_enriched_slice.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(dst_ptr, src_ptr, nbytes, self.stream.cu_stream())
.context("grad_h_enriched copy")?;
// 5. K iterations: per-position forward (CfC + heads + BCE) +
// backward through heads + CfC. CfC/heads grads accumulate
// into shared buffers; grad_x_k goes into slot k of
// grad_h_enriched_seq.
let mut total_loss = 0.0f32;
let mut n_supervised = 0usize;
let h_enriched_slice_bytes = HIDDEN_DIM * std::mem::size_of::<f32>();
for k in 0..self.cfg.seq_len {
// Position k's labels — skip if ALL horizons are NaN (no
// supervision at this position).
let labels_k = &labels_per_position[k];
if labels_k.iter().all(|v| !v.is_finite()) {
continue;
}
// Slice h_enriched_seq[0, k, :] into self.snap_feat_d? No —
// we use the raw CudaSlice slot at the right offset via a
// small DtoD copy into a single-position scratch. Actually
// cleaner: compute the offset pointer and pass it directly
// to the cfc_step kernel.
// Allocate per-position scratch x_k and copy from
// h_enriched_seq[k] (DtoD).
let mut x_k = self.stream
.alloc_zeros::<f32>(HIDDEN_DIM)
.context("x_k alloc")?;
unsafe {
let (src_base, _g1) = h_enriched_seq.cuda_data().device_ptr(&self.stream);
let src_off = src_base + (k * h_enriched_slice_bytes) as u64;
let (dst_ptr, _g2) = x_k.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(
dst_ptr, src_off, h_enriched_slice_bytes, self.stream.cu_stream(),
)
.context("x_k DtoD copy")?;
}
// CfC forward (h_old already zero from step 3).
{
let mut launch = self.stream.launch_builder(&self.step_fn);
launch
.arg(&self.w_in_d).arg(&self.w_rec_d).arg(&self.b_d).arg(&self.tau_d)
.arg(&x_k).arg(&self.h_old_d)
.arg(&dt_s).arg(&n_in_i).arg(&n_hid_i)
.arg(&mut self.h_new_d);
unsafe { launch.launch(cfg_cfc).context("cfc fwd k")?; }
}
// Heads forward.
{
let mut launch = self.stream.launch_builder(&self.heads_fn);
launch.arg(&self.heads_w_d).arg(&self.heads_b_d)
.arg(&self.h_new_d).arg(&mut self.probs_d);
unsafe { launch.launch(cfg_heads).context("heads fwd k")?; }
}
self.stream.synchronize().context("fwd sync k")?;
// BCE: skip horizons with NaN labels by zeroing their grad
// contribution (BCE kernel currently treats them numerically;
// we mask after the fact).
let probs_host = download(&self.stream, &self.probs_d)?;
let mut labels_clean = [0.5f32; N_HORIZONS]; // 0.5 = no-op for BCE if mask zeros grad
let mut mask = [0.0f32; N_HORIZONS];
let mut n_valid_k = 0;
for h in 0..N_HORIZONS {
if labels_k[h].is_finite() {
labels_clean[h] = labels_k[h];
mask[h] = 1.0;
n_valid_k += 1;
}
}
if n_valid_k == 0 { continue; }
let bce_out = bce_multi_horizon_loss_and_grad_gpu(
&self.dev,
&BceInput {
probs: probs_host.clone(),
labels: labels_clean.to_vec(),
n_horizons: N_HORIZONS,
n_pos: 1,
},
)?;
// Mask grad_probs at NaN-label horizons + count only valid loss.
let mut grad_probs_masked = bce_out.grad_probs.clone();
let mut loss_k = 0.0f32;
for h in 0..N_HORIZONS {
if mask[h] == 0.0 {
grad_probs_masked[h] = 0.0;
} else {
// Recover per-horizon BCE: -y*log(p) - (1-y)*log(1-p).
let p = probs_host[h].clamp(1e-7, 1.0 - 1e-7);
let y = labels_clean[h];
loss_k += -(y * p.ln() + (1.0 - y) * (1.0 - p).ln());
}
}
total_loss += loss_k / n_valid_k as f32;
n_supervised += 1;
// Heads backward: accumulates grad_heads_w/b via += into
// self.grad_heads_w_d / self.grad_heads_b_d.
let grad_probs_d = upload(&self.stream, &grad_probs_masked)?;
let probs_d_alias = upload(&self.stream, &probs_host)?;
{
let mut launch = self.stream.launch_builder(&self.heads_bwd_fn);
launch
.arg(&self.heads_w_d).arg(&probs_d_alias)
.arg(&self.h_new_d).arg(&grad_probs_d)
.arg(&mut self.grad_heads_w_d).arg(&mut self.grad_heads_b_d)
.arg(&mut self.grad_h_new_d);
unsafe { launch.launch(cfg_heads).context("heads bwd k")?; }
}
// CfC backward: accumulates grad_w_in/grad_w_rec/grad_b via
// += into shared buffers; emits grad_x_k → slot k of
// grad_h_enriched_seq.
{
let mut launch = self.stream.launch_builder(&self.step_bwd_fn);
launch
.arg(&self.w_in_d).arg(&self.w_rec_d).arg(&self.b_d).arg(&self.tau_d)
.arg(&x_k).arg(&self.h_old_d).arg(&self.grad_h_new_d)
.arg(&dt_s).arg(&n_in_i).arg(&n_hid_i)
.arg(&mut self.grad_w_in_d).arg(&mut self.grad_w_rec_d)
.arg(&mut self.grad_b_d).arg(&mut self.grad_h_old_d)
.arg(&mut self.grad_x_d);
unsafe { launch.launch(cfg_cfc_bwd).context("cfc bwd k")?; }
}
// Write grad_x_k into grad_h_enriched_seq[0, k, :].
unsafe {
let (src_ptr, _g1) = self.grad_x_d.device_ptr(&self.stream);
let (dst_base, _g2) = grad_h_enriched_seq.data_mut().device_ptr_mut(&self.stream);
let dst_off = dst_base + (k * h_enriched_slice_bytes) as u64;
cudarc::driver::result::memcpy_dtod_async(
dst_off, src_ptr, h_enriched_slice_bytes, self.stream.cu_stream(),
)
.context("grad_h_enriched_seq slot copy")?;
}
}
let grad_h_enriched_tensor =
GpuTensor::new(grad_h_enriched_slice, vec![1, HIDDEN_DIM])
.map_err(|e| anyhow::anyhow!("grad_h_enriched as GpuTensor: {e}"))?;
self.stream.synchronize().context("per-position bwd sync")?;
if n_supervised == 0 {
// No supervised positions — skip the optimizer step.
return Ok(0.0);
}
// 6. Mamba2 backward — single call consumes the full per-step
// grad_h_enriched_seq, unrolling SSM gradients through all K
// timesteps.
let mamba2_grads = self
.mamba2
.backward_from_h_enriched(&cache, &grad_h_enriched_tensor)
.context("mamba2 backward_from_h_enriched")?;
.backward_from_h_enriched_seq(&cache, &grad_h_enriched_seq)
.context("mamba2 backward_from_h_enriched_seq")?;
// 9. Apply AdamW updates on all 6 param groups.
// 7. Apply AdamW updates on all 6 param groups.
self.opt_w_in.step(&mut self.w_in_d, &self.grad_w_in_d)?;
self.opt_w_rec.step(&mut self.w_rec_d, &self.grad_w_rec_d)?;
self.opt_b.step(&mut self.b_d, &self.grad_b_d)?;
@@ -431,7 +554,7 @@ impl PerceptionTrainer {
.step(&mut self.mamba2, &mamba2_grads)
.context("mamba2 AdamW step")?;
Ok(loss)
Ok(total_loss / n_supervised as f32)
}
pub fn last_probs(&self) -> Result<[f32; N_HORIZONS]> {

View File

@@ -14,7 +14,11 @@ fn test_device() -> MlDevice {
MlDevice::cuda(0).expect("CUDA 0 required for ml-alpha tests")
}
fn synthetic_seq(seq_len: usize, mut prev_mid: f32, mut ts_ns: u64) -> (Vec<Mbp10RawInput>, [f32; 5]) {
fn synthetic_seq(
seq_len: usize,
mut prev_mid: f32,
mut ts_ns: u64,
) -> (Vec<Mbp10RawInput>, Vec<[f32; 5]>) {
let mut out = Vec::with_capacity(seq_len);
for k in 0..seq_len {
let next_mid = prev_mid + 0.25;
@@ -31,10 +35,7 @@ fn synthetic_seq(seq_len: usize, mut prev_mid: f32, mut ts_ns: u64) -> (Vec<Mbp1
let prev_ts = ts_ns;
ts_ns += 20_000_000;
out.push(Mbp10RawInput {
bid_px,
bid_sz,
ask_px,
ask_sz,
bid_px, bid_sz, ask_px, ask_sz,
prev_mid,
trade_signed_vol: 1.0,
trade_count: 1,
@@ -44,7 +45,11 @@ fn synthetic_seq(seq_len: usize, mut prev_mid: f32, mut ts_ns: u64) -> (Vec<Mbp1
prev_mid = next_mid;
let _ = k;
}
(out, [1.0, 1.0, 1.0, 1.0, 1.0])
// Per-position labels: every position knows the next K snapshots all
// move up (synthetic monotone ramp), so label = 1.0 for every horizon
// at every position. Drives the trainer to learn "always predict 1".
let labels = vec![[1.0; 5]; seq_len];
(out, labels)
}
#[test]
@@ -73,7 +78,7 @@ fn stacked_trainer_loss_shrinks_on_constant_signal() {
let mut prev_mid = 5500.0_f32;
for _ in 0..8 {
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
let l = trainer.step(&seq, &labels).expect("step warm");
let l = trainer.step(&seq, labels.as_slice()).expect("step warm");
initial_total += l;
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
ts = seq.last().unwrap().ts_ns;
@@ -86,7 +91,7 @@ fn stacked_trainer_loss_shrinks_on_constant_signal() {
let mut window_count = 0usize;
for step_idx in 0..250 {
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
let l = trainer.step(&seq, &labels).expect("train step");
let l = trainer.step(&seq, labels.as_slice()).expect("train step");
window_loss += l;
window_count += 1;
if step_idx % 50 == 49 {
@@ -106,7 +111,7 @@ fn stacked_trainer_loss_shrinks_on_constant_signal() {
let mut final_total = 0.0_f32;
for _ in 0..8 {
let (seq, labels) = synthetic_seq(cfg.seq_len, prev_mid, ts);
let l = trainer.step(&seq, &labels).expect("step final");
let l = trainer.step(&seq, labels.as_slice()).expect("step final");
final_total += l;
prev_mid = 0.5 * (seq.last().unwrap().bid_px[0] + seq.last().unwrap().ask_px[0]);
ts = seq.last().unwrap().ts_ns;