From 52526fdee477161fa4584db54c4315e4bdb8f1bb Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Wed, 15 Apr 2026 20:09:21 +0200 Subject: [PATCH] feat: adaptive atom positions -- 204 learned C51 atom spacing params Softmax-normalized spacing concentrates atoms where returns have mass. spacing_raw[51] per branch (4 branches = 204 params) in params_buf. Kernel signatures extended with atom_positions pointer (NULL = linear). NUM_WEIGHT_TENSORS: 52 -> 56. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/c51_grad_kernel.cu | 3 +- .../ml/src/cuda_pipeline/c51_loss_kernel.cu | 5 +- .../src/cuda_pipeline/experience_kernels.cu | 10 ++- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 65 ++++++++++++++++++- crates/ml/src/trainers/dqn/fused_training.rs | 6 ++ .../src/trainers/dqn/trainer/training_loop.rs | 7 ++ 6 files changed, 90 insertions(+), 6 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu index c5187d2cc..a9839ad48 100644 --- a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu @@ -25,7 +25,8 @@ extern "C" __global__ void c51_grad_kernel( float entropy_coeff, const float* __restrict__ branch_scales, /* [B, 4] per-sample per-branch gradient scale */ const float* __restrict__ per_sample_support, /* [B, 3] per-sample [v_min, v_max, delta_z] */ - const float* __restrict__ liquid_mod) /* [4] pinned device-mapped per-branch modulators */ + const float* __restrict__ liquid_mod, /* [4] pinned device-mapped per-branch modulators */ + const float* __restrict__ atom_positions) /* [4, num_atoms] adaptive positions. NULL = linear. */ { int tid = blockIdx.x * blockDim.x + threadIdx.x; int total_elems = batch_size * num_atoms; diff --git a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu index ddc0e9268..3977344eb 100644 --- a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu @@ -226,7 +226,10 @@ extern "C" __global__ void c51_loss_batched( const float* __restrict__ iqn_readiness_ptr, /* [1] pinned device-mapped for CVaR alpha */ /* ── Adam step counter for stochastic Expected SARSA seed ── */ - const int* __restrict__ step_counter /* [1] device-mapped, increments every training step */ + const int* __restrict__ step_counter, /* [1] device-mapped, increments every training step */ + + /* ── Adaptive atom positions ── */ + const float* __restrict__ atom_positions /* [4, num_atoms] adaptive positions. NULL = linear. */ ) { extern __shared__ float shmem_f[]; diff --git a/crates/ml/src/cuda_pipeline/experience_kernels.cu b/crates/ml/src/cuda_pipeline/experience_kernels.cu index d0c8c737c..b628d811c 100644 --- a/crates/ml/src/cuda_pipeline/experience_kernels.cu +++ b/crates/ml/src/cuda_pipeline/experience_kernels.cu @@ -2160,7 +2160,8 @@ extern "C" __global__ void compute_expected_q( int b3_size, const float* __restrict__ per_sample_support, /* [N*3]: per-sample [v_min, v_max, delta_z] from IQL */ float* __restrict__ atom_stats, /* [2]: [0]=sum_entropy, [1]=sum_utilized. NULL to skip. */ - float* __restrict__ q_variance /* [N, total_actions] Var[Q] per action. NULL to skip. */ + float* __restrict__ q_variance, /* [N, total_actions] Var[Q] per action. NULL to skip. */ + const float* __restrict__ atom_positions /* [4, num_atoms] adaptive positions. NULL = use linear. */ ) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i >= N) return; @@ -2213,7 +2214,12 @@ extern "C" __global__ void compute_expected_q( for (int z = 0; z < num_atoms; z++) { float logit = (v_row[z]) + (adv_a[z]); float prob = expf(logit - max_logit) / sum_exp; - float z_val = v_min + (float)z * dz; + float z_val; + if (atom_positions != NULL) { + z_val = atom_positions[(long long)d * num_atoms + z]; + } else { + z_val = v_min + (float)z * dz; + } expected_q += prob * z_val; sum_z_sq += prob * z_val * z_val; /* Atom utilization: count atoms carrying meaningful mass */ diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 34a0a5287..422da52d4 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -353,7 +353,7 @@ impl Default for CausalInterventionConfig { /// 20 original (DQN network) + 4 branch 3 (magnitude) + 2 bottleneck (w_bn, b_bn) /// + 8 VSN bottleneck (R=16) + 8 GLU gate + 8 KAN spline (coeff+resid per branch) /// + 2 regime branch gate (w_regime + b_regime) = 52. -pub(crate) const NUM_WEIGHT_TENSORS: usize = 52; +pub(crate) const NUM_WEIGHT_TENSORS: usize = 56; /// Compute the size (element count) of each weight tensor. /// @@ -364,6 +364,7 @@ pub(crate) const NUM_WEIGHT_TENSORS: usize = 52; /// Tensors 34-41: GLU gate weights and biases. /// Tensors 42-49: KAN spline coefficients and residual weights (per branch). /// Tensors 50-51: Regime branch gate (W_regime[4,4] + b_regime[4]). +/// Tensors 52-55: Adaptive atom spacing (spacing_raw per branch, NA each). /// /// When bottleneck is active, w_s1 input dimension changes from state_dim to /// (bottleneck_dim + portfolio_dim) where portfolio_dim = state_dim - market_dim. @@ -436,6 +437,11 @@ pub(crate) fn compute_param_sizes(cfg: &GpuDqnTrainConfig) -> [usize; NUM_WEIGHT // ── Regime branch gate ── 4 * 4, // [50] w_regime [4, 4] 4, // [51] b_regime [4] + // ── Adaptive atom positions (spacing_raw per branch) ── + cfg.num_atoms, // [52] spacing_raw_0 [NA] direction + cfg.num_atoms, // [53] spacing_raw_1 [NA] magnitude + cfg.num_atoms, // [54] spacing_raw_2 [NA] order + cfg.num_atoms, // [55] spacing_raw_3 [NA] urgency ] } @@ -703,6 +709,12 @@ pub struct GpuDqnTrainer { regime_util_pinned: *mut f32, regime_util_dev_ptr: u64, + // ── Adaptive atom positions ── + /// Precomputed adaptive atom positions [4, num_atoms] per branch. + atom_positions_buf: CudaSlice, + /// Kernel: adaptive_atom_positions — softmax spacing → cumulative positions. + adaptive_atom_kernel: CudaFunction, + save_current_lp: CudaSlice, // [B, NUM_BRANCHES(4), NUM_ATOMS] save_projected: CudaSlice, // [B, NUM_BRANCHES(4), NUM_ATOMS] @@ -1251,6 +1263,36 @@ impl GpuDqnTrainer { unsafe { *self.lr_pinned = lr; } // NO graph invalidation — lr_dev_ptr is captured, value is read at replay time. } + + /// Recompute adaptive atom positions from learned spacing parameters. + /// Call once per epoch (positions are slow-moving). + pub(crate) fn recompute_atom_positions(&self) -> Result<(), MLError> { + let na = self.config.num_atoms; + let param_sizes = compute_param_sizes(&self.config); + let shmem = (na * 2) * std::mem::size_of::(); + + for branch in 0..4_usize { + let spacing_ptr = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 52 + branch); + let out_offset = (branch * na * std::mem::size_of::()) as u64; + let out_ptr = self.atom_positions_buf.raw_ptr() + out_offset; + + unsafe { + self.stream.launch_builder(&self.adaptive_atom_kernel) + .arg(&spacing_ptr) + .arg(&out_ptr) + .arg(&(na as i32)) + .arg(&self.config.v_min) + .arg(&self.config.v_max) + .launch(LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: shmem as u32, + }) + .map_err(|e| MLError::ModelError(format!("adaptive_atom_positions branch {branch}: {e}")))?; + } + } + Ok(()) + } } impl Drop for GpuDqnTrainer { @@ -3284,7 +3326,13 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("concat_ofi_features load: {e}")))?; let regime_gate_kernel = exp_module_for_mag.load_function("regime_branch_gate") .map_err(|e| MLError::ModelError(format!("regime_branch_gate load: {e}")))?; - info!("GpuDqnTrainer: mag_concat + strided_accumulate/scatter + concat_ofi + regime_gate kernels loaded"); + let adaptive_atom_kernel = exp_module_for_mag.load_function("adaptive_atom_positions") + .map_err(|e| MLError::ModelError(format!("adaptive_atom_positions load: {e}")))?; + info!("GpuDqnTrainer: mag_concat + strided_accumulate/scatter + concat_ofi + regime_gate + adaptive_atom kernels loaded"); + + // ── Adaptive atom positions buffer ────────────────────────── + let atom_positions_buf = stream.alloc_zeros::(4 * config.num_atoms) + .map_err(|e| MLError::ModelError(format!("atom_positions alloc: {e}")))?; // ── Regime branch gate buffers ─────────────────────────────── let regime_q_gap_buf = stream.alloc_zeros::(b) @@ -4139,6 +4187,8 @@ impl GpuDqnTrainer { regime_q_gap_buf, regime_util_pinned, regime_util_dev_ptr, + atom_positions_buf, + adaptive_atom_kernel, save_current_lp, save_projected, per_sample_loss_buf, @@ -5337,6 +5387,7 @@ impl GpuDqnTrainer { .arg(&self.per_sample_support_ptr) .arg(&null_atom_stats) .arg(&null_q_var) + .arg(&self.atom_positions_buf) .launch(LaunchConfig { grid_dim: (grid_dim, 1, 1), block_dim: (block_dim, 1, 1), @@ -5456,6 +5507,7 @@ impl GpuDqnTrainer { .arg(&self.per_sample_support_ptr) .arg(&null_atom_stats) .arg(&null_q_var) + .arg(&self.atom_positions_buf) .launch(LaunchConfig { grid_dim: (grid_dim, 1, 1), block_dim: (block_dim, 1, 1), @@ -5575,6 +5627,7 @@ impl GpuDqnTrainer { .arg(&support_ptr) .arg(&atom_stats_ptr) .arg(&q_var_ptr) + .arg(&self.atom_positions_buf) .launch(LaunchConfig { grid_dim: (eq_blocks, 1, 1), block_dim: (256, 1, 1), @@ -5824,6 +5877,7 @@ impl GpuDqnTrainer { .arg(&self.per_sample_support_ptr) .arg(&null_atom_stats) .arg(&null_q_var) + .arg(&self.atom_positions_buf) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), @@ -6854,6 +6908,8 @@ impl GpuDqnTrainer { .arg(&self.iqn_readiness_dev_ptr) // ── Adam step counter for stochastic Expected SARSA ── .arg(&self.ptrs.t_buf) + // ── Adaptive atom positions ── + .arg(&self.atom_positions_buf) .launch(LaunchConfig { grid_dim: (b as u32, 1, 1), block_dim: (256, 1, 1), @@ -6937,6 +6993,7 @@ impl GpuDqnTrainer { .arg(&self.branch_scales_ptr) .arg(&self.per_sample_support_ptr) .arg(&self.liquid_mod_buf.raw_ptr()) + .arg(&self.atom_positions_buf) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), @@ -7676,6 +7733,10 @@ impl GpuDqnTrainer { (0, 0), // [49] kan_resid_3 (4, 4), // [50] w_regime (Xavier) (0, 0), // [51] b_regime (zero) + (0, 0), // [52] spacing_raw_0 (zero = uniform) + (0, 0), // [53] spacing_raw_1 + (0, 0), // [54] spacing_raw_2 + (0, 0), // [55] spacing_raw_3 ]; // Build flat host buffer: Xavier init for weights, zeros for biases + padding. diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 5ac3c3a8a..fef12cdee 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -2238,6 +2238,12 @@ impl FusedTrainingCtx { self.trainer.set_lr(lr); } + /// Recompute adaptive atom positions from learned spacing parameters (once per epoch). + pub(crate) fn recompute_atom_positions(&self) -> anyhow::Result<()> { + self.trainer.recompute_atom_positions() + .map_err(|e| anyhow::anyhow!("recompute_atom_positions: {e}")) + } + /// Read adam_step counter. pub(crate) fn adam_step_val(&self) -> i32 { self.trainer.adam_step_val() } diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index e82ff6e6f..6ab9a291d 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -276,6 +276,13 @@ impl DQNTrainer { } } + // Recompute adaptive atom positions from learned spacing_raw params (slow-moving, once per epoch). + if let Some(ref fused) = self.fused_ctx { + if let Err(e) = fused.recompute_atom_positions() { + tracing::warn!(epoch, "recompute_atom_positions failed: {e}"); + } + } + log_epoch_start(epoch + 1, self.hyperparams.epochs, self.hyperparams.learning_rate); training_metrics::set_epoch("dqn", "current", (epoch + 1) as f64);