feat(dqn): pre-upload training data to GPU, use cached targets in hot loop
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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<DqnGpuData>,
|
||||
}
|
||||
|
||||
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() {
|
||||
|
||||
Reference in New Issue
Block a user