diff --git a/crates/ml/src/trainers/dqn/trainer.rs b/crates/ml/src/trainers/dqn/trainer.rs index 58a631c26..0ecaf5534 100644 --- a/crates/ml/src/trainers/dqn/trainer.rs +++ b/crates/ml/src/trainers/dqn/trainer.rs @@ -14,6 +14,7 @@ use risk::drawdown_monitor::DrawdownMonitor; use risk::safety::position_limiter::HybridPositionLimiter; use risk::safety::PositionLimiterConfig; use num_traits::ToPrimitive; +use crate::cuda_pipeline::DqnGpuData; use rust_decimal::Decimal; use tokio::sync::RwLock; use tracing::{debug, info, warn}; @@ -192,6 +193,9 @@ pub struct DQNTrainer { /// Current effective batch size (may be reduced by OOM recovery) current_batch_size: usize, + + /// Pre-uploaded GPU training data (set once, reused across epochs) + gpu_data: Option, } impl std::fmt::Debug for DQNTrainer { @@ -738,6 +742,9 @@ impl DQNTrainer { // OOM recovery: track effective batch size current_batch_size: initial_batch_size, + + // GPU pipeline: pre-uploaded training data (initialized lazily at first epoch) + gpu_data: None, }) } @@ -1282,6 +1289,24 @@ impl DQNTrainer { // Within a single training run, portfolio compounds across ALL epochs. let epoch_start = std::time::Instant::now(); + + // Phase 1: Pre-upload training data to GPU (once, reused across epochs) + if self.gpu_data.is_none() { + match DqnGpuData::upload(&training_data, &self.device) { + Ok(gpu_data) => { + info!("GPU data pre-uploaded: {} bars x {} features ({:.1} MB)", + gpu_data.num_bars, + gpu_data.feature_dim, + (gpu_data.num_bars * (51 + 4) * 4) as f64 / 1_048_576.0 + ); + self.gpu_data = Some(gpu_data); + } + Err(e) => { + debug!("GPU data pre-upload skipped (CPU fallback): {}", e); + } + } + } + let mut epoch_loss = 0.0; let mut epoch_q_value = 0.0; let mut epoch_gradient_norm = 0.0; @@ -1344,34 +1369,45 @@ impl DQNTrainer { for (idx_in_batch, &i) in batch_indices.iter().enumerate() { let state = &states[idx_in_batch]; let action = actions[idx_in_batch]; - let target = &training_data[i].1; - // Extract close prices for next state calculation + // Use pre-uploaded GPU targets if available, otherwise fall back to CPU // WAVE 3 BUG FIX: Use raw prices (indices 2,3) for barrier tracker, preprocessed (indices 0,1) for rewards - let current_close_raw = if target.len() >= 4 { - target[2] // Raw price for barrier tracker - } else if target.len() >= 2 { - target[0] // Fallback to preprocessed for old data - } else { - training_data[i].0[3] - }; - let next_close_raw = if target.len() >= 4 { - target[3] // Raw price for barrier tracker - } else if target.len() >= 2 { - target[1] // Fallback to preprocessed for old data - } else { - current_close_raw - }; - let current_close = if target.len() >= 2 { - target[0] // Preprocessed for reward calculation - } else { - training_data[i].0[3] - }; - let next_close = if target.len() >= 2 { - target[1] // Preprocessed for reward calculation - } else { - current_close - }; + let (current_close_raw, next_close_raw, current_close, next_close) = + if let Some(ref gpu_data) = self.gpu_data { + let t = gpu_data.bar_target_values(i)?; + let cc = t[0] as f64; + let nc = t[1] as f64; + let ccr = if t[2] != 0.0 { t[2] as f64 } else { cc }; + let ncr = if t[3] != 0.0 { t[3] as f64 } else { cc }; + (ccr, ncr, cc, nc) + } else { + let target = &training_data[i].1; + let current_close_raw = if target.len() >= 4 { + target[2] // Raw price for barrier tracker + } else if target.len() >= 2 { + target[0] // Fallback to preprocessed for old data + } else { + training_data[i].0[3] + }; + let next_close_raw = if target.len() >= 4 { + target[3] // Raw price for barrier tracker + } else if target.len() >= 2 { + target[1] // Fallback to preprocessed for old data + } else { + current_close_raw + }; + let current_close = if target.len() >= 2 { + target[0] // Preprocessed for reward calculation + } else { + training_data[i].0[3] + }; + let next_close = if target.len() >= 2 { + target[1] // Preprocessed for reward calculation + } else { + current_close + }; + (current_close_raw, next_close_raw, current_close, next_close) + }; // Get next state (BUG #42: made mutable to update portfolio features after trade execution) let mut next_state = if i + 1 < training_data.len() {