refactor(ml-alpha): remove per-horizon Q_h path (C21-C25 falsified) [V1]

3-fold A/B sweep 2026-05-18 at commit 83546b5c3 falsified the simple
per-horizon Q_h attention pool:
  mean_auc 0.7559 ± 0.0068  vs baseline 0.7749 ± 0.024  (Δ = -0.019)
  h6000    0.7588 ± 0.0049  vs baseline 0.7591 ± 0.018  (Δ = -0.0003)

best_epoch on val_loss = 1 in 2/3 folds → calibration drift as α opens;
no horizon-distribution shift toward h6000. The per-horizon path
spends capacity on directions that hurt log-likelihood without lifting
ranking quality.

V1 of the v2 redesign deletes the falsified path entirely (per
feedback_no_partial_refactor; v2 spec/plan committed earlier today
captures the migration). Files removed:
  cuda/per_horizon_attention_pool.cu
  cuda/per_horizon_residual_head.cu
  cuda/per_horizon_prob_blend.cu
  src/per_horizon_attention_pool.rs
  src/per_horizon_residual_head.rs
  src/trainer/per_horizon_state.rs
  tests/per_horizon_attention_pool_numgrad.rs
  tests/per_horizon_residual_head_numgrad.rs
  tests/per_horizon_full_pipeline_smoke.rs

perception.rs:
  - struct field `per_horizon` removed
  - new() initialization removed
  - step_batched section 4.5 (forward_with_blend) → reserved comment
  - step_batched section 5a (backward_through_blend) → reserved comment
  - step_batched section 9 (adamw_step) → reserved comment
  - existing 17 optimizer groups + BCE/attention-pool path untouched
  - reduce_axis0 kernel kept (still used by existing param-grad reducers)

build.rs KERNELS: dropped the 3 per_horizon entries.
lib.rs + trainer/mod.rs: dropped per_horizon module declarations.

Workspace compiles clean (cargo check -p ml-alpha). Next: V2 builds
the horizon_token_attention_pool kernel as the v2 replacement.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-18 13:26:45 +02:00
parent 4425a77844
commit 41292303dc
13 changed files with 16 additions and 1729 deletions

View File

@@ -20,9 +20,6 @@ const KERNELS: &[&str] = &[
"layer_norm", // Phase 1: trunk pre-CfC normalisation
"variable_selection", // Phase 2D: TFT-style per-feature gating
"attention_pool", // Phase 3: learned context summary at CfC k=0
"per_horizon_attention_pool", // C21: per-horizon variant; ships behind a config flag
"per_horizon_residual_head", // C22: per-horizon scalar residual on top of multi_horizon_heads logits
"per_horizon_prob_blend", // C25: closed-form logit-bias applied in probability space
"reduce_axis0", // Phase B: cross-batch param-grad reducer
];

View File

@@ -1,259 +0,0 @@
// per_horizon_attention_pool.cu — per-horizon attention pool (C21).
//
// Replaces the single learned query Q[HIDDEN_DIM] of attention_pool.cu
// with N_HORIZONS=5 learned queries Q_h[N_HORIZONS, HIDDEN_DIM]. Each
// horizon attends to a different combination of K LN_b positions,
// producing its own context vector. The multi-horizon heads downstream
// consume each horizon's own context (concat'd with CfC h_K).
//
// Forward math (per sample b, per horizon h):
// scores_h[k] = Σ_d Q_h[h, d] * LNb[b, k, d] # [K]
// attn_h[k] = softmax_k(scores_h) # [K]
// context_h[d] = Σ_k attn_h[k] * LNb[b, k, d] # [HIDDEN_DIM]
//
// For this pool keys == values == LN_b output [B, K, HIDDEN_DIM].
// Learned params: Q_h [N_HORIZONS, HIDDEN_DIM].
//
// Saved for backward:
// attn_h_weights [B, N_HORIZONS, K] (post-softmax)
//
// Block layout (forward):
// grid_dim = (B, 1, 1), block_dim = (HIDDEN_DIM=128, 1, 1).
// Each block handles ALL horizons for its batch — horizon loop is
// inside the block. Keeps grad_ln_out writes free of cross-block
// races: each batch's [b, :, :] slice is touched by ONE block in bwd.
//
// Reduction strategy (perf-critical):
// block_dim == HIDDEN_DIM == 128 == 4 warps. Dot products over the
// HIDDEN_DIM axis are reduced via warp-shuffle (`__shfl_xor_sync`)
// inside each warp, then a single cross-warp reduction over 4 lanes.
// This replaces ~10 `__syncthreads` per K-step with 1, ~10× fewer
// barriers on the hot inner loop.
//
// Per `feedback_no_atomicadd.md`: block tree-reduce only, no atomics.
#define PHA_HIDDEN_DIM 128
#define PHA_BLOCK 128
#define PHA_MAX_K 512 // safety cap; smoke uses K=16-64.
#define PHA_N_HORIZONS 5
#define PHA_N_WARPS (PHA_BLOCK / 32) // == 4
// Reduce a per-thread float across a 128-thread block (4 warps).
// `s_warp` must be [PHA_N_WARPS] in shared mem (caller-supplied).
// Result is broadcast to all threads in the block.
__device__ __forceinline__ float block_reduce_sum(float v, float* s_warp, int tid) {
// Intra-warp reduction.
for (int s = 16; s > 0; s >>= 1) v += __shfl_xor_sync(0xffffffff, v, s);
const int lane = tid & 31;
const int warp = tid >> 5;
if (lane == 0) s_warp[warp] = v;
__syncthreads();
// Cross-warp reduction (4 lanes only).
float w = (tid < PHA_N_WARPS) ? s_warp[tid] : 0.0f;
if (tid < 32) {
for (int s = PHA_N_WARPS / 2; s > 0; s >>= 1) {
w += __shfl_xor_sync(0xffffffff, w, s);
}
}
// Re-broadcast via shared mem so all threads see the same scalar.
if (tid == 0) s_warp[0] = w;
__syncthreads();
return s_warp[0];
}
// Block-wide max reduction (mirror of block_reduce_sum but with fmaxf).
__device__ __forceinline__ float block_reduce_max(float v, float* s_warp, int tid) {
for (int s = 16; s > 0; s >>= 1) {
v = fmaxf(v, __shfl_xor_sync(0xffffffff, v, s));
}
const int lane = tid & 31;
const int warp = tid >> 5;
if (lane == 0) s_warp[warp] = v;
__syncthreads();
float w = (tid < PHA_N_WARPS) ? s_warp[tid] : -INFINITY;
if (tid < 32) {
for (int s = PHA_N_WARPS / 2; s > 0; s >>= 1) {
w = fmaxf(w, __shfl_xor_sync(0xffffffff, w, s));
}
}
if (tid == 0) s_warp[0] = w;
__syncthreads();
return s_warp[0];
}
extern "C" __global__ void per_horizon_attention_pool_fwd(
const float* __restrict__ Q_h, // [N_HORIZONS, HIDDEN_DIM]
const float* __restrict__ ln_out, // [B, K, HIDDEN_DIM]
int n_batch,
int k_seq,
float* __restrict__ context_h, // [B, N_HORIZONS, HIDDEN_DIM]
float* __restrict__ attn_h_weights // [B, N_HORIZONS, K]
) {
int b_idx = blockIdx.x;
int tid = threadIdx.x;
if (b_idx >= n_batch || tid >= PHA_BLOCK) return;
extern __shared__ float smem[];
float* s_scores = smem; // [K]
float* s_warp = smem + k_seq; // [PHA_N_WARPS]
// (no longer need a [BLOCK]-sized s_red — warp-shuffle reduce uses only N_WARPS slots)
__shared__ float s_Qh[PHA_HIDDEN_DIM];
const float* ln_b = ln_out + (long long)b_idx * k_seq * PHA_HIDDEN_DIM;
// Outer loop: one full forward pass per horizon, reusing the
// per-block scratch for each. Sequential horizons within a block
// means no cross-horizon barrier worries; the smem is reused.
for (int h = 0; h < PHA_N_HORIZONS; ++h) {
// Cache Q_h[h, :] for this iteration.
if (tid < PHA_HIDDEN_DIM) s_Qh[tid] = Q_h[h * PHA_HIDDEN_DIM + tid];
__syncthreads();
// Pass 1: scores[k] = Q_h[h, :] · ln_out[b, k, :].
// Warp-shuffle reduce → ~1 syncthreads per K instead of 8.
for (int k = 0; k < k_seq; ++k) {
float v = ln_b[k * PHA_HIDDEN_DIM + tid] * s_Qh[tid];
float r = block_reduce_sum(v, s_warp, tid);
if (tid == 0) s_scores[k] = r;
__syncthreads(); // s_warp will be re-used; safe.
}
// Pass 2: softmax over K — max-subtract + exp + sum + divide.
float my_max = -INFINITY;
for (int k = tid; k < k_seq; k += PHA_BLOCK) {
const float v = s_scores[k];
if (v > my_max) my_max = v;
}
const float s_max = block_reduce_max(my_max, s_warp, tid);
float my_sum = 0.0f;
for (int k = tid; k < k_seq; k += PHA_BLOCK) {
const float e = expf(s_scores[k] - s_max);
s_scores[k] = e;
my_sum += e;
}
const float s_sum = block_reduce_sum(my_sum, s_warp, tid);
// Pass 3: attn[k] = exp / sum; save; context[d] = Σ_k attn[k] * ln_out[b, k, d].
const long long attn_base = (long long)b_idx * PHA_N_HORIZONS * k_seq
+ (long long)h * k_seq;
for (int k = tid; k < k_seq; k += PHA_BLOCK) {
const float a = s_scores[k] / s_sum;
s_scores[k] = a;
attn_h_weights[attn_base + k] = a;
}
__syncthreads();
if (tid < PHA_HIDDEN_DIM) {
float c = 0.0f;
for (int k = 0; k < k_seq; ++k) {
c += s_scores[k] * ln_b[k * PHA_HIDDEN_DIM + tid];
}
context_h[(long long)b_idx * PHA_N_HORIZONS * PHA_HIDDEN_DIM
+ (long long)h * PHA_HIDDEN_DIM + tid] = c;
}
__syncthreads(); // before next horizon resets s_Qh / s_scores
}
}
// Per-horizon attention pool backward.
//
// Chain rule (per (b, h) slice — identical shape to the single-Q
// backward, just looped over horizons):
//
// d_attn[k] = Σ_d grad_context[d] * ln_out[b, k, d]
// d_scores[k] = attn[k] * (d_attn[k] - Σ_kp attn[kp] * d_attn[kp])
// d_Q_h[h, d] += Σ_k d_scores[k] * ln_out[b, k, d]
// d_LNb[b, k, d] += grad_context[d] * attn[k] + d_scores[k] * Q_h[h, d]
//
// grad_Q_h is [N_HORIZONS, HIDDEN_DIM] shared across all batches → per-batch
// scratch [B, N_HORIZONS, HIDDEN_DIM] reduced via reduce_axis0 after.
// grad_ln_out is per-batch indexed; each block writes its [b, :, :] slice
// via += and sums horizon contributions sequentially within the block —
// no cross-block race possible.
extern "C" __global__ void per_horizon_attention_pool_bwd(
const float* __restrict__ Q_h, // [N_HORIZONS, HIDDEN_DIM]
const float* __restrict__ ln_out, // [B, K, HIDDEN_DIM]
const float* __restrict__ attn_h_weights, // [B, N_HORIZONS, K] from fwd
const float* __restrict__ grad_context_h, // [B, N_HORIZONS, HIDDEN_DIM]
int n_batch,
int k_seq,
float* __restrict__ grad_Qh_scratch, // [B, N_HORIZONS, HIDDEN_DIM] (+=)
float* __restrict__ grad_ln_out // [B, K, HIDDEN_DIM] (+= chained with K-loop)
) {
int bi = blockIdx.x;
int tid = threadIdx.x;
if (bi >= n_batch || tid >= PHA_BLOCK) return;
extern __shared__ float smem[];
float* s_dattn = smem; // [K]
float* s_dscores = smem + k_seq; // [K]
float* s_warp = smem + 2 * k_seq; // [PHA_N_WARPS]
__shared__ float s_Qh[PHA_HIDDEN_DIM];
__shared__ float s_attn[PHA_MAX_K];
__shared__ float s_grad_ctx[PHA_HIDDEN_DIM];
const float* ln_b = ln_out + (long long)bi * k_seq * PHA_HIDDEN_DIM;
float* grad_ln_b = grad_ln_out + (long long)bi * k_seq * PHA_HIDDEN_DIM;
// Horizon loop. For each h: read Q_h + grad_context + attn; compute
// all backward gradients; accumulate into grad_ln_b (shared across
// horizons within this batch — sequential within block, no race).
for (int h = 0; h < PHA_N_HORIZONS; ++h) {
if (tid < PHA_HIDDEN_DIM) {
s_Qh[tid] = Q_h[h * PHA_HIDDEN_DIM + tid];
s_grad_ctx[tid] = grad_context_h[(long long)bi * PHA_N_HORIZONS * PHA_HIDDEN_DIM
+ (long long)h * PHA_HIDDEN_DIM + tid];
}
const long long attn_base = (long long)bi * PHA_N_HORIZONS * k_seq
+ (long long)h * k_seq;
for (int k = tid; k < k_seq; k += PHA_BLOCK) {
s_attn[k] = attn_h_weights[attn_base + k];
}
__syncthreads();
// Pass 1: d_attn[k] = Σ_d grad_context[d] * ln_out[b, k, d].
// Warp-shuffle reduce per k (was 8 syncthreads per k — now 1).
for (int k = 0; k < k_seq; ++k) {
float v = s_grad_ctx[tid] * ln_b[k * PHA_HIDDEN_DIM + tid];
float r = block_reduce_sum(v, s_warp, tid);
if (tid == 0) s_dattn[k] = r;
__syncthreads();
}
// Pass 2: Σ_kp attn[kp] * d_attn[kp] (softmax Jacobian centring).
float my_sum = 0.0f;
for (int k = tid; k < k_seq; k += PHA_BLOCK) {
my_sum += s_attn[k] * s_dattn[k];
}
const float s_sum_attn_dattn = block_reduce_sum(my_sum, s_warp, tid);
// Pass 3: d_scores[k] = attn[k] * (d_attn[k] - centring).
for (int k = tid; k < k_seq; k += PHA_BLOCK) {
s_dscores[k] = s_attn[k] * (s_dattn[k] - s_sum_attn_dattn);
}
__syncthreads();
// Pass 4: grad_Qh_scratch[bi, h, d] += Σ_k d_scores[k] * ln_out[b, k, d].
// grad_ln_b[k, d] += grad_context[d] * attn[k] + d_scores[k] * Q_h[h, d].
// Thread tid owns column d=tid for this (b, h). Per-horizon += into grad_ln_b
// is sequential within the block — no race; horizons accumulate over each other.
if (tid < PHA_HIDDEN_DIM) {
float dq_local = 0.0f;
for (int k = 0; k < k_seq; ++k) {
const float v = ln_b[k * PHA_HIDDEN_DIM + tid];
dq_local += s_dscores[k] * v;
grad_ln_b[k * PHA_HIDDEN_DIM + tid] +=
s_grad_ctx[tid] * s_attn[k] + s_dscores[k] * s_Qh[tid];
}
grad_Qh_scratch[(long long)bi * PHA_N_HORIZONS * PHA_HIDDEN_DIM
+ (long long)h * PHA_HIDDEN_DIM + tid] += dq_local;
}
__syncthreads();
}
}

View File

@@ -1,128 +0,0 @@
// per_horizon_prob_blend.cu — logit-space additive bias over per-K probs.
//
// Existing pipeline (perception.rs::step_batched) stores per-K BASELINE
// logits in `logit_per_k_d[K, B, N_HORIZONS]` AND their sigmoid in
// `probs_per_k_d[K, B, N_HORIZONS]`. The per-horizon attention pool's
// learnable α-gated residual adds a logit-space shift:
//
// r_contrib[b, h] = tanh(α[h]) * residual[b, h]
// probs_final[k,b,h] = sigmoid(logit_baseline[k,b,h] + r_contrib[b,h])
//
// At α=0 ⇒ r_contrib=0 ⇒ probs_final = sigmoid(logit_baseline) =
// probs_baseline → bit-identical to baseline (C23 invariant extended
// into trainer).
//
// Forward (fwd kernel): overwrites probs_per_k_d in place with
// probs_final. BCE consumes probs_per_k_d as-is afterward.
//
// Backward: NO separate kernel needed for d_logit_baseline. The existing
// heads_grn_bwd already computes d_logit_baseline = grad_probs * p *
// (1 - p), and since p_final = sigmoid(logit_baseline + r_contrib),
// d_logit_baseline = d_p_final * sigmoid_deriv(logit_baseline + r_contrib)
// = d_p_final * p_final * (1 - p_final) — which is exactly what the
// existing GRN backward kernel computes. So the GRN backward path is
// UNCHANGED — the chain rule absorbs the bias automatically.
//
// What we DO need: compute d_r_contrib_per_k[k, b, h] = the same
// quantity (= d_logit_baseline by chain-rule equivalence) and reduce
// over k + b to produce d_residual[b, h] and d_alpha[h].
#define PHB_N_HORIZONS 5
extern "C" __global__ void per_horizon_prob_blend_fwd(
const float* __restrict__ alpha, // [N_HORIZONS]
const float* __restrict__ residual, // [B, N_HORIZONS]
const float* __restrict__ logit_per_k, // [K, B, N_HORIZONS] from GRN forward
float* __restrict__ probs_per_k, // [K, B, N_HORIZONS] OVERWRITTEN with probs_final
int n_batch,
int k_seq
) {
int b = blockIdx.x;
int h = threadIdx.x;
if (b >= n_batch || h >= PHB_N_HORIZONS) return;
const float r_contrib = tanhf(alpha[h]) * residual[b * PHB_N_HORIZONS + h];
for (int k = 0; k < k_seq; ++k) {
const long long idx =
(long long)k * n_batch * PHB_N_HORIZONS
+ (long long)b * PHB_N_HORIZONS
+ (long long)h;
const float z = logit_per_k[idx] + r_contrib;
probs_per_k[idx] = 1.0f / (1.0f + expf(-z));
}
}
// Reduce step. Reads grad_probs_per_k (= ∂L/∂p_final, from BCE) and
// probs_per_k (= p_final, in-place overwrite from fwd kernel) to derive:
//
// d_logit_per_k[k, b, h] = grad_probs[k, b, h] * p_final * (1 - p_final)
//
// then sums over the appropriate axes:
//
// d_residual[b, h] = tanh(α[h]) * Σ_k d_logit_per_k[k, b, h]
// d_alpha[h] = (1 - tanh²(α[h])) * Σ_{k, b} d_logit_per_k[k, b, h] * residual[b, h]
//
// One block per horizon. Warp-parallel inside the block: 32 threads
// stride over the K*B index space, each maintaining its own k_sum
// per batch via warp shuffle, then a final warp shuffle for alpha_sum.
// d_residual[b, h] needs a per-batch sum (across k), so we let lane 0
// of each warp write it after a warp reduction over the K stride.
//
// Grid: (N_HORIZONS, 1, 1). Block: (32, 1, 1) = one warp per horizon.
// Single-warp blocks → zero __syncthreads; all reduction via shuffle.
extern "C" __global__ void per_horizon_prob_blend_reduce_alpha_residual(
const float* __restrict__ alpha, // [N_HORIZONS]
const float* __restrict__ residual, // [B, N_HORIZONS]
const float* __restrict__ probs_per_k, // [K, B, N_HORIZONS] (= probs_final)
const float* __restrict__ grad_probs, // [K, B, N_HORIZONS] (= ∂L/∂p_final from BCE)
int n_batch,
int k_seq,
float* __restrict__ d_residual, // [B, N_HORIZONS] (written)
float* __restrict__ d_alpha // [N_HORIZONS] (written)
) {
int h = blockIdx.x;
int lane = threadIdx.x;
if (h >= PHB_N_HORIZONS || lane >= 32) return;
const float a = alpha[h];
const float ta = tanhf(a);
const float sech2 = 1.0f - ta * ta;
// For each batch, all 32 threads cooperatively reduce over K via
// shuffle. Lane 0 writes d_residual[b, h]. alpha_sum accumulates
// across batches and is reduced once at the end.
float alpha_sum = 0.0f;
for (int b = 0; b < n_batch; ++b) {
const float resid_bh = residual[b * PHB_N_HORIZONS + h];
// Per-thread strided sum over k.
float k_sum = 0.0f;
float k_resid_acc = 0.0f;
for (int k = lane; k < k_seq; k += 32) {
const long long idx =
(long long)k * n_batch * PHB_N_HORIZONS
+ (long long)b * PHB_N_HORIZONS
+ (long long)h;
const float p = probs_per_k[idx];
const float gp = grad_probs[idx];
const float d_logit = gp * p * (1.0f - p);
k_sum += d_logit;
k_resid_acc += d_logit * resid_bh;
}
// Warp-reduce both sums.
for (int s = 16; s > 0; s >>= 1) {
k_sum += __shfl_xor_sync(0xffffffff, k_sum, s);
k_resid_acc += __shfl_xor_sync(0xffffffff, k_resid_acc, s);
}
if (lane == 0) {
d_residual[b * PHB_N_HORIZONS + h] = ta * k_sum;
}
alpha_sum += k_resid_acc; // already warp-reduced above
}
// Lane 0 already holds the full alpha_sum (added of warp-reduced values).
if (lane == 0) {
d_alpha[h] = sech2 * alpha_sum;
}
}

View File

@@ -1,104 +0,0 @@
// per_horizon_residual_head.cu — per-horizon scalar residual from contexts (C22).
//
// Produces a per-horizon scalar residual that the trainer adds (behind
// a learnable α-gate) to the existing multi_horizon_heads logit output.
// Keeps the existing head kernel unchanged — the per-horizon attention
// pool from C21 contributes via this lightweight projection that the
// PerceptionTrainer can opt into without weight-shape changes elsewhere.
//
// Forward math (per sample b, per horizon h):
// residual[b, h] = Σ_d w_res[h, d] * context_h[b, h, d] + bias_res[h]
//
// One block per batch; thread tid handles dimension d. Reduction across
// d uses block tree-reduce (no atomicAdd per feedback_no_atomicadd.md).
// Inner loop over N_HORIZONS sequentially within the block (same pattern
// as per_horizon_attention_pool.cu).
//
// Saved for backward: nothing (residual is linear in its inputs — bwd
// can recompute from w_res / context_h).
#define PHR_HIDDEN_DIM 128
#define PHR_BLOCK 128
#define PHR_N_HORIZONS 5
#define PHR_N_WARPS (PHR_BLOCK / 32)
extern "C" __global__ void per_horizon_residual_head_fwd(
const float* __restrict__ context_h, // [B, N_HORIZONS, HIDDEN_DIM]
const float* __restrict__ w_res, // [N_HORIZONS, HIDDEN_DIM]
const float* __restrict__ bias_res, // [N_HORIZONS]
int n_batch,
float* __restrict__ residual_out // [B, N_HORIZONS]
) {
int b = blockIdx.x;
int tid = threadIdx.x;
if (b >= n_batch || tid >= PHR_BLOCK) return;
// Warp-shuffle reduce: 1 partial per warp.
__shared__ float s_warp[PHR_N_WARPS];
for (int h = 0; h < PHR_N_HORIZONS; ++h) {
// Per-thread partial: w_res[h, tid] * context_h[b, h, tid].
float v = (tid < PHR_HIDDEN_DIM)
? w_res[h * PHR_HIDDEN_DIM + tid]
* context_h[(long long)b * PHR_N_HORIZONS * PHR_HIDDEN_DIM
+ (long long)h * PHR_HIDDEN_DIM + tid]
: 0.0f;
// Intra-warp reduction via shuffle (zero barriers).
for (int s = 16; s > 0; s >>= 1) v += __shfl_xor_sync(0xffffffff, v, s);
if ((tid & 31) == 0) s_warp[tid >> 5] = v;
__syncthreads();
// Cross-warp reduce: WHOLE warp 0 must call __shfl_xor_sync with
// mask=0xffffffff (any divergence within the warp is UB and hangs
// on Ampere/Ada). Inactive lanes contribute 0.
float w = (tid < PHR_N_WARPS) ? s_warp[tid] : 0.0f;
if (tid < 32) {
for (int s = PHR_N_WARPS / 2; s > 0; s >>= 1) {
w += __shfl_xor_sync(0xffffffff, w, s);
}
if (tid == 0) {
residual_out[(long long)b * PHR_N_HORIZONS + h] = w + bias_res[h];
}
}
__syncthreads();
}
}
// Backward:
// d_w_res[h, d] += Σ_b context_h[b, h, d] * grad_residual[b, h]
// d_bias_res[h] += Σ_b grad_residual[b, h]
// d_context_h[b,h,d] = w_res[h, d] * grad_residual[b, h]
//
// One block per batch; thread tid owns column d. Writes:
// d_context_h[b, h, d] — sole writer per (b, h, d). No race.
// d_w_res_scratch[b, h, d] — per-block scratch; host reduce_axis0 collapses
// across batches to recover the shared [N_HORIZONS, HIDDEN_DIM] grad.
// d_bias_res_scratch[b, h] — per-block scratch; same reduction pattern.
extern "C" __global__ void per_horizon_residual_head_bwd(
const float* __restrict__ context_h, // [B, N_HORIZONS, HIDDEN_DIM]
const float* __restrict__ w_res, // [N_HORIZONS, HIDDEN_DIM]
const float* __restrict__ grad_residual, // [B, N_HORIZONS]
int n_batch,
float* __restrict__ d_w_res_scratch, // [B, N_HORIZONS, HIDDEN_DIM] (+=)
float* __restrict__ d_bias_res_scratch, // [B, N_HORIZONS] (+=)
float* __restrict__ d_context_h // [B, N_HORIZONS, HIDDEN_DIM] (+= chained)
) {
int b = blockIdx.x;
int tid = threadIdx.x;
if (b >= n_batch || tid >= PHR_BLOCK) return;
for (int h = 0; h < PHR_N_HORIZONS; ++h) {
const float g = grad_residual[(long long)b * PHR_N_HORIZONS + h];
if (tid < PHR_HIDDEN_DIM) {
const long long ctx_idx = (long long)b * PHR_N_HORIZONS * PHR_HIDDEN_DIM
+ (long long)h * PHR_HIDDEN_DIM + tid;
const float ctx = context_h[ctx_idx];
const float w = w_res[h * PHR_HIDDEN_DIM + tid];
d_w_res_scratch[ctx_idx] += g * ctx;
d_context_h[ctx_idx] += g * w;
}
if (tid == 0) {
d_bias_res_scratch[(long long)b * PHR_N_HORIZONS + h] += g;
}
}
}

View File

@@ -30,8 +30,6 @@ pub mod data;
pub mod eval;
pub mod heads;
pub mod isv;
pub mod per_horizon_attention_pool;
pub mod per_horizon_residual_head;
pub mod pinned;
pub mod pinned_mem;
pub mod trainer;

View File

@@ -1,136 +0,0 @@
//! Per-horizon attention pool kernel host wrapper (C21).
//!
//! See `crates/ml-alpha/cuda/per_horizon_attention_pool.cu` for kernel
//! math + the design spec at
//! `docs/superpowers/specs/2026-05-18-per-horizon-attention-pool-design.md`.
//!
//! v1 of this binding: standalone forward + backward callable with
//! cudarc CudaSlices, used by the numgrad parity test in
//! `tests/per_horizon_attention_pool_numgrad.rs`. Wiring into
//! `PerceptionTrainer`'s captured graph is a follow-up commit gated on
//! the numgrad parity check passing here.
use anyhow::{Context, Result};
use cudarc::driver::{
CudaContext, CudaFunction, CudaModule, CudaSlice, CudaStream, LaunchConfig, PushKernelArg,
};
use std::sync::Arc;
pub const PHA_HIDDEN_DIM: usize = 128;
pub const PHA_BLOCK: usize = 128;
pub const PHA_N_HORIZONS: usize = 5;
/// Warp-shuffle reduce footprint inside the block: 1 partial per warp.
pub const PHA_N_WARPS: usize = PHA_BLOCK / 32;
const CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/per_horizon_attention_pool.cubin"));
pub struct PerHorizonAttentionPool {
_module: Arc<CudaModule>,
fwd_fn: CudaFunction,
bwd_fn: CudaFunction,
stream: Arc<CudaStream>,
}
impl PerHorizonAttentionPool {
pub fn new(ctx: &Arc<CudaContext>, stream: Arc<CudaStream>) -> Result<Self> {
let module = ctx
.load_cubin(CUBIN.to_vec())
.context("load per_horizon_attention_pool cubin")?;
let fwd_fn = module
.load_function("per_horizon_attention_pool_fwd")
.context("load per_horizon_attention_pool_fwd")?;
let bwd_fn = module
.load_function("per_horizon_attention_pool_bwd")
.context("load per_horizon_attention_pool_bwd")?;
Ok(Self {
_module: module,
fwd_fn,
bwd_fn,
stream,
})
}
/// Forward. Inputs/outputs are flat row-major CudaSlices:
/// q_h: [N_HORIZONS, HIDDEN_DIM]
/// ln_out: [B, K, HIDDEN_DIM]
/// context_h: [B, N_HORIZONS, HIDDEN_DIM] (written)
/// attn_h: [B, N_HORIZONS, K] (written, post-softmax)
pub fn forward(
&self,
q_h: &CudaSlice<f32>,
ln_out: &CudaSlice<f32>,
n_batch: i32,
k_seq: i32,
context_h: &mut CudaSlice<f32>,
attn_h: &mut CudaSlice<f32>,
) -> Result<()> {
// Smem layout: s_scores[K] + s_warp[N_WARPS]. Warp-shuffle reduce
// removed the prior [BLOCK] + [HIDDEN_DIM] regions.
let smem_bytes =
((k_seq as usize + PHA_N_WARPS) * std::mem::size_of::<f32>()) as u32;
let cfg = LaunchConfig {
grid_dim: (n_batch as u32, 1, 1),
block_dim: (PHA_BLOCK as u32, 1, 1),
shared_mem_bytes: smem_bytes,
};
let mut launch = self.stream.launch_builder(&self.fwd_fn);
unsafe {
launch
.arg(q_h)
.arg(ln_out)
.arg(&n_batch)
.arg(&k_seq)
.arg(context_h)
.arg(attn_h)
.launch(cfg)
.context("per_horizon_attention_pool_fwd")?;
}
// No synchronize: same-stream issue order is sufficient and
// synchronize is illegal during CUDA Graph capture.
Ok(())
}
/// Backward. Adds gradients into `grad_qh_scratch` + `grad_ln_out`
/// (both `+=`). `grad_qh_scratch` is per-block scratch of shape
/// `[B, N_HORIZONS, HIDDEN_DIM]`; reduce across the batch axis
/// host-side via `reduce_axis0` to recover the shared `grad_Q_h`.
pub fn backward(
&self,
q_h: &CudaSlice<f32>,
ln_out: &CudaSlice<f32>,
attn_h: &CudaSlice<f32>,
grad_context_h: &CudaSlice<f32>,
n_batch: i32,
k_seq: i32,
grad_qh_scratch: &mut CudaSlice<f32>,
grad_ln_out: &mut CudaSlice<f32>,
) -> Result<()> {
// Smem layout: s_dattn[K] + s_dscores[K] + s_warp[N_WARPS].
// Warp-shuffle reduce removed the prior [BLOCK] region.
let smem_bytes =
((2 * k_seq as usize + PHA_N_WARPS) * std::mem::size_of::<f32>()) as u32;
let cfg = LaunchConfig {
grid_dim: (n_batch as u32, 1, 1),
block_dim: (PHA_BLOCK as u32, 1, 1),
shared_mem_bytes: smem_bytes,
};
let mut launch = self.stream.launch_builder(&self.bwd_fn);
unsafe {
launch
.arg(q_h)
.arg(ln_out)
.arg(attn_h)
.arg(grad_context_h)
.arg(&n_batch)
.arg(&k_seq)
.arg(grad_qh_scratch)
.arg(grad_ln_out)
.launch(cfg)
.context("per_horizon_attention_pool_bwd")?;
}
// No synchronize: same-stream issue order is sufficient and
// synchronize is illegal during CUDA Graph capture.
Ok(())
}
}

View File

@@ -1,115 +0,0 @@
//! Per-horizon residual head kernel host wrapper (C22).
//!
//! See `crates/ml-alpha/cuda/per_horizon_residual_head.cu` for kernel
//! math. v1 of this binding: standalone forward + backward, validated
//! via numgrad parity test. Trainer integration (gating the residual
//! addition behind a learnable α-scalar) is a follow-up commit.
use anyhow::{Context, Result};
use cudarc::driver::{
CudaContext, CudaFunction, CudaModule, CudaSlice, CudaStream, LaunchConfig, PushKernelArg,
};
use std::sync::Arc;
pub const PHR_HIDDEN_DIM: usize = 128;
pub const PHR_BLOCK: usize = 128;
pub const PHR_N_HORIZONS: usize = 5;
const CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/per_horizon_residual_head.cubin"));
pub struct PerHorizonResidualHead {
_module: Arc<CudaModule>,
fwd_fn: CudaFunction,
bwd_fn: CudaFunction,
stream: Arc<CudaStream>,
}
impl PerHorizonResidualHead {
pub fn new(ctx: &Arc<CudaContext>, stream: Arc<CudaStream>) -> Result<Self> {
let module = ctx
.load_cubin(CUBIN.to_vec())
.context("load per_horizon_residual_head cubin")?;
let fwd_fn = module
.load_function("per_horizon_residual_head_fwd")
.context("load per_horizon_residual_head_fwd")?;
let bwd_fn = module
.load_function("per_horizon_residual_head_bwd")
.context("load per_horizon_residual_head_bwd")?;
Ok(Self {
_module: module,
fwd_fn,
bwd_fn,
stream,
})
}
/// Forward. context_h: [B, N_HORIZONS, HIDDEN_DIM].
/// w_res: [N_HORIZONS, HIDDEN_DIM]. bias_res: [N_HORIZONS].
/// residual_out: [B, N_HORIZONS] (written).
pub fn forward(
&self,
context_h: &CudaSlice<f32>,
w_res: &CudaSlice<f32>,
bias_res: &CudaSlice<f32>,
n_batch: i32,
residual_out: &mut CudaSlice<f32>,
) -> Result<()> {
let cfg = LaunchConfig {
grid_dim: (n_batch as u32, 1, 1),
block_dim: (PHR_BLOCK as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.fwd_fn);
unsafe {
launch
.arg(context_h)
.arg(w_res)
.arg(bias_res)
.arg(&n_batch)
.arg(residual_out)
.launch(cfg)
.context("per_horizon_residual_head_fwd")?;
}
// No synchronize: same-stream issue order is sufficient and
// synchronize is illegal during CUDA Graph capture.
Ok(())
}
/// Backward. d_w_res_scratch + d_bias_res_scratch are per-block
/// scratch tensors; reduce across batch with `reduce_axis0` to get
/// the shared `d_w_res[N_HORIZONS, HIDDEN_DIM]` + `d_bias[N_HORIZONS]`.
/// d_context_h is per-batch indexed; updates +=.
pub fn backward(
&self,
context_h: &CudaSlice<f32>,
w_res: &CudaSlice<f32>,
grad_residual: &CudaSlice<f32>,
n_batch: i32,
d_w_res_scratch: &mut CudaSlice<f32>,
d_bias_res_scratch: &mut CudaSlice<f32>,
d_context_h: &mut CudaSlice<f32>,
) -> Result<()> {
let cfg = LaunchConfig {
grid_dim: (n_batch as u32, 1, 1),
block_dim: (PHR_BLOCK as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.bwd_fn);
unsafe {
launch
.arg(context_h)
.arg(w_res)
.arg(grad_residual)
.arg(&n_batch)
.arg(d_w_res_scratch)
.arg(d_bias_res_scratch)
.arg(d_context_h)
.launch(cfg)
.context("per_horizon_residual_head_bwd")?;
}
// No synchronize: same-stream issue order is sufficient and
// synchronize is illegal during CUDA Graph capture.
Ok(())
}
}

View File

@@ -3,5 +3,4 @@
pub mod loss;
pub mod optim;
pub mod per_horizon_state;
pub mod perception;

View File

@@ -1,409 +0,0 @@
//! Per-horizon attention pool trainer state (C24).
//!
//! Bundles the device buffers, kernel bindings, gradient scratch, and
//! AdamW optimizer instances for the per-horizon attention-pool path
//! defined in C21 (`per_horizon_attention_pool`) + C22
//! (`per_horizon_residual_head`). Lives on PerceptionTrainer as one
//! cohesive field instead of fifteen scattered ones.
//!
//! The α-gate is **zero-initialised**, which (per the C23
//! `alpha_zero_init_is_identity_to_baseline` invariant) makes this
//! state's contribution to per-batch logits bit-identical to baseline
//! at the start of training. AdamW updates may push α away from 0 if
//! the residual signal correlates with loss gradient.
use anyhow::{Context, Result};
use cudarc::driver::{
CudaFunction, CudaModule, CudaSlice, CudaStream, LaunchConfig, PushKernelArg,
};
use ml_core::device::MlDevice;
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
use std::sync::Arc;
use crate::heads::{HIDDEN_DIM, N_HORIZONS};
use crate::per_horizon_attention_pool::PerHorizonAttentionPool;
use crate::per_horizon_residual_head::PerHorizonResidualHead;
use crate::trainer::optim::AdamW;
const PROB_BLEND_CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/per_horizon_prob_blend.cubin"));
const REDUCE_AXIS0_CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/reduce_axis0.cubin"));
pub struct PerHorizonTrainState {
// Kernel bindings.
pub pool: PerHorizonAttentionPool,
pub head: PerHorizonResidualHead,
// Learnable parameters.
pub q_h_d: CudaSlice<f32>, // [N_HORIZONS, HIDDEN_DIM] — attention queries
pub w_res_d: CudaSlice<f32>, // [N_HORIZONS, HIDDEN_DIM] — residual head weights
pub bias_res_d: CudaSlice<f32>, // [N_HORIZONS]
pub alpha_d: CudaSlice<f32>, // [N_HORIZONS] — learnable gate scalars (init 0)
// Forward intermediates (per-batch).
pub context_d: CudaSlice<f32>, // [B, N_HORIZONS, HIDDEN_DIM]
pub attn_weights_d: CudaSlice<f32>, // [B, N_HORIZONS, K]
pub residual_d: CudaSlice<f32>, // [B, N_HORIZONS]
// Backward grads — per-batch scratch reduced to shared via reduce_axis0.
pub grad_q_h_scratch_d: CudaSlice<f32>, // [B, N_HORIZONS, HIDDEN_DIM]
pub grad_w_res_scratch_d: CudaSlice<f32>, // [B, N_HORIZONS, HIDDEN_DIM]
pub grad_bias_res_scratch_d: CudaSlice<f32>, // [B, N_HORIZONS]
pub grad_alpha_d: CudaSlice<f32>, // [N_HORIZONS] — host-reduced
pub grad_context_d: CudaSlice<f32>, // [B, N_HORIZONS, HIDDEN_DIM]
pub grad_residual_d: CudaSlice<f32>, // [B, N_HORIZONS]
// Reduced (shared) grad buffers — what AdamW reads.
pub grad_q_h_d: CudaSlice<f32>, // [N_HORIZONS, HIDDEN_DIM]
pub grad_w_res_d: CudaSlice<f32>, // [N_HORIZONS, HIDDEN_DIM]
pub grad_bias_res_d: CudaSlice<f32>, // [N_HORIZONS]
// Optimizers.
pub opt_q_h: AdamW,
pub opt_w_res: AdamW,
pub opt_bias_res: AdamW,
pub opt_alpha: AdamW,
// Dimensions captured at construction.
pub n_batch: usize,
pub k_seq: usize,
stream: Arc<CudaStream>,
// C25 prob-blend kernel handles + per-K scratch (baseline probs stash +
// ∂L/∂r_contrib per K position).
_prob_blend_module: Arc<CudaModule>,
prob_blend_fwd_fn: CudaFunction,
prob_blend_reduce_fn: CudaFunction,
/// d_α [N_HORIZONS] — what AdamW updates for α. Written by the
/// reduce kernel after BCE backward has populated grad_probs_per_k.
pub d_alpha_reduced_d: CudaSlice<f32>,
// reduce_axis0 GPU kernel — collapses [B, n_tail] per-batch
// gradient scratch into shared [n_tail] grads. Must be GPU
// (capture-safe); host-side reduction is forbidden during graph
// capture per pearl_no_host_branches_in_captured_graph.
_reduce_axis0_module: Arc<CudaModule>,
reduce_axis0_fn: CudaFunction,
}
impl PerHorizonTrainState {
/// Construct + initialise.
/// `lr` matches the existing attn_q AdamW LR. `seed` drives Xavier
/// init for Q_h and small init for w_res; alpha + bias stay at 0.
pub fn new(dev: &MlDevice, n_batch: usize, k_seq: usize, lr: f32, seed: u64) -> Result<Self> {
anyhow::ensure!(n_batch > 0 && k_seq > 0, "n_batch and k_seq must be > 0");
let stream = dev.cuda_stream().context("phts stream")?.clone();
let ctx = dev.cuda_context().context("phts ctx")?;
let pool = PerHorizonAttentionPool::new(ctx, stream.clone())
.context("PerHorizonAttentionPool")?;
let head = PerHorizonResidualHead::new(ctx, stream.clone())
.context("PerHorizonResidualHead")?;
let prob_blend_module = ctx
.load_cubin(PROB_BLEND_CUBIN.to_vec())
.context("load per_horizon_prob_blend cubin")?;
let prob_blend_fwd_fn = prob_blend_module
.load_function("per_horizon_prob_blend_fwd")
.context("per_horizon_prob_blend_fwd")?;
let prob_blend_reduce_fn = prob_blend_module
.load_function("per_horizon_prob_blend_reduce_alpha_residual")
.context("per_horizon_prob_blend_reduce_alpha_residual")?;
let reduce_axis0_module = ctx
.load_cubin(REDUCE_AXIS0_CUBIN.to_vec())
.context("load reduce_axis0 cubin")?;
let reduce_axis0_fn = reduce_axis0_module
.load_function("reduce_axis0")
.context("reduce_axis0 symbol")?;
let mut rng = ChaCha8Rng::seed_from_u64(seed);
// Q_h: Xavier-style init at 1/sqrt(HIDDEN_DIM). Per the spec
// §5 "Q_h initialisation" decision — start with the same scale
// as the existing Q to give every horizon comparable variance
// at step 0.
let qh_scale = (1.0_f32 / HIDDEN_DIM as f32).sqrt();
let q_h_init: Vec<f32> = (0..N_HORIZONS * HIDDEN_DIM)
.map(|_| rng.gen_range(-qh_scale..qh_scale))
.collect();
let q_h_d = upload(&stream, &q_h_init)?;
// w_res: small init so the residual contribution starts subtle.
// The α-gate already zeros the contribution at step 0, so this
// is really a "post-gate-opening starting magnitude" knob.
let w_res_scale = 0.1_f32 * qh_scale;
let w_res_init: Vec<f32> = (0..N_HORIZONS * HIDDEN_DIM)
.map(|_| rng.gen_range(-w_res_scale..w_res_scale))
.collect();
let w_res_d = upload(&stream, &w_res_init)?;
let bias_res_d = stream.alloc_zeros::<f32>(N_HORIZONS).context("bias_res alloc")?;
// α = 0 ⇒ tanh(0) = 0 ⇒ residual contribution = 0 at step 0.
// Per C23 invariant: bit-identical to baseline until training
// moves α away from 0.
let alpha_d = stream.alloc_zeros::<f32>(N_HORIZONS).context("alpha alloc")?;
let context_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS * HIDDEN_DIM)
.context("context alloc")?;
let attn_weights_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS * k_seq)
.context("attn alloc")?;
let residual_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS)
.context("residual alloc")?;
let grad_q_h_scratch_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS * HIDDEN_DIM)?;
let grad_w_res_scratch_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS * HIDDEN_DIM)?;
let grad_bias_res_scratch_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS)?;
let grad_alpha_d = stream.alloc_zeros::<f32>(N_HORIZONS)?;
let grad_context_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS * HIDDEN_DIM)?;
let grad_residual_d = stream.alloc_zeros::<f32>(n_batch * N_HORIZONS)?;
let grad_q_h_d = stream.alloc_zeros::<f32>(N_HORIZONS * HIDDEN_DIM)?;
let grad_w_res_d = stream.alloc_zeros::<f32>(N_HORIZONS * HIDDEN_DIM)?;
let grad_bias_res_d = stream.alloc_zeros::<f32>(N_HORIZONS)?;
let opt_q_h = AdamW::new(dev, N_HORIZONS * HIDDEN_DIM, lr).context("opt_q_h")?;
let opt_w_res = AdamW::new(dev, N_HORIZONS * HIDDEN_DIM, lr).context("opt_w_res")?;
let opt_bias_res = AdamW::new(dev, N_HORIZONS, lr).context("opt_bias_res")?;
// α gets a slightly lower LR — it controls gate sensitivity and
// we want it to ramp slowly. 1/4 of the param LR is a defensible
// default; can be tuned per spec §5 open question 2.
let opt_alpha = AdamW::new(dev, N_HORIZONS, lr * 0.25).context("opt_alpha")?;
let d_alpha_reduced_d = stream
.alloc_zeros::<f32>(N_HORIZONS)
.context("d_alpha_reduced alloc")?;
Ok(Self {
pool, head,
q_h_d, w_res_d, bias_res_d, alpha_d,
context_d, attn_weights_d, residual_d,
grad_q_h_scratch_d, grad_w_res_scratch_d, grad_bias_res_scratch_d,
grad_alpha_d, grad_context_d, grad_residual_d,
grad_q_h_d, grad_w_res_d, grad_bias_res_d,
opt_q_h, opt_w_res, opt_bias_res, opt_alpha,
n_batch, k_seq,
stream,
_prob_blend_module: prob_blend_module,
prob_blend_fwd_fn,
prob_blend_reduce_fn,
d_alpha_reduced_d,
_reduce_axis0_module: reduce_axis0_module,
reduce_axis0_fn,
})
}
/// C25 forward: per-horizon attention pool → residual head → in-place
/// logit-bias rewrite of `probs_per_k` using `logit_per_k` as the
/// baseline (saved by the GRN forward).
/// At α=0 (init), r_contrib = 0 → probs unchanged → bit-identical
/// to baseline (extension of the C23 identity invariant into the
/// trainer hot loop).
///
/// `ln_b_out` is the perception trainer's LN_b output `[B, K, HIDDEN_DIM]`
/// (same tensor that feeds the existing single-Q attention pool).
/// `logit_per_k` is `[K, B, N_HORIZONS]` from GRN forward.
/// `probs_per_k` is `[K, B, N_HORIZONS]` — overwritten in-place
/// with `sigmoid(logit_per_k + tanh(α) * residual)`.
pub fn forward_with_blend(
&mut self,
ln_b_out: &CudaSlice<f32>,
logit_per_k: &CudaSlice<f32>,
probs_per_k: &mut CudaSlice<f32>,
) -> Result<()> {
let n = self.n_batch as i32;
let k = self.k_seq as i32;
// 1. Per-horizon attention pool → context_h, attn_weights.
self.pool.forward(&self.q_h_d, ln_b_out, n, k,
&mut self.context_d, &mut self.attn_weights_d)?;
// 2. Residual head → residual[B, N_HORIZONS].
self.head.forward(&self.context_d, &self.w_res_d, &self.bias_res_d, n,
&mut self.residual_d)?;
// 3. In-place logit-bias rewrite: probs_per_k ← sigmoid(logit_per_k + r).
let cfg = LaunchConfig {
grid_dim: (self.n_batch as u32, 1, 1),
block_dim: (N_HORIZONS as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.prob_blend_fwd_fn);
unsafe {
launch
.arg(&self.alpha_d)
.arg(&self.residual_d)
.arg(logit_per_k)
.arg(probs_per_k)
.arg(&n)
.arg(&k)
.launch(cfg)
.context("per_horizon_prob_blend_fwd")?;
}
// No synchronize: same-stream issue order is sufficient and
// synchronize is illegal during CUDA Graph capture.
Ok(())
}
/// C25 backward. Called AFTER BCE has populated `grad_probs_per_k`
/// (= ∂L/∂p_final).
///
/// Key insight: because the bias is additive INSIDE the sigmoid,
/// ∂L/∂logit_baseline = ∂L/∂r_contrib = ∂L/∂p_final * p_final * (1 - p_final).
/// The existing GRN backward already computes the LHS via the same
/// formula (using probs_per_k = p_final), so the GRN backward path
/// is UNTOUCHED. We just read the same gradient ourselves and reduce
/// to get d_residual + d_alpha.
///
/// `probs_per_k` is the trainer's probs buffer (post-blend, == p_final).
/// `grad_probs_per_k` is from BCE backward.
/// `ln_b_out` + `grad_ln_b_out` are LN_b's fwd output and accumulating
/// gradient — the attention pool bwd `+=`'s onto grad_ln_b_out.
pub fn backward_through_blend(
&mut self,
probs_per_k: &CudaSlice<f32>,
grad_probs_per_k: &CudaSlice<f32>,
ln_b_out: &CudaSlice<f32>,
grad_ln_b_out: &mut CudaSlice<f32>,
) -> Result<()> {
let n = self.n_batch as i32;
let k = self.k_seq as i32;
// 1. Reduce: compute d_logit_baseline = grad_probs * p * (1-p)
// inline + reduce over k for d_residual, over k+b for d_alpha.
// One warp per horizon: 32 threads cooperatively reduce via shuffle.
let cfg_per_h = LaunchConfig {
grid_dim: (N_HORIZONS as u32, 1, 1),
block_dim: (32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.prob_blend_reduce_fn);
unsafe {
launch
.arg(&self.alpha_d)
.arg(&self.residual_d)
.arg(probs_per_k)
.arg(grad_probs_per_k)
.arg(&n)
.arg(&k)
.arg(&mut self.grad_residual_d)
.arg(&mut self.d_alpha_reduced_d)
.launch(cfg_per_h)
.context("per_horizon_prob_blend_reduce_alpha_residual")?;
}
// No synchronize: same-stream issue order suffices and
// synchronize is illegal during CUDA Graph capture.
// 2. Residual-head bwd: consumes d_residual → produces
// d_w_res_scratch, d_bias_res_scratch, d_context.
self.head.backward(
&self.context_d, &self.w_res_d, &self.grad_residual_d, n,
&mut self.grad_w_res_scratch_d,
&mut self.grad_bias_res_scratch_d,
&mut self.grad_context_d,
)?;
// 3. Attention-pool bwd: consumes d_context → produces
// d_q_h_scratch + adds to grad_ln_b_out.
self.pool.backward(
&self.q_h_d, ln_b_out, &self.attn_weights_d, &self.grad_context_d,
n, k,
&mut self.grad_q_h_scratch_d, grad_ln_b_out,
)?;
// 4. Reduce per-batch scratches to shared grads.
self.reduce_per_batch_scratches_to_shared()?;
Ok(())
}
/// Reduce per-batch scratches into shared grad buffers using the
/// `reduce_axis0` GPU kernel. Capture-safe: no host allocations,
/// no D↔H copies, no host-side arithmetic. One kernel launch per
/// scratch (3 total): grad_q_h, grad_w_res, grad_bias_res.
fn reduce_per_batch_scratches_to_shared(&mut self) -> Result<()> {
let n_qh = (N_HORIZONS * HIDDEN_DIM) as i32;
let n_bres = N_HORIZONS as i32;
let n_batch_i = self.n_batch as i32;
// grad_q_h_scratch [B, n_qh] → grad_q_h_d [n_qh]
{
let cfg = LaunchConfig {
grid_dim: (n_qh as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn);
launch
.arg(&self.grad_q_h_scratch_d)
.arg(&n_batch_i)
.arg(&n_qh)
.arg(&mut self.grad_q_h_d);
unsafe { launch.launch(cfg).context("reduce_axis0 grad_q_h")?; }
}
// grad_w_res_scratch [B, n_qh] → grad_w_res_d [n_qh]
{
let cfg = LaunchConfig {
grid_dim: (n_qh as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn);
launch
.arg(&self.grad_w_res_scratch_d)
.arg(&n_batch_i)
.arg(&n_qh)
.arg(&mut self.grad_w_res_d);
unsafe { launch.launch(cfg).context("reduce_axis0 grad_w_res")?; }
}
// grad_bias_res_scratch [B, n_bres] → grad_bias_res_d [n_bres]
{
let cfg = LaunchConfig {
grid_dim: (n_bres as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn);
launch
.arg(&self.grad_bias_res_scratch_d)
.arg(&n_batch_i)
.arg(&n_bres)
.arg(&mut self.grad_bias_res_d);
unsafe { launch.launch(cfg).context("reduce_axis0 grad_bias_res")?; }
}
Ok(())
}
/// AdamW step for all four per-horizon param groups. Called after
/// `backward_through_blend` populates the shared grad buffers.
pub fn adamw_step(&mut self) -> Result<()> {
self.opt_q_h.step(&mut self.q_h_d, &self.grad_q_h_d)?;
self.opt_w_res.step(&mut self.w_res_d, &self.grad_w_res_d)?;
self.opt_bias_res.step(&mut self.bias_res_d, &self.grad_bias_res_d)?;
self.opt_alpha.step(&mut self.alpha_d, &self.d_alpha_reduced_d)?;
Ok(())
}
/// Zero all gradient scratch buffers between training steps.
/// Uses device-side `memset_zeros` (capture-safe). Host-zero +
/// memcpy_htod is forbidden during CUDA Graph capture.
pub fn zero_grads(&mut self) -> Result<()> {
let stream = &self.stream;
stream.memset_zeros(&mut self.grad_q_h_scratch_d)?;
stream.memset_zeros(&mut self.grad_w_res_scratch_d)?;
stream.memset_zeros(&mut self.grad_bias_res_scratch_d)?;
stream.memset_zeros(&mut self.grad_alpha_d)?;
stream.memset_zeros(&mut self.grad_context_d)?;
stream.memset_zeros(&mut self.grad_residual_d)?;
stream.memset_zeros(&mut self.grad_q_h_d)?;
stream.memset_zeros(&mut self.grad_w_res_d)?;
stream.memset_zeros(&mut self.grad_bias_res_d)?;
stream.memset_zeros(&mut self.d_alpha_reduced_d)?;
Ok(())
}
}
fn upload(stream: &std::sync::Arc<cudarc::driver::CudaStream>, host: &[f32]) -> Result<CudaSlice<f32>> {
let mut buf = stream.alloc_zeros::<f32>(host.len()).context("upload alloc")?;
stream.memcpy_htod(host, &mut buf).context("upload memcpy")?;
Ok(buf)
}

View File

@@ -375,14 +375,6 @@ pub struct PerceptionTrainer {
attn_bwd_fn: CudaFunction,
_attn_module: Arc<CudaModule>,
/// Per-horizon attention pool + residual head trainer state (C24).
/// Lives behind a single field rather than scattered alongside the
/// other attention pool fields above; α-gate is zero-initialised so
/// the contribution is bit-identical to baseline at step 0 (proven
/// by the C23 alpha_zero_init_is_identity_to_baseline test).
/// Forward/backward integration into step_batched is C25.
pub per_horizon: crate::trainer::per_horizon_state::PerHorizonTrainState,
// ── K-loop parallelization (Phase B) ──
// Per-batch grad scratch buffers for cfc_step_backward_batched.
// Zeroed once per training step; the K-loop's 64 bwd calls accumulate
@@ -811,17 +803,13 @@ impl PerceptionTrainer {
// Phase B: attn pool per-batch grad scratch.
let attn_grad_q_scratch_d = stream.alloc_zeros::<f32>(cfg.n_batch * HIDDEN_DIM)?;
// C24: per-horizon attention pool trainer state. α-init=0 makes
// this contribution bit-identical to baseline at step 0 (proven
// by C23 alpha_zero_init_is_identity_to_baseline). Forward/
// backward wiring into step_batched is C25.
let per_horizon = crate::trainer::per_horizon_state::PerHorizonTrainState::new(
dev,
cfg.n_batch,
cfg.seq_len,
cfg.lr_cfc,
cfg.seed.wrapping_add(0xA110_C00A),
)?;
// V1 (2026-05-18): per-horizon Q_h trainer state (C24/C25) removed.
// A/B sweep at commit 83546b5c3 falsified the approach (mean_auc
// -0.019 vs baseline; h6000 essentially tied). v2 design with
// horizon-token K-prepend + inverted attention + regime-MoE +
// Kendall σ-BCE + L2 anchor replaces it. See
// docs/superpowers/specs/2026-05-18-ml-alpha-v2-multi-horizon-design.md
// and the V1-V13 plan.
let k = cfg.seq_len;
Ok(Self {
@@ -878,8 +866,6 @@ impl PerceptionTrainer {
attn_fwd_fn,
attn_bwd_fn,
_attn_module: attn_module,
// C24: per-horizon attention pool trainer state.
per_horizon,
// Phase B: cfc per-batch grad scratch + reducer.
cfc_grad_w_in_scratch_d,
cfc_grad_w_rec_scratch_d,
@@ -1665,20 +1651,9 @@ impl PerceptionTrainer {
}
drop((_g_hpk, _g_probs, _g_z1, _g_a1, _g_z2, _g_gate, _g_main, _g_logit));
// ── 4.5. C25 per-horizon attention pool forward + prob-blend.
// Reads ln_out_d + logit_per_k_d (baseline GRN logits)
// and overwrites probs_per_k_d in place with
// sigmoid(logit_baseline + tanh(α) * residual).
// At α=0 (init): r_contrib=0, probs unchanged → baseline
// training is bit-identical to pre-C25 (per the C23
// alpha_zero_init_is_identity invariant extended into
// the trainer hot loop).
self.per_horizon.zero_grads()?;
self.per_horizon.forward_with_blend(
&self.ln_out_d,
&self.logit_per_k_d,
&mut self.probs_per_k_d,
)?;
// ── 4.5. (v2 reserved) — per-horizon Q_h path removed at V1.
// The v2 design (horizon-token K-prepend + inverted attn +
// regime-MoE) wires in here in commits V9/V10.
// ── 5. Fused multi-horizon BCE over the full [K*B, N_HORIZONS]
// grid. The kernel doesn't distinguish position-vs-batch
@@ -1703,22 +1678,8 @@ impl PerceptionTrainer {
unsafe { launch.launch(bce_cfg).context("bce launch")?; }
}
// ── 5a. C25 per-horizon attention pool backward.
// Reads probs_per_k_d (now post-blend, = p_final) +
// grad_probs_per_k_d (= ∂L/∂p_final from BCE) and:
// 1. Computes d_residual + d_α via per-horizon reduce
// (chain rule equivalence makes the existing GRN
// backward unmodified — see kernel header).
// 2. Runs residual_head bwd → d_w_res, d_bias_res, d_context.
// 3. Runs attention_pool bwd → d_q_h, += grad_h_enriched_seq_d.
// AdamW updates for the four per-horizon param groups
// land in section 9 alongside the existing 17 groups.
self.per_horizon.backward_through_blend(
&self.probs_per_k_d,
&self.grad_probs_per_k_d,
&self.ln_out_d,
self.grad_h_enriched_seq_d.data_mut(),
)?;
// ── 5a. (v2 reserved) — per-horizon Q_h backward removed at V1.
// v2 backward path will live here in commit V10.
// ── 5b. ISV-driven per-horizon EMA + lambda. Updates the EMA
// of unweighted per-horizon BCE and emits a clamped
@@ -2233,11 +2194,10 @@ impl PerceptionTrainer {
self.opt_vsn_b.step(&mut self.vsn_b_d, &self.grad_vsn_b_d)?;
self.opt_attn_q.step(&mut self.attn_q_d, &self.grad_attn_q_d)?;
// C25: per-horizon attention pool AdamW step (4 param groups:
// q_h, w_res, bias_res, α). At α=0 init + with backward-flow
// proving residual signal correlates with loss gradient, α
// will move away from 0 if the per-horizon path is useful.
self.per_horizon.adamw_step()?;
// (v2 reserved) — per-horizon AdamW step removed at V1; v2's six
// optimizer groups (horizon_tokens, Q_inv, w_fuse + b_fuse,
// moe_gate, experts, log_sigma) will land in commit V10.
// GPU-resident grad-clip + AdamW (zero host roundtrips).
// Replaces step_from_buffers which did 9× memcpy_dtoh per step.
self.mamba2_adamw

View File

@@ -1,152 +0,0 @@
//! Numerical-gradient parity check for the per-horizon attention pool
//! kernel (C21). Same approach as the VSN numgrad test
//! (Phase 2D.2 / task #191): compute analytical gradients via the
//! backward kernel, then perturb each input element by ±eps, run the
//! forward kernel twice, and verify the central-difference matches
//! within tolerance.
//!
//! Two grads are checked:
//! d_Q_h [N_HORIZONS, HIDDEN_DIM] — query weights
//! d_LNb [B, K, HIDDEN_DIM] — input gradient
//!
//! Loss function is the sum over context_h (so d_context = 1 for every
//! output element), which gives a simple analytical reduction.
use anyhow::Result;
use cudarc::driver::CudaSlice;
use ml_alpha::per_horizon_attention_pool::{PerHorizonAttentionPool, PHA_HIDDEN_DIM, PHA_N_HORIZONS};
use ml_core::device::MlDevice;
use rand::SeedableRng;
use rand::Rng;
use rand_chacha::ChaCha8Rng;
const B: usize = 2;
const K: usize = 8; // small K → fast numgrad
fn try_dev() -> Option<MlDevice> {
match MlDevice::cuda(0) {
Ok(d) => Some(d),
Err(e) => {
eprintln!("skipping: cuda device unavailable ({e})");
None
}
}
}
fn alloc_and_upload(stream: &std::sync::Arc<cudarc::driver::CudaStream>, host: &[f32]) -> CudaSlice<f32> {
let mut buf = stream.alloc_zeros::<f32>(host.len()).expect("alloc");
stream.memcpy_htod(host, &mut buf).expect("htod");
buf
}
fn download(stream: &std::sync::Arc<cudarc::driver::CudaStream>, src: &CudaSlice<f32>) -> Vec<f32> {
let mut out = vec![0.0f32; src.len()];
stream.memcpy_dtoh(src, out.as_mut_slice()).expect("dtoh");
out
}
fn run_forward_loss(
pool: &PerHorizonAttentionPool,
stream: &std::sync::Arc<cudarc::driver::CudaStream>,
q_h_host: &[f32],
ln_host: &[f32],
) -> f32 {
let q_h = alloc_and_upload(stream, q_h_host);
let ln = alloc_and_upload(stream, ln_host);
let mut ctx = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM).unwrap();
let mut attn = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * K).unwrap();
pool.forward(&q_h, &ln, B as i32, K as i32, &mut ctx, &mut attn).unwrap();
let ctx_host = download(stream, &ctx);
ctx_host.iter().sum()
}
#[test]
#[ignore = "requires CUDA"]
fn forward_then_backward_matches_central_difference() -> Result<()> {
let Some(dev) = try_dev() else { return Ok(()); };
let ctx = dev.cuda_context()?.clone();
let stream = dev.cuda_stream()?.clone();
let pool = PerHorizonAttentionPool::new(&ctx, stream.clone())?;
let mut rng = ChaCha8Rng::seed_from_u64(0xCAFE_F00D);
let q_h_host: Vec<f32> = (0..PHA_N_HORIZONS * PHA_HIDDEN_DIM)
.map(|_| rng.gen_range(-0.1..0.1)).collect();
let ln_host: Vec<f32> = (0..B * K * PHA_HIDDEN_DIM)
.map(|_| rng.gen_range(-0.5..0.5)).collect();
// Analytical: backward with grad_context = 1 (since loss = Σ context).
let q_h_d = alloc_and_upload(&stream, &q_h_host);
let ln_d = alloc_and_upload(&stream, &ln_host);
let mut ctx_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
let mut attn_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * K)?;
pool.forward(&q_h_d, &ln_d, B as i32, K as i32, &mut ctx_d, &mut attn_d)?;
let grad_ctx_host = vec![1.0f32; B * PHA_N_HORIZONS * PHA_HIDDEN_DIM];
let grad_ctx_d = alloc_and_upload(&stream, &grad_ctx_host);
let mut grad_qh_scratch_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
let mut grad_ln_d = stream.alloc_zeros::<f32>(B * K * PHA_HIDDEN_DIM)?;
pool.backward(
&q_h_d, &ln_d, &attn_d, &grad_ctx_d,
B as i32, K as i32,
&mut grad_qh_scratch_d, &mut grad_ln_d,
)?;
let grad_qh_scratch = download(&stream, &grad_qh_scratch_d);
let grad_ln = download(&stream, &grad_ln_d);
// Reduce grad_qh_scratch over batch axis → [N_HORIZONS, HIDDEN_DIM].
let mut grad_qh_analytical = vec![0.0f32; PHA_N_HORIZONS * PHA_HIDDEN_DIM];
for b in 0..B {
for i in 0..PHA_N_HORIZONS * PHA_HIDDEN_DIM {
grad_qh_analytical[i] += grad_qh_scratch[b * PHA_N_HORIZONS * PHA_HIDDEN_DIM + i];
}
}
let eps = 1e-2f32;
let tol = 5e-2f32; // 5% rel-tol, plus abs floor 5e-3 for tiny grads
// 1) Probe d_Q_h at a few random indices.
let q_h_idxs: Vec<usize> = (0..8)
.map(|_| rng.gen_range(0..PHA_N_HORIZONS * PHA_HIDDEN_DIM))
.collect();
for &idx in &q_h_idxs {
let mut q_h_plus = q_h_host.clone();
q_h_plus[idx] += eps;
let mut q_h_minus = q_h_host.clone();
q_h_minus[idx] -= eps;
let l_plus = run_forward_loss(&pool, &stream, &q_h_plus, &ln_host);
let l_minus = run_forward_loss(&pool, &stream, &q_h_minus, &ln_host);
let fd = (l_plus - l_minus) / (2.0 * eps);
let analytical = grad_qh_analytical[idx];
let abs_err = (fd - analytical).abs();
let rel_err = abs_err / fd.abs().max(1e-6);
let pass = abs_err < 5e-3 || rel_err < tol;
assert!(
pass,
"d_Q_h[{idx}]: analytical={analytical:.6} fd={fd:.6} abs_err={abs_err:.6} rel_err={rel_err:.4}"
);
}
// 2) Probe d_LNb at a few random indices.
let ln_idxs: Vec<usize> = (0..8)
.map(|_| rng.gen_range(0..B * K * PHA_HIDDEN_DIM))
.collect();
for &idx in &ln_idxs {
let mut ln_plus = ln_host.clone();
ln_plus[idx] += eps;
let mut ln_minus = ln_host.clone();
ln_minus[idx] -= eps;
let l_plus = run_forward_loss(&pool, &stream, &q_h_host, &ln_plus);
let l_minus = run_forward_loss(&pool, &stream, &q_h_host, &ln_minus);
let fd = (l_plus - l_minus) / (2.0 * eps);
let analytical = grad_ln[idx];
let abs_err = (fd - analytical).abs();
let rel_err = abs_err / fd.abs().max(1e-6);
let pass = abs_err < 5e-3 || rel_err < tol;
assert!(
pass,
"d_LNb[{idx}]: analytical={analytical:.6} fd={fd:.6} abs_err={abs_err:.6} rel_err={rel_err:.4}"
);
}
Ok(())
}

View File

@@ -1,201 +0,0 @@
//! End-to-end smoke test for the per-horizon attention pool +
//! residual head + α-gate composition (C23).
//!
//! Pipeline:
//! LNb [B, K, HIDDEN_DIM]
//! → per_horizon_attention_pool_fwd → context_h [B, N_HORIZONS, HIDDEN_DIM]
//! → per_horizon_residual_head_fwd → residual [B, N_HORIZONS]
//! → final_logit[h] = baseline_logit[h] + tanh(α[h]) * residual[h]
//!
//! Verifies:
//! 1. α = 0 → final_logit == baseline_logit (bit-equal) — proves the
//! additive integration is identity at zero-init, so adopting this
//! path can't regress the existing baseline.
//! 2. α ≠ 0 → final_logit ≠ baseline_logit AND no NaN/Inf leaks.
//! 3. Backward: residual_head_bwd's d_context_h feeds the attention
//! pool bwd; full gradient chain finite + non-zero where expected.
use anyhow::Result;
use cudarc::driver::CudaSlice;
use ml_alpha::per_horizon_attention_pool::{
PerHorizonAttentionPool, PHA_HIDDEN_DIM, PHA_N_HORIZONS,
};
use ml_alpha::per_horizon_residual_head::PerHorizonResidualHead;
use ml_core::device::MlDevice;
use rand::Rng;
use rand::SeedableRng;
use rand_chacha::ChaCha8Rng;
const B: usize = 4;
const K: usize = 16;
fn try_dev() -> Option<MlDevice> {
match MlDevice::cuda(0) {
Ok(d) => Some(d),
Err(e) => {
eprintln!("skipping: cuda device unavailable ({e})");
None
}
}
}
fn alloc_upload(
stream: &std::sync::Arc<cudarc::driver::CudaStream>,
host: &[f32],
) -> CudaSlice<f32> {
let mut buf = stream.alloc_zeros::<f32>(host.len()).expect("alloc");
stream.memcpy_htod(host, &mut buf).expect("htod");
buf
}
fn download(stream: &std::sync::Arc<cudarc::driver::CudaStream>, src: &CudaSlice<f32>) -> Vec<f32> {
let mut out = vec![0.0f32; src.len()];
stream.memcpy_dtoh(src, out.as_mut_slice()).expect("dtoh");
out
}
#[test]
#[ignore = "requires CUDA"]
fn alpha_zero_init_is_identity_to_baseline() -> Result<()> {
let Some(dev) = try_dev() else { return Ok(()); };
let ctx = dev.cuda_context()?.clone();
let stream = dev.cuda_stream()?.clone();
let pool = PerHorizonAttentionPool::new(&ctx, stream.clone())?;
let head = PerHorizonResidualHead::new(&ctx, stream.clone())?;
let mut rng = ChaCha8Rng::seed_from_u64(0xA11A_0BEEF);
let ln_host: Vec<f32> = (0..B * K * PHA_HIDDEN_DIM).map(|_| rng.gen_range(-0.5..0.5)).collect();
let q_h_host: Vec<f32> = (0..PHA_N_HORIZONS * PHA_HIDDEN_DIM).map(|_| rng.gen_range(-0.1..0.1)).collect();
let w_res_host: Vec<f32> = (0..PHA_N_HORIZONS * PHA_HIDDEN_DIM).map(|_| rng.gen_range(-0.1..0.1)).collect();
let bias_res_host: Vec<f32> = (0..PHA_N_HORIZONS).map(|_| rng.gen_range(-0.05..0.05)).collect();
let baseline_logit_host: Vec<f32> = (0..B * PHA_N_HORIZONS).map(|_| rng.gen_range(-2.0..2.0)).collect();
// α = 0 — final logit MUST equal baseline.
let alpha_zero = vec![0.0f32; PHA_N_HORIZONS];
let ln_d = alloc_upload(&stream, &ln_host);
let q_h_d = alloc_upload(&stream, &q_h_host);
let w_res_d = alloc_upload(&stream, &w_res_host);
let b_res_d = alloc_upload(&stream, &bias_res_host);
let mut ctx_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
let mut attn_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * K)?;
pool.forward(&q_h_d, &ln_d, B as i32, K as i32, &mut ctx_d, &mut attn_d)?;
let mut residual_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS)?;
head.forward(&ctx_d, &w_res_d, &b_res_d, B as i32, &mut residual_d)?;
// Host-side gate combination: final[b,h] = baseline[b,h] + tanh(α[h]) * residual[b,h].
let residual_host = download(&stream, &residual_d);
let mut final_logit = vec![0.0f32; B * PHA_N_HORIZONS];
for b in 0..B {
for h in 0..PHA_N_HORIZONS {
let idx = b * PHA_N_HORIZONS + h;
let gate = alpha_zero[h].tanh();
final_logit[idx] = baseline_logit_host[idx] + gate * residual_host[idx];
}
}
for i in 0..B * PHA_N_HORIZONS {
assert_eq!(
final_logit[i].to_bits(),
baseline_logit_host[i].to_bits(),
"alpha=0 at idx {i}: final {} ≠ baseline {}",
final_logit[i], baseline_logit_host[i]
);
}
Ok(())
}
#[test]
#[ignore = "requires CUDA"]
fn alpha_nonzero_changes_output_and_grads_flow_end_to_end() -> Result<()> {
let Some(dev) = try_dev() else { return Ok(()); };
let ctx = dev.cuda_context()?.clone();
let stream = dev.cuda_stream()?.clone();
let pool = PerHorizonAttentionPool::new(&ctx, stream.clone())?;
let head = PerHorizonResidualHead::new(&ctx, stream.clone())?;
let mut rng = ChaCha8Rng::seed_from_u64(0xBABE_F00D);
let ln_host: Vec<f32> = (0..B * K * PHA_HIDDEN_DIM).map(|_| rng.gen_range(-0.5..0.5)).collect();
let q_h_host: Vec<f32> = (0..PHA_N_HORIZONS * PHA_HIDDEN_DIM).map(|_| rng.gen_range(-0.1..0.1)).collect();
let w_res_host: Vec<f32> = (0..PHA_N_HORIZONS * PHA_HIDDEN_DIM).map(|_| rng.gen_range(-0.1..0.1)).collect();
let bias_res_host: Vec<f32> = (0..PHA_N_HORIZONS).map(|_| rng.gen_range(-0.05..0.05)).collect();
let baseline_logit_host: Vec<f32> = (0..B * PHA_N_HORIZONS).map(|_| rng.gen_range(-2.0..2.0)).collect();
// Non-zero alpha — different per horizon.
let alpha = [0.5f32, -0.3, 0.2, -0.1, 0.4];
let ln_d = alloc_upload(&stream, &ln_host);
let q_h_d = alloc_upload(&stream, &q_h_host);
let w_res_d = alloc_upload(&stream, &w_res_host);
let b_res_d = alloc_upload(&stream, &bias_res_host);
let mut ctx_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
let mut attn_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * K)?;
pool.forward(&q_h_d, &ln_d, B as i32, K as i32, &mut ctx_d, &mut attn_d)?;
let mut residual_d = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS)?;
head.forward(&ctx_d, &w_res_d, &b_res_d, B as i32, &mut residual_d)?;
let residual_host = download(&stream, &residual_d);
let mut final_logit = vec![0.0f32; B * PHA_N_HORIZONS];
let mut differences_seen = 0;
for b in 0..B {
for h in 0..PHA_N_HORIZONS {
let idx = b * PHA_N_HORIZONS + h;
let gate = alpha[h].tanh();
final_logit[idx] = baseline_logit_host[idx] + gate * residual_host[idx];
assert!(final_logit[idx].is_finite(), "non-finite final[{idx}]");
if (final_logit[idx] - baseline_logit_host[idx]).abs() > 1e-6 {
differences_seen += 1;
}
}
}
assert!(
differences_seen > 0,
"α non-zero but final logit unchanged from baseline — residual contributing nothing"
);
// Now run the full backward chain. Loss = Σ final_logit, so
// d_final[b,h] = 1
// d_residual[b,h] = tanh(α[h])
// d_α[h] = (1 - tanh(α[h])²) * Σ_b residual[b, h]
// d_baseline[b,h] = 1
let mut d_residual_host = vec![0.0f32; B * PHA_N_HORIZONS];
for b in 0..B {
for h in 0..PHA_N_HORIZONS {
d_residual_host[b * PHA_N_HORIZONS + h] = alpha[h].tanh();
}
}
let d_residual_d = alloc_upload(&stream, &d_residual_host);
let mut d_w_res_scratch = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
let mut d_b_res_scratch = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS)?;
let mut d_ctx = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
head.backward(
&ctx_d, &w_res_d, &d_residual_d, B as i32,
&mut d_w_res_scratch, &mut d_b_res_scratch, &mut d_ctx,
)?;
let mut d_q_h_scratch = stream.alloc_zeros::<f32>(B * PHA_N_HORIZONS * PHA_HIDDEN_DIM)?;
let mut d_ln = stream.alloc_zeros::<f32>(B * K * PHA_HIDDEN_DIM)?;
pool.backward(
&q_h_d, &ln_d, &attn_d, &d_ctx,
B as i32, K as i32,
&mut d_q_h_scratch, &mut d_ln,
)?;
let d_q_h_host = download(&stream, &d_q_h_scratch);
let d_ln_host = download(&stream, &d_ln);
// No NaN/Inf in any gradient buffer.
for &g in &d_q_h_host { assert!(g.is_finite(), "d_q_h scratch has non-finite"); }
for &g in &d_ln_host { assert!(g.is_finite(), "d_ln has non-finite"); }
// At least one non-zero gradient (else gate path didn't actually trigger).
let any_nonzero_q = d_q_h_host.iter().any(|&g| g.abs() > 1e-8);
let any_nonzero_ln = d_ln_host.iter().any(|&g| g.abs() > 1e-8);
assert!(any_nonzero_q, "d_q_h all zero — backward didn't flow through attention");
assert!(any_nonzero_ln, "d_ln all zero — backward didn't flow back to LN_b input");
Ok(())
}

View File

@@ -1,163 +0,0 @@
//! Numerical-gradient parity check for the per-horizon residual head
//! kernel (C22). Mirrors the C21 numgrad pattern.
//!
//! Three grads checked:
//! d_w_res [N_HORIZONS, HIDDEN_DIM]
//! d_bias_res [N_HORIZONS]
//! d_context_h [B, N_HORIZONS, HIDDEN_DIM]
//!
//! Loss = Σ residual_out (so d_residual = 1 everywhere).
use anyhow::Result;
use cudarc::driver::CudaSlice;
use ml_alpha::per_horizon_residual_head::{PerHorizonResidualHead, PHR_HIDDEN_DIM, PHR_N_HORIZONS};
use ml_core::device::MlDevice;
use rand::Rng;
use rand::SeedableRng;
use rand_chacha::ChaCha8Rng;
const B: usize = 3;
fn try_dev() -> Option<MlDevice> {
match MlDevice::cuda(0) {
Ok(d) => Some(d),
Err(e) => {
eprintln!("skipping: cuda device unavailable ({e})");
None
}
}
}
fn alloc_upload(stream: &std::sync::Arc<cudarc::driver::CudaStream>, host: &[f32]) -> CudaSlice<f32> {
let mut buf = stream.alloc_zeros::<f32>(host.len()).expect("alloc");
stream.memcpy_htod(host, &mut buf).expect("htod");
buf
}
fn download(stream: &std::sync::Arc<cudarc::driver::CudaStream>, src: &CudaSlice<f32>) -> Vec<f32> {
let mut out = vec![0.0f32; src.len()];
stream.memcpy_dtoh(src, out.as_mut_slice()).expect("dtoh");
out
}
fn forward_loss(
head: &PerHorizonResidualHead,
stream: &std::sync::Arc<cudarc::driver::CudaStream>,
context: &[f32],
w_res: &[f32],
bias: &[f32],
) -> f32 {
let ctx_d = alloc_upload(stream, context);
let w_d = alloc_upload(stream, w_res);
let b_d = alloc_upload(stream, bias);
let mut out_d = stream.alloc_zeros::<f32>(B * PHR_N_HORIZONS).unwrap();
head.forward(&ctx_d, &w_d, &b_d, B as i32, &mut out_d).unwrap();
download(stream, &out_d).iter().sum()
}
#[test]
#[ignore = "requires CUDA"]
fn forward_then_backward_matches_central_difference() -> Result<()> {
let Some(dev) = try_dev() else { return Ok(()); };
let ctx_h = dev.cuda_context()?.clone();
let stream = dev.cuda_stream()?.clone();
let head = PerHorizonResidualHead::new(&ctx_h, stream.clone())?;
let mut rng = ChaCha8Rng::seed_from_u64(0xBEEF_F00D);
let context_host: Vec<f32> = (0..B * PHR_N_HORIZONS * PHR_HIDDEN_DIM)
.map(|_| rng.gen_range(-0.5..0.5)).collect();
let w_host: Vec<f32> = (0..PHR_N_HORIZONS * PHR_HIDDEN_DIM)
.map(|_| rng.gen_range(-0.1..0.1)).collect();
let bias_host: Vec<f32> = (0..PHR_N_HORIZONS)
.map(|_| rng.gen_range(-0.05..0.05)).collect();
// Analytical backward with grad_residual = 1 everywhere.
let ctx_d = alloc_upload(&stream, &context_host);
let w_d = alloc_upload(&stream, &w_host);
let b_d = alloc_upload(&stream, &bias_host);
let mut out_d = stream.alloc_zeros::<f32>(B * PHR_N_HORIZONS)?;
head.forward(&ctx_d, &w_d, &b_d, B as i32, &mut out_d)?;
let grad_res_host = vec![1.0f32; B * PHR_N_HORIZONS];
let grad_res_d = alloc_upload(&stream, &grad_res_host);
let mut d_w_scratch_d = stream.alloc_zeros::<f32>(B * PHR_N_HORIZONS * PHR_HIDDEN_DIM)?;
let mut d_b_scratch_d = stream.alloc_zeros::<f32>(B * PHR_N_HORIZONS)?;
let mut d_ctx_d = stream.alloc_zeros::<f32>(B * PHR_N_HORIZONS * PHR_HIDDEN_DIM)?;
head.backward(
&ctx_d, &w_d, &grad_res_d, B as i32,
&mut d_w_scratch_d, &mut d_b_scratch_d, &mut d_ctx_d,
)?;
let d_w_scratch = download(&stream, &d_w_scratch_d);
let d_b_scratch = download(&stream, &d_b_scratch_d);
let d_ctx = download(&stream, &d_ctx_d);
// Reduce d_w_scratch across batch → [N_HORIZONS, HIDDEN_DIM].
let mut d_w_analytical = vec![0.0f32; PHR_N_HORIZONS * PHR_HIDDEN_DIM];
for b in 0..B {
for i in 0..PHR_N_HORIZONS * PHR_HIDDEN_DIM {
d_w_analytical[i] += d_w_scratch[b * PHR_N_HORIZONS * PHR_HIDDEN_DIM + i];
}
}
// Reduce d_b_scratch across batch → [N_HORIZONS].
let mut d_b_analytical = vec![0.0f32; PHR_N_HORIZONS];
for b in 0..B {
for h in 0..PHR_N_HORIZONS {
d_b_analytical[h] += d_b_scratch[b * PHR_N_HORIZONS + h];
}
}
let eps = 1e-2f32;
let tol = 5e-2f32;
// 1) Probe d_w_res at random indices.
for _ in 0..8 {
let idx = rng.gen_range(0..PHR_N_HORIZONS * PHR_HIDDEN_DIM);
let mut wp = w_host.clone(); wp[idx] += eps;
let mut wm = w_host.clone(); wm[idx] -= eps;
let lp = forward_loss(&head, &stream, &context_host, &wp, &bias_host);
let lm = forward_loss(&head, &stream, &context_host, &wm, &bias_host);
let fd = (lp - lm) / (2.0 * eps);
let an = d_w_analytical[idx];
let abs_err = (fd - an).abs();
let rel_err = abs_err / fd.abs().max(1e-6);
assert!(
abs_err < 5e-3 || rel_err < tol,
"d_w_res[{idx}]: analytical={an:.6} fd={fd:.6} abs={abs_err:.6} rel={rel_err:.4}"
);
}
// 2) Probe d_bias_res at every horizon.
for h in 0..PHR_N_HORIZONS {
let mut bp = bias_host.clone(); bp[h] += eps;
let mut bm = bias_host.clone(); bm[h] -= eps;
let lp = forward_loss(&head, &stream, &context_host, &w_host, &bp);
let lm = forward_loss(&head, &stream, &context_host, &w_host, &bm);
let fd = (lp - lm) / (2.0 * eps);
let an = d_b_analytical[h];
let abs_err = (fd - an).abs();
let rel_err = abs_err / fd.abs().max(1e-6);
assert!(
abs_err < 5e-3 || rel_err < tol,
"d_bias_res[{h}]: analytical={an:.6} fd={fd:.6} abs={abs_err:.6} rel={rel_err:.4}"
);
}
// 3) Probe d_context_h at random indices.
for _ in 0..8 {
let idx = rng.gen_range(0..B * PHR_N_HORIZONS * PHR_HIDDEN_DIM);
let mut cp = context_host.clone(); cp[idx] += eps;
let mut cm = context_host.clone(); cm[idx] -= eps;
let lp = forward_loss(&head, &stream, &cp, &w_host, &bias_host);
let lm = forward_loss(&head, &stream, &cm, &w_host, &bias_host);
let fd = (lp - lm) / (2.0 * eps);
let an = d_ctx[idx];
let abs_err = (fd - an).abs();
let rel_err = abs_err / fd.abs().max(1e-6);
assert!(
abs_err < 5e-3 || rel_err < tol,
"d_context_h[{idx}]: analytical={an:.6} fd={fd:.6} abs={abs_err:.6} rel={rel_err:.4}"
);
}
Ok(())
}