feat: magnitude branch forward/backward uses [h_s2; Q_dir] conditioning

Forward: branch 1 FC GEMM reads [B, SH2+3] from mag_concat_buf.
All 3 forward paths (online_raw, online_f32, target_raw) updated.
Concat uses previous step's direction logits (one-step lag, negligible).

Backward: branch 1 dW uses wider X, dX writes to d_mag_concat_buf.
strided_accumulate extracts first SH2 columns into d_h_s2 with beta=1.
d_Q_dir (last 3 columns) discarded — detached, no backprop to direction.

Fixes backward goff_b_b1fc offset: uses ah*(sh2+3) matching compute_param_sizes.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-14 22:48:22 +02:00
parent cf009ad49e
commit 85b5f15520
4 changed files with 146 additions and 26 deletions

View File

@@ -242,6 +242,8 @@ impl CublasBackwardSet {
}
// Branch FC layers: (adv_h, shared_h2)
add_fc_shapes(&mut unique_shapes, ah, sh2);
// Branch 1 FC (magnitude) wider input: (adv_h, shared_h2 + 3)
add_fc_shapes(&mut unique_shapes, ah, sh2 + 3);
// Value output: (NA, value_h)
add_fc_shapes(&mut unique_shapes, na, vh);
// Value FC: (value_h, shared_h2)
@@ -719,6 +721,9 @@ impl CublasBackwardSet {
// When non-zero, computes d_loss/d_bn_concat (upstream gradient for bottleneck).
// When 0, dX is not computed (states not trainable — no bottleneck).
s1_dx_output: u64,
// Magnitude branch conditioning: saved forward concat and dX output
mag_concat_ptr: u64, // [B, SH2+3] saved forward concat (for dW)
d_mag_concat_ptr: u64, // [B, SH2+3] dX output for branch 1 (caller accumulates)
) -> Result<(), MLError> {
let b = self.batch_size;
let na = self.num_atoms;
@@ -768,7 +773,7 @@ impl CublasBackwardSet {
let goff_w_b0out: u64 = goff_b_b0fc + ah * f32;
let goff_b_b0out: u64 = goff_w_b0out + b0 * na64 * ah * f32;
let goff_w_b1fc: u64 = goff_b_b0out + b0 * na64 * f32;
let goff_b_b1fc: u64 = goff_w_b1fc + ah * sh2 * f32;
let goff_b_b1fc: u64 = goff_w_b1fc + ah * (sh2 + 3) * f32; // magnitude uses wider input: SH2+3
let goff_w_b1out: u64 = goff_b_b1fc + ah * f32;
let goff_b_b1out: u64 = goff_w_b1out + b1 * na64 * ah * f32;
let goff_w_b2fc: u64 = goff_b_b1out + b1 * na64 * f32;
@@ -869,32 +874,55 @@ impl CublasBackwardSet {
b * self.adv_h,
)?;
// dW for branch FC: dW[AH, SH2] += dY^T @ h_s2
// (dX is computed separately below to allow accumulation)
// dW for branch FC: dW += dY^T @ X
// Branch 1 (magnitude) uses wider concat input [B, SH2+3] for dW.
let (fc_input, fc_in_dim) = if d == 1 && mag_concat_ptr != 0 {
(mag_concat_ptr, self.shared_h2 + 3)
} else {
(save_h_s2, self.shared_h2)
};
self.launch_dw_only(
stream,
scratch_d_h_b[d], // dY [B, AH] — f32
save_h_s2, // X [B, SH2]
grad_buf_base + goff_w_bfc[d], // dW [AH, SH2]
fc_input, // X [B, fc_in_dim]
grad_buf_base + goff_w_bfc[d], // dW [AH, fc_in_dim]
grad_buf_base + goff_b_bfc[d], // db [AH]
self.adv_h, // out_dim
self.shared_h2, // in_dim
fc_in_dim, // in_dim
b,
)?;
// dX for branch FC: d_h_s2 += dY @ W_bdk_fc
// Use beta = if d==0 { 0.0 } else { 1.0 } to accumulate branches.
let beta_s2 = if d == 0 { 0.0_f32 } else { 1.0_f32 };
self.launch_dx_only(
stream,
scratch_d_h_b[d], // dY [B, AH] — f32
w_fc, // W [AH, SH2]
scratch_d_h_s2, // dX [B, SH2] — f32
self.adv_h, // out_dim
self.shared_h2, // in_dim
b,
beta_s2,
)?;
// Branch 1 (magnitude): dX writes to d_mag_concat [B, SH2+3].
// The caller accumulates first SH2 columns into d_h_s2 via strided_accumulate.
if d == 1 && mag_concat_ptr != 0 {
// Magnitude: dX = dY @ W_b1fc → [B, SH2+3], write to d_mag_concat
self.launch_dx_only(
stream,
scratch_d_h_b[d], // dY [B, AH] — f32
w_fc, // W [AH, SH2+3]
d_mag_concat_ptr, // dX [B, SH2+3] — f32
self.adv_h, // out_dim
self.shared_h2 + 3, // in_dim
b,
0.0_f32, // overwrite (fresh buffer)
)?;
// Caller will call accumulate_d_h_s2_from_concat(d_mag_concat, d_h_s2, B, 1.0)
// after backward_full returns.
} else {
// Other branches: accumulate into d_h_s2 directly.
let beta_s2 = if d == 0 { 0.0_f32 } else { 1.0_f32 };
self.launch_dx_only(
stream,
scratch_d_h_b[d], // dY [B, AH] — f32
w_fc, // W [AH, SH2]
scratch_d_h_s2, // dX [B, SH2] — f32
self.adv_h, // out_dim
self.shared_h2, // in_dim
b,
beta_s2,
)?;
}
}
// ══════════════════════════════════════════════════════════════════

View File

@@ -145,6 +145,10 @@ pub struct CublasGemmSet {
branch_2_size: usize,
branch_3_size: usize,
/// Magnitude branch (d==1) FC input dimension: shared_h2 + 3
/// (concat of h_s2 and Q_dir). Used when mag_concat_ptr != 0.
mag_concat_dim: usize,
// ── Multi-stream branch dispatch ──
/// 4 forked CUDA streams for parallel advantage branch execution.
/// Each branch (exposure, order, urgency, ...) submits GEMMs to its own stream,
@@ -261,6 +265,8 @@ impl CublasGemmSet {
unique_shapes.push((num_atoms, batch_size, value_h, value_h));
// h_bd (×4): M=adv_h, N=batch, K=shared_h2, ldb=shared_h2
unique_shapes.push((adv_h, batch_size, shared_h2, shared_h2));
// h_bd magnitude (d==1) wider input: M=adv_h, N=batch, K=shared_h2+3, ldb=shared_h2+3
unique_shapes.push((adv_h, batch_size, shared_h2 + 3, shared_h2 + 3));
// adv_logits (×4): M=branch_k*num_atoms, N=batch, K=adv_h, ldb=adv_h
for &bs in &branch_sizes {
unique_shapes.push((bs * num_atoms, batch_size, adv_h, adv_h));
@@ -281,6 +287,7 @@ impl CublasGemmSet {
(shared_h2, batch_size, shared_h1, shared_h1), // h_s2
(value_h, batch_size, shared_h2, shared_h2), // h_v
(adv_h, batch_size, shared_h2, shared_h2), // h_bd (×4)
(adv_h, batch_size, shared_h2 + 3, shared_h2 + 3), // h_bd magnitude wider input
];
let mut gemm_cache_relu_bias = HashMap::new();
for &(n, b, k, ldb) in &relu_bias_shapes {
@@ -310,6 +317,7 @@ impl CublasGemmSet {
branch_1_size,
branch_2_size,
branch_3_size,
mag_concat_dim: shared_h2 + 3,
branch_streams,
_branch_workspace_bufs: branch_workspace_bufs,
branch_workspace_ptrs,
@@ -375,6 +383,7 @@ impl CublasGemmSet {
h_s1_ptr: u64, h_s2_ptr: u64, h_v_ptr: u64,
h_b0_ptr: u64, h_b1_ptr: u64, h_b2_ptr: u64, h_b3_ptr: u64,
v_logits_ptr: u64, b_logits_ptr: u64,
mag_concat_ptr: u64,
) -> Result<(), MLError> {
let b = self.batch_size;
@@ -450,10 +459,18 @@ impl CublasGemmSet {
// Try fused GEMM+bias+ReLU epilogue, fall back to separate kernels.
let bws = self.branch_workspace_ptrs[d];
let bwss = self.handle.lt_workspace_size;
if self.sgemm_f32_fused_relu_bias(bs, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1],
self.adv_h, b, self.shared_h2, self.shared_h2, bws, bwss, "h_bd").is_err()
// Branch 1 (magnitude) uses wider [h_s2; Q_dir] concat input when available.
let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 {
(mag_concat_ptr, self.mag_concat_dim)
} else {
(h_s2_ptr, self.shared_h2)
};
if self.sgemm_f32_fused_relu_bias(bs, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1],
self.adv_h, b, fc_k, fc_k, bws, bwss, "h_bd").is_err()
{
self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, d, "h_bd")?;
self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, d, "h_bd")?;
self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
}
@@ -478,7 +495,13 @@ impl CublasGemmSet {
let n_d = branch_sizes[d];
let w_fc_idx = branch_w_base[d];
self.sgemm_f32(stream, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "h_bd")?;
let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 {
(mag_concat_ptr, self.mag_concat_dim)
} else {
(h_s2_ptr, self.shared_h2)
};
self.sgemm_f32(stream, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, "h_bd")?;
self.launch_add_bias_relu_f32_raw(stream, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
let adv_out_ptr = b_logits_ptr + logit_byte_offset;
@@ -507,6 +530,7 @@ impl CublasGemmSet {
h_s1_ptr: u64, h_s2_ptr: u64, h_v_ptr: u64, // F32 activation buffers
h_b0_ptr: u64, h_b1_ptr: u64, h_b2_ptr: u64, h_b3_ptr: u64,
v_logits_ptr: u64, b_logits_ptr: u64, // F32 output logits
mag_concat_ptr: u64,
) -> Result<(), MLError> {
let b = self.batch_size;
@@ -547,7 +571,14 @@ impl CublasGemmSet {
.map_err(|e| MLError::ModelError(format!("f32 branch {d} wait trunk: {e}")))?;
// Per-branch workspace: eliminates contention between parallel branch streams.
self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, d, "f32_h_bd")?;
// Branch 1 (magnitude) uses wider [h_s2; Q_dir] concat input when available.
let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 {
(mag_concat_ptr, self.mag_concat_dim)
} else {
(h_s2_ptr, self.shared_h2)
};
self.sgemm_f32_branch(bs, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, d, "f32_h_bd")?;
self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
let adv_out_ptr = b_logits_ptr + logit_byte_offset;
@@ -570,7 +601,13 @@ impl CublasGemmSet {
let n_d = branch_sizes[d];
let w_fc_idx = branch_w_base[d];
self.sgemm_f32(stream, w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d], self.adv_h, b, self.shared_h2, "f32_h_bd")?;
let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 {
(mag_concat_ptr, self.mag_concat_dim)
} else {
(h_s2_ptr, self.shared_h2)
};
self.sgemm_f32(stream, w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d], self.adv_h, b, fc_k, "f32_h_bd")?;
self.launch_add_bias_relu_f32_raw(stream, branch_h_ptrs[d], w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
let adv_out_ptr = b_logits_ptr + logit_byte_offset;
@@ -614,6 +651,7 @@ impl CublasGemmSet {
raw_bf16_ptr(save_h_b3, stream),
raw_f32_ptr(v_logits_buf, stream),
raw_f32_ptr(b_logits_buf, stream),
0u64, // mag_concat_ptr: not used in CudaSlice wrapper path
)
}
@@ -681,6 +719,7 @@ impl CublasGemmSet {
h_s1_ptr: u64, h_s2_ptr: u64, h_v_ptr: u64,
h_b0_ptr: u64, h_b1_ptr: u64, h_b2_ptr: u64, h_b3_ptr: u64,
v_logits_ptr: u64, b_logits_ptr: u64,
mag_concat_ptr: u64,
) -> Result<(), MLError> {
let b = self.batch_size;
@@ -725,8 +764,15 @@ impl CublasGemmSet {
.map_err(|e| MLError::ModelError(format!("tg branch {d} wait trunk: {e}")))?;
// Hidden layer — per-branch workspace eliminates contention.
self.sgemm_f32_branch(bs, tg_w_ptrs[w_fc_idx], h_s2_ptr, branch_h_ptrs[d],
self.adv_h, b, self.shared_h2, d, "tg_h_bd")?;
// Branch 1 (magnitude) uses wider [h_s2; Q_dir] concat input when available.
let (fc_input, fc_k) = if d == 1 && mag_concat_ptr != 0 {
(mag_concat_ptr, self.mag_concat_dim)
} else {
(h_s2_ptr, self.shared_h2)
};
self.sgemm_f32_branch(bs, tg_w_ptrs[w_fc_idx], fc_input, branch_h_ptrs[d],
self.adv_h, b, fc_k, d, "tg_h_bd")?;
self.launch_add_bias_relu_f32_raw(bs, branch_h_ptrs[d], tg_w_ptrs[w_fc_idx + 1], self.adv_h, b)?;
// Output layer: f32 cublasLtMatmul + f32 bias — per-branch workspace.
@@ -779,6 +825,7 @@ impl CublasGemmSet {
raw_f32_ptr(tg_h_b3_scratch, stream),
raw_f32_ptr(tg_v_logits_buf, stream),
raw_f32_ptr(tg_b_logits_buf, stream),
0u64, // mag_concat_ptr: not used in CudaSlice wrapper path
)
}

View File

@@ -1967,8 +1967,20 @@ impl GpuDqnTrainer {
scratch_d_h_s2, scratch_d_h_s1, scratch_d_h_v,
&[scratch_d_h_b0, scratch_d_h_b1, scratch_d_h_b2, scratch_d_h_b3],
0, // CQL backward: no bottleneck dX needed (separate gradient budget)
self.ptrs.mag_concat_buf,
self.ptrs.d_mag_concat_buf,
).map_err(|e| MLError::ModelError(format!("CQL backward_full: {e}")))?;
// CQL: accumulate magnitude branch dX into scratch_d_h_s2
if self.ptrs.mag_concat_buf != 0 {
self.accumulate_d_h_s2_from_concat(
self.ptrs.d_mag_concat_buf,
scratch_d_h_s2,
self.config.batch_size,
1.0, // beta=1: d==0 already wrote to scratch_d_h_s2
)?;
}
Ok(true)
}
@@ -3808,6 +3820,7 @@ impl GpuDqnTrainer {
h_s1, h_s2, h_v, h_b0, h_b1, h_b2, h_b3,
self.ptrs.on_next_v_logits_buf,
self.ptrs.on_next_b_logits_buf,
0u64, // mag_concat_ptr: causal intervention uses unconditioned forward
)?;
// Step 4: GPU reduction — |Q_orig - Q_interv|² → sensitivity[k]
@@ -4288,11 +4301,14 @@ impl GpuDqnTrainer {
fn replay_forward_ungraphed(&self) -> Result<(), MLError> {
let param_sizes = compute_param_sizes(&self.config);
let on_w_ptrs = f32_weight_ptrs_from_base(self.ptrs.params_ptr, &param_sizes);
// Build mag_concat from previous step's logits (one-step lag)
self.launch_mag_concat(self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf)?;
self.cublas_forward.forward_online_raw(
&self.stream, self.ptrs.states_buf, &on_w_ptrs,
self.ptrs.save_h_s1, self.ptrs.save_h_s2, self.ptrs.save_h_v,
self.ptrs.save_h_b0, self.ptrs.save_h_b1, self.ptrs.save_h_b2, self.ptrs.save_h_b3,
self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf,
self.ptrs.mag_concat_buf,
)
}
@@ -4348,11 +4364,14 @@ impl GpuDqnTrainer {
// Pre-step-0: no graph captured. Run ungraphed forward.
let param_sizes = compute_param_sizes(&self.config);
let on_w_ptrs = f32_weight_ptrs_from_base(self.ptrs.params_ptr, &param_sizes);
// Build mag_concat from previous step's logits (zeros at step 0)
self.launch_mag_concat(self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf)?;
self.cublas_forward.forward_online_raw(
&self.stream, self.ptrs.states_buf, &on_w_ptrs,
self.ptrs.save_h_s1, self.ptrs.save_h_s2, self.ptrs.save_h_v,
self.ptrs.save_h_b0, self.ptrs.save_h_b1, self.ptrs.save_h_b2, self.ptrs.save_h_b3,
self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf,
self.ptrs.mag_concat_buf,
)?;
}
@@ -4987,6 +5006,7 @@ impl GpuDqnTrainer {
self.ptrs.on_next_h_b_scratch, self.ptrs.on_next_h_b_scratch,
self.ptrs.on_next_h_b_scratch, self.ptrs.on_next_h_b_scratch,
self.ptrs.on_next_v_logits_buf, self.ptrs.on_next_b_logits_buf,
0u64, // mag_concat_ptr: DDQN argmax pass — no magnitude conditioning
)?;
Ok(())
}
@@ -5223,6 +5243,13 @@ impl GpuDqnTrainer {
self.ptrs.states_buf
};
// ── Build mag_concat from previous step's direction logits (one-step lag).
// At step 0, logit buffers are zero-init → mag_concat = [h_s2; 0,0,0].
// From step 1+, uses the previous step's direction Q-values.
// This call is captured in the CUDA graph — on replay it reads the logits
// from the previous graph execution (still in the buffer).
self.launch_mag_concat(self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf)?;
// ── Pass 1: Online forward on STATES (graph-safe: CachedPtrs, no device_ptr)
// When bottleneck is active, s1_input_ptr = bn_concat (compressed features).
// When disabled, s1_input_ptr = states_buf (original features).
@@ -5231,6 +5258,7 @@ impl GpuDqnTrainer {
self.ptrs.save_h_s1, self.ptrs.save_h_s2, self.ptrs.save_h_v,
self.ptrs.save_h_b0, self.ptrs.save_h_b1, self.ptrs.save_h_b2, self.ptrs.save_h_b3,
self.ptrs.on_v_logits_buf, self.ptrs.on_b_logits_buf,
self.ptrs.mag_concat_buf,
)?;
// ── #21 Stochastic depth: scale hidden activations by per-layer mask ──
@@ -5271,6 +5299,9 @@ impl GpuDqnTrainer {
// ── Pass 2: Target forward on NEXT_STATES — main stream (graph-safe)
// No event sync needed — Pass 3 (Double DQN) is submitted separately
// via submit_forward_ops_ddqn() and captured on its own stream/graph.
// Build target mag_concat from previous step's target logits (one-step lag).
// Reuses the same mag_concat_buf — target forward runs after online forward.
self.launch_mag_concat(self.ptrs.tg_v_logits_buf, self.ptrs.tg_b_logits_buf)?;
cublas.forward_target_raw(
&self.stream, self.ptrs.next_states_buf, &tg_w_ptrs,
self.ptrs.tg_h_s1_scratch, self.ptrs.tg_h_s2_buf,
@@ -5278,6 +5309,7 @@ impl GpuDqnTrainer {
self.ptrs.tg_h_b0_scratch, self.ptrs.tg_h_b1_scratch,
self.ptrs.tg_h_b2_scratch, self.ptrs.tg_h_b3_scratch,
self.ptrs.tg_v_logits_buf, self.ptrs.tg_b_logits_buf,
self.ptrs.mag_concat_buf,
)?;
Ok(())
@@ -5888,8 +5920,20 @@ impl GpuDqnTrainer {
d_h_v_ptr,
&[d_h_b0_ptr, d_h_b1_ptr, d_h_b2_ptr, d_h_b3_ptr],
s1_dx_output,
self.ptrs.mag_concat_buf,
self.ptrs.d_mag_concat_buf,
)?;
// Accumulate magnitude branch dX (first SH2 columns) into d_h_s2
if self.ptrs.mag_concat_buf != 0 {
self.accumulate_d_h_s2_from_concat(
self.ptrs.d_mag_concat_buf,
d_h_s2_ptr,
self.config.batch_size,
1.0, // beta=1: d==0 already wrote to d_h_s2 in backward_full
)?;
}
// ── #31 Bottleneck backward: chain rule through tanh + GEMM ──
// After backward_full, d_bn_concat_buf has dL/d(bn_concat) [B, concat_dim] f32.
// Chain rule: dL/d_bn = dL/d_concat[:, :bn_dim] * tanh'(bn_raw)

View File

@@ -1935,6 +1935,7 @@ impl GpuExperienceCollector {
self.exp_h_b0_f32.raw_ptr(), self.exp_h_b1_f32.raw_ptr(),
self.exp_h_b2_f32.raw_ptr(), self.exp_h_b3_f32.raw_ptr(),
self.exp_v_logits.raw_ptr(), self.exp_b_logits.raw_ptr(),
0u64, // mag_concat_ptr: experience collector uses unconditioned forward
)?;
}