Files
foxhunt/crates/ml-alpha/cuda/aux_trunk.cu
jgrusewski 21e7dfd63c feat(aux-supervision): wire AuxTrunk + AuxHeads + Huber loss into PerceptionTrainer (B5)
Wires the aux supervision path parallel to BCE:
- AuxTrunk (64-hidden single-bucket CfC) consumes the same encoder output
  as the main BCE trunk
- AuxHeads (linear regression long/short) maps aux_trunk output to
  per-(direction, horizon) predicted outcomes
- AuxHuberLoss supervises against D-style labels from MultiHorizonLoader

Backward path with asymmetric stop-grad at encoder boundary:
- aux_trunk gets gradient signal into its OWN params at all times
- aux_trunk's encoder-boundary gradient is INITIALLY blocked
  (stop_grad_aux_to_encoder = true)
- Conditional lift per E3 design: if aux_huber_ema < 0.4 AND
  aux_dir_acc_ema > 0.85 within 200 steps, lift the stop-grad
- When lifted, aux_vec_add kernel folds aux's grad_x into the main
  grad_h_enriched_seq slot (element-wise += per feedback_no_atomicadd)

ISV signals added: aux_huber_per_h, aux_dir_acc_per_h (per pearl).

Per-trunk scratch + reduced grad buffers (no Adam state sharing per
pearl_adam_normalizes_loss_weights — opt_aux is its own Adam group).

New helper kernel cuda/aux_vec_add.cu: position-local dst += src for
the asymmetric stop-grad lift accumulation.

New synthetic test stacked_trainer_aux_supervision_converges_on_constant_signal
validates end-to-end:
  aux_huber_ema_per_h    = [0.087, 0.087, 0.087]  (converged)
  aux_dir_acc_ema_per_h  = [1.0, 1.0, 1.0]        (perfect on constant)
  stop_grad_aux_to_encoder = false                (lift fired)

All 5 stacked_trainer tests pass on RTX 3050 (lib still converges, no
regression from parallel aux wiring).

Not yet consumed by decision policy (B7) — aux output flows through
training only.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-05-22 09:45:29 +02:00

254 lines
12 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// aux_trunk.cu — single-bucket CfC step (forward + backward) at AUX_HIDDEN=64.
//
// Per `pearl_separate_aux_trunk_when_shared_starves`: this is the kernel
// behind a SECOND, smaller CfC trunk used for outcome-supervision (D-labels).
// It runs in parallel with the main per-branch CfC trunk on the same encoder
// output, with INDEPENDENT parameters (no shared w_in/w_rec/b/tau) so that
// BCE gradient flow on the main trunk does not starve the aux supervision.
//
// Design contract:
// • Single bucket — no per-horizon channel splitting, no
// channels_in_bucket lookup, no bucket_dim_k. Aux supervision is
// per-K at the head (B4 wiring), not at the τ-state.
// • Hidden dim is AUX_HIDDEN=64 (half of main trunk's HIDDEN_DIM=128).
// • Algorithm matches `cfc_step_per_branch.cu`: continuous-time CfC
// decay update `h_new = h_old·decay + (1-decay)·tanh(pre)`, with
// `decay = exp(-dt / max(tau, 1e-6))`.
// • Input feature dim is parameterised (`feat_dim` runtime argument)
// so a single cubin handles both raw-snap input (FEATURE_DIM=40) and
// post-encoder input (HIDDEN_DIM=128). The host wrapper passes the
// correct dim; the kernel iterates `feat_dim` in the input MVM.
//
// Per `feedback_no_atomicadd.md`: no atomicAdd anywhere. Each thread
// writes to its own (batch, c) slot in forward, and to its own per-batch
// grad slice in backward. Cross-batch reduction is the caller's
// responsibility via the existing `reduce_axis0` infrastructure.
//
// Per `pearl_cooperative_staging_eliminates_redundant_reads`: the per-
// batch `x` and `h_old` rows are SHARED across all output channels in a
// block. Stage them into shared memory once at block entry so the K-loop
// reads from smem instead of issuing AUX_HIDDEN × redundant DRAM reads.
//
// Per `pearl_no_host_branches_in_captured_graph`: no host branching;
// the kernel takes feat_dim as a runtime arg but does not consult any
// host-side state during execution.
//
// Per `feedback_nvidia_grade_perf_for_kernels.md`: warp-uniform branches
// (the only conditional is `batch < B`, identical for every thread in a
// block since batch == blockIdx.x), no atomicAdd, coalesced loads via
// cooperative staging.
#define AUX_HIDDEN 64
#define MAX_FEAT_DIM 128 // upper bound for shared-memory staging; covers
// both FEATURE_DIM=40 and HIDDEN_DIM=128 callers
// ─────────────────────────────────────────────────────────────────────
// aux_trunk_fwd: forward pass.
//
// Launch:
// grid = (B, 1, 1)
// block = (AUX_HIDDEN = 64, 1, 1)
// shared_mem_bytes = (feat_dim + AUX_HIDDEN) * sizeof(float)
//
// Per (batch, c) thread computes:
// pre = b[c] + Σ_k W_in[c, k] * x[batch, k] + Σ_k W_rec[c, k] * h_old[batch, k]
// decay = exp(-dt / max(tau[c], 1e-6))
// h_new[batch, c] = h_old[batch, c] * decay + (1 - decay) * tanh(pre)
//
// W_in is [AUX_HIDDEN × feat_dim], W_rec is [AUX_HIDDEN × AUX_HIDDEN].
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void aux_trunk_fwd(
const float* __restrict__ w_in, // [AUX_HIDDEN × feat_dim]
const float* __restrict__ w_rec, // [AUX_HIDDEN × AUX_HIDDEN]
const float* __restrict__ b, // [AUX_HIDDEN]
const float* __restrict__ tau, // [AUX_HIDDEN]
const float* __restrict__ x, // [B × feat_dim]
const float* __restrict__ h_old, // [B × AUX_HIDDEN]
float dt,
int B,
int feat_dim,
float* __restrict__ h_new // [B × AUX_HIDDEN]
) {
extern __shared__ float smem[];
float* x_local = smem; // [feat_dim]
float* h_old_local = smem + feat_dim; // [AUX_HIDDEN]
int batch = blockIdx.x;
int c = threadIdx.x;
if (batch >= B) return;
// Cooperative staging of per-batch x[batch, *] and h_old[batch, *].
// Each thread loads ceil(feat_dim / AUX_HIDDEN) = up to ~2 x slots
// and exactly one h_old slot (since blockDim.x == AUX_HIDDEN).
for (int i = c; i < feat_dim; i += AUX_HIDDEN) {
x_local[i] = x[batch * feat_dim + i];
}
h_old_local[c] = h_old[batch * AUX_HIDDEN + c];
__syncthreads();
// pre = b[c] + Σ_k W_in[c, k] * x_local[k] + Σ_k W_rec[c, k] * h_old_local[k]
float pre = b[c];
for (int k = 0; k < feat_dim; ++k) {
pre += w_in[c * feat_dim + k] * x_local[k];
}
for (int k = 0; k < AUX_HIDDEN; ++k) {
pre += w_rec[c * AUX_HIDDEN + k] * h_old_local[k];
}
float decay = expf(-dt / fmaxf(tau[c], 1e-6f));
h_new[batch * AUX_HIDDEN + c] =
h_old_local[c] * decay + (1.0f - decay) * tanhf(pre);
}
// ─────────────────────────────────────────────────────────────────────
// aux_trunk_bwd: backward pass.
//
// Launch:
// grid = (B, 1, 1)
// block = (AUX_HIDDEN = 64, 1, 1)
// shared_mem_bytes = (feat_dim + AUX_HIDDEN) * sizeof(float)
//
// Per-batch grad slices accumulation semantics (SDD-3 Layer B5 callers):
// * grad_w_in / grad_w_rec / grad_b / grad_tau — `+=` per K-loop launch.
// Multiple K-iterations sum into the same per-batch scratch (each step
// has different x/h_old → different contribution), and reduce_axis0
// at end of step collapses across B.
// * grad_h_old / grad_x — `=` overwrite (consumed immediately by the
// caller after each per-step launch: grad_h_old carries into k-1,
// grad_x is accumulated into the encoder grad slot or discarded).
// Caller must zero the four param-grad scratches at step start (the
// trainer does this alongside the rest of its per-step scratch memsets).
//
// grad shapes:
// grad_w_in : [B × AUX_HIDDEN × feat_dim] per-batch scratch
// grad_w_rec : [B × AUX_HIDDEN × AUX_HIDDEN] per-batch scratch
// grad_b : [B × AUX_HIDDEN] per-batch scratch
// grad_tau : [B × AUX_HIDDEN] per-batch scratch
// grad_h_old : [B × AUX_HIDDEN] per-batch (decay
// contribution only;
// cross-channel reduction
// deferred to caller if
// unrolled K>1)
// grad_x : [B × feat_dim] per-batch (reduced
// across c via shared-
// memory tree-reduce)
//
// grad_x cross-channel reduction is done inside this kernel via a
// block-resident shared-memory tree-reduce (no atomicAdd per
// `feedback_no_atomicadd`).
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void aux_trunk_bwd(
const float* __restrict__ w_in, // [AUX_HIDDEN × feat_dim]
const float* __restrict__ w_rec, // [AUX_HIDDEN × AUX_HIDDEN]
const float* __restrict__ b, // [AUX_HIDDEN]
const float* __restrict__ tau, // [AUX_HIDDEN]
const float* __restrict__ x, // [B × feat_dim]
const float* __restrict__ h_old, // [B × AUX_HIDDEN]
const float* __restrict__ grad_h_new, // [B × AUX_HIDDEN]
float dt,
int B,
int feat_dim,
float* __restrict__ grad_w_in, // [B × AUX_HIDDEN × feat_dim]
float* __restrict__ grad_w_rec, // [B × AUX_HIDDEN × AUX_HIDDEN]
float* __restrict__ grad_b, // [B × AUX_HIDDEN]
float* __restrict__ grad_tau, // [B × AUX_HIDDEN]
float* __restrict__ grad_h_old, // [B × AUX_HIDDEN]
float* __restrict__ grad_x // [B × feat_dim]
) {
// Shared layout:
// [0 .. feat_dim) : x_local
// [feat_dim .. feat_dim+AH) : h_old_local
// [feat_dim+AH .. feat_dim+AH+AH) : d_pre_smem (per-channel d_pre
// for cross-channel grad_x reduce)
extern __shared__ float smem[];
float* x_local = smem;
float* h_old_local = smem + feat_dim;
float* d_pre_smem = smem + feat_dim + AUX_HIDDEN;
int batch = blockIdx.x;
int c = threadIdx.x;
if (batch >= B) return;
// Cooperative staging — identical pattern to forward.
for (int i = c; i < feat_dim; i += AUX_HIDDEN) {
x_local[i] = x[batch * feat_dim + i];
}
h_old_local[c] = h_old[batch * AUX_HIDDEN + c];
__syncthreads();
// Recompute forward pre + tanh to derive d_pre.
float pre = b[c];
for (int k = 0; k < feat_dim; ++k) {
pre += w_in[c * feat_dim + k] * x_local[k];
}
for (int k = 0; k < AUX_HIDDEN; ++k) {
pre += w_rec[c * AUX_HIDDEN + k] * h_old_local[k];
}
const float tau_eps = 1e-6f;
const float tau_raw = tau[c];
const float tau_eff = fmaxf(tau_raw, tau_eps);
const float decay = expf(-dt / tau_eff);
const float s = tanhf(pre);
const float dh = grad_h_new[batch * AUX_HIDDEN + c];
const float d_pre = dh * (1.0f - decay) * (1.0f - s * s);
const float d_decay = dh * (h_old_local[c] - s);
// Per-batch grad-scratch writes. Each thread (batch, c) is the sole
// writer of grad_*[batch, c, ...] — no race, no atomicAdd.
// Param-grad scratches use `+=` so the trainer's reverse-K loop sums
// contributions across all K positions (each with different x/h_old
// through the same params).
grad_b[batch * AUX_HIDDEN + c] += d_pre;
// grad_tau receives 0 when tau hits the clamp floor
// (strict d(max)/d(tau) = 0 below tau_eps).
const float gate = (tau_raw > tau_eps) ? 1.0f : 0.0f;
grad_tau[batch * AUX_HIDDEN + c] +=
gate * d_decay * decay * dt / (tau_eff * tau_eff);
// grad_w_in[batch, c, k] += d_pre * x_local[k]
for (int k = 0; k < feat_dim; ++k) {
const long long off =
(long long)batch * AUX_HIDDEN * feat_dim
+ (long long)c * feat_dim + k;
grad_w_in[off] += d_pre * x_local[k];
}
// grad_w_rec[batch, c, k] += d_pre * h_old_local[k]
for (int k = 0; k < AUX_HIDDEN; ++k) {
const long long off =
(long long)batch * AUX_HIDDEN * AUX_HIDDEN
+ (long long)c * AUX_HIDDEN + k;
grad_w_rec[off] += d_pre * h_old_local[k];
}
// grad_h_old[batch, c] direct decay contribution. OVERWRITE — the
// trainer's per-step DtoD copy reads this immediately into the next
// K iteration's carry buffer; no cross-K accumulation here.
// (Cross-channel term Σ_j d_pre[j] * W_rec[j, c] is left for the
// caller to add if BPTT through h_old is needed — for K=1 per-step
// aux supervision this direct term is sufficient.)
grad_h_old[batch * AUX_HIDDEN + c] = dh * decay;
// grad_x[batch, k] = Σ_c d_pre[c] * W_in[c, k]
// Cross-channel reduction via shared memory: stash d_pre to smem,
// then have each thread compute its assigned k-slice by summing
// across c. This avoids atomicAdd while keeping the reduction
// block-local. (Per pearl_no_atomicadd + the block tree-reduce
// pattern used in reduce_axis0.cu.)
d_pre_smem[c] = d_pre;
__syncthreads();
// Each of the AUX_HIDDEN threads handles a strided subset of feat_dim
// output positions. With feat_dim ≤ MAX_FEAT_DIM = 128 and
// AUX_HIDDEN = 64, each thread covers at most 2 k positions.
for (int k = c; k < feat_dim; k += AUX_HIDDEN) {
float acc = 0.0f;
for (int j = 0; j < AUX_HIDDEN; ++j) {
acc += d_pre_smem[j] * w_in[j * feat_dim + k];
}
grad_x[batch * feat_dim + k] = acc;
}
}