diff --git a/crates/ml/tests/aux_trunk_oracle_tests.rs b/crates/ml/tests/aux_trunk_oracle_tests.rs index 326ec0d73..88c1d32e0 100644 --- a/crates/ml/tests/aux_trunk_oracle_tests.rs +++ b/crates/ml/tests/aux_trunk_oracle_tests.rs @@ -1292,4 +1292,660 @@ mod gpu { "EMA blend (h=3.0, α={ALPHA}): ISV[449] = {slot_b:.6} (expected {expected:.6})" ); } + + // ════════════════════════════════════════════════════════════════════ + // SP14 Layer C Phase C.9 synthetic smoke (2026-05-08). + // + // Verifies the full aux trunk learning pipeline: + // + // - aux_trunk_learns_synthetic_uptrend — 100 steps of forward + + // backward + Adam on a synthetic all-UP batch; asserts CE loss < + // 0.1 and dir_acc > 0.95 at convergence. + // + // Architectural intent: if this test FAILS (loss stays near ln(2) ≈ + // 0.693 = random baseline), the aux trunk gradient chain is broken and + // L40S dispatch must be blocked until the root cause is fixed. See + // `pearl_separate_aux_trunk_when_shared_starves.md` — aux-loss + // flatline at random baseline is a structural symptom, not a + // hyperparameter problem. + // + // Plan: docs/superpowers/plans/2026-05-07-sp14-layer-c-separate-aux-trunk.md §C.9 + // ════════════════════════════════════════════════════════════════════ + + const AUX_HEADS_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/aux_heads_kernel.cubin")); + const DQN_UTILITY_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/dqn_utility_kernels.cubin")); + + struct SmokeKernels { + trunk_fwd: CudaFunction, + trunk_bwd_dh_pre: CudaFunction, + trunk_bwd_dW: CudaFunction, + trunk_bwd_db: CudaFunction, + head_fwd: CudaFunction, + head_loss: CudaFunction, + head_bwd: CudaFunction, + grad_norm: CudaFunction, + grad_norm_final: CudaFunction, + adam: CudaFunction, + } + + fn load_smoke_kernels(stream: &Arc) -> SmokeKernels { + let trunk_fwd_mod = stream.context() + .load_cubin(AUX_TRUNK_FORWARD_CUBIN.to_vec()) + .expect("load trunk_fwd cubin"); + let trunk_bwd_mod = stream.context() + .load_cubin(AUX_TRUNK_BACKWARD_CUBIN.to_vec()) + .expect("load trunk_bwd cubin"); + let heads_mod = stream.context() + .load_cubin(AUX_HEADS_CUBIN.to_vec()) + .expect("load aux_heads cubin"); + let util_mod = stream.context() + .load_cubin(DQN_UTILITY_CUBIN.to_vec()) + .expect("load dqn_utility cubin"); + + SmokeKernels { + trunk_fwd: trunk_fwd_mod.load_function("aux_trunk_forward").expect("aux_trunk_forward"), + trunk_bwd_dh_pre: trunk_bwd_mod.load_function("aux_trunk_bwd_dh_pre").expect("bwd_dh_pre"), + trunk_bwd_dW: trunk_bwd_mod.load_function("aux_trunk_bwd_dW_reduce").expect("bwd_dW"), + trunk_bwd_db: trunk_bwd_mod.load_function("aux_trunk_bwd_db_reduce").expect("bwd_db"), + head_fwd: heads_mod.load_function("aux_next_bar_forward").expect("head_fwd"), + head_loss: heads_mod.load_function("aux_next_bar_loss_reduce").expect("head_loss"), + head_bwd: heads_mod.load_function("aux_next_bar_backward").expect("head_bwd"), + grad_norm: util_mod.load_function("dqn_grad_norm_kernel").expect("grad_norm"), + grad_norm_final: util_mod.load_function("dqn_grad_norm_finalize").expect("grad_norm_fin"), + adam: util_mod.load_function("dqn_adam_update_kernel").expect("adam"), + } + } + + /// Synthetic learning smoke: 100 forward + backward + Adam steps on an + /// all-UP batch. Asserts CE loss < 0.1 and dir_acc > 0.95. + /// + /// Topology (small dims for speed on RTX 3050): + /// B=16, ENC=32, H1_trunk=32, H2_trunk=16, SH2=32 + /// H_HEAD=32 (compile-time in aux_heads_kernel.cu), K=2 + /// + /// Only the 6 aux trunk weight tensors are updated via Adam. The 4 aux + /// head params are updated via simple SGD on the host (test harness only) + /// by reading the per-sample partial gradients through mapped-pinned memory + /// and summing them in Rust. This is valid test orchestration — not a + /// CPU compute path in production. + /// + /// Failure condition: loss > 0.1 after 100 steps → gradient chain broken. + #[test] + #[ignore = "requires GPU"] + fn aux_trunk_learns_synthetic_uptrend() { + const B: usize = 16; + const ENC: usize = 32; // encoder_out_dim fed into trunk + const H1: usize = 32; // trunk hidden layer 1 + const H2: usize = 16; // trunk hidden layer 2 + const SH2: usize = 32; // trunk output dim = head input dim + const H_HEAD: usize = 32; // AUX_HIDDEN_DIM compile-time in aux_heads_kernel.cu + const K: usize = 2; // binary: class 0 = DOWN, class 1 = UP + const STEPS: usize = 100; + const LR_TRUNK: f32 = 3e-3; // aggressive LR for fast convergence on clean signal + const LR_HEAD: f32 = 3e-3; + const LOSS_THRESHOLD: f32 = 0.1; + const DIR_ACC_THRESHOLD: f32 = 0.95; + + let stream = make_test_stream(); + let k = load_smoke_kernels(&stream); + + let mut rng: u32 = 0xDEAD_BEEF; + + // ── Allocate aux trunk params + m/v ───────────────────────────── + // w1 [ENC, H1], b1 [H1], w2 [H1, H2], b2 [H2], w3 [H2, SH2], b3 [SH2] + let sizes_trunk: [usize; 6] = [ + ENC * H1, H1, H1 * H2, H2, H2 * SH2, SH2, + ]; + let total_trunk: usize = sizes_trunk.iter().sum(); + + macro_rules! xavier_buf { + ($n:expr, $fan:expr) => {{ + let buf = unsafe { MappedF32Buffer::new($n) }.expect("alloc"); + let host = unsafe { std::slice::from_raw_parts_mut(buf.host_ptr, $n) }; + let bound = (2.0_f32 / $fan as f32).sqrt(); + for v in host.iter_mut() { + *v = lcg_next_signed(&mut rng) * bound; + } + buf + }}; + } + macro_rules! zero_buf { + ($n:expr) => {{ + let buf = unsafe { MappedF32Buffer::new($n) }.expect("alloc"); + let host = unsafe { std::slice::from_raw_parts_mut(buf.host_ptr, $n) }; + for v in host.iter_mut() { *v = 0.0; } + buf + }}; + } + + // Trunk weights — Xavier init + let w1 = xavier_buf!(ENC * H1, ENC + H1); + let b1 = zero_buf!(H1); + let w2 = xavier_buf!(H1 * H2, H1 + H2); + let b2 = zero_buf!(H2); + let w3 = xavier_buf!(H2 * SH2, H2 + SH2); + let b3 = zero_buf!(SH2); + // Adam m/v for trunk — zero init + let w1_m = zero_buf!(ENC * H1); let w1_v = zero_buf!(ENC * H1); + let b1_m = zero_buf!(H1); let b1_v = zero_buf!(H1); + let w2_m = zero_buf!(H1 * H2); let w2_v = zero_buf!(H1 * H2); + let b2_m = zero_buf!(H2); let b2_v = zero_buf!(H2); + let w3_m = zero_buf!(H2 * SH2); let w3_v = zero_buf!(H2 * SH2); + let b3_m = zero_buf!(SH2); let b3_v = zero_buf!(SH2); + // Adam grad buffers for trunk — zero init + let w1_g = zero_buf!(ENC * H1); + let b1_g = zero_buf!(H1); + let w2_g = zero_buf!(H1 * H2); + let b2_g = zero_buf!(H2); + let w3_g = zero_buf!(H2 * SH2); + let b3_g = zero_buf!(SH2); + // Adam utility buffers for trunk (block_sums/grad_norm_buf replaced by flat_grad path below) + let _norm_blocks = (total_trunk + 255) / 256; + let _grad_norm_buf = zero_buf!(1); + let wd_mask = { + let buf = unsafe { MappedF32Buffer::new(total_trunk) }.expect("alloc"); + let host = unsafe { std::slice::from_raw_parts_mut(buf.host_ptr, total_trunk) }; + for v in host.iter_mut() { *v = 0.0; } // no weight decay + buf + }; + let nan_flags = zero_buf!(2); // diag_slot=0, tiny buf + let engage_buf = zero_buf!(1); // unused (engage_buf_offset=-1) + // Adam step counter + LR + clip — mapped-pinned scalars + let adam_t_buf = { + let buf = unsafe { MappedF32Buffer::new(1) }.expect("alloc t_buf"); + // We'll write an i32 through the f32 mapped-pinned buffer host_ptr. + // Safe because i32 and f32 are same size and we treat it as raw bytes. + buf + }; + let lr_buf = { + let buf = unsafe { MappedF32Buffer::new(1) }.expect("alloc lr"); + let host = unsafe { std::slice::from_raw_parts_mut(buf.host_ptr, 1) }; + host[0] = LR_TRUNK; + buf + }; + let clip_buf = { + let buf = unsafe { MappedF32Buffer::new(1) }.expect("alloc clip"); + let host = unsafe { std::slice::from_raw_parts_mut(buf.host_ptr, 1) }; + host[0] = 10.0_f32; // generous clip + buf + }; + + // ── Allocate aux head params (trained by host-side SGD) ────────── + // wh1 [H_HEAD, SH2], bh1 [H_HEAD], wh2 [K, H_HEAD], bh2 [K] + let wh1 = xavier_buf!(H_HEAD * SH2, H_HEAD + SH2); + let bh1 = zero_buf!(H_HEAD); + let wh2 = xavier_buf!(K * H_HEAD, K + H_HEAD); + let bh2 = zero_buf!(K); + + // ── Allocate forward scratch buffers ───────────────────────────── + let x_in = zero_buf!(B * ENC); + let h_aux1 = zero_buf!(B * H1); + let h_aux2 = zero_buf!(B * H2); + let h_s2_aux = zero_buf!(B * SH2); + let hidden_out = zero_buf!(B * H_HEAD); + let logits_out = zero_buf!(B * K); + let softmax_out = zero_buf!(B * K); + let loss_out = zero_buf!(1); + let valid_count = zero_buf!(1); + let dh_s2_aux_out = zero_buf!(B * SH2); + + // ── Allocate head backward partial-grad buffers ────────────────── + // Per-sample partials — summed on host, applied as SGD + let dWh1_partial = zero_buf!(B * H_HEAD * SH2); + let dbh1_partial = zero_buf!(B * H_HEAD); + let dWh2_partial = zero_buf!(B * K * H_HEAD); + let dbh2_partial = zero_buf!(B * K); + + // ── Allocate labels buffer (all UP = label 1) ──────────────────── + // Labels are i32; reuse a MappedF32Buffer for the raw pointer. + let labels_buf = { + let buf = unsafe { MappedF32Buffer::new(B) }.expect("alloc labels"); + let host = unsafe { + std::slice::from_raw_parts_mut(buf.host_ptr as *mut i32, B) + }; + for v in host.iter_mut() { *v = 1; } // all UP + buf + }; + + // ── Synthetic encoder output: constant signal ──────────────────── + // All samples share the same non-zero encoder output so the trunk + // can learn a direction-invariant mapping to the "UP" class. + { + let host = unsafe { std::slice::from_raw_parts_mut(x_in.host_ptr, B * ENC) }; + for (i, v) in host.iter_mut().enumerate() { + // Monotone: feature[j] = +1/(j+1) — same for all samples. + let j = i % ENC; + *v = 0.5_f32 / (j as f32 + 1.0); + } + } + + let mut final_loss = f32::NAN; + let mut final_acc = f32::NAN; + + for step in 0..STEPS { + // ── Step counter (i32 written through f32 mapped-pinned buf) ── + let t_val: i32 = (step + 1) as i32; + unsafe { + *(adam_t_buf.host_ptr as *mut i32) = t_val; + } + + // ── 1. Aux trunk forward ────────────────────────────────────── + let b_i = B as i32; + let enc_i = ENC as i32; + let h1_i = H1 as i32; + let h2_i = H2 as i32; + let sh2_i = SH2 as i32; + let smem_trunk = ((H1 + H2) * std::mem::size_of::()) as u32; + unsafe { + stream.launch_builder(&k.trunk_fwd) + .arg(&x_in.dev_ptr) + .arg(&w1.dev_ptr).arg(&b1.dev_ptr) + .arg(&w2.dev_ptr).arg(&b2.dev_ptr) + .arg(&w3.dev_ptr).arg(&b3.dev_ptr) + .arg(&h_aux1.dev_ptr) + .arg(&h_aux2.dev_ptr) + .arg(&h_s2_aux.dev_ptr) + .arg(&b_i).arg(&enc_i).arg(&h1_i).arg(&h2_i).arg(&sh2_i) + .launch(LaunchConfig { + grid_dim: (B as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_trunk, + }) + .expect("trunk_fwd launch"); + } + stream.synchronize().expect("sync trunk_fwd"); + + // ── 2. Aux head forward ─────────────────────────────────────── + let _hh_i = H_HEAD as i32; // H_HEAD hard-coded in aux_heads_kernel.cu + let k_i = K as i32; + let smem_head = ((H_HEAD + K) * std::mem::size_of::()) as u32; + unsafe { + stream.launch_builder(&k.head_fwd) + .arg(&h_s2_aux.dev_ptr) + .arg(&wh1.dev_ptr).arg(&bh1.dev_ptr) + .arg(&wh2.dev_ptr).arg(&bh2.dev_ptr) + .arg(&b_i).arg(&sh2_i).arg(&k_i) + .arg(&hidden_out.dev_ptr) + .arg(&logits_out.dev_ptr) + .arg(&softmax_out.dev_ptr) + .launch(LaunchConfig { + grid_dim: (B as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_head, + }) + .expect("head_fwd launch"); + } + stream.synchronize().expect("sync head_fwd"); + + // ── 3. CE loss ──────────────────────────────────────────────── + let smem_loss = (2 * AUX_TRUNK_BLOCK as usize * std::mem::size_of::()) as u32; + unsafe { + stream.launch_builder(&k.head_loss) + .arg(&softmax_out.dev_ptr) + .arg(&labels_buf.dev_ptr) + .arg(&b_i).arg(&k_i) + .arg(&loss_out.dev_ptr) + .arg(&valid_count.dev_ptr) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_loss, + }) + .expect("head_loss launch"); + } + stream.synchronize().expect("sync head_loss"); + + // ── 4. Head backward (per-sample partial grads + dh_s2_aux) ── + let smem_bwd = ((2 * H_HEAD + K) * std::mem::size_of::()) as u32; + unsafe { + // Zero dh_s2_aux_out before backward (accumulate fresh each step) + let dh_host = std::slice::from_raw_parts_mut(dh_s2_aux_out.host_ptr, B * SH2); + for v in dh_host.iter_mut() { *v = 0.0; } + stream.launch_builder(&k.head_bwd) + .arg(&h_s2_aux.dev_ptr) + .arg(&wh1.dev_ptr) + .arg(&wh2.dev_ptr) + .arg(&hidden_out.dev_ptr) + .arg(&softmax_out.dev_ptr) + .arg(&labels_buf.dev_ptr) + .arg(&valid_count.dev_ptr) + .arg(&b_i).arg(&sh2_i).arg(&k_i) + .arg(&dWh1_partial.dev_ptr).arg(&dbh1_partial.dev_ptr) + .arg(&dWh2_partial.dev_ptr).arg(&dbh2_partial.dev_ptr) + .arg(&dh_s2_aux_out.dev_ptr) + .launch(LaunchConfig { + grid_dim: (B as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_bwd, + }) + .expect("head_bwd launch"); + } + stream.synchronize().expect("sync head_bwd"); + + // ── 5. Host-side SGD on head params ────────────────────────── + // Read per-sample partials via mapped-pinned, sum → SGD update. + // This is test-harness orchestration — the production aux head + // trainer does not execute on the CPU. + { + // dWh1: sum over B dim, shape [H_HEAD, SH2] + let dwh1_h = unsafe { std::slice::from_raw_parts(dWh1_partial.host_ptr, B * H_HEAD * SH2) }; + let wh1_h = unsafe { std::slice::from_raw_parts_mut(wh1.host_ptr, H_HEAD * SH2) }; + for idx in 0..H_HEAD * SH2 { + let mut g = 0.0_f32; + for b_idx in 0..B { g += dwh1_h[b_idx * H_HEAD * SH2 + idx]; } + wh1_h[idx] -= LR_HEAD * g; + } + // dbh1 + let dbh1_h = unsafe { std::slice::from_raw_parts(dbh1_partial.host_ptr, B * H_HEAD) }; + let bh1_h = unsafe { std::slice::from_raw_parts_mut(bh1.host_ptr, H_HEAD) }; + for idx in 0..H_HEAD { + let mut g = 0.0_f32; + for b_idx in 0..B { g += dbh1_h[b_idx * H_HEAD + idx]; } + bh1_h[idx] -= LR_HEAD * g; + } + // dWh2: shape [K, H_HEAD] + let dwh2_h = unsafe { std::slice::from_raw_parts(dWh2_partial.host_ptr, B * K * H_HEAD) }; + let wh2_h = unsafe { std::slice::from_raw_parts_mut(wh2.host_ptr, K * H_HEAD) }; + for idx in 0..K * H_HEAD { + let mut g = 0.0_f32; + for b_idx in 0..B { g += dwh2_h[b_idx * K * H_HEAD + idx]; } + wh2_h[idx] -= LR_HEAD * g; + } + // dbh2 + let dbh2_h = unsafe { std::slice::from_raw_parts(dbh2_partial.host_ptr, B * K) }; + let bh2_h = unsafe { std::slice::from_raw_parts_mut(bh2.host_ptr, K) }; + for idx in 0..K { + let mut g = 0.0_f32; + for b_idx in 0..B { g += dbh2_h[b_idx * K + idx]; } + bh2_h[idx] -= LR_HEAD * g; + } + } + + // ── 6. Aux trunk backward ───────────────────────────────────── + // + // Kernel signatures (from aux_trunk_backward_kernel.cu): + // + // aux_trunk_bwd_dh_pre(d_logits[B,SH2], w3[H2,SH2], w2[H1,H2], + // h_aux1[B,H1], h_aux2[B,H2], + // dh_aux2_pre_out[B,H2], dh_aux1_pre_out[B,H1], + // B, H1, H2, AUX_HIDDEN_DIM=SH2) + // shmem: H2 * sizeof(f32) (sh_dh2_pre cache) + // grid: (B,1,1) + // + // d_logits here = dh_s2_aux_out from head backward (upstream grad). + // AUX_HIDDEN_DIM = SH2 in the kernel (the trunk output / head input dim). + // + // aux_trunk_bwd_dW_reduce(A[B,Krows], B_grad[B,Jcols], dW_out[Krows,Jcols], + // B, Krows, Jcols) + // shmem: 256 * sizeof(f32) grid: (Krows*Jcols, 1, 1) + // Called 3×: + // dW3: A=h_aux2[B,H2], B_grad=dh_s2_aux_out[B,SH2], dW3[H2,SH2] + // dW2: A=h_aux1[B,H1], B_grad=dh_pre2[B,H2], dW2[H1,H2] + // dW1: A=x_in [B,ENC], B_grad=dh_pre1[B,H1], dW1[ENC,H1] + // + // aux_trunk_bwd_db_reduce(B_grad[B,Jcols], db_out[Jcols], B, Jcols) + // shmem: 256 * sizeof(f32) grid: (Jcols, 1, 1) + // Called 3×: db3(SH2), db2(H2), db1(H1) + + // Phase A: dh_pre (layer 2 and 1 pre-activation grads). + let dh_pre1 = zero_buf!(B * H1); // dh_aux1_pre_out [B, H1] + let dh_pre2 = zero_buf!(B * H2); // dh_aux2_pre_out [B, H2] + // shmem = H2 floats (sh_dh2_pre cache used inside kernel) + let smem_dh = (H2 * std::mem::size_of::()) as u32; + let sh2_i_aux_hidden = sh2_i; // AUX_HIDDEN_DIM = SH2 in this test + unsafe { + stream.launch_builder(&k.trunk_bwd_dh_pre) + .arg(&dh_s2_aux_out.dev_ptr) // d_logits [B, SH2] + .arg(&w3.dev_ptr) // w3 [H2, SH2] + .arg(&w2.dev_ptr) // w2 [H1, H2] + .arg(&h_aux1.dev_ptr) // h_aux1 [B, H1] + .arg(&h_aux2.dev_ptr) // h_aux2 [B, H2] + .arg(&dh_pre2.dev_ptr) // dh_aux2_pre_out [B, H2] + .arg(&dh_pre1.dev_ptr) // dh_aux1_pre_out [B, H1] + .arg(&b_i) // B + .arg(&h1_i) // H1 + .arg(&h2_i) // H2 + .arg(&sh2_i_aux_hidden) // AUX_HIDDEN_DIM (=SH2 here) + .launch(LaunchConfig { + grid_dim: (B as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_dh, + }) + .expect("bwd_dh_pre launch"); + } + stream.synchronize().expect("sync bwd_dh_pre"); + + // Phase B: dW reduce — 3 separate launches (generic kernel, one call per layer). + // shmem = 256 * sizeof(f32) for the block-tree-reduce. + let smem_dW = (256 * std::mem::size_of::()) as u32; + let smem_db = (256 * std::mem::size_of::()) as u32; + unsafe { + // dW3: A=h_aux2[B,H2], B_grad=dh_s2_aux_out[B,SH2], dW3[H2,SH2] + stream.launch_builder(&k.trunk_bwd_dW) + .arg(&h_aux2.dev_ptr) + .arg(&dh_s2_aux_out.dev_ptr) + .arg(&w3_g.dev_ptr) + .arg(&b_i).arg(&h2_i).arg(&sh2_i) + .launch(LaunchConfig { + grid_dim: ((H2 * SH2) as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_dW, + }) + .expect("bwd_dW3 launch"); + stream.synchronize().expect("sync bwd_dW3"); + + // dW2: A=h_aux1[B,H1], B_grad=dh_pre2[B,H2], dW2[H1,H2] + stream.launch_builder(&k.trunk_bwd_dW) + .arg(&h_aux1.dev_ptr) + .arg(&dh_pre2.dev_ptr) + .arg(&w2_g.dev_ptr) + .arg(&b_i).arg(&h1_i).arg(&h2_i) + .launch(LaunchConfig { + grid_dim: ((H1 * H2) as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_dW, + }) + .expect("bwd_dW2 launch"); + stream.synchronize().expect("sync bwd_dW2"); + + // dW1: A=x_in[B,ENC], B_grad=dh_pre1[B,H1], dW1[ENC,H1] + stream.launch_builder(&k.trunk_bwd_dW) + .arg(&x_in.dev_ptr) + .arg(&dh_pre1.dev_ptr) + .arg(&w1_g.dev_ptr) + .arg(&b_i).arg(&enc_i).arg(&h1_i) + .launch(LaunchConfig { + grid_dim: ((ENC * H1) as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_dW, + }) + .expect("bwd_dW1 launch"); + stream.synchronize().expect("sync bwd_dW1"); + } + + // Phase C: db reduce — 3 separate launches. + unsafe { + // db3: B_grad=dh_s2_aux_out[B,SH2], db3[SH2] + stream.launch_builder(&k.trunk_bwd_db) + .arg(&dh_s2_aux_out.dev_ptr) + .arg(&b3_g.dev_ptr) + .arg(&b_i).arg(&sh2_i) + .launch(LaunchConfig { + grid_dim: (SH2 as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_db, + }) + .expect("bwd_db3 launch"); + stream.synchronize().expect("sync bwd_db3"); + + // db2: B_grad=dh_pre2[B,H2], db2[H2] + stream.launch_builder(&k.trunk_bwd_db) + .arg(&dh_pre2.dev_ptr) + .arg(&b2_g.dev_ptr) + .arg(&b_i).arg(&h2_i) + .launch(LaunchConfig { + grid_dim: (H2 as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_db, + }) + .expect("bwd_db2 launch"); + stream.synchronize().expect("sync bwd_db2"); + + // db1: B_grad=dh_pre1[B,H1], db1[H1] + stream.launch_builder(&k.trunk_bwd_db) + .arg(&dh_pre1.dev_ptr) + .arg(&b1_g.dev_ptr) + .arg(&b_i).arg(&h1_i) + .launch(LaunchConfig { + grid_dim: (H1 as u32, 1, 1), + block_dim: (AUX_TRUNK_BLOCK, 1, 1), + shared_mem_bytes: smem_db, + }) + .expect("bwd_db1 launch"); + stream.synchronize().expect("sync bwd_db1"); + } + + // ── 7. Grad norm for trunk (Phase 1: standalone per-tensor) ─── + // Concatenate all trunk grad buffers logically via flattened ptrs. + // We do ONE grad_norm over the concatenated logical gradient space. + // Since grad bufs are NOT contiguous, we use the dqn_grad_norm_kernel + // individually per tensor, accumulate block_sums, then finalize once. + // + // Simpler approach: since total_trunk is small (≤ 5K), a single + // per-elem flat reduction over a scratch grad buffer is fine. + // We write a flat copy to a contiguous scratch, run grad_norm once, + // then run 6 Adam launches. + let flat_grad = { + let buf = unsafe { MappedF32Buffer::new(total_trunk) }.expect("flat_grad"); + let host = unsafe { std::slice::from_raw_parts_mut(buf.host_ptr, total_trunk) }; + let g_bufs: [(&MappedF32Buffer, usize); 6] = [ + (&w1_g, ENC*H1), (&b1_g, H1), (&w2_g, H1*H2), + (&b2_g, H2), (&w3_g, H2*SH2), (&b3_g, SH2), + ]; + let mut off = 0; + for (buf_ref, n) in g_bufs { + let src = unsafe { std::slice::from_raw_parts(buf_ref.host_ptr, n) }; + host[off..off+n].copy_from_slice(src); + off += n; + } + buf + }; + let flat_norm_blocks = (total_trunk + 255) / 256; + let flat_block_sums = zero_buf!(flat_norm_blocks); + let flat_grad_norm = zero_buf!(1); + { + let total_i = total_trunk as i32; + let num_b_i = flat_norm_blocks as i32; + unsafe { + stream.launch_builder(&k.grad_norm) + .arg(&flat_grad.dev_ptr) + .arg(&flat_block_sums.dev_ptr) + .arg(&total_i) + .launch(LaunchConfig { + grid_dim: (flat_norm_blocks as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .expect("grad_norm launch"); + stream.synchronize().expect("sync grad_norm"); + stream.launch_builder(&k.grad_norm_final) + .arg(&flat_block_sums.dev_ptr) + .arg(&flat_grad_norm.dev_ptr) + .arg(&num_b_i) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .expect("grad_norm_final launch"); + } + stream.synchronize().expect("sync grad_norm_final"); + } + + // ── 8. Adam update for each trunk tensor ────────────────────── + let trunk_tensors: [(&MappedF32Buffer, &MappedF32Buffer, &MappedF32Buffer, &MappedF32Buffer, usize); 6] = [ + (&w1, &w1_g, &w1_m, &w1_v, ENC*H1), + (&b1, &b1_g, &b1_m, &b1_v, H1), + (&w2, &w2_g, &w2_m, &w2_v, H1*H2), + (&b2, &b2_g, &b2_m, &b2_v, H2), + (&w3, &w3_g, &w3_m, &w3_v, H2*SH2), + (&b3, &b3_g, &b3_m, &b3_v, SH2), + ]; + for (p_buf, g_buf, m_buf, v_buf, n) in &trunk_tensors { + let n_i = *n as i32; + let blocks = (n + 255) / 256; + let beta1: f32 = 0.9; + let beta2: f32 = 0.999; + let eps: f32 = 1e-8; + let wd: f32 = 0.0; + let wc_max: f32 = 0.0; // disabled + let l1_end: i32 = 0; + let l1_lam: f32 = 0.0; + let diag: i32 = 0; + let engage_off: i32 = -1; // disabled + unsafe { + stream.launch_builder(&k.adam) + .arg(&p_buf.dev_ptr) + .arg(&g_buf.dev_ptr) + .arg(&m_buf.dev_ptr) + .arg(&v_buf.dev_ptr) + .arg(&flat_grad_norm.dev_ptr) + .arg(&lr_buf.dev_ptr) + .arg(&beta1) + .arg(&beta2) + .arg(&eps) + .arg(&wd) + .arg(&clip_buf.dev_ptr) + .arg(&adam_t_buf.dev_ptr) + .arg(&n_i) + .arg(&wd_mask.dev_ptr) + .arg(&l1_end) + .arg(&l1_lam) + .arg(&wc_max) + .arg(&nan_flags.dev_ptr) + .arg(&diag) + .arg(&engage_buf.dev_ptr) + .arg(&engage_off) + .launch(LaunchConfig { + grid_dim: (blocks as u32, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .expect("adam trunk launch"); + } + } + stream.synchronize().expect("sync adam"); + + // ── 9. Read loss + compute dir_acc for this step ────────────── + let loss_h = unsafe { std::slice::from_raw_parts(loss_out.host_ptr, 1) }; + let sm_h = unsafe { std::slice::from_raw_parts(softmax_out.host_ptr, B * K) }; + let mut correct = 0usize; + for b_idx in 0..B { + // argmax of softmax row: UP (class 1) predicted? + let p0 = sm_h[b_idx * K]; + let p1 = sm_h[b_idx * K + 1]; + if p1 >= p0 { correct += 1; } + } + let acc = correct as f32 / B as f32; + final_loss = loss_h[0]; + final_acc = acc; + + if step % 20 == 0 || step == STEPS - 1 { + eprintln!("step {:3}: CE loss = {:.4}, dir_acc = {:.3}", step + 1, final_loss, acc); + } + } + + assert!( + final_loss < LOSS_THRESHOLD, + "C.9 smoke: aux trunk CE loss = {final_loss:.4} after {STEPS} steps (expected < {LOSS_THRESHOLD}). \ + If loss is near ln(2)≈0.693, gradient chain is broken — diagnose before L40S dispatch.", + ); + assert!( + final_acc >= DIR_ACC_THRESHOLD, + "C.9 smoke: aux trunk dir_acc = {final_acc:.3} after {STEPS} steps (expected ≥ {DIR_ACC_THRESHOLD}).", + ); + eprintln!("C.9 smoke PASS: CE loss = {final_loss:.4} < {LOSS_THRESHOLD}, dir_acc = {final_acc:.3} ≥ {DIR_ACC_THRESHOLD}"); + } } diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index be5d08d3c..e5a85501e 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -2,6 +2,66 @@ **Status:** Populated during Plan 1 Task 6 (A.5 orphan audit). Updated on every commit per Invariant 7. +## 2026-05-08 — SP14 Layer C Phases C.8 + C.9: ISV-driven Adam hyperparams (already complete) + synthetic smoke test + +### C.8 — ISV-driven aux trunk Adam β1/β2/ε/LR/grad-clip + +**Status: already complete in C.5a (commit c90de9859).** No new code required. + +During review for C.8, all five ISV reads for the aux trunk Adam update were found already present in `launch_aux_trunk_adam_update` (implemented in C.5a when the launcher was first wired): + +```rust +// gpu_dqn_trainer.rs — launch_aux_trunk_adam_update +let beta1 = self.read_isv_signal_at(AUX_TRUNK_BETA1_INDEX).max(0.5_f32).min(0.9999_f32); +let beta2 = self.read_isv_signal_at(AUX_TRUNK_BETA2_INDEX).max(0.9_f32).min(0.99999_f32); +let epsilon = self.read_isv_signal_at(AUX_TRUNK_EPS_INDEX).max(1e-12_f32); +// LR + grad-clip written to mapped-pinned scalars per-step: +let aux_lr = self.read_isv_signal_at(AUX_TRUNK_LR_INDEX); +let aux_clip = self.read_isv_signal_at(AUX_TRUNK_GRAD_CLIP_INDEX); +``` + +All 5 fold-boundary StateResetRegistry default-write arms (`sp14_c_aux_trunk_lr`, `sp14_c_aux_trunk_beta1`, `sp14_c_aux_trunk_beta2`, `sp14_c_aux_trunk_eps`, `sp14_c_aux_trunk_grad_clip`) were also already present in `training_loop.rs` (committed with C.5a). ISV defaults: LR=1e-4, β1=0.9, β2=0.999, ε=1e-8, clip=1.0. + +The natural implementation sequence in C.5a bundled the hyperparameter wiring with the Adam launch infrastructure. This is the correct outcome per `feedback_wire_everything_up` — no orphan launchers without hyperparameter reads. + +### C.9 — Synthetic-data smoke: aux trunk gradient chain verification + +Adds `aux_trunk_learns_synthetic_uptrend` to `crates/ml/tests/aux_trunk_oracle_tests.rs`. + +**Topology** (small for RTX 3050 speed): B=16, ENC=32, H1=32, H2=16, SH2=32, H_HEAD=32, K=2, STEPS=100. + +**Signal**: All samples have constant uptrend encoder output `x_in[b, j] = 0.5 / (j+1)`. All labels = 1 (UP). A converging aux trunk must learn to map this pattern to P(UP) → 1. + +**Update scheme**: +- Trunk params (w1/b1/w2/b2/w3/b3): GPU Adam via `dqn_adam_update_kernel`, LR=3e-3. +- Head params (wh1/bh1/wh2/bh2): host-side SGD reading per-sample partial grads through mapped-pinned memory. Valid test orchestration — not a production compute path. + +**Backward kernel invocations** (correct signatures confirmed against `.cu` source): +``` +// Phase A: dh_pre — shmem = H2 floats (sh_dh2_pre cache) +trunk_bwd_dh_pre(dh_s2_aux_out, w3, w2, h_aux1, h_aux2, + dh_pre2[B,H2], dh_pre1[B,H1], + B, H1, H2, SH2) // SH2 = AUX_HIDDEN_DIM + +// Phase B: dW_reduce — 3 launches, shmem=256 floats each +dW3: (h_aux2[B,H2], dh_s2_aux_out[B,SH2], w3_g[H2,SH2], B, H2, SH2) +dW2: (h_aux1[B,H1], dh_pre2[B,H2], w2_g[H1,H2], B, H1, H2) +dW1: (x_in [B,ENC], dh_pre1[B,H1], w1_g[ENC,H1], B, ENC, H1) + +// Phase C: db_reduce — 3 launches, shmem=256 floats each +db3: (dh_s2_aux_out[B,SH2], b3_g[SH2], B, SH2) +db2: (dh_pre2[B,H2], b2_g[H2], B, H2) +db1: (dh_pre1[B,H1], b1_g[H1], B, H1) +``` + +**Pass criteria**: CE loss < 0.1 AND dir_acc ≥ 0.95 after 100 steps. Near-random baseline (ln(2)≈0.693) indicates a broken gradient chain — L40S dispatch is blocked until the root cause is diagnosed (see `pearl_separate_aux_trunk_when_shared_starves.md`). + +**Run command**: `SQLX_OFFLINE=true CUDA_COMPUTE_CAP=86 cargo test -p ml --test aux_trunk_oracle_tests --release -- --ignored aux_trunk_learns_synthetic_uptrend --nocapture` + +**Test result** (RTX 3050 Ti, sm_86): pending GPU execution. Compile clean. + +--- + ## 2026-05-08 — SP14 Layer C Phase C.4: aux trunk backward kernel + gradient check + stop-grad invariant test Phase C.4 of `docs/superpowers/plans/2026-05-07-sp14-layer-c-separate-aux-trunk.md`. Additive commit — backward kernel lands here; wire-up into the collector backward chain + Adam updates land atomically in C.5 per `feedback_no_partial_refactor`.