From ff7073480201e8a520880019ebc47ea2cd052365 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Fri, 22 May 2026 09:11:02 +0200 Subject: [PATCH] feat(aux-heads): linear regression heads + Huber loss for aux supervision (B4) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds the aux heads that map AuxTrunk's output (64-dim) to per-(direction, horizon) predicted outcomes, plus the Huber loss kernel that supervises them against the D-style labels generated in B1. Files (1212 LOC): - cuda/aux_heads.cu (197 LOC): fused fwd+bwd for linear projection h_aux[B, 64] → y_long_hat[B, N_HORIZONS] + y_short_hat[B, N_HORIZONS]. Single launch for both directions, per-batch grad scratch + caller-side reduce_axis0, cooperative h_aux staging in shmem. - cuda/aux_loss.cu (143 LOC): Huber loss fwd+bwd with NaN-masking. Per direction call; returns Σ Huber + valid_count separately so caller picks reduction policy. - src/aux_heads.rs (412 LOC): AuxHeads + AuxHuberLoss wrappers, weight structs, Xavier init under scoped_init_seed. - tests/aux_heads.rs (460 LOC): 4 #[ignore]'d GPU oracle tests (fwd_matches_naive, bwd_finite_diff_matches_bias, huber_loss_matches_ naive, huber_loss_nan_mask_does_not_propagate). 4/4 PASS on RTX 3050. - build.rs: KERNELS += aux_heads, aux_loss; cache-bust bumped to v17. - src/lib.rs: pub mod aux_heads. Design decisions (full notes in subagent report): - Linear heads (no GRN) — D-labels already encode asymmetric loss-aversion - Per-direction Huber launches (cleaner per-direction telemetry) - NaN-masking in loss kernel (single NaN label can't poison batch grad) - Unreduced sum + valid_count separately (caller picks mean policy) Not yet wired into trainer (B5) or used in policy (B7). Co-Authored-By: Claude Opus 4.7 (1M context) --- crates/ml-alpha/build.rs | 4 +- crates/ml-alpha/cuda/aux_heads.cu | 197 ++++++++++++ crates/ml-alpha/cuda/aux_loss.cu | 143 +++++++++ crates/ml-alpha/src/aux_heads.rs | 412 ++++++++++++++++++++++++++ crates/ml-alpha/src/lib.rs | 1 + crates/ml-alpha/tests/aux_heads.rs | 460 +++++++++++++++++++++++++++++ 6 files changed, 1216 insertions(+), 1 deletion(-) create mode 100644 crates/ml-alpha/cuda/aux_heads.cu create mode 100644 crates/ml-alpha/cuda/aux_loss.cu create mode 100644 crates/ml-alpha/src/aux_heads.rs create mode 100644 crates/ml-alpha/tests/aux_heads.rs diff --git a/crates/ml-alpha/build.rs b/crates/ml-alpha/build.rs index 7d38b55c2..89102a5c7 100644 --- a/crates/ml-alpha/build.rs +++ b/crates/ml-alpha/build.rs @@ -28,9 +28,11 @@ const KERNELS: &[&str] = &[ "cfc_step_per_branch", // Per-horizon CfC Phase 2: fused per-(batch, branch) fwd + bwd over [25,25,25,25,28] buckets "heads_block_diagonal_fwd", // Per-horizon CfC Phase 2: heads w_skip projection with compact ragged storage (640→128 floats) "aux_trunk", // SDD-3 Layer B3: smaller single-bucket CfC trunk (AUX_HIDDEN=64) for outcome-supervision (D-labels) + "aux_heads", // SDD-3 Layer B4: per-direction linear regression heads on AuxTrunk output (long + short, N_AUX_HORIZONS each) + "aux_loss", // SDD-3 Layer B4: Huber loss + grad for aux trade-outcome regression targets (NaN-masked) ]; -// Cache bust v16 (2026-05-21): ALPHA fix — channels_in_bucket_kernel + zero_off_bucket_kernel replace tau_reorder_kernel. +// Cache bust v17 (2026-05-22): SDD-3 Layer B4 — aux_heads + aux_loss cubins for trade-outcome regression. fn main() { println!("cargo:rerun-if-changed=build.rs"); diff --git a/crates/ml-alpha/cuda/aux_heads.cu b/crates/ml-alpha/cuda/aux_heads.cu new file mode 100644 index 000000000..5a2844089 --- /dev/null +++ b/crates/ml-alpha/cuda/aux_heads.cu @@ -0,0 +1,197 @@ +// aux_heads.cu — linear regression heads on the AuxTrunk output. +// +// Per `pearl_separate_aux_trunk_when_shared_starves` (SDD-3 Layer B3): +// the aux trunk runs alongside the main per-branch CfC trunk and feeds +// its 64-dim hidden state to a pair of LINEAR projection heads: +// +// y_long_hat [B, N_AUX_HORIZONS] = h_aux [B, AUX_HIDDEN] @ W_long^T + b_long +// y_short_hat[B, N_AUX_HORIZONS] = h_aux [B, AUX_HIDDEN] @ W_short^T + b_short +// +// where (W_long, b_long, W_short, b_short) are independent per-direction +// parameters. No activation — these are regression outputs supervised by +// the D-style asymmetric-loss-aversion labels generated in +// `multi_horizon_labels::generate_outcome_labels_d` (SDD-3 Layer B1) and +// scored by the Huber kernel in `aux_loss.cu` (this same task). +// +// Design constraints (per CLAUDE.md + repo memory): +// * `feedback_no_atomicadd.md` — no atomicAdd. Each thread is +// the sole writer of the slot(s) it touches; cross-batch reduction +// of grad_w / grad_b is deferred to the caller (via the existing +// `reduce_axis0` infrastructure that `cfc/aux_trunk.rs` already +// uses for its own per-batch grad scratch). +// * `feedback_no_nvrtc.md` — pre-compiled cubin via build.rs. +// * `feedback_nvidia_grade_perf_for_kernels.md` — warp-uniform launch, +// no host branches, coalesced loads via cooperative h_aux staging +// (one shared-memory pass per block — `pearl_cooperative_staging_eliminates_redundant_reads`). +// * `pearl_no_host_branches_in_captured_graph` — kernel takes no flags +// that affect control flow. +// +// Block layout: one block per batch sample, AUX_HIDDEN=64 threads per +// block. Forward issues a single matmul per direction with thread role +// = output channel k. Backward uses thread role = hidden unit c so that +// grad_h_aux (which sums over k) is the sole-writer responsibility of +// thread c; the parameter-grad writes (grad_w[k, c], grad_b[k]) then +// loop k inside each thread c. + +#define AUX_HIDDEN 64 +#define N_AUX_HORIZONS 3 +#define N_DIRS 2 // long, short — laid out as two contiguous blocks + + +// ───────────────────────────────────────────────────────────────────── +// aux_heads_fwd: per-direction linear projection, batched. +// +// Launch: +// grid = (B, 1, 1) +// block = (AUX_HIDDEN = 64, 1, 1) +// shared_mem_bytes = AUX_HIDDEN * sizeof(float) (h_aux row stage) +// +// Outputs are written in one kernel launch: +// y_long_hat [batch, k] = b_long[k] + Σ_c W_long [k, c] * h_aux[batch, c] +// y_short_hat[batch, k] = b_short[k] + Σ_c W_short[k, c] * h_aux[batch, c] +// +// Layout: W_long, W_short are row-major [N_AUX_HORIZONS, AUX_HIDDEN]. +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void aux_heads_fwd( + const float* __restrict__ w_long, // [N_AUX_HORIZONS × AUX_HIDDEN] + const float* __restrict__ b_long, // [N_AUX_HORIZONS] + const float* __restrict__ w_short, // [N_AUX_HORIZONS × AUX_HIDDEN] + const float* __restrict__ b_short, // [N_AUX_HORIZONS] + const float* __restrict__ h_aux, // [B × AUX_HIDDEN] + int B, + float* __restrict__ y_long_hat, // [B × N_AUX_HORIZONS] + float* __restrict__ y_short_hat // [B × N_AUX_HORIZONS] +) { + extern __shared__ float s_h[]; // [AUX_HIDDEN] + + const int batch = blockIdx.x; + const int c = threadIdx.x; // 0..AUX_HIDDEN-1 + + if (batch >= B) return; + + // Cooperative staging of h_aux[batch, *] into shared mem so every + // thread's later dot-product reads from smem rather than re-fetching + // the same row N_DIRS × N_AUX_HORIZONS times from DRAM. + s_h[c] = h_aux[batch * AUX_HIDDEN + c]; + __syncthreads(); + + // Compute the 6 outputs (2 dirs × 3 horizons) via 6 threads (k=0..5). + // Remaining 58 threads sit idle this phase — block is intentionally + // sized to AUX_HIDDEN so the *backward* kernel (which is the heavier + // pass) gets full occupancy on h_aux/grad_h_aux. The forward's idle + // lanes are a constant per-block overhead, not a per-batch cost. + if (c < N_DIRS * N_AUX_HORIZONS) { + const int dir = c / N_AUX_HORIZONS; // 0 = long, 1 = short + const int k = c % N_AUX_HORIZONS; + + const float* w = (dir == 0) ? w_long : w_short; + const float* bv = (dir == 0) ? b_long : b_short; + + float acc = bv[k]; + #pragma unroll + for (int ci = 0; ci < AUX_HIDDEN; ++ci) { + acc += w[k * AUX_HIDDEN + ci] * s_h[ci]; + } + float* out = (dir == 0) ? y_long_hat : y_short_hat; + out[batch * N_AUX_HORIZONS + k] = acc; + } +} + + +// ───────────────────────────────────────────────────────────────────── +// aux_heads_bwd: backward through the two per-direction linear heads. +// +// Launch: +// grid = (B, 1, 1) +// block = (AUX_HIDDEN = 64, 1, 1) +// shared_mem_bytes = AUX_HIDDEN * sizeof(float) (h_aux row stage) +// +// Grad shapes (per-batch scratch; OVERWRITE semantics — caller reduces +// axis 0 via existing `reduce_axis0` infrastructure): +// grad_w_long : [B × N_AUX_HORIZONS × AUX_HIDDEN] +// grad_b_long : [B × N_AUX_HORIZONS] +// grad_w_short : [B × N_AUX_HORIZONS × AUX_HIDDEN] +// grad_b_short : [B × N_AUX_HORIZONS] +// grad_h_aux : [B × AUX_HIDDEN] — to be added by the +// trainer to the aux trunk's grad_h_new (cross-batch +// reduction NOT needed — h_aux is per-batch). +// +// Math (per batch, per direction d ∈ {long, short}): +// grad_w_d[k, c] = grad_y_d_hat[k] * h_aux[c] +// grad_b_d[k] = grad_y_d_hat[k] +// grad_h_aux[c] += Σ_k grad_y_long_hat [k] * W_long [k, c] +// + Σ_k grad_y_short_hat[k] * W_short[k, c] +// +// Thread role: c = threadIdx.x is the hidden-unit dim. Each thread c +// is the sole writer of: +// * grad_w_{long,short}[batch, k, c] for all k (loops k internally) +// * grad_h_aux[batch, c] +// And the first N_AUX_HORIZONS threads additionally write grad_b for +// each direction. +// +// No atomicAdd; no cross-thread reduction needed for grad_w / grad_h_aux +// because each thread owns disjoint output slots (per +// `feedback_no_atomicadd.md`). +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void aux_heads_bwd( + const float* __restrict__ w_long, // [N_AUX_HORIZONS × AUX_HIDDEN] + const float* __restrict__ w_short, // [N_AUX_HORIZONS × AUX_HIDDEN] + const float* __restrict__ h_aux, // [B × AUX_HIDDEN] + const float* __restrict__ grad_y_long_hat, // [B × N_AUX_HORIZONS] + const float* __restrict__ grad_y_short_hat, // [B × N_AUX_HORIZONS] + int B, + float* __restrict__ grad_w_long, // [B × N_AUX_HORIZONS × AUX_HIDDEN] + float* __restrict__ grad_b_long, // [B × N_AUX_HORIZONS] + float* __restrict__ grad_w_short, // [B × N_AUX_HORIZONS × AUX_HIDDEN] + float* __restrict__ grad_b_short, // [B × N_AUX_HORIZONS] + float* __restrict__ grad_h_aux // [B × AUX_HIDDEN] +) { + extern __shared__ float s_h[]; // [AUX_HIDDEN] + + const int batch = blockIdx.x; + const int c = threadIdx.x; // 0..AUX_HIDDEN-1 + + if (batch >= B) return; + + // Stage h_aux[batch, *] and the per-direction grad_y vectors. The + // grad_y vectors are small (N_AUX_HORIZONS = 3 each) so we keep them + // in shared via the first N_AUX_HORIZONS threads. + s_h[c] = h_aux[batch * AUX_HIDDEN + c]; + + __shared__ float s_gyl[N_AUX_HORIZONS]; + __shared__ float s_gys[N_AUX_HORIZONS]; + if (c < N_AUX_HORIZONS) { + s_gyl[c] = grad_y_long_hat [batch * N_AUX_HORIZONS + c]; + s_gys[c] = grad_y_short_hat[batch * N_AUX_HORIZONS + c]; + } + __syncthreads(); + + // grad_b_{long,short}[batch, k] = grad_y_{long,short}_hat[batch, k] + // Sole writer: thread c == k (for k in [0, N_AUX_HORIZONS)). + if (c < N_AUX_HORIZONS) { + grad_b_long [batch * N_AUX_HORIZONS + c] = s_gyl[c]; + grad_b_short[batch * N_AUX_HORIZONS + c] = s_gys[c]; + } + + // For thread c (the hidden-unit dim) — write grad_w[batch, k, c] for + // each k, and accumulate grad_h_aux[batch, c] across both directions. + const float h_c = s_h[c]; + float gh_acc = 0.0f; + #pragma unroll + for (int k = 0; k < N_AUX_HORIZONS; ++k) { + const float gyl_k = s_gyl[k]; + const float gys_k = s_gys[k]; + + // grad_w_d[batch, k, c] = grad_y_d_hat[batch, k] * h_aux[batch, c] + const long long off = + (long long)batch * N_AUX_HORIZONS * AUX_HIDDEN + + (long long)k * AUX_HIDDEN + c; + grad_w_long [off] = gyl_k * h_c; + grad_w_short[off] = gys_k * h_c; + + // grad_h_aux[batch, c] += grad_y_d_hat[k] * W_d[k, c] for both dirs. + gh_acc += gyl_k * w_long [k * AUX_HIDDEN + c]; + gh_acc += gys_k * w_short[k * AUX_HIDDEN + c]; + } + grad_h_aux[batch * AUX_HIDDEN + c] = gh_acc; +} diff --git a/crates/ml-alpha/cuda/aux_loss.cu b/crates/ml-alpha/cuda/aux_loss.cu new file mode 100644 index 000000000..c51b1c294 --- /dev/null +++ b/crates/ml-alpha/cuda/aux_loss.cu @@ -0,0 +1,143 @@ +// aux_loss.cu — Huber loss kernel for the aux trade-outcome heads. +// +// Per E2's design memo (referenced from `pearl_separate_aux_trunk_when_shared_starves`): +// the aux head predicts continuous-valued trade-outcome targets +// (D-style asymmetric loss-aversion labels from +// `multi_horizon_labels::generate_outcome_labels_d`). Huber is the +// canonical regression loss when the target distribution has heavy +// tails — squared near zero (smooth gradient), absolute far from zero +// (bounded gradient magnitude). δ = 1.0 by default per E2. +// +// Formulation (per element): +// r = y_hat - y_true +// L_huber = (|r| <= δ) ? 0.5 * r² : δ * (|r| - 0.5 * δ) +// dL/dy_hat = (|r| <= δ) ? r : δ * sign(r) +// +// NaN handling: y_true may be NaN at positions where the D-label +// generator dropped the target (edge/invalid samples). We MUST mask +// these out — the loss contribution is 0 AND the grad is 0 so the +// downstream `aux_heads_bwd` propagates zero through the head and the +// `aux_trunk_bwd` sees no spurious gradient. +// +// Reduction discipline (per `feedback_no_atomicadd.md`): +// * Per-thread strided accumulation into registers. +// * Warp-shuffle reduce (`__shfl_xor_sync`) to one float per warp. +// * One `__syncthreads` to gather warp partials in shared memory. +// * Final warp-0 reduction to total loss / total valid count. +// * No atomicAdd anywhere. +// +// Per-direction split: this kernel processes ONE direction's tensor at +// a time. The trainer calls it twice (once for long, once for short) +// and sums the two scalars into the total aux loss. Keeping the kernel +// per-direction keeps the block layout simple and lets the per-direction +// telemetry (mean Huber, valid count) emerge for free. +// +// Block layout: grid = (1), block = (BLOCK = 256). Single block handles +// the full [B × N_AUX_HORIZONS] grid (max ~32 × 3 = 96, easily within +// stride reach). Mirrors the pattern in `bce_loss_multi_horizon.cu`. + +#define AUX_LOSS_BLOCK 256 +#define AUX_LOSS_WARPS (AUX_LOSS_BLOCK / 32) // == 8 + + +__device__ __forceinline__ float warp_reduce_sum_f32(float v) { + #pragma unroll + for (int s = 16; s > 0; s >>= 1) { + v += __shfl_xor_sync(0xffffffff, v, s); + } + return v; +} + +__device__ __forceinline__ int warp_reduce_sum_i32(int v) { + #pragma unroll + for (int s = 16; s > 0; s >>= 1) { + v += __shfl_xor_sync(0xffffffff, v, s); + } + return v; +} + + +// ───────────────────────────────────────────────────────────────────── +// aux_huber_loss_fwd_bwd: +// Fused Huber-loss forward + per-element gradient writeback for ONE +// direction (long or short). Inputs/outputs are [B × N_AUX_HORIZONS] +// layout. NaN-masked positions contribute zero loss and write zero +// gradient (no contamination of downstream consumers). +// +// Outputs: +// loss_out [1] — Σ_valid L_huber (UNREDUCED; +// caller divides by valid_count +// if a mean is desired). +// grad_y_hat [B × N_AUX_HORIZONS] — dL/dy_hat per element. +// valid_count_out [1] — #{(b, k) : isfinite(y_true)}. +// +// `loss_out` is intentionally a SUM, not a mean, so the caller controls +// the reduction policy (mean over valid, mean over total elements, etc.) +// without making the kernel re-runnable for different conventions. +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void aux_huber_loss_fwd_bwd( + const float* __restrict__ y_hat, // [B × N_AUX_HORIZONS] + const float* __restrict__ y_true, // [B × N_AUX_HORIZONS] (NaN-maskable) + float delta, // Huber transition point (default 1.0) + int n_total, // = B × N_AUX_HORIZONS + float* __restrict__ loss_out, // [1] — Σ L (unreduced) + float* __restrict__ grad_y_hat, // [B × N_AUX_HORIZONS] + int* __restrict__ valid_count_out // [1] +) { + const int tid = threadIdx.x; + const int lane = tid & 31; + const int warp = tid >> 5; + + // Per-thread strided accumulation into registers. + float local_loss = 0.0f; + int local_valid = 0; + + for (int i = tid; i < n_total; i += AUX_LOSS_BLOCK) { + const float y_t = y_true[i]; + if (!isfinite(y_t)) { + grad_y_hat[i] = 0.0f; + continue; + } + const float y_p = y_hat[i]; + const float r = y_p - y_t; + const float ar = fabsf(r); + + float L, g; + if (ar <= delta) { + L = 0.5f * r * r; + g = r; + } else { + L = delta * (ar - 0.5f * delta); + // δ * sign(r). sign uses copysign to keep g in the same sign + // bucket as r (and to give 0 when r == 0, though the |r| > δ + // branch already excludes r == 0). + g = copysignf(delta, r); + } + local_loss += L; + local_valid += 1; + grad_y_hat[i] = g; + } + + // Warp-shuffle reduce. + float warp_loss = warp_reduce_sum_f32(local_loss); + int warp_valid = warp_reduce_sum_i32(local_valid); + + __shared__ float s_loss[AUX_LOSS_WARPS]; + __shared__ int s_valid[AUX_LOSS_WARPS]; + if (lane == 0) { + s_loss[warp] = warp_loss; + s_valid[warp] = warp_valid; + } + __syncthreads(); + + if (warp == 0) { + float v_l = (lane < AUX_LOSS_WARPS) ? s_loss[lane] : 0.0f; + int v_v = (lane < AUX_LOSS_WARPS) ? s_valid[lane] : 0; + v_l = warp_reduce_sum_f32(v_l); + v_v = warp_reduce_sum_i32(v_v); + if (lane == 0) { + loss_out[0] = v_l; + valid_count_out[0] = v_v; + } + } +} diff --git a/crates/ml-alpha/src/aux_heads.rs b/crates/ml-alpha/src/aux_heads.rs new file mode 100644 index 000000000..797549e4f --- /dev/null +++ b/crates/ml-alpha/src/aux_heads.rs @@ -0,0 +1,412 @@ +//! SDD-3 Layer B4 — Aux heads (linear regression) + Huber loss. +//! +//! Sits on top of [`crate::cfc::AuxTrunk`] and predicts continuous-valued +//! trade-outcome targets emitted by [`crate::multi_horizon_labels::generate_outcome_labels_d`] +//! (the D-style asymmetric loss-aversion labels from SDD-3 Layer B1). +//! +//! ## Architecture +//! +//! Two INDEPENDENT per-direction linear heads: +//! +//! ```text +//! y_long_hat [B, N] = h_aux [B, H] @ W_long^T + b_long +//! y_short_hat[B, N] = h_aux [B, H] @ W_short^T + b_short +//! ``` +//! +//! where `H = AUX_HIDDEN = 64` and `N = N_AUX_HORIZONS = 3`. No activation +//! — these are regression outputs supervised by Huber loss with NaN masking. +//! +//! ## Why linear (not GRN) +//! +//! Per E2's design memo: the asymmetric loss-aversion labels are +//! already non-linear functions of price + drawdown over the K-horizon; +//! the head's role is to predict a calibrated scalar, not to re-derive +//! the asymmetry. A linear projection keeps the head's gradient signal +//! interpretable and avoids the extra-parameter risk of the main BCE +//! head's GRN structure. +//! +//! ## Initialisation +//! +//! Per [`pearl_scoped_init_seed_for_reproducibility`]: +//! `AuxHeads::new` installs a [`scoped_init_seed`] guard so the Xavier +//! draws are deterministic given the seed. +//! +//! * `w_long`, `w_short`: Xavier uniform `[-sqrt(6/(H+N)), +]` +//! * `b_long`, `b_short`: zeros +//! +//! ## Constraints honoured +//! +//! * `feedback_no_atomicadd.md` — backward writes per-batch grad +//! scratch; cross-batch reduction is the caller's responsibility +//! via the existing `reduce_axis0` infrastructure. +//! * `feedback_no_htod_htoh_only_mapped_pinned.md` — uploads go +//! through mapped-pinned staging. +//! * `feedback_no_nvrtc.md` — pre-compiled cubin via `build.rs`. + +use std::sync::Arc; + +use anyhow::{Context, Result}; +use cudarc::driver::{ + CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut, LaunchConfig, PushKernelArg, +}; +use ml_core::cuda_autograd::init::scoped_init_seed; +use ml_core::device::MlDevice; +use rand::{Rng, SeedableRng}; +use rand_chacha::ChaCha8Rng; + +use crate::cfc::AUX_HIDDEN; +use crate::pinned_mem::MappedF32Buffer; + +const AUX_HEADS_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/aux_heads.cubin")); +const AUX_LOSS_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/aux_loss.cubin")); + +/// Number of aux-horizon regression outputs per direction. Matches the +/// main BCE head's [`crate::heads::N_HORIZONS`] (= 3) by construction: +/// the D-label generator emits one outcome per (direction, horizon) +/// across the same horizon grid as the BCE supervision so the two +/// heads can share the loader pipeline. +pub const N_AUX_HORIZONS: usize = crate::heads::N_HORIZONS; + +/// Default Huber transition point per E2's recommendation. The +/// `aux_huber_loss_gpu` wrapper exposes it as a runtime argument so +/// downstream callers can sweep without rebuilding the cubin. +pub const DEFAULT_HUBER_DELTA: f32 = 1.0; + +/// Host-side weight storage for [`AuxHeads`] (checkpoint round-trip). +#[derive(Clone, Debug)] +pub struct AuxHeadsWeights { + pub w_long: Vec, // [N_AUX_HORIZONS × AUX_HIDDEN] + pub b_long: Vec, // [N_AUX_HORIZONS] + pub w_short: Vec, // [N_AUX_HORIZONS × AUX_HIDDEN] + pub b_short: Vec, // [N_AUX_HORIZONS] +} + +/// Construction config for [`AuxHeads`]. +#[derive(Clone, Debug)] +pub struct AuxHeadsConfig { + /// Random seed for Xavier init. Reproducibility-critical per + /// `pearl_scoped_init_seed_for_reproducibility`. + pub seed: u64, +} + +/// Aux-direction regression heads. Owns device weights + cached cubin +/// function handles for forward / backward. +pub struct AuxHeads { + cfg: AuxHeadsConfig, + stream: Arc, + + _module: Arc, + pub fwd_fn: CudaFunction, + pub bwd_fn: CudaFunction, + + pub w_long_d: CudaSlice, + pub b_long_d: CudaSlice, + pub w_short_d: CudaSlice, + pub b_short_d: CudaSlice, +} + +/// Wrapper around the Huber loss cubin. Cached once per trainer so the +/// per-step launch path doesn't reload the module. +pub struct AuxHuberLoss { + stream: Arc, + _module: Arc, + pub fn_: CudaFunction, +} + +impl AuxHeads { + pub fn new(dev: &MlDevice, cfg: AuxHeadsConfig) -> Result { + let stream: Arc = dev.cuda_stream().context("aux_heads stream")?.clone(); + let ctx = dev.cuda_context().context("aux_heads ctx")?; + let module = ctx + .load_cubin(AUX_HEADS_CUBIN.to_vec()) + .context("load aux_heads cubin")?; + let fwd_fn = module + .load_function("aux_heads_fwd") + .context("load aux_heads_fwd")?; + let bwd_fn = module + .load_function("aux_heads_bwd") + .context("load aux_heads_bwd")?; + + let _seed_guard = scoped_init_seed(cfg.seed); + let mut r = ChaCha8Rng::seed_from_u64(cfg.seed); + let scale = (6.0_f32 / (AUX_HIDDEN + N_AUX_HORIZONS) as f32).sqrt(); + let w_long: Vec = (0..N_AUX_HORIZONS * AUX_HIDDEN) + .map(|_| r.gen_range(-scale..scale)) + .collect(); + let w_short: Vec = (0..N_AUX_HORIZONS * AUX_HIDDEN) + .map(|_| r.gen_range(-scale..scale)) + .collect(); + let b_long: Vec = vec![0.0; N_AUX_HORIZONS]; + let b_short: Vec = vec![0.0; N_AUX_HORIZONS]; + + let w_long_d = upload(&stream, &w_long)?; + let b_long_d = upload(&stream, &b_long)?; + let w_short_d = upload(&stream, &w_short)?; + let b_short_d = upload(&stream, &b_short)?; + + Ok(Self { + cfg, + stream, + _module: module, + fwd_fn, + bwd_fn, + w_long_d, + b_long_d, + w_short_d, + b_short_d, + }) + } + + pub fn config(&self) -> &AuxHeadsConfig { + &self.cfg + } + + pub fn stream(&self) -> &Arc { + &self.stream + } + + /// Download the head weights into host vectors. Used by checkpoint + /// save and unit tests. + pub fn download_weights(&self) -> Result { + let mut w_long = vec![0.0_f32; self.w_long_d.len()]; + let mut b_long = vec![0.0_f32; self.b_long_d.len()]; + let mut w_short = vec![0.0_f32; self.w_short_d.len()]; + let mut b_short = vec![0.0_f32; self.b_short_d.len()]; + self.stream + .memcpy_dtoh(&self.w_long_d, w_long.as_mut_slice()) + .context("aux_heads dtoh w_long")?; + self.stream + .memcpy_dtoh(&self.b_long_d, b_long.as_mut_slice()) + .context("aux_heads dtoh b_long")?; + self.stream + .memcpy_dtoh(&self.w_short_d, w_short.as_mut_slice()) + .context("aux_heads dtoh w_short")?; + self.stream + .memcpy_dtoh(&self.b_short_d, b_short.as_mut_slice()) + .context("aux_heads dtoh b_short")?; + Ok(AuxHeadsWeights { + w_long, + b_long, + w_short, + b_short, + }) + } +} + +impl AuxHuberLoss { + pub fn new(dev: &MlDevice) -> Result { + let stream: Arc = dev.cuda_stream().context("aux_loss stream")?.clone(); + let ctx = dev.cuda_context().context("aux_loss ctx")?; + let module = ctx + .load_cubin(AUX_LOSS_CUBIN.to_vec()) + .context("load aux_loss cubin")?; + let fn_ = module + .load_function("aux_huber_loss_fwd_bwd") + .context("load aux_huber_loss_fwd_bwd")?; + Ok(Self { + stream, + _module: module, + fn_, + }) + } + + pub fn stream(&self) -> &Arc { + &self.stream + } +} + +// ── Launch wrappers ────────────────────────────────────────────────── + +/// Launch the fused aux-heads forward kernel. +/// +/// Buffer contract: +/// * `h_aux_d` : `[B × AUX_HIDDEN]` — aux trunk hidden state +/// * `y_long_hat_d` : `[B × N_AUX_HORIZONS]` — long predictions (overwrite) +/// * `y_short_hat_d` : `[B × N_AUX_HORIZONS]` — short predictions (overwrite) +/// +/// Shared memory layout: `AUX_HIDDEN * sizeof(float)` covering the +/// cooperative-staged `h_aux` row. +#[allow(clippy::too_many_arguments)] +pub fn aux_heads_fwd_gpu( + stream: &Arc, + func: &CudaFunction, + w_long_d: &CudaSlice, + b_long_d: &CudaSlice, + w_short_d: &CudaSlice, + b_short_d: &CudaSlice, + h_aux_d: &CudaSlice, + b_sz: i32, + y_long_hat_d: &mut CudaSlice, + y_short_hat_d: &mut CudaSlice, +) -> Result<()> { + let b_sz_u = b_sz as usize; + debug_assert_eq!(w_long_d.len(), N_AUX_HORIZONS * AUX_HIDDEN); + debug_assert_eq!(b_long_d.len(), N_AUX_HORIZONS); + debug_assert_eq!(w_short_d.len(), N_AUX_HORIZONS * AUX_HIDDEN); + debug_assert_eq!(b_short_d.len(), N_AUX_HORIZONS); + debug_assert_eq!(h_aux_d.len(), b_sz_u * AUX_HIDDEN); + debug_assert_eq!(y_long_hat_d.len(), b_sz_u * N_AUX_HORIZONS); + debug_assert_eq!(y_short_hat_d.len(), b_sz_u * N_AUX_HORIZONS); + + let cfg = LaunchConfig { + grid_dim: (b_sz as u32, 1, 1), + block_dim: (AUX_HIDDEN as u32, 1, 1), + shared_mem_bytes: (AUX_HIDDEN * std::mem::size_of::()) as u32, + }; + let mut launch = stream.launch_builder(func); + launch + .arg(w_long_d) + .arg(b_long_d) + .arg(w_short_d) + .arg(b_short_d) + .arg(h_aux_d) + .arg(&b_sz) + .arg(y_long_hat_d) + .arg(y_short_hat_d); + unsafe { + launch.launch(cfg).context("aux_heads_fwd launch")?; + } + Ok(()) +} + +/// Launch the fused aux-heads backward kernel. +/// +/// Per-batch grad scratch is OVERWRITTEN by this kernel; the caller is +/// responsible for the cross-batch reduction (e.g. `reduce_axis0`). +/// `grad_h_aux_d` is per-batch and does NOT require cross-batch reduction +/// — it feeds back into the aux-trunk's `grad_h_new` as a per-sample +/// signal. +/// +/// Shared memory layout: `AUX_HIDDEN * sizeof(float)` covering the +/// cooperative-staged `h_aux` row. +#[allow(clippy::too_many_arguments)] +pub fn aux_heads_bwd_gpu( + stream: &Arc, + func: &CudaFunction, + w_long_d: &CudaSlice, + w_short_d: &CudaSlice, + h_aux_d: &CudaSlice, + grad_y_long_hat_d: &CudaSlice, + grad_y_short_hat_d: &CudaSlice, + b_sz: i32, + grad_w_long_d: &mut CudaSlice, + grad_b_long_d: &mut CudaSlice, + grad_w_short_d: &mut CudaSlice, + grad_b_short_d: &mut CudaSlice, + grad_h_aux_d: &mut CudaSlice, +) -> Result<()> { + let b_sz_u = b_sz as usize; + debug_assert_eq!(w_long_d.len(), N_AUX_HORIZONS * AUX_HIDDEN); + debug_assert_eq!(w_short_d.len(), N_AUX_HORIZONS * AUX_HIDDEN); + debug_assert_eq!(h_aux_d.len(), b_sz_u * AUX_HIDDEN); + debug_assert_eq!(grad_y_long_hat_d.len(), b_sz_u * N_AUX_HORIZONS); + debug_assert_eq!(grad_y_short_hat_d.len(), b_sz_u * N_AUX_HORIZONS); + debug_assert_eq!(grad_w_long_d.len(), b_sz_u * N_AUX_HORIZONS * AUX_HIDDEN); + debug_assert_eq!(grad_b_long_d.len(), b_sz_u * N_AUX_HORIZONS); + debug_assert_eq!(grad_w_short_d.len(), b_sz_u * N_AUX_HORIZONS * AUX_HIDDEN); + debug_assert_eq!(grad_b_short_d.len(), b_sz_u * N_AUX_HORIZONS); + debug_assert_eq!(grad_h_aux_d.len(), b_sz_u * AUX_HIDDEN); + + let cfg = LaunchConfig { + grid_dim: (b_sz as u32, 1, 1), + block_dim: (AUX_HIDDEN as u32, 1, 1), + shared_mem_bytes: (AUX_HIDDEN * std::mem::size_of::()) as u32, + }; + let mut launch = stream.launch_builder(func); + launch + .arg(w_long_d) + .arg(w_short_d) + .arg(h_aux_d) + .arg(grad_y_long_hat_d) + .arg(grad_y_short_hat_d) + .arg(&b_sz) + .arg(grad_w_long_d) + .arg(grad_b_long_d) + .arg(grad_w_short_d) + .arg(grad_b_short_d) + .arg(grad_h_aux_d); + unsafe { + launch.launch(cfg).context("aux_heads_bwd launch")?; + } + Ok(()) +} + +/// Launch the fused Huber loss forward + backward kernel for ONE +/// direction. The kernel writes: +/// * `loss_out_d[0]` — Σ_valid Huber loss (unreduced sum; +/// caller divides by valid_count if a mean +/// is desired). +/// * `grad_y_hat_d` — per-element dL/dy_hat (NaN-masked zero). +/// * `valid_count_out_d[0]` — #{(b, k) : isfinite(y_true)}. +/// +/// Call twice — once for long, once for short — and sum the two loss +/// scalars to get the total aux loss. The per-direction split keeps the +/// telemetry (mean Huber per direction, valid count per direction) +/// available for free without extra kernels. +#[allow(clippy::too_many_arguments)] +pub fn aux_huber_loss_gpu( + stream: &Arc, + func: &CudaFunction, + y_hat_d: &CudaSlice, + y_true_d: &CudaSlice, + delta: f32, + n_total: i32, + loss_out_d: &mut CudaSlice, + grad_y_hat_d: &mut CudaSlice, + valid_count_out_d: &mut CudaSlice, +) -> Result<()> { + let n_total_u = n_total as usize; + debug_assert_eq!(y_hat_d.len(), n_total_u); + debug_assert_eq!(y_true_d.len(), n_total_u); + debug_assert_eq!(grad_y_hat_d.len(), n_total_u); + debug_assert_eq!(loss_out_d.len(), 1); + debug_assert_eq!(valid_count_out_d.len(), 1); + + const AUX_LOSS_BLOCK: u32 = 256; + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (AUX_LOSS_BLOCK, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = stream.launch_builder(func); + launch + .arg(y_hat_d) + .arg(y_true_d) + .arg(&delta) + .arg(&n_total) + .arg(loss_out_d) + .arg(grad_y_hat_d) + .arg(valid_count_out_d); + unsafe { + launch.launch(cfg).context("aux_huber_loss launch")?; + } + Ok(()) +} + +// ── pinned-staging upload helper (mirrors cfc/aux_trunk.rs::upload) ── + +fn upload(stream: &Arc, host: &[f32]) -> Result> { + let n = host.len(); + let staging = unsafe { MappedF32Buffer::new(n) } + .map_err(|e| anyhow::anyhow!("aux_heads upload staging: {e}"))?; + staging.write_from_slice(host); + let mut dst = stream + .alloc_zeros::(n) + .context("aux_heads upload alloc")?; + if n > 0 { + let nbytes = n * std::mem::size_of::(); + unsafe { + let (dst_ptr, _g) = dst.device_ptr_mut(stream); + cudarc::driver::result::memcpy_dtod_async( + dst_ptr, + staging.dev_ptr, + nbytes, + stream.cu_stream(), + ) + .context("aux_heads upload DtoD")?; + } + } + Ok(dst) +} diff --git a/crates/ml-alpha/src/lib.rs b/crates/ml-alpha/src/lib.rs index 6e9be402e..03adf6afe 100644 --- a/crates/ml-alpha/src/lib.rs +++ b/crates/ml-alpha/src/lib.rs @@ -25,6 +25,7 @@ // Task 10+: pub mod trainer; // Task 12: pub mod data; // Task 16+: pub mod gate; +pub mod aux_heads; pub mod cfc; pub mod data; pub mod eval; diff --git a/crates/ml-alpha/tests/aux_heads.rs b/crates/ml-alpha/tests/aux_heads.rs new file mode 100644 index 000000000..d98fc5563 --- /dev/null +++ b/crates/ml-alpha/tests/aux_heads.rs @@ -0,0 +1,460 @@ +//! GPU oracle tests for `aux_heads.cu` + `aux_loss.cu` (SDD-3 Layer B4). +//! +//! Run with: +//! SQLX_OFFLINE=true cargo test -p ml-alpha --test aux_heads \ +//! -- --ignored --nocapture +//! +//! These tests are `#[ignore]` because they require a CUDA device. They +//! drive the cubins directly through the host wrappers in +//! `crates/ml-alpha/src/aux_heads.rs` and compare against a deterministic +//! CPU naive reference (per `feedback_no_cpu_test_fallbacks.md` the +//! production path stays GPU-only; the CPU implementation here exists +//! ONLY as a unit-test oracle). + +use anyhow::Result; +use ml_alpha::aux_heads::{ + aux_heads_bwd_gpu, aux_heads_fwd_gpu, aux_huber_loss_gpu, AuxHeads, AuxHeadsConfig, + AuxHuberLoss, DEFAULT_HUBER_DELTA, N_AUX_HORIZONS, +}; +use ml_alpha::cfc::AUX_HIDDEN; +use ml_core::device::MlDevice; + +fn upload( + stream: &std::sync::Arc, + host: &[f32], +) -> Result> { + let mut d = stream.alloc_zeros::(host.len())?; + stream.memcpy_htod(host, &mut d)?; + Ok(d) +} + +fn upload_i32( + stream: &std::sync::Arc, + host: &[i32], +) -> Result> { + let mut d = stream.alloc_zeros::(host.len())?; + stream.memcpy_htod(host, &mut d)?; + Ok(d) +} + +fn download( + stream: &std::sync::Arc, + d: &cudarc::driver::CudaSlice, +) -> Result> { + let mut h = vec![0.0_f32; d.len()]; + stream.memcpy_dtoh(d, &mut h)?; + Ok(h) +} + +fn download_i32( + stream: &std::sync::Arc, + d: &cudarc::driver::CudaSlice, +) -> Result> { + let mut h = vec![0_i32; d.len()]; + stream.memcpy_dtoh(d, &mut h)?; + Ok(h) +} + +/// Forward shape correctness + value match against a naive CPU oracle. +/// +/// Validates: +/// * Output shapes ([B × N_AUX_HORIZONS] per direction). +/// * Per-element value within tolerance (5e-4) of the naive reference. +/// * Long and short heads use INDEPENDENT parameter sets — flipping a +/// single long weight does not perturb a short output and vice versa. +#[test] +#[ignore = "requires CUDA"] +fn fwd_matches_naive_reference() -> Result<()> { + let dev = MlDevice::cuda(0)?; + let stream = dev.cuda_stream()?.clone(); + let heads = AuxHeads::new(&dev, AuxHeadsConfig { seed: 0xAE_AB_BEEF })?; + + let b_sz: usize = 4; + let h_aux: Vec = (0..b_sz * AUX_HIDDEN) + .map(|i| ((i as f32) * 0.013).sin() * 0.6) + .collect(); + let h_aux_d = upload(&stream, &h_aux)?; + let mut y_long_d = stream.alloc_zeros::(b_sz * N_AUX_HORIZONS)?; + let mut y_short_d = stream.alloc_zeros::(b_sz * N_AUX_HORIZONS)?; + + aux_heads_fwd_gpu( + &stream, + &heads.fwd_fn, + &heads.w_long_d, + &heads.b_long_d, + &heads.w_short_d, + &heads.b_short_d, + &h_aux_d, + b_sz as i32, + &mut y_long_d, + &mut y_short_d, + )?; + stream.synchronize()?; + + let y_long = download(&stream, &y_long_d)?; + let y_short = download(&stream, &y_short_d)?; + let w = heads.download_weights()?; + + for batch in 0..b_sz { + for k in 0..N_AUX_HORIZONS { + let mut expected_long = w.b_long[k]; + let mut expected_short = w.b_short[k]; + for c in 0..AUX_HIDDEN { + expected_long += w.w_long[k * AUX_HIDDEN + c] * h_aux[batch * AUX_HIDDEN + c]; + expected_short += w.w_short[k * AUX_HIDDEN + c] * h_aux[batch * AUX_HIDDEN + c]; + } + let got_long = y_long[batch * N_AUX_HORIZONS + k]; + let got_short = y_short[batch * N_AUX_HORIZONS + k]; + assert!( + (got_long - expected_long).abs() < 5e-4, + "y_long[{},{}] = {got_long} vs expected {expected_long} \ + (diff {})", + batch, + k, + (got_long - expected_long).abs(), + ); + assert!( + (got_short - expected_short).abs() < 5e-4, + "y_short[{},{}] = {got_short} vs expected {expected_short} \ + (diff {})", + batch, + k, + (got_short - expected_short).abs(), + ); + } + } + + Ok(()) +} + +/// Backward grad shapes finite, no NaN/Inf, and bias-grad finite-difference +/// matches within tolerance against the implicit loss +/// L = Σ grad_y_*_hat[batch, k] * y_*_hat[batch, k] (linear in y_hat so +/// dL/dW and dL/db are closed-form and stable under FD). +/// +/// Also asserts the `grad_h_aux` per-batch slot is non-zero and finite, +/// which is the load-bearing signal that flows back into `aux_trunk_bwd`. +#[test] +#[ignore = "requires CUDA"] +fn bwd_finite_diff_matches_bias_sample() -> Result<()> { + let dev = MlDevice::cuda(0)?; + let stream = dev.cuda_stream()?.clone(); + let heads = AuxHeads::new(&dev, AuxHeadsConfig { seed: 0xBEEF_FEED })?; + + let b_sz: usize = 3; + let h_aux: Vec = (0..b_sz * AUX_HIDDEN) + .map(|i| ((i as f32) * 0.011).cos() * 0.4) + .collect(); + // Structured per-element grad (non-zero everywhere) so every grad + // path exercises a non-trivial accumulation. + let grad_y_long: Vec = (0..b_sz * N_AUX_HORIZONS) + .map(|i| ((i as f32) * 0.17).sin() * 0.5 + 0.1) + .collect(); + let grad_y_short: Vec = (0..b_sz * N_AUX_HORIZONS) + .map(|i| ((i as f32) * 0.23).cos() * 0.4 - 0.05) + .collect(); + + let h_aux_d = upload(&stream, &h_aux)?; + let grad_y_long_d = upload(&stream, &grad_y_long)?; + let grad_y_short_d = upload(&stream, &grad_y_short)?; + + let mut grad_w_long_d = + stream.alloc_zeros::(b_sz * N_AUX_HORIZONS * AUX_HIDDEN)?; + let mut grad_b_long_d = stream.alloc_zeros::(b_sz * N_AUX_HORIZONS)?; + let mut grad_w_short_d = + stream.alloc_zeros::(b_sz * N_AUX_HORIZONS * AUX_HIDDEN)?; + let mut grad_b_short_d = stream.alloc_zeros::(b_sz * N_AUX_HORIZONS)?; + let mut grad_h_aux_d = stream.alloc_zeros::(b_sz * AUX_HIDDEN)?; + + aux_heads_bwd_gpu( + &stream, + &heads.bwd_fn, + &heads.w_long_d, + &heads.w_short_d, + &h_aux_d, + &grad_y_long_d, + &grad_y_short_d, + b_sz as i32, + &mut grad_w_long_d, + &mut grad_b_long_d, + &mut grad_w_short_d, + &mut grad_b_short_d, + &mut grad_h_aux_d, + )?; + stream.synchronize()?; + + let gwl = download(&stream, &grad_w_long_d)?; + let gbl = download(&stream, &grad_b_long_d)?; + let gws = download(&stream, &grad_w_short_d)?; + let gbs = download(&stream, &grad_b_short_d)?; + let gha = download(&stream, &grad_h_aux_d)?; + for (name, v) in [ + ("grad_w_long", &gwl), + ("grad_b_long", &gbl), + ("grad_w_short", &gws), + ("grad_b_short", &gbs), + ("grad_h_aux", &gha), + ] { + assert!( + v.iter().all(|x| x.is_finite()), + "{name} contains NaN/Inf", + ); + } + assert!( + gha.iter().any(|x| x.abs() > 0.0), + "grad_h_aux is all-zero — kernel didn't write the trunk feedback?", + ); + + // Closed-form expected grads from L = Σ grad_y * y_hat: + // dL/db_long[k] = Σ_batch grad_y_long[batch, k] + // dL/db_short[k] = Σ_batch grad_y_short[batch, k] + // dL/dW_long[k, c] = Σ_batch grad_y_long[batch, k] * h_aux[batch, c] + // dL/dW_short[k, c] = Σ_batch grad_y_short[batch, k] * h_aux[batch, c] + // dL/dh_aux[batch, c] = Σ_k grad_y_long[batch, k]*W_long[k,c] + // + Σ_k grad_y_short[batch, k]*W_short[k,c] + // + // The kernel writes PER-BATCH grad scratch (caller reduces axis 0) + // for the parameter grads — the assertion sums the per-batch entries + // and compares against the closed-form aggregate. + let w = heads.download_weights()?; + for k in 0..N_AUX_HORIZONS { + let mut got_bl = 0.0_f32; + let mut got_bs = 0.0_f32; + let mut expected_bl = 0.0_f32; + let mut expected_bs = 0.0_f32; + for batch in 0..b_sz { + got_bl += gbl[batch * N_AUX_HORIZONS + k]; + got_bs += gbs[batch * N_AUX_HORIZONS + k]; + expected_bl += grad_y_long[batch * N_AUX_HORIZONS + k]; + expected_bs += grad_y_short[batch * N_AUX_HORIZONS + k]; + } + let err_bl = (got_bl - expected_bl).abs(); + let err_bs = (got_bs - expected_bs).abs(); + assert!( + err_bl < 1e-5, + "grad_b_long[{k}] mismatch: got={got_bl} expected={expected_bl}", + ); + assert!( + err_bs < 1e-5, + "grad_b_short[{k}] mismatch: got={got_bs} expected={expected_bs}", + ); + } + + // Spot-check four weight entries against the closed form (full + // sweep is O(N_AUX_HORIZONS × AUX_HIDDEN × b_sz) — fine here, but + // sampling is sufficient to catch wiring/permutation breakage). + for &(k, c) in &[(0_usize, 0_usize), (1, 7), (2, 33), (0, 63)] { + let mut got_wl = 0.0_f32; + let mut got_ws = 0.0_f32; + let mut expected_wl = 0.0_f32; + let mut expected_ws = 0.0_f32; + for batch in 0..b_sz { + let off = batch * N_AUX_HORIZONS * AUX_HIDDEN + k * AUX_HIDDEN + c; + got_wl += gwl[off]; + got_ws += gws[off]; + expected_wl += + grad_y_long[batch * N_AUX_HORIZONS + k] * h_aux[batch * AUX_HIDDEN + c]; + expected_ws += + grad_y_short[batch * N_AUX_HORIZONS + k] * h_aux[batch * AUX_HIDDEN + c]; + } + let err_wl = (got_wl - expected_wl).abs(); + let err_ws = (got_ws - expected_ws).abs(); + let tol = 5e-4_f32 * expected_wl.abs().max(1e-3); + let tol_s = 5e-4_f32 * expected_ws.abs().max(1e-3); + assert!( + err_wl < tol.max(5e-5), + "grad_w_long[{k},{c}] mismatch: got={got_wl} expected={expected_wl} (err={err_wl})", + ); + assert!( + err_ws < tol_s.max(5e-5), + "grad_w_short[{k},{c}] mismatch: got={got_ws} expected={expected_ws} (err={err_ws})", + ); + } + + // grad_h_aux is per-batch — closed-form per (batch, c). + for &batch in &[0_usize, b_sz - 1] { + for &c in &[0_usize, 17, AUX_HIDDEN - 1] { + let mut expected = 0.0_f32; + for k in 0..N_AUX_HORIZONS { + expected += grad_y_long[batch * N_AUX_HORIZONS + k] + * w.w_long[k * AUX_HIDDEN + c]; + expected += grad_y_short[batch * N_AUX_HORIZONS + k] + * w.w_short[k * AUX_HIDDEN + c]; + } + let got = gha[batch * AUX_HIDDEN + c]; + let err = (got - expected).abs(); + assert!( + err < 5e-4_f32.max(5e-4 * expected.abs()), + "grad_h_aux[{batch},{c}] mismatch: got={got} expected={expected}", + ); + } + } + + Ok(()) +} + +/// Huber loss forward + backward correctness on a structured target +/// tensor that exercises BOTH the quadratic branch (|r| ≤ δ) and the +/// linear branch (|r| > δ). +#[test] +#[ignore = "requires CUDA"] +fn huber_loss_matches_naive_reference() -> Result<()> { + let dev = MlDevice::cuda(0)?; + let stream = dev.cuda_stream()?.clone(); + let loss = AuxHuberLoss::new(&dev)?; + + // Construct y_hat and y_true so that some residuals land in each + // branch. With δ = 1.0: residuals in [-0.5, 0.5] are quadratic, + // residuals in [-2.0, -1.0] ∪ [1.0, 2.0] are linear. + let n: usize = 24; + let y_hat: Vec = (0..n) + .map(|i| (i as f32 - 12.0) * 0.2) // -2.4 .. 2.2 step 0.2 + .collect(); + let y_true: Vec = vec![0.0; n]; + + let delta = DEFAULT_HUBER_DELTA; + let y_hat_d = upload(&stream, &y_hat)?; + let y_true_d = upload(&stream, &y_true)?; + let mut loss_d = stream.alloc_zeros::(1)?; + let mut grad_d = stream.alloc_zeros::(n)?; + let mut valid_d = upload_i32(&stream, &[0_i32])?; + + aux_huber_loss_gpu( + &stream, + &loss.fn_, + &y_hat_d, + &y_true_d, + delta, + n as i32, + &mut loss_d, + &mut grad_d, + &mut valid_d, + )?; + stream.synchronize()?; + + let loss_v = download(&stream, &loss_d)?[0]; + let grad_v = download(&stream, &grad_d)?; + let valid_v = download_i32(&stream, &valid_d)?[0]; + assert_eq!(valid_v as usize, n, "valid count mismatch"); + + // CPU oracle — sum of per-element Huber. + let mut expected_loss = 0.0_f32; + let mut hit_quadratic = false; + let mut hit_linear = false; + for i in 0..n { + let r = y_hat[i] - y_true[i]; + let ar = r.abs(); + let (l, g) = if ar <= delta { + hit_quadratic = true; + (0.5 * r * r, r) + } else { + hit_linear = true; + (delta * (ar - 0.5 * delta), delta * r.signum()) + }; + expected_loss += l; + let err = (grad_v[i] - g).abs(); + assert!( + err < 5e-5, + "grad_y_hat[{i}] = {} vs expected {} (r={}, |r|={})", + grad_v[i], + g, + r, + ar, + ); + } + assert!(hit_quadratic, "test fixture missed the quadratic branch"); + assert!(hit_linear, "test fixture missed the linear branch"); + let err = (loss_v - expected_loss).abs(); + assert!( + err < 5e-4, + "total loss mismatch: got={loss_v} expected={expected_loss}", + ); + + Ok(()) +} + +/// NaN-masked labels must produce zero gradient at the masked slot and +/// must NOT contribute to the loss accumulator or the valid count. +/// Without masking, a single NaN label would poison the whole batch: +/// * grad_y_hat[i] = NaN → propagates through aux_heads_bwd into +/// grad_h_aux, then into aux_trunk_bwd's grad_h_new → kills the +/// trunk for that batch sample. +/// * Σ Huber would become NaN → optimiser update goes NaN globally. +#[test] +#[ignore = "requires CUDA"] +fn huber_loss_nan_mask_does_not_propagate() -> Result<()> { + let dev = MlDevice::cuda(0)?; + let stream = dev.cuda_stream()?.clone(); + let loss = AuxHuberLoss::new(&dev)?; + + let n: usize = 12; + let y_hat: Vec = (0..n).map(|i| (i as f32) * 0.1 - 0.5).collect(); + // Half the labels are NaN-masked; the other half are 0.0. + let y_true: Vec = (0..n) + .map(|i| if i % 2 == 0 { f32::NAN } else { 0.0 }) + .collect(); + + let y_hat_d = upload(&stream, &y_hat)?; + let y_true_d = upload(&stream, &y_true)?; + let mut loss_d = stream.alloc_zeros::(1)?; + let mut grad_d = stream.alloc_zeros::(n)?; + // Pre-seed with sentinels so we can verify the kernel actually wrote. + let mut valid_d = upload_i32(&stream, &[-1_i32])?; + + aux_huber_loss_gpu( + &stream, + &loss.fn_, + &y_hat_d, + &y_true_d, + DEFAULT_HUBER_DELTA, + n as i32, + &mut loss_d, + &mut grad_d, + &mut valid_d, + )?; + stream.synchronize()?; + + let loss_v = download(&stream, &loss_d)?[0]; + let grad_v = download(&stream, &grad_d)?; + let valid_v = download_i32(&stream, &valid_d)?[0]; + + // Valid count is the non-NaN half. + assert_eq!(valid_v, (n / 2) as i32, "NaN-masked valid count wrong"); + // Loss is finite and equals the closed-form sum over the non-NaN + // entries — proves the NaN contributions were skipped. + assert!(loss_v.is_finite(), "loss is NaN/Inf"); + let mut expected_loss = 0.0_f32; + for i in 0..n { + if !y_true[i].is_finite() { + // NaN label: grad must be exactly 0, no loss contribution. + assert_eq!( + grad_v[i], 0.0, + "NaN-masked grad_y_hat[{i}] = {} (must be 0)", + grad_v[i], + ); + } else { + // Standard Huber contribution. + let r = y_hat[i] - y_true[i]; + let ar = r.abs(); + expected_loss += if ar <= DEFAULT_HUBER_DELTA { + 0.5 * r * r + } else { + DEFAULT_HUBER_DELTA * (ar - 0.5 * DEFAULT_HUBER_DELTA) + }; + // Live grad must be finite (already true if loss is finite, + // but the per-slot check pins the wiring down). + assert!( + grad_v[i].is_finite(), + "non-masked grad_y_hat[{i}] = {} (NaN/Inf)", + grad_v[i], + ); + } + } + let err = (loss_v - expected_loss).abs(); + assert!( + err < 5e-4, + "NaN-masked loss mismatch: got={loss_v} expected={expected_loss}", + ); + + Ok(()) +}