feat(isv): recursive confidence head — model predicts own TD-error
h_s2 → sigmoid(w_conf @ h + b_conf) → predicted_error [B]. MSE loss vs lagged_td_error (0.01× weight). Backward accumulates into trunk gradient + conf weight gradients in main grad_buf. Creates self-improvement loop: model learns to predict when wrong. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -4993,3 +4993,53 @@ extern "C" __global__ void branch_confidence_routing(
|
||||
branch_offset += A_d;
|
||||
}
|
||||
}
|
||||
|
||||
/* ================================================================== */
|
||||
/* Kernel: recursive_confidence_forward — predict own TD-error */
|
||||
/* ================================================================== */
|
||||
extern "C" __global__ void recursive_confidence_forward(
|
||||
const float* __restrict__ h_s2, /* [B, SH2] */
|
||||
const float* __restrict__ w_conf, /* [SH2] */
|
||||
const float* __restrict__ b_conf, /* [1] */
|
||||
float* __restrict__ predicted_error, /* [B] */
|
||||
int B, int SH2
|
||||
) {
|
||||
int i = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (i >= B) return;
|
||||
|
||||
const float* h = h_s2 + (long long)i * SH2;
|
||||
float val = b_conf[0];
|
||||
for (int k = 0; k < SH2; k++)
|
||||
val += w_conf[k] * h[k];
|
||||
predicted_error[i] = 1.0f / (1.0f + expf(-val));
|
||||
}
|
||||
|
||||
/* ================================================================== */
|
||||
/* Kernel: recursive_confidence_backward — MSE grad into trunk */
|
||||
/* ================================================================== */
|
||||
extern "C" __global__ void recursive_confidence_backward(
|
||||
const float* __restrict__ h_s2, /* [B, SH2] */
|
||||
const float* __restrict__ predicted_error, /* [B] */
|
||||
const float* __restrict__ lagged_td_error, /* [1] pinned — target */
|
||||
const float* __restrict__ w_conf, /* [SH2] */
|
||||
float* __restrict__ d_w_conf, /* [SH2] gradient accumulator */
|
||||
float* __restrict__ d_b_conf, /* [1] gradient accumulator */
|
||||
float* __restrict__ d_h_s2, /* [B, SH2] trunk gradient (accumulate) */
|
||||
int B, int SH2,
|
||||
float loss_weight /* 0.01 */
|
||||
) {
|
||||
int i = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (i >= B) return;
|
||||
|
||||
float pred = predicted_error[i];
|
||||
float target = lagged_td_error[0];
|
||||
float d_loss = loss_weight * 2.0f * (pred - target) / (float)B;
|
||||
float d_sigmoid = d_loss * pred * (1.0f - pred);
|
||||
|
||||
const float* h = h_s2 + (long long)i * SH2;
|
||||
atomicAdd(d_b_conf, d_sigmoid);
|
||||
for (int k = 0; k < SH2; k++) {
|
||||
atomicAdd(&d_w_conf[k], d_sigmoid * h[k]);
|
||||
atomicAdd(&d_h_s2[(long long)i * SH2 + k], d_sigmoid * w_conf[k]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1344,6 +1344,10 @@ pub struct GpuDqnTrainer {
|
||||
gamma_mod_buf: CudaSlice<f32>, // [1]
|
||||
gamma_buf: CudaSlice<f32>, // [B] per-sample effective gamma
|
||||
predicted_error_buf: CudaSlice<f32>, // [B] recursive confidence output
|
||||
|
||||
// ── Recursive confidence kernels ──
|
||||
recursive_conf_fwd_kernel: CudaFunction,
|
||||
recursive_conf_bwd_kernel: CudaFunction,
|
||||
}
|
||||
|
||||
impl GpuDqnTrainer {
|
||||
@@ -2068,6 +2072,70 @@ impl GpuDqnTrainer {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Recursive confidence forward: predict own TD-error from h_s2.
|
||||
/// h_s2 → sigmoid(w_conf @ h + b_conf) → predicted_error [B].
|
||||
pub(crate) fn launch_recursive_confidence_forward(&self, batch_size: usize) -> Result<(), MLError> {
|
||||
let param_sizes = compute_param_sizes(&self.config);
|
||||
let w_conf = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 76);
|
||||
let b_conf = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 77);
|
||||
let blocks = ((batch_size as u32 + 255) / 256).max(1);
|
||||
let b_i32 = batch_size as i32;
|
||||
let sh2 = self.config.shared_h2 as i32;
|
||||
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.recursive_conf_fwd_kernel)
|
||||
.arg(&self.save_h_s2)
|
||||
.arg(&w_conf)
|
||||
.arg(&b_conf)
|
||||
.arg(&self.predicted_error_buf)
|
||||
.arg(&b_i32)
|
||||
.arg(&sh2)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("recursive_confidence_forward: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Recursive confidence backward: MSE loss gradient into trunk + conf weight gradients.
|
||||
/// Accumulates into grad_buf (same buffer Adam reads) and bw_d_h_s2 trunk gradient.
|
||||
pub(crate) fn launch_recursive_confidence_backward(&self, batch_size: usize) -> Result<(), MLError> {
|
||||
let param_sizes = compute_param_sizes(&self.config);
|
||||
let w_conf = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 76);
|
||||
// Gradient accumulators for w_conf and b_conf in main grad_buf
|
||||
let d_w_conf = self.ptrs.grad_buf + padded_byte_offset(¶m_sizes, 76);
|
||||
let d_b_conf = self.ptrs.grad_buf + padded_byte_offset(¶m_sizes, 77);
|
||||
|
||||
let blocks = ((batch_size as u32 + 255) / 256).max(1);
|
||||
let b_i32 = batch_size as i32;
|
||||
let sh2 = self.config.shared_h2 as i32;
|
||||
let loss_weight = 0.01_f32;
|
||||
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.recursive_conf_bwd_kernel)
|
||||
.arg(&self.save_h_s2)
|
||||
.arg(&self.predicted_error_buf)
|
||||
.arg(&self.lagged_td_error_dev_ptr)
|
||||
.arg(&w_conf)
|
||||
.arg(&d_w_conf)
|
||||
.arg(&d_b_conf)
|
||||
.arg(&self.bw_d_h_s2)
|
||||
.arg(&b_i32)
|
||||
.arg(&sh2)
|
||||
.arg(&loss_weight)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("recursive_confidence_backward: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Broadcast base_gamma * gamma_mod[0] into gamma_buf [B] for per-sample C51 loss.
|
||||
pub(crate) fn fill_gamma_buf(&self) -> Result<(), MLError> {
|
||||
let base_gamma = self.adaptive_gamma.powi(self.config.n_steps as i32);
|
||||
@@ -4173,6 +4241,10 @@ impl GpuDqnTrainer {
|
||||
.map_err(|e| MLError::ModelError(format!("isv_forward load: {e}")))?;
|
||||
let fill_gamma_buf_kernel = exp_module_for_mag.load_function("fill_gamma_buf")
|
||||
.map_err(|e| MLError::ModelError(format!("fill_gamma_buf load: {e}")))?;
|
||||
let recursive_conf_fwd_kernel = exp_module_for_mag.load_function("recursive_confidence_forward")
|
||||
.map_err(|e| MLError::ModelError(format!("recursive_confidence_forward load: {e}")))?;
|
||||
let recursive_conf_bwd_kernel = exp_module_for_mag.load_function("recursive_confidence_backward")
|
||||
.map_err(|e| MLError::ModelError(format!("recursive_confidence_backward load: {e}")))?;
|
||||
info!("GpuDqnTrainer: mag_concat + strided_accumulate/scatter + concat_ofi + regime_gate + adaptive_atom + atom_grad + q_anchor + regime_dropout + G5/G6/G10/G12 + risk_budget + isv_signal_update + isv_forward + fill_gamma_buf kernels loaded");
|
||||
|
||||
// ── G5: Epistemic-gated magnitude — pinned var_ema threshold ─
|
||||
@@ -5673,6 +5745,8 @@ impl GpuDqnTrainer {
|
||||
gamma_mod_buf,
|
||||
gamma_buf,
|
||||
predicted_error_buf,
|
||||
recursive_conf_fwd_kernel,
|
||||
recursive_conf_bwd_kernel,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -6914,6 +6988,9 @@ impl GpuDqnTrainer {
|
||||
// ISV forward: encoder MLP → branch gate + gamma mod
|
||||
self.launch_isv_forward()?;
|
||||
|
||||
// Recursive confidence: predict own TD-error from h_s2
|
||||
self.launch_recursive_confidence_forward(batch_size)?;
|
||||
|
||||
// Risk budget forward: h_s2 → risk_budget R ∈ (0,1) before Q-value computation
|
||||
self.risk_budget_forward(batch_size)?;
|
||||
|
||||
|
||||
@@ -1373,6 +1373,11 @@ impl FusedTrainingCtx {
|
||||
}
|
||||
}
|
||||
|
||||
// Recursive confidence backward: MSE grad into trunk + conf weight gradients.
|
||||
// Must run before Adam (which reads grad_buf for the parameter update).
|
||||
self.trainer.launch_recursive_confidence_backward(self.batch_size)
|
||||
.map_err(|e| anyhow::anyhow!("Recursive confidence backward: {e}"))?;
|
||||
|
||||
// Regime-adaptive PER scaling.
|
||||
self.trainer.regime_scale_td_errors()
|
||||
.map_err(|e| anyhow::anyhow!("Regime PER scaling: {e}"))?;
|
||||
|
||||
Reference in New Issue
Block a user