diff --git a/crates/ml/src/trainers/ppo.rs b/crates/ml/src/trainers/ppo.rs index eb6eb9c29..d9848ccd7 100644 --- a/crates/ml/src/trainers/ppo.rs +++ b/crates/ml/src/trainers/ppo.rs @@ -7,7 +7,7 @@ //! - Hyperparameter configuration //! checkpoint management, and comprehensive metrics reporting. -use ml_core::cuda_autograd::GpuTensor; +use ml_core::cuda_autograd::GpuVarStore; use std::collections::VecDeque; use std::path::{Path, PathBuf}; use std::sync::Arc; @@ -211,7 +211,7 @@ pub struct PpoTrainingMetrics { pub struct PpoTrainer { model: Arc>, hyperparams: PpoHyperparameters, - device: Device, + device: MlDevice, checkpoint_dir: PathBuf, state_dim: usize, /// Number of parallel environments (None or Some(1) = standard, Some(n>1) = vectorized) @@ -387,13 +387,9 @@ impl PpoTrainer { &mut self, data: &[([f64; 42], Vec)], ) -> Result<(), MLError> { - let cuda_dev = match &self.device { - MlDevice::Cuda { stream: _, context: _ } => d, - MlDevice::Cpu => { - return Err(MLError::ConfigError("CUDA required for PPO set_raw_market_data".to_owned())); - } - }; - let stream = cuda_dev.cuda_stream(); + let stream = self.device.cuda_stream() + .map_err(|_| MLError::ConfigError("CUDA required for PPO set_raw_market_data".to_owned()))? + .clone(); let num_bars = data.len(); // Upload features [num_bars * 42] @@ -493,16 +489,15 @@ impl PpoTrainer { let avg_spread = self.hyperparams.avg_spread as f32; let cash_reserve_pct = self.hyperparams.cash_reserve_pct as f32 / 100.0; match (|| -> Result<_, MLError> { - let cuda_device = match &self.device { - MlDevice::Cuda { stream: _, context: _ } => d, - MlDevice::Cpu => return Err(MLError::ModelError("Not a CUDA device".into())), - }; - let stream = cuda_device.cuda_stream(); + let stream = self.device.cuda_stream() + .map_err(|_| MLError::ModelError("Not a CUDA device".into()))? + .clone(); + let curiosity_placeholder = GpuVarStore::new(stream.clone()); crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector::new( stream, actor_vars, critic_vars, - &GpuVarStore::new(), // curiosity placeholder — PPO curiosity integration is future work + &curiosity_placeholder, // curiosity placeholder — PPO curiosity integration is future work initial_capital, avg_spread, cash_reserve_pct, @@ -841,8 +836,7 @@ impl PpoTrainer { } /// Compute additional metrics (KL divergence, explained variance, etc.) - /// Uses GPU tensor ops to replace 5 sequential CPU passes + residuals allocation - /// with fused tensor operations and single scalar extractions. + /// Uses O(n) CPU arithmetic on host-resident data — zero GPU round-trips. fn compute_metrics( &self, batch: &TrajectoryBatch, @@ -856,26 +850,10 @@ impl PpoTrainer { let values = &batch.values; let rewards = &batch.rewards; - // Use tensor ops to compute explained variance, mean_reward, std_reward - // in fused passes instead of 5 sequential CPU loops + residuals Vec allocation. + // CPU-resident O(n) computation — data is already on host (downloaded from + // GPU kernel via PpoExperienceBatch). No GPU upload/download round-trip. let (explained_variance, mean_reward, std_reward) = - self.compute_metrics_tensors(returns, values, rewards) - .unwrap_or_else(|_| { - // Fallback to CPU if tensor ops fail - let mean_ret = returns.iter().sum::() / returns.len().max(1) as f32; - let var_ret = returns.iter().map(|r| (r - mean_ret).powi(2)).sum::() - / returns.len().max(1) as f32; - let mean_res: f32 = returns.iter().zip(values.iter()) - .map(|(r, v)| r - v).sum::() / returns.len().max(1) as f32; - let var_res: f32 = returns.iter().zip(values.iter()) - .map(|(r, v)| ((r - v) - mean_res).powi(2)).sum::() - / returns.len().max(1) as f32; - let ev = if var_ret > 0.0 { 1.0 - var_res / var_ret } else { 0.0 }; - let mr = rewards.iter().sum::() / rewards.len().max(1) as f32; - let sr = (rewards.iter().map(|r| (r - mr).powi(2)).sum::() - / rewards.len().max(1) as f32).sqrt(); - (ev, mr, sr) - }); + Self::compute_metrics_cpu(returns, values, rewards); // Compute real entropy from the epoch action distribution let entropy = { @@ -906,109 +884,68 @@ impl PpoTrainer { Ok((kl_div, explained_variance, mean_reward, std_reward, entropy)) } - /// Tensor-based computation of explained variance, mean reward, and std reward. - /// Replaces 5 sequential CPU passes with fused GPU/CPU tensor operations. - fn compute_metrics_tensors( - &self, + /// Compute explained variance, mean reward, and std reward from CPU-resident slices. + /// + /// Data is already on CPU (downloaded from GPU kernel via `PpoExperienceBatch`). + /// Pure CPU O(n) arithmetic avoids the GPU round-trip that the old Tensor-based + /// implementation incurred (3 uploads + 8 kernel launches + 4 scalar readbacks). + fn compute_metrics_cpu( returns: &[f32], values: &[f32], rewards: &[f32], - ) -> Result<(f32, f32, f32), MLError> { + ) -> (f32, f32, f32) { if returns.is_empty() || values.is_empty() || rewards.is_empty() { - return Ok((0.0, 0.0, 0.0)); + return (0.0, 0.0, 0.0); } - let device = &self.device; + let n = returns.len() as f32; - // Explained variance: 1 - Var(returns - values) / Var(returns) - let returns_t = Tensor::from_slice(returns, returns.len(), device) - .map_err(|e| MLError::ModelError(e.to_string()))?; - let values_t = Tensor::from_slice(values, values.len(), device) - .map_err(|e| MLError::ModelError(e.to_string()))?; + // Pass 1: accumulate sums for mean(returns), mean(residuals), mean(rewards) + let mut sum_ret = 0.0_f64; + let mut sum_res = 0.0_f64; + let mut sum_rew = 0.0_f64; + for i in 0..returns.len() { + let r = returns.get(i).copied().unwrap_or(0.0) as f64; + let v = values.get(i).copied().unwrap_or(0.0) as f64; + sum_ret += r; + sum_res += r - v; + // Rewards slice may be shorter than returns; guard with .get() + sum_rew += rewards.get(i).copied().unwrap_or(0.0) as f64; + } + let mean_ret = sum_ret / n as f64; + let mean_res = sum_res / n as f64; + let n_rew = rewards.len().max(1) as f64; + let mean_rew = sum_rew / n_rew; - // Var(returns): E[(x - E[x])^2] — keep as Tensor (no sync) - let returns_mean = returns_t.mean_all() - .map_err(|e| MLError::ModelError(e.to_string()))?; - let returns_centered = returns_t.broadcast_sub(&returns_mean) - .map_err(|e| MLError::ModelError(e.to_string()))?; - let var_returns_t = returns_centered.sqr() - .and_then(|s| s.mean_all()) - .map_err(|e| MLError::ModelError(e.to_string()))?; + // Pass 2: accumulate variance numerators + let mut var_ret_acc = 0.0_f64; + let mut var_res_acc = 0.0_f64; + let mut var_rew_acc = 0.0_f64; + for i in 0..returns.len() { + let r = returns.get(i).copied().unwrap_or(0.0) as f64; + let v = values.get(i).copied().unwrap_or(0.0) as f64; + let d_ret = r - mean_ret; + let d_res = (r - v) - mean_res; + var_ret_acc += d_ret * d_ret; + var_res_acc += d_res * d_res; + } + for i in 0..rewards.len() { + let d = rewards.get(i).copied().unwrap_or(0.0) as f64 - mean_rew; + var_rew_acc += d * d; + } + let var_ret = var_ret_acc / n as f64; + let var_res = var_res_acc / n as f64; + let var_rew = var_rew_acc / n_rew; - // Var(residuals): E[((r-v) - E[r-v])^2] — keep as Tensor (no sync) - let residuals_t = returns_t.sub(&values_t) - .map_err(|e| MLError::ModelError(e.to_string()))?; - let res_mean = residuals_t.mean_all() - .map_err(|e| MLError::ModelError(e.to_string()))?; - let res_centered = residuals_t.broadcast_sub(&res_mean) - .map_err(|e| MLError::ModelError(e.to_string()))?; - let var_residuals_t = res_centered.sqr() - .and_then(|s| s.mean_all()) - .map_err(|e| MLError::ModelError(e.to_string()))?; - - // Reward statistics via tensor ops — GPU-native mean centering (no scalar round-trip) - let rewards_t = Tensor::from_slice(rewards, rewards.len(), device) - .map_err(|e| MLError::ModelError(e.to_string()))?; - let mean_reward_t = rewards_t.mean_all() - .map_err(|e| MLError::ModelError(e.to_string()))?; - let rewards_centered = rewards_t.broadcast_sub(&mean_reward_t) - .map_err(|e| MLError::ModelError(e.to_string()))?; - let var_reward_t = rewards_centered.sqr() - .and_then(|s| s.mean_all()) - .map_err(|e| MLError::ModelError(e.to_string()))?; - - // Individual scalar readbacks (4 × .to_scalar, no bulk .to_vec1) - let var_returns = var_returns_t.to_scalar::() - .map_err(|e| MLError::ModelError(e.to_string()))?; - let var_residuals = var_residuals_t.to_scalar::() - .map_err(|e| MLError::ModelError(e.to_string()))?; - let mean_reward = mean_reward_t.to_scalar::() - .map_err(|e| MLError::ModelError(e.to_string()))?; - let var_reward = var_reward_t.to_scalar::() - .map_err(|e| MLError::ModelError(e.to_string()))?; - - let explained_variance = if var_returns > 0.0 { - 1.0 - var_residuals / var_returns + let explained_variance = if var_ret > 0.0 { + (1.0 - var_res / var_ret) as f32 } else { - 0.0 + 0.0_f32 }; - let std_reward = var_reward.sqrt(); + let mean_reward = mean_rew as f32; + let std_reward = (var_rew.sqrt()) as f32; - Ok((explained_variance, mean_reward, std_reward)) - } - - /// Sample action via GPU-side Gumbel-max trick. - /// - /// Returns `(action_index, log_prob)` — single scalar extraction from GPU. - /// Gumbel-max: argmax(log(p) - log(-log(U))) is equivalent to categorical sampling - /// but keeps the computation on the GPU, avoiding a full probability vector transfer. - fn sample_action(&self, probs: &Tensor) -> Result<(usize, f32), MLError> { - let flat_probs = probs.flatten_all()?; - let num_actions = flat_probs.elem_count(); - - // Gumbel-max sampling on GPU: argmax(log(p + eps) + gumbel_noise) - let eps_tensor = (flat_probs.ones_like()? * 1e-8)?; - let safe_probs = (flat_probs + eps_tensor)?; - let log_probs = safe_probs.log()?; - let uniform = Tensor::rand(0_f32, 1_f32, (num_actions,), log_probs.device())?; - let gumbel = uniform.log()?.neg()?.log()?.neg()?; - let perturbed = (log_probs.clone() + gumbel)?; - let action_idx_t = perturbed.argmax(0)?; - - // Single scalar GPU→CPU sync for the action index - let action_idx = action_idx_t - .to_dtype(ml_core::native_types::NativeDType::F32)? - .to_scalar::() - .unwrap_or(0) as usize; - let action_idx = action_idx.min(num_actions.saturating_sub(1)); - - // Extract log_prob for the selected action — single scalar sync - let log_prob = log_probs - .i(action_idx)? - .to_scalar::() - .unwrap_or(-10.0); - - Ok((action_idx, log_prob)) + (explained_variance, mean_reward, std_reward) } /// Compute reward based on actual PnL from price movements