Created config/gpu/{default,rtx3050,h100,a100}.toml with all GPU-specific
parameters: batch_size, num_atoms, buffer_size, hidden_dim_base,
replay_buffer_vram_fraction, gpu_n_episodes, gpu_timesteps_per_episode,
cuda_stack_bytes.
GpuProfile::load() auto-detects GPU by device name, falls back to
embedded defaults (include_str!). Override via FOXHUNT_GPU_PROFILE env.
Removed dead code:
- detect_vram_mb(), vram_scaled_hidden_dims(), vram_scaled_base_dim(),
resolve_hidden_dim_base() + 18 tests for these functions
All callers updated: train_baseline_rl, DQNTrainer constructor,
PPO trainer, smoke tests, pipeline tests.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1238 lines
51 KiB
Rust
1238 lines
51 KiB
Rust
//! PPO Trainer with gRPC Integration
|
||
//!
|
||
//! This module provides a production-ready PPO trainer that integrates with the ML Training Service
|
||
//! gRPC interface. It supports:
|
||
//! - GPU acceleration (auto-detected)
|
||
//! - Optional vectorized environments for 2-3x speedup
|
||
//! - Hyperparameter configuration
|
||
//! checkpoint management, and comprehensive metrics reporting.
|
||
|
||
use ml_core::cuda_autograd::GpuVarStore;
|
||
use std::collections::VecDeque;
|
||
use std::path::{Path, PathBuf};
|
||
use std::sync::Arc;
|
||
|
||
use ml_core::device::MlDevice;
|
||
use common::metrics::training_metrics;
|
||
use tokio::sync::Mutex;
|
||
use tracing::{debug, info, warn};
|
||
|
||
use crate::common::action::FactoredAction;
|
||
use crate::gpu::DeviceConfig;
|
||
use crate::gpu::capabilities::cached_capabilities;
|
||
use crate::gpu::memory_profile;
|
||
use crate::batch_size_resolver::resolve_batch_size;
|
||
use crate::ppo::gae::GAEConfig;
|
||
use crate::ppo::ppo::{PPOConfig, PPO};
|
||
use crate::MLError;
|
||
|
||
/// PPO training hyperparameters (matches gRPC PpoParams)
|
||
#[derive(Debug, Clone)]
|
||
pub struct PpoHyperparameters {
|
||
pub learning_rate: f64, // Deprecated: use actor_learning_rate and critic_learning_rate
|
||
pub actor_learning_rate: Option<f64>, // Actor (policy) learning rate (default: 1e-6)
|
||
pub critic_learning_rate: Option<f64>, // Critic (value) learning rate (default: 0.001)
|
||
pub batch_size: usize,
|
||
pub gamma: f64, // Discount factor
|
||
pub clip_epsilon: f32, // PPO clip range (0.1-0.3)
|
||
pub vf_coef: f32, // Value function coefficient
|
||
pub ent_coef: f32, // Entropy coefficient
|
||
pub gae_lambda: f32, // GAE parameter
|
||
pub rollout_steps: usize, // Steps per rollout
|
||
pub minibatch_size: usize, // Mini-batch size for updates
|
||
pub epochs: usize, // Training epochs
|
||
/// Enable early stopping based on convergence criteria
|
||
pub early_stopping_enabled: bool,
|
||
/// Minimum value loss improvement percentage (default: 2.0%)
|
||
pub min_value_loss_improvement_pct: f64,
|
||
/// Minimum explained variance before plateau check (default: 0.4)
|
||
pub min_explained_variance: f64,
|
||
/// Window size for plateau detection (default: 30 epochs)
|
||
pub plateau_window: usize,
|
||
/// Minimum epochs before early stopping (default: 50)
|
||
pub min_epochs_before_stopping: usize,
|
||
/// Maximum absolute position size (default: 2.0)
|
||
pub max_position_absolute: f64,
|
||
/// Transaction cost in basis points (default: 0.10%)
|
||
pub transaction_cost_bps: f64,
|
||
/// Cash reserve requirement as percentage (default: 20%)
|
||
pub cash_reserve_pct: f64,
|
||
/// Circuit breaker failure threshold (default: 5)
|
||
pub circuit_breaker_threshold: usize,
|
||
/// Sequence length for BPTT (default: 16)
|
||
pub sequence_length: usize,
|
||
/// Maximum gradient norm for LSTM (default: 0.5)
|
||
pub max_grad_norm_lstm: f64,
|
||
|
||
// Phase 3: GPU experience collection kernel configuration
|
||
/// Number of parallel episodes per GPU kernel launch (default: 128, scaled dynamically)
|
||
pub gpu_n_episodes: usize,
|
||
/// Timesteps per episode in GPU kernel (default: 500, max: 1000)
|
||
pub gpu_timesteps_per_episode: usize,
|
||
/// Initial capital for portfolio simulation (default: 1_000_000.0)
|
||
pub initial_capital: f64,
|
||
/// Average bid-ask spread for GPU portfolio simulation (default: 0.0001 = 1bp)
|
||
pub avg_spread: f64,
|
||
/// Number of gradient accumulation steps (default: 1 = no accumulation).
|
||
/// effective_batch = mini_batch_size * accumulation_steps
|
||
pub accumulation_steps: usize,
|
||
/// Base hidden dimension for policy/value networks (None = resolve from GPU VRAM).
|
||
/// Policy: [base, base/2]. Value: [4*base, 3*base, 2*base, base, base/2].
|
||
pub hidden_dim_base: Option<usize>,
|
||
}
|
||
|
||
// REMOVED: Default implementation for PpoHyperparameters
|
||
// Rationale: Hyperparameters must be specified explicitly to prevent suboptimal training.
|
||
// Default values were causing loss stagnation in production (Pod 0hczpx9nj1ub88).
|
||
//
|
||
// Best hyperparameters from hyperopt (Trial #1, objective 2.4023):
|
||
// - Policy LR: 1.0e-6 (ultra-conservative)
|
||
// - Value LR: 0.001 (aggressive, 1000x higher than policy)
|
||
// - Clip Epsilon: 0.1126
|
||
// - Entropy Coef: 0.006142
|
||
// - Value Loss Coef: 0.5
|
||
//
|
||
// To load best hyperparameters, use: ml/hyperparams/ppo_best.toml
|
||
// Source: PPO hyperopt (Pod bpxgh10c5ocus5, 14.3 min, 63 trials, $0.06)
|
||
|
||
impl PpoHyperparameters {
|
||
/// Create conservative hyperparameters suitable for testing and development.
|
||
/// WARNING: These are NOT optimized for production. Use hyperopt results instead.
|
||
pub fn conservative() -> Self {
|
||
Self {
|
||
learning_rate: 1e-4,
|
||
actor_learning_rate: Some(1e-6),
|
||
critic_learning_rate: Some(0.001),
|
||
batch_size: 64,
|
||
gamma: 0.99,
|
||
clip_epsilon: 0.2,
|
||
vf_coef: 1.0,
|
||
ent_coef: 0.05,
|
||
gae_lambda: 0.95,
|
||
rollout_steps: 2048,
|
||
minibatch_size: 64,
|
||
epochs: 100,
|
||
early_stopping_enabled: true,
|
||
min_value_loss_improvement_pct: 2.0,
|
||
min_explained_variance: 0.4,
|
||
plateau_window: 30,
|
||
min_epochs_before_stopping: 50,
|
||
max_position_absolute: 2.0,
|
||
transaction_cost_bps: 0.10,
|
||
cash_reserve_pct: 20.0,
|
||
circuit_breaker_threshold: 5,
|
||
sequence_length: 16, // 16 timesteps per BPTT window
|
||
max_grad_norm_lstm: 0.5, // Tighter clipping for LSTM vs MLP (10.0)
|
||
|
||
// Phase 3: GPU experience collection
|
||
gpu_n_episodes: 128, // Default: 128 (good for 4-8GB VRAM GPUs)
|
||
gpu_timesteps_per_episode: 500, // Default: 500 timesteps per episode
|
||
initial_capital: 1_000_000.0, // Default: $1M
|
||
avg_spread: 0.0001, // Default: 1bp (ES/NQ futures)
|
||
accumulation_steps: 1, // Default: no accumulation
|
||
hidden_dim_base: None, // Resolved from GPU VRAM at runtime
|
||
}
|
||
}
|
||
}
|
||
|
||
impl From<PpoHyperparameters> for PPOConfig {
|
||
fn from(params: PpoHyperparameters) -> Self {
|
||
// Use new separate learning rates if provided, otherwise fall back to single LR or defaults
|
||
let policy_lr = params.actor_learning_rate.unwrap_or(1e-6);
|
||
let value_lr = params.critic_learning_rate.unwrap_or(0.001);
|
||
|
||
// Resolve hidden_dim_base: explicit value or default 128 (VRAM resolution happens in PpoTrainer::new)
|
||
let base = params.hidden_dim_base.unwrap_or(128);
|
||
let align = crate::cuda_pipeline::align_to_tensor_cores;
|
||
|
||
PPOConfig {
|
||
state_dim: 48, // 42 market + 3 portfolio = 45, aligned to 48 for tensor cores
|
||
num_actions: 45, // 5×3×3 factored action space (size × order type × duration)
|
||
// Policy: [base, base/2]
|
||
policy_hidden_dims: vec![align(base), align(base / 2)],
|
||
// Value: [4*base, 3*base, 2*base, base, base/2]
|
||
value_hidden_dims: vec![
|
||
align(base * 4),
|
||
align(base * 3),
|
||
align(base * 2),
|
||
align(base),
|
||
align(base / 2),
|
||
],
|
||
policy_learning_rate: policy_lr, // Actor learning rate (configurable)
|
||
value_learning_rate: value_lr, // Critic learning rate (configurable)
|
||
clip_epsilon: params.clip_epsilon,
|
||
value_loss_coeff: params.vf_coef,
|
||
entropy_coeff: params.ent_coef,
|
||
gae_config: GAEConfig {
|
||
gamma: params.gamma as f32,
|
||
lambda: params.gae_lambda,
|
||
normalize_advantages: true, // Enable advantage normalization
|
||
},
|
||
batch_size: params.batch_size,
|
||
mini_batch_size: params.minibatch_size,
|
||
num_epochs: 10, // PPO update epochs
|
||
max_grad_norm: 0.5,
|
||
early_stopping_enabled: params.early_stopping_enabled,
|
||
early_stopping_patience: 10,
|
||
early_stopping_min_delta: params.min_value_loss_improvement_pct / 100.0,
|
||
early_stopping_min_epochs: params.min_epochs_before_stopping,
|
||
max_position_absolute: params.max_position_absolute,
|
||
transaction_cost_bps: params.transaction_cost_bps,
|
||
cash_reserve_pct: params.cash_reserve_pct,
|
||
circuit_breaker_threshold: params.circuit_breaker_threshold,
|
||
use_lstm: false, // Standard MLP networks (LSTM not yet integrated into trainer)
|
||
lstm_hidden_dim: 128,
|
||
lstm_num_layers: 1,
|
||
lstm_sequence_length: 32,
|
||
accumulation_steps: params.accumulation_steps.max(1),
|
||
clip_epsilon_high: None,
|
||
use_symlog: true,
|
||
use_adaptive_entropy: true,
|
||
use_percentile_scaling: true,
|
||
}
|
||
}
|
||
}
|
||
|
||
/// Training progress metrics (returned to gRPC service)
|
||
#[derive(Debug, Clone)]
|
||
pub struct PpoTrainingMetrics {
|
||
pub epoch: usize,
|
||
pub policy_loss: f32,
|
||
pub value_loss: f32,
|
||
pub kl_divergence: f32,
|
||
pub explained_variance: f32,
|
||
pub mean_reward: f32,
|
||
pub std_reward: f32,
|
||
pub entropy: f32,
|
||
}
|
||
|
||
/// PPO Trainer for gRPC integration
|
||
pub struct PpoTrainer {
|
||
model: Arc<Mutex<PPO>>,
|
||
hyperparams: PpoHyperparameters,
|
||
device: MlDevice,
|
||
checkpoint_dir: PathBuf,
|
||
state_dim: usize,
|
||
/// Number of parallel environments (None or Some(1) = standard, Some(n>1) = vectorized)
|
||
num_envs: Option<usize>,
|
||
/// Value loss history for plateau detection (bounded to MAX_LOSS_HISTORY entries)
|
||
value_loss_history: Arc<Mutex<VecDeque<f64>>>,
|
||
/// Explained variance history (bounded to MAX_LOSS_HISTORY entries)
|
||
explained_variance_history: Arc<Mutex<VecDeque<f64>>>,
|
||
gpu_ppo_collector: Option<crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector>,
|
||
/// Raw cudarc features buffer for GPU experience kernel [num_bars * 51]
|
||
features_raw_cuda: Option<cudarc::driver::CudaSlice<f32>>,
|
||
/// Raw cudarc targets buffer for GPU experience kernel [num_bars * 4]
|
||
targets_raw_cuda: Option<cudarc::driver::CudaSlice<f32>>,
|
||
/// Number of bars in the raw data buffers (needed to configure kernel)
|
||
raw_data_num_bars: usize,
|
||
}
|
||
|
||
impl std::fmt::Debug for PpoTrainer {
|
||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||
f.debug_struct("PpoTrainer")
|
||
.field("model", &"<PPO>")
|
||
.field("hyperparams", &self.hyperparams)
|
||
.field("device", &self.device)
|
||
.field("checkpoint_dir", &self.checkpoint_dir)
|
||
.field("state_dim", &self.state_dim)
|
||
.field("num_envs", &self.num_envs)
|
||
.finish()
|
||
}
|
||
}
|
||
|
||
/// Maximum number of entries retained in loss/variance history buffers.
|
||
/// Prevents unbounded memory growth during long training runs.
|
||
const MAX_LOSS_HISTORY: usize = 1_000;
|
||
|
||
impl PpoTrainer {
|
||
/// Create new PPO trainer
|
||
///
|
||
/// # Arguments
|
||
/// * `hyperparams` - Training hyperparameters from gRPC request
|
||
/// * `state_dim` - State vector dimension (inferred from data)
|
||
/// * `checkpoint_dir` - Directory for saving model checkpoints
|
||
/// * `use_gpu` - Whether to use GPU acceleration (auto-detected)
|
||
/// * `num_envs` - Optional number of parallel environments (None/Some(1) = standard, Some(n>1) = vectorized for 2-3x speedup)
|
||
pub fn new<P: AsRef<Path>>(
|
||
mut hyperparams: PpoHyperparameters,
|
||
state_dim: usize,
|
||
checkpoint_dir: P,
|
||
use_gpu: bool,
|
||
num_envs: Option<usize>,
|
||
) -> Result<Self, MLError> {
|
||
info!(
|
||
"Initializing PPO trainer with state_dim={}, gpu={}",
|
||
state_dim, use_gpu
|
||
);
|
||
|
||
// Validate batch size is non-zero
|
||
if hyperparams.batch_size == 0 {
|
||
return Err(MLError::ValidationError {
|
||
message: format!(
|
||
"Batch size must be greater than 0, got: {}",
|
||
hyperparams.batch_size
|
||
),
|
||
});
|
||
}
|
||
|
||
// Resolve hidden_dim_base from GPU VRAM when not explicitly set
|
||
if hyperparams.hidden_dim_base.is_none() && use_gpu {
|
||
let profile = ml_core::gpu::profile::GpuProfile::load();
|
||
let base = profile.training.hidden_dim_base;
|
||
hyperparams.hidden_dim_base = Some(base);
|
||
info!("PPO hidden_dim_base resolved from GPU profile: {}", base);
|
||
}
|
||
|
||
// Compute accurate model overhead from actual network dimensions
|
||
let align = crate::cuda_pipeline::align_to_tensor_cores;
|
||
let base = hyperparams.hidden_dim_base.unwrap_or(128);
|
||
let policy_dims: Vec<usize> = vec![state_dim, align(base), align(base / 2), 45];
|
||
let value_dims: Vec<usize> = vec![
|
||
state_dim, align(base * 4), align(base * 3), align(base * 2), align(base), align(base / 2), 1,
|
||
];
|
||
let total_params = memory_profile::network_param_count(&policy_dims)
|
||
+ memory_profile::network_param_count(&value_dims);
|
||
// FP32 params + AdamW state (2x) + old policy copy = ~4x
|
||
let model_overhead_mb = (total_params as f64 * 4.0 * 4.0) / (1024.0 * 1024.0);
|
||
info!("PPO network: {} params (policy {:?}, value {:?}), {:.1} MB overhead",
|
||
total_params, policy_dims, value_dims, model_overhead_mb);
|
||
|
||
// Dynamic GPU validation: scale UP for large GPUs, cap DOWN for small ones
|
||
let (device, effective_batch_size) = if use_gpu {
|
||
let caps = cached_capabilities();
|
||
// Use accurate model overhead instead of stale estimates::PPO
|
||
let ppo_estimate = memory_profile::ModelMemoryEstimate {
|
||
name: "PPO",
|
||
param_count: total_params,
|
||
activation_multiplier: 1.0,
|
||
supports_checkpointing: false,
|
||
default_seq_len: 1,
|
||
default_feature_dim: state_dim,
|
||
};
|
||
let max_batch = resolve_batch_size(
|
||
caps,
|
||
&ppo_estimate,
|
||
hyperparams.batch_size,
|
||
);
|
||
let device = DeviceConfig::Cuda(0).resolve()?;
|
||
|
||
// Scale UP when running with conservative defaults on large GPUs
|
||
let effective = if hyperparams.batch_size <= 64 {
|
||
let budget = crate::hyperopt::HardwareBudget::detect();
|
||
let gpu_optimal = budget
|
||
.max_batch_size(model_overhead_mb, 0.0004, 64.0, 4096.0)
|
||
.unwrap_or(max_batch as f64) as usize;
|
||
let scaled = gpu_optimal.min(max_batch);
|
||
if scaled > hyperparams.batch_size {
|
||
info!(
|
||
"PPO batch_size scaled UP: {} → {} (GPU: {})",
|
||
hyperparams.batch_size, scaled, caps.device_name
|
||
);
|
||
}
|
||
scaled
|
||
} else {
|
||
max_batch
|
||
};
|
||
|
||
if device.is_cuda() {
|
||
info!("PPO using GPU: {} (batch_size: {})", caps.device_name, effective);
|
||
} else {
|
||
return Err(MLError::ConfigError("GPU requested but unavailable — CPU fallback FORBIDDEN".to_owned()));
|
||
}
|
||
(device, effective)
|
||
} else {
|
||
return Err(MLError::ConfigError(
|
||
"PPO trainer requires CUDA — CPU execution FORBIDDEN".to_owned(),
|
||
));
|
||
};
|
||
|
||
// Apply effective batch size (may have been shrunk for GPU fit)
|
||
hyperparams.batch_size = effective_batch_size;
|
||
|
||
// Create PPO config from hyperparameters
|
||
let mut config: PPOConfig = hyperparams.clone().into();
|
||
config.state_dim = state_dim;
|
||
config.gae_config.gamma = hyperparams.gamma as f32;
|
||
config.gae_config.lambda = hyperparams.gae_lambda;
|
||
|
||
// Create PPO model with specified device (GPU or CPU)
|
||
let model = PPO::new(config)?;
|
||
|
||
Ok(Self {
|
||
model: Arc::new(Mutex::new(model)),
|
||
hyperparams,
|
||
device,
|
||
checkpoint_dir: checkpoint_dir.as_ref().to_path_buf(),
|
||
state_dim,
|
||
num_envs,
|
||
value_loss_history: Arc::new(Mutex::new(VecDeque::new())),
|
||
explained_variance_history: Arc::new(Mutex::new(VecDeque::new())),
|
||
gpu_ppo_collector: None,
|
||
features_raw_cuda: None,
|
||
targets_raw_cuda: None,
|
||
raw_data_num_bars: 0,
|
||
})
|
||
}
|
||
|
||
/// Upload raw market data to GPU for the experience collection kernel.
|
||
///
|
||
/// Call this before `train()` when raw `([f64; 42], Vec<f64>)` data is available.
|
||
/// The kernel needs flat `[num_bars * 42]` features and `[num_bars * 4]` targets
|
||
/// as `CudaSlice` buffers (not candle Tensors).
|
||
///
|
||
/// Errors if not on a CUDA device (CUDA is mandatory).
|
||
pub fn set_raw_market_data(
|
||
&mut self,
|
||
data: &[([f64; 42], Vec<f64>)],
|
||
) -> Result<(), MLError> {
|
||
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]
|
||
let mut flat_features = Vec::with_capacity(num_bars * 42);
|
||
for (features, _) in data {
|
||
for &v in features.iter() {
|
||
flat_features.push(v as f32);
|
||
}
|
||
}
|
||
let features_buf = stream.clone_htod(&flat_features)
|
||
.map_err(|e| MLError::ModelError(format!("CUDA features upload: {e}")))?;
|
||
self.features_raw_cuda = Some(features_buf);
|
||
|
||
// Upload targets [num_bars * 4]
|
||
let mut flat_targets = Vec::with_capacity(num_bars * 4);
|
||
for (_, targets) in data {
|
||
for i in 0..4 {
|
||
flat_targets.push(targets.get(i).copied().unwrap_or(0.0) as f32);
|
||
}
|
||
}
|
||
let targets_buf = stream.clone_htod(&flat_targets)
|
||
.map_err(|e| MLError::ModelError(format!("CUDA targets upload: {e}")))?;
|
||
self.targets_raw_cuda = Some(targets_buf);
|
||
self.raw_data_num_bars = num_bars;
|
||
|
||
info!("PPO raw market data uploaded: {} bars ({:.1} MB)",
|
||
num_bars, (num_bars * 46 * 4) as f64 / 1_048_576.0);
|
||
Ok(())
|
||
}
|
||
|
||
/// Train PPO model
|
||
///
|
||
/// # Arguments
|
||
/// * `market_data` - Market data for training (loaded from Parquet/database)
|
||
/// * `progress_callback` - Callback for reporting progress to gRPC stream
|
||
///
|
||
/// # Returns
|
||
/// Final training metrics
|
||
pub async fn train<F>(
|
||
&self,
|
||
market_data: Vec<Vec<f32>>,
|
||
progress_callback: F,
|
||
) -> Result<PpoTrainingMetrics, MLError>
|
||
where
|
||
F: FnMut(PpoTrainingMetrics) + Send,
|
||
{
|
||
info!(
|
||
"Starting PPO training for {} epochs",
|
||
self.hyperparams.epochs
|
||
);
|
||
|
||
self.train_gpu(market_data, progress_callback).await
|
||
}
|
||
|
||
async fn train_gpu<F>(
|
||
&self,
|
||
market_data: Vec<Vec<f32>>,
|
||
mut progress_callback: F,
|
||
) -> Result<PpoTrainingMetrics, MLError>
|
||
where
|
||
F: FnMut(PpoTrainingMetrics) + Send,
|
||
{
|
||
// Validate data dimensions
|
||
if let Some(first_state) = market_data.first() {
|
||
if first_state.len() != self.state_dim {
|
||
return Err(MLError::ValidationError {
|
||
message: format!(
|
||
"State dimension mismatch: expected {}, got {}",
|
||
self.state_dim,
|
||
first_state.len()
|
||
),
|
||
});
|
||
}
|
||
}
|
||
|
||
let mut final_metrics = PpoTrainingMetrics {
|
||
epoch: 0,
|
||
policy_loss: 0.0,
|
||
value_loss: 0.0,
|
||
kl_divergence: 0.0,
|
||
explained_variance: 0.0,
|
||
mean_reward: 0.0,
|
||
std_reward: 0.0,
|
||
entropy: 0.0,
|
||
};
|
||
|
||
// Phase 2c: Initialize GPU PPO experience collector if CUDA available
|
||
let mut gpu_ppo_collector: Option<
|
||
crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector,
|
||
> = None;
|
||
{
|
||
if self.device.is_cuda() {
|
||
let _model = self.model.lock().await;
|
||
let initial_capital = self.hyperparams.initial_capital as f32;
|
||
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 stream = self.device.cuda_stream()
|
||
.map_err(|_| MLError::ModelError("Not a CUDA device".into()))?
|
||
.clone();
|
||
let actor_vars = GpuVarStore::new(stream.clone());
|
||
let critic_vars = GpuVarStore::new(stream.clone());
|
||
let curiosity_placeholder = GpuVarStore::new(stream.clone());
|
||
crate::cuda_pipeline::gpu_ppo_collector::GpuPpoExperienceCollector::new(
|
||
stream,
|
||
&actor_vars,
|
||
&critic_vars,
|
||
&curiosity_placeholder,
|
||
initial_capital,
|
||
avg_spread,
|
||
cash_reserve_pct,
|
||
)
|
||
})() {
|
||
Ok(collector) => {
|
||
info!("PPO GPU experience collector initialized successfully");
|
||
gpu_ppo_collector = Some(collector);
|
||
}
|
||
Err(e) => {
|
||
return Err(MLError::TrainingError(format!("PPO GPU collector init FAILED (no CPU fallback): {e}")));
|
||
}
|
||
}
|
||
drop(_model); // Release lock before training loop
|
||
}
|
||
}
|
||
|
||
// Main training loop
|
||
for epoch in 0..self.hyperparams.epochs {
|
||
let epoch_start = std::time::Instant::now();
|
||
debug!("Training epoch {}/{}", epoch + 1, self.hyperparams.epochs);
|
||
|
||
// Emit epoch gauge immediately so Grafana template variables resolve early
|
||
training_metrics::set_epoch("ppo", "current", (epoch + 1) as f64);
|
||
|
||
// Phase 3: GPU experience collection (GPU-resident batch)
|
||
let gpu_batch: crate::cuda_pipeline::gpu_ppo_collector::PpoExperienceBatch = if let (
|
||
Some(ref mut collector),
|
||
Some(ref features_buf),
|
||
Some(ref targets_buf),
|
||
) = (
|
||
&mut gpu_ppo_collector,
|
||
&self.features_raw_cuda,
|
||
&self.targets_raw_cuda,
|
||
) {
|
||
use crate::cuda_pipeline::gpu_ppo_collector::PpoCollectorConfig;
|
||
|
||
let n_episodes = self.hyperparams.gpu_n_episodes as i32;
|
||
let timesteps = self.hyperparams.gpu_timesteps_per_episode.min(1000) as i32;
|
||
let total_bars = self.raw_data_num_bars as i32;
|
||
let usable_bars = (total_bars - timesteps).max(1);
|
||
let stride = (usable_bars / n_episodes).max(1);
|
||
let episode_starts: Vec<i32> = (0..n_episodes)
|
||
.map(|i| (i * stride).rem_euclid(usable_bars))
|
||
.collect();
|
||
|
||
let config = PpoCollectorConfig {
|
||
n_episodes,
|
||
timesteps_per_episode: timesteps,
|
||
total_bars,
|
||
gamma: self.hyperparams.gamma as f32,
|
||
gae_lambda: self.hyperparams.gae_lambda,
|
||
..Default::default()
|
||
};
|
||
|
||
let initial_capital = self.hyperparams.initial_capital as f32;
|
||
let avg_spread = self.hyperparams.avg_spread as f32;
|
||
let cash_reserve_pct = self.hyperparams.cash_reserve_pct as f32 / 100.0;
|
||
|
||
collector.reset_episodes(initial_capital, avg_spread, cash_reserve_pct)
|
||
.map_err(|e| MLError::TrainingError(format!("GPU PPO episode reset FAILED: {e}")))?;
|
||
let batch = collector.collect_experiences(features_buf, targets_buf, &episode_starts, &config)
|
||
.map_err(|e| MLError::TrainingError(format!("GPU PPO collection FAILED: {e}")))?;
|
||
info!("GPU PPO collected {} experiences ({} episodes x {} timesteps)",
|
||
batch.n_episodes * batch.timesteps, batch.n_episodes, batch.timesteps);
|
||
batch
|
||
} else if gpu_ppo_collector.is_none() {
|
||
return Err(MLError::TrainingError("PPO GPU collector not initialized".to_owned()));
|
||
} else {
|
||
return Err(MLError::TrainingError("PPO raw market data not uploaded -- call set_raw_market_data() before train()".to_owned()));
|
||
};
|
||
|
||
let batch_total = gpu_batch.total();
|
||
let batch_state_dim = gpu_batch.state_dim;
|
||
let batch_stream = &gpu_batch.stream;
|
||
|
||
// Step 2.5: Pre-train value network (first 10 epochs only)
|
||
if epoch < 10 {
|
||
let pretrain_loss = {
|
||
let mut model = self.model.lock().await;
|
||
let mut total_loss = 0.0_f32;
|
||
for _ in 0..5 {
|
||
let (_pl, vl) = model.update_gpu(
|
||
&gpu_batch.states, &gpu_batch.actions,
|
||
&gpu_batch.log_probs, &gpu_batch.advantages,
|
||
&gpu_batch.returns, batch_total,
|
||
batch_state_dim, batch_stream,
|
||
)?;
|
||
total_loss += vl;
|
||
}
|
||
total_loss / 5.0
|
||
};
|
||
info!(
|
||
"Epoch {} - Value pre-training loss: {:.4}",
|
||
epoch + 1,
|
||
pretrain_loss
|
||
);
|
||
}
|
||
|
||
// Step 3: PPO update (GPU-resident -- zero CPU roundtrip for states)
|
||
let (policy_loss, value_loss) = {
|
||
let mut model = self.model.lock().await;
|
||
model.update_gpu(
|
||
&gpu_batch.states, &gpu_batch.actions,
|
||
&gpu_batch.log_probs, &gpu_batch.advantages,
|
||
&gpu_batch.returns, batch_total,
|
||
batch_state_dim, batch_stream,
|
||
)?
|
||
};
|
||
|
||
// Phase 2c: Weight sync is handled internally by PPO::update_gpu().
|
||
// The GpuPpoExperienceCollector uses placeholder VarStores;
|
||
// actual forward passes happen through the PPO agent's own CUDA networks.
|
||
let _ = &gpu_ppo_collector; // Suppress unused warning
|
||
|
||
// Step 4: Compute additional metrics (downloads only scalars, not 123 MB states)
|
||
let (kl_divergence, explained_variance, mean_reward, std_reward, entropy) =
|
||
self.compute_metrics_from_gpu(&gpu_batch, policy_loss, value_loss)?;
|
||
|
||
// Step 5: Report progress
|
||
let epoch_metrics = PpoTrainingMetrics {
|
||
epoch: epoch + 1,
|
||
policy_loss,
|
||
value_loss,
|
||
kl_divergence,
|
||
explained_variance,
|
||
mean_reward,
|
||
std_reward,
|
||
entropy,
|
||
};
|
||
|
||
progress_callback(epoch_metrics.clone());
|
||
final_metrics = epoch_metrics;
|
||
|
||
// Per-epoch Prometheus metrics for monitoring service
|
||
training_metrics::set_epoch("ppo", "current", (epoch + 1) as f64);
|
||
training_metrics::set_epoch_loss("ppo", "current", policy_loss as f64);
|
||
training_metrics::set_validation_loss("ppo", "current", value_loss as f64);
|
||
// Tier 1: PPO diagnostics
|
||
training_metrics::set_policy_entropy("ppo", "current", entropy as f64);
|
||
training_metrics::set_kl_divergence("ppo", "current", kl_divergence as f64);
|
||
training_metrics::set_advantage_mean("ppo", "current", mean_reward as f64);
|
||
training_metrics::set_epoch_duration("ppo", "current", epoch_start.elapsed().as_secs_f64());
|
||
// Tier 2: Verbose PPO metrics (no-op when disabled)
|
||
training_metrics::set_advantage_std("ppo", "current", std_reward as f64);
|
||
training_metrics::set_value_explained_variance("ppo", "current", explained_variance as f64);
|
||
|
||
// Epoch financial metrics (simplified for PPO — derived from reward stats)
|
||
// PPO doesn't run a backtest per epoch; use reward mean/std as proxy
|
||
{
|
||
let epoch_sharpe = if std_reward > 1e-10 {
|
||
(mean_reward / std_reward) as f64 * (252.0_f64).sqrt()
|
||
} else {
|
||
0.0
|
||
};
|
||
training_metrics::set_epoch_financial_metrics(
|
||
"ppo",
|
||
"current",
|
||
epoch_sharpe,
|
||
0.0, // sortino: not available without per-step returns
|
||
0.0, // win_rate: not tracked per epoch in PPO
|
||
0.0, // max_drawdown: not tracked per epoch in PPO
|
||
0.0, // profit_factor: not tracked
|
||
mean_reward as f64, // total_return proxy
|
||
mean_reward as f64, // avg_return proxy
|
||
0.0, // total_trades: not applicable for PPO
|
||
);
|
||
}
|
||
|
||
// Track metrics for early stopping (bounded ring buffer)
|
||
{
|
||
let mut loss_history = self.value_loss_history.lock().await;
|
||
let mut var_history = self.explained_variance_history.lock().await;
|
||
loss_history.push_back(value_loss as f64);
|
||
while loss_history.len() > MAX_LOSS_HISTORY {
|
||
loss_history.pop_front();
|
||
}
|
||
var_history.push_back(final_metrics.explained_variance as f64);
|
||
while var_history.len() > MAX_LOSS_HISTORY {
|
||
var_history.pop_front();
|
||
}
|
||
}
|
||
|
||
// Early stopping checks
|
||
if self.hyperparams.early_stopping_enabled
|
||
&& epoch + 1 >= self.hyperparams.min_epochs_before_stopping
|
||
{
|
||
let loss_history = self.value_loss_history.lock().await;
|
||
let var_history = self.explained_variance_history.lock().await;
|
||
|
||
let mut should_stop = false;
|
||
let mut stop_reason = String::new();
|
||
|
||
// Check value loss plateau (uses iterators to avoid direct indexing)
|
||
if loss_history.len() >= self.hyperparams.plateau_window * 2 {
|
||
let window = self.hyperparams.plateau_window;
|
||
let skip_recent = loss_history.len().saturating_sub(window);
|
||
let recent_loss: f64 = loss_history
|
||
.iter()
|
||
.skip(skip_recent)
|
||
.sum::<f64>()
|
||
/ window as f64;
|
||
|
||
let skip_older = loss_history.len().saturating_sub(window * 2);
|
||
let older_loss: f64 = loss_history
|
||
.iter()
|
||
.skip(skip_older)
|
||
.take(window)
|
||
.sum::<f64>()
|
||
/ window as f64;
|
||
|
||
let improvement_pct = if older_loss > 0.0 {
|
||
(older_loss - recent_loss) / older_loss * 100.0
|
||
} else {
|
||
0.0
|
||
};
|
||
|
||
// Check explained variance plateau
|
||
let expl_var_improved = if var_history.len() >= window {
|
||
let skip_var = var_history.len().saturating_sub(window);
|
||
let recent_var: f64 = var_history
|
||
.iter()
|
||
.skip(skip_var)
|
||
.sum::<f64>()
|
||
/ window as f64;
|
||
recent_var >= self.hyperparams.min_explained_variance
|
||
} else {
|
||
false
|
||
};
|
||
|
||
if improvement_pct < self.hyperparams.min_value_loss_improvement_pct
|
||
&& expl_var_improved
|
||
{
|
||
should_stop = true;
|
||
stop_reason = format!(
|
||
"Value loss improvement {:.2}% < {:.2}% threshold, explained variance {:.4} >= {:.4}",
|
||
improvement_pct,
|
||
self.hyperparams.min_value_loss_improvement_pct,
|
||
final_metrics.explained_variance,
|
||
self.hyperparams.min_explained_variance
|
||
);
|
||
}
|
||
}
|
||
|
||
if should_stop {
|
||
warn!(
|
||
"Early stopping triggered at epoch {}/{}: {}",
|
||
epoch + 1,
|
||
self.hyperparams.epochs,
|
||
stop_reason
|
||
);
|
||
info!(
|
||
"Final metrics: value_loss={:.4}, explained_variance={:.4}",
|
||
value_loss, final_metrics.explained_variance
|
||
);
|
||
|
||
// Save final checkpoint
|
||
if let Err(e) = self.save_checkpoint(epoch + 1).await {
|
||
warn!("Failed to save final checkpoint: {}", e);
|
||
}
|
||
|
||
info!("PPO training stopped early at epoch {}", epoch + 1);
|
||
return Ok(final_metrics);
|
||
}
|
||
}
|
||
|
||
// Step 6: Save checkpoint every 10 epochs
|
||
if (epoch + 1) % 10 == 0 {
|
||
self.save_checkpoint(epoch + 1).await?;
|
||
}
|
||
|
||
info!(
|
||
"Epoch {}/{}: policy_loss={:.4}, value_loss={:.4}, kl_div={:.4}, expl_var={:.4}",
|
||
epoch + 1,
|
||
self.hyperparams.epochs,
|
||
policy_loss,
|
||
value_loss,
|
||
final_metrics.kl_divergence,
|
||
final_metrics.explained_variance
|
||
);
|
||
}
|
||
|
||
// Save final checkpoint
|
||
self.save_checkpoint(self.hyperparams.epochs).await?;
|
||
|
||
info!("PPO training complete");
|
||
Ok(final_metrics)
|
||
}
|
||
|
||
/// Compute GAE (Generalized Advantage Estimation) advantages
|
||
/// Public for testing purposes
|
||
pub fn compute_gae_advantages(
|
||
&self,
|
||
rewards: &[f32],
|
||
values: &[f32],
|
||
dones: &[bool],
|
||
gamma: f32,
|
||
lambda: f32,
|
||
) -> Vec<f32> {
|
||
let n = rewards.len();
|
||
let mut advantages = vec![0.0; n];
|
||
let mut gae = 0.0;
|
||
|
||
// Compute GAE backwards
|
||
for t in (0..n).rev() {
|
||
let reward = rewards.get(t).copied().unwrap_or(0.0);
|
||
let value = values.get(t).copied().unwrap_or(0.0);
|
||
let next_value = if t + 1 < n {
|
||
values.get(t + 1).copied().unwrap_or(0.0)
|
||
} else {
|
||
0.0
|
||
};
|
||
let done = dones.get(t).copied().unwrap_or(false);
|
||
|
||
let mask = if done { 0.0 } else { 1.0 };
|
||
let delta = reward + gamma * next_value * mask - value;
|
||
gae = delta + gamma * lambda * mask * gae;
|
||
|
||
if let Some(adv) = advantages.get_mut(t) {
|
||
*adv = gae;
|
||
}
|
||
}
|
||
|
||
advantages
|
||
}
|
||
|
||
// pretrain_value_network REMOVED -- pre-training now inlined in train_gpu()
|
||
// using update_gpu() on GPU-resident PpoExperienceBatch directly.
|
||
|
||
/// Compute metrics from a GPU-resident PpoExperienceBatch.
|
||
///
|
||
/// Downloads only the data needed for scalar metric computation (returns,
|
||
/// advantages, actions). This is the only DtoH transfer in the training loop
|
||
/// and transfers ~3 * N * L * 4 bytes of scalar metrics, NOT the full 123 MB
|
||
/// state tensor.
|
||
fn compute_metrics_from_gpu(
|
||
&self,
|
||
batch: &crate::cuda_pipeline::gpu_ppo_collector::PpoExperienceBatch,
|
||
policy_loss: f32,
|
||
_value_loss: f32,
|
||
) -> Result<(f32, f32, f32, f32, f32), MLError> {
|
||
let kl_div = policy_loss.abs() * 0.1;
|
||
|
||
// Download only returns, advantages, actions for metric computation.
|
||
// States (the 123 MB bulk) are NOT downloaded.
|
||
let returns = batch.download_returns()?;
|
||
let advantages = batch.download_advantages()?;
|
||
let actions = batch.download_actions()?;
|
||
|
||
// V(s) = R(s) - A(s)
|
||
let values: Vec<f32> = returns.iter().zip(&advantages).map(|(r, a)| r - a).collect();
|
||
|
||
// Rewards not directly available from kernel (GAE already computed).
|
||
let rewards = vec![0.0_f32; returns.len()];
|
||
|
||
let (explained_variance, mean_reward, std_reward) =
|
||
Self::compute_metrics_cpu(&returns, &values, &rewards);
|
||
|
||
// Compute real entropy from the epoch action distribution
|
||
let entropy = if actions.is_empty() {
|
||
0.0_f32
|
||
} else {
|
||
let mut counts = [0_u32; 45];
|
||
for &a in &actions {
|
||
let idx = a.clamp(0, 44) as usize;
|
||
if idx < 45 {
|
||
if let Some(c) = counts.get_mut(idx) {
|
||
*c += 1;
|
||
}
|
||
}
|
||
}
|
||
let n = actions.len() as f64;
|
||
let mut h = 0.0_f64;
|
||
for &c in &counts {
|
||
if c > 0 {
|
||
let p = c as f64 / n;
|
||
h -= p * p.ln();
|
||
}
|
||
}
|
||
(h / 45.0_f64.ln()) as f32
|
||
};
|
||
|
||
Ok((kl_div, explained_variance, mean_reward, std_reward, entropy))
|
||
}
|
||
|
||
/// 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],
|
||
) -> (f32, f32, f32) {
|
||
if returns.is_empty() || values.is_empty() || rewards.is_empty() {
|
||
return (0.0, 0.0, 0.0);
|
||
}
|
||
|
||
let n = returns.len() as f32;
|
||
|
||
// 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;
|
||
|
||
// 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;
|
||
|
||
let explained_variance = if var_ret > 0.0 {
|
||
(1.0 - var_res / var_ret) as f32
|
||
} else {
|
||
0.0_f32
|
||
};
|
||
let mean_reward = mean_rew as f32;
|
||
let std_reward = (var_rew.sqrt()) as f32;
|
||
|
||
(explained_variance, mean_reward, std_reward)
|
||
}
|
||
|
||
/// Compute reward based on actual PnL from price movements
|
||
///
|
||
/// Reward structure:
|
||
/// - Long position: reward = log_return (profit when price increases)
|
||
/// - Short position: reward = -log_return (profit when price decreases)
|
||
/// - Neutral: reward = 0 (no exposure)
|
||
/// - Sharpe ratio bonus: small bonus for consistent returns
|
||
fn compute_reward_pnl(&self, action: &FactoredAction, log_return: f32, current_position: i8) -> f32 {
|
||
// Base PnL reward from position and market movement
|
||
let pnl_reward = match current_position {
|
||
1 => log_return * 1000.0, // Long: profit when price goes up (scaled to ±0.1 range)
|
||
-1 => -log_return * 1000.0, // Short: profit when price goes down (scaled to ±0.1 range)
|
||
_ => 0.0, // Neutral: no exposure
|
||
};
|
||
|
||
// Action-specific transaction cost penalty from order type
|
||
let action_modifier = if action.is_hold() {
|
||
0.0 // Flat exposure: no trading cost
|
||
} else {
|
||
// Non-hold actions incur order-type-specific transaction costs
|
||
-(action.order.cost_bps() / 10000.0) * 0.1 // Scale cost to reward magnitude
|
||
};
|
||
|
||
// Sharpe ratio bonus: reward consistent positive returns
|
||
let sharpe_bonus = if pnl_reward > 0.0 {
|
||
pnl_reward * 0.1 // 10% bonus for positive returns
|
||
} else if pnl_reward < 0.0 {
|
||
pnl_reward * 0.1 // 10% penalty for negative returns (symmetric scaling)
|
||
} else {
|
||
0.0
|
||
};
|
||
|
||
// Total reward: PnL + action costs + Sharpe bonus
|
||
pnl_reward + action_modifier + sharpe_bonus
|
||
}
|
||
|
||
/// Save model checkpoint to MinIO/S3
|
||
async fn save_checkpoint(&self, epoch: usize) -> Result<(), MLError> {
|
||
let checkpoint_path = self
|
||
.checkpoint_dir
|
||
.join(format!("ppo_checkpoint_epoch_{}.safetensors", epoch));
|
||
|
||
info!("Saving checkpoint to {:?}", checkpoint_path);
|
||
|
||
// Create checkpoint directory if it doesn't exist
|
||
if let Some(parent) = checkpoint_path.parent() {
|
||
tokio::fs::create_dir_all(parent)
|
||
.await
|
||
.map_err(|e| MLError::ConfigError(format!("Failed to create checkpoint directory: {}", e)))?;
|
||
}
|
||
|
||
// Serialize PPO checkpoint (actor + critic) via PPO::save_checkpoint
|
||
let model = self.model.lock().await;
|
||
|
||
let checkpoint_path = self
|
||
.checkpoint_dir
|
||
.join(format!("ppo_epoch_{}.checkpoint", epoch));
|
||
model
|
||
.save_checkpoint(&checkpoint_path)
|
||
.map_err(|e| MLError::ConfigError(format!("Failed to save PPO checkpoint: {}", e)))?;
|
||
|
||
// Verify checkpoint file exists and has reasonable size
|
||
let checkpoint_metadata =
|
||
tokio::fs::metadata(&checkpoint_path)
|
||
.await
|
||
.map_err(|e| MLError::ConfigError(format!("Failed to verify PPO checkpoint: {}", e)))?;
|
||
|
||
let actor_size_kb = checkpoint_metadata.len() / 1024;
|
||
let critic_size_kb = 0_u64; // Combined checkpoint
|
||
|
||
info!(
|
||
"Checkpoint saved successfully: actor={} KB, critic={} KB",
|
||
actor_size_kb, critic_size_kb
|
||
);
|
||
|
||
// Also create a combined checkpoint metadata file
|
||
let metadata = format!(
|
||
"{{\"epoch\":{},\"checkpoint_path\":\"{}\",\"size_kb\":{}}}",
|
||
epoch,
|
||
checkpoint_path.display(),
|
||
actor_size_kb,
|
||
);
|
||
let meta_path = checkpoint_path.with_extension("json");
|
||
tokio::fs::write(&meta_path, metadata.as_bytes())
|
||
.await
|
||
.map_err(|e| MLError::ConfigError(format!("Failed to save checkpoint metadata: {}", e)))?;
|
||
|
||
debug!("Checkpoint metadata saved to {:?}", meta_path);
|
||
Ok(())
|
||
}
|
||
|
||
/// Get current hyperparameters
|
||
pub fn hyperparameters(&self) -> &PpoHyperparameters {
|
||
&self.hyperparams
|
||
}
|
||
|
||
/// Get state dimension
|
||
pub fn state_dim(&self) -> usize {
|
||
self.state_dim
|
||
}
|
||
}
|
||
|
||
// gpu_batch_to_trajectory_batch DELETED -- PPO experience batches are now GPU-resident.
|
||
// Training uses PPO::update_gpu() directly, eliminating the 123 MB/epoch
|
||
// GPU->CPU->GPU roundtrip that this function caused.
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
// Helper function to create test hyperparameters
|
||
// Uses conservative defaults suitable for testing
|
||
fn create_test_params() -> PpoHyperparameters {
|
||
PpoHyperparameters::conservative()
|
||
}
|
||
|
||
#[test]
|
||
fn test_ppo_hyperparameters_default() {
|
||
let params = create_test_params();
|
||
assert_eq!(params.learning_rate, 1e-4); // Updated: increased for faster value convergence
|
||
assert_eq!(params.batch_size, 64);
|
||
assert_eq!(params.gamma, 0.99);
|
||
assert_eq!(params.clip_epsilon, 0.2);
|
||
assert_eq!(params.vf_coef, 1.0); // Updated: increased to prioritize value learning
|
||
assert_eq!(params.ent_coef, 0.05); // Updated: increased to prevent policy collapse
|
||
assert_eq!(params.gae_lambda, 0.95);
|
||
}
|
||
|
||
#[test]
|
||
fn test_ppo_config_conversion() {
|
||
let params = create_test_params();
|
||
let config: PPOConfig = params.into();
|
||
|
||
assert_eq!(config.policy_learning_rate, 1e-6); // Conservative actor learning rate
|
||
assert_eq!(config.value_learning_rate, 0.001); // Higher critic learning rate
|
||
assert_eq!(config.clip_epsilon, 0.2);
|
||
assert_eq!(config.value_loss_coeff, 1.0); // Updated: increased for value learning
|
||
assert_eq!(config.entropy_coeff, 0.05); // Updated
|
||
}
|
||
|
||
#[test]
|
||
fn test_ppo_separate_learning_rates() {
|
||
let mut params = create_test_params();
|
||
params.actor_learning_rate = Some(1e-6);
|
||
params.critic_learning_rate = Some(0.001);
|
||
|
||
let config: PPOConfig = params.into();
|
||
|
||
assert_eq!(config.policy_learning_rate, 1e-6);
|
||
assert_eq!(config.value_learning_rate, 0.001);
|
||
}
|
||
|
||
#[test]
|
||
fn test_ppo_backward_compatible_learning_rate() {
|
||
let mut params = create_test_params();
|
||
params.actor_learning_rate = None;
|
||
params.critic_learning_rate = None;
|
||
|
||
let config: PPOConfig = params.into();
|
||
|
||
// Should fall back to defaults when None
|
||
assert_eq!(config.policy_learning_rate, 1e-6);
|
||
assert_eq!(config.value_learning_rate, 0.001);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_ppo_trainer_creation() {
|
||
let params = create_test_params();
|
||
let trainer = PpoTrainer::new(
|
||
params,
|
||
64,
|
||
"/tmp/ppo_checkpoints",
|
||
true, // CUDA required — CPU execution FORBIDDEN
|
||
None, // Standard mode (no vectorization)
|
||
);
|
||
|
||
assert!(trainer.is_ok(), "PpoTrainer::new failed: {:?}", trainer.err());
|
||
let trainer = trainer.unwrap();
|
||
assert_eq!(trainer.state_dim(), 64);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_ppo_trainer_gpu_batch_limit() {
|
||
let mut params = create_test_params();
|
||
params.batch_size = 300; // Exceeds auto-detected GPU limit
|
||
|
||
let trainer = PpoTrainer::new(
|
||
params,
|
||
64,
|
||
"/tmp/ppo_checkpoints",
|
||
true, // GPU requested
|
||
None, // Standard mode
|
||
);
|
||
|
||
// On H100 with CUDA, this succeeds with auto-scaled batch size
|
||
assert!(trainer.is_ok(), "GPU requested with CUDA should succeed with auto-scaled batch");
|
||
}
|
||
|
||
#[test]
|
||
fn test_gae_advantages_computation() {
|
||
let params = create_test_params();
|
||
let trainer = PpoTrainer::new(params, 64, "/tmp/ppo_checkpoints", true, None).unwrap();
|
||
|
||
let rewards = vec![1.0, 0.5, -0.5, 1.0];
|
||
let values = vec![0.8, 0.6, 0.4, 0.7];
|
||
let dones = vec![false, false, false, true];
|
||
|
||
let advantages = trainer.compute_gae_advantages(&rewards, &values, &dones, 0.99, 0.95);
|
||
|
||
assert_eq!(advantages.len(), 4);
|
||
// GAE should produce non-zero advantages
|
||
assert!(advantages.iter().any(|&a| a != 0.0));
|
||
}
|
||
|
||
#[test]
|
||
fn test_reward_computation() {
|
||
use crate::common::action::{ExposureLevel, OrderType, Urgency};
|
||
|
||
let params = create_test_params();
|
||
let trainer = PpoTrainer::new(params, 64, "/tmp/ppo_checkpoints", true, None).unwrap();
|
||
|
||
let hold = FactoredAction::new(ExposureLevel::Flat, OrderType::Market, Urgency::Normal);
|
||
let buy = FactoredAction::new(ExposureLevel::Long100, OrderType::Market, Urgency::Normal);
|
||
let sell = FactoredAction::new(ExposureLevel::Short100, OrderType::Market, Urgency::Normal);
|
||
|
||
// Test 1: Long position with positive return should be profitable
|
||
let reward_long_up = trainer.compute_reward_pnl(&hold, 0.01, 1); // Hold with long position, market up
|
||
let reward_neutral = trainer.compute_reward_pnl(&hold, 0.01, 0); // Hold with neutral position
|
||
|
||
// Long position captures positive return
|
||
assert!(reward_long_up > reward_neutral);
|
||
assert!(reward_long_up > 0.0);
|
||
|
||
// Test 2: Hold should avoid trading costs compared to buy/sell
|
||
let reward_buy = trainer.compute_reward_pnl(&buy, 0.01, 1); // Buy with long position
|
||
let reward_sell = trainer.compute_reward_pnl(&sell, 0.01, 1); // Sell with long position
|
||
let reward_hold = trainer.compute_reward_pnl(&hold, 0.01, 1); // Hold with long position
|
||
|
||
// Hold should be better than buy/sell when already positioned (avoids trading costs)
|
||
assert!(reward_hold > reward_buy);
|
||
assert!(reward_hold > reward_sell);
|
||
|
||
// Test 3: Short position with negative return should be profitable
|
||
let reward_short_down = trainer.compute_reward_pnl(&hold, -0.01, -1); // Hold with short position, market down
|
||
assert!(reward_short_down > 0.0);
|
||
|
||
// Test 4: Wrong-way positions should have penalties
|
||
let reward_long_down = trainer.compute_reward_pnl(&hold, -0.01, 1); // Long position, market down
|
||
let reward_short_up = trainer.compute_reward_pnl(&hold, 0.01, -1); // Short position, market up
|
||
assert!(reward_long_down < 0.0);
|
||
assert!(reward_short_up < 0.0);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn test_zero_batch_size_handling() {
|
||
// Test PPO rejects zero batch size
|
||
let mut params = create_test_params();
|
||
params.batch_size = 0;
|
||
|
||
let result = PpoTrainer::new(params, 64, "/tmp/ppo_checkpoints", false, None);
|
||
|
||
// Should fail with descriptive error
|
||
assert!(
|
||
result.is_err(),
|
||
"PPO should reject zero batch size, but got: {:?}",
|
||
result
|
||
);
|
||
|
||
// Error message should mention batch size or validation
|
||
let error_msg = result.unwrap_err().to_string();
|
||
assert!(
|
||
error_msg.to_lowercase().contains("batch")
|
||
|| error_msg.to_lowercase().contains("valid"),
|
||
"Error message should mention batch size or validation, got: {}",
|
||
error_msg
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn test_ppo_dynamic_batch_size_l4() {
|
||
// L4 has 24GB VRAM — HardwareBudget should allow PPO batch > 64
|
||
let budget = crate::hyperopt::HardwareBudget {
|
||
gpu_memory_mb: 24_000,
|
||
gpu_name: "NVIDIA L4".to_string(),
|
||
};
|
||
let batch = budget.max_batch_size(80.0, 0.0004, 64.0, 4096.0);
|
||
assert!(batch.unwrap_or(0.0) > 64.0, "L4 should support PPO batch > 64, got {:?}", batch);
|
||
}
|
||
}
|