#![allow( clippy::assertions_on_constants, clippy::assertions_on_result_states, clippy::clone_on_copy, clippy::decimal_literal_representation, clippy::doc_markdown, clippy::empty_line_after_doc_comments, clippy::field_reassign_with_default, clippy::get_unwrap, clippy::identity_op, clippy::inconsistent_digit_grouping, clippy::indexing_slicing, clippy::integer_division, clippy::len_zero, clippy::let_underscore_must_use, clippy::manual_div_ceil, clippy::manual_let_else, clippy::manual_range_contains, clippy::modulo_arithmetic, clippy::needless_range_loop, clippy::non_ascii_literal, clippy::redundant_clone, clippy::shadow_reuse, clippy::shadow_same, clippy::shadow_unrelated, clippy::single_match_else, clippy::str_to_string, clippy::string_slice, clippy::tests_outside_test_module, clippy::too_many_lines, clippy::unnecessary_wraps, clippy::unseparated_literal_suffix, clippy::use_debug, clippy::useless_vec, clippy::wildcard_enum_match_arm, clippy::else_if_without_else, clippy::expect_used, clippy::missing_const_for_fn, clippy::similar_names, clippy::type_complexity, clippy::collapsible_else_if, clippy::doc_lazy_continuation, clippy::items_after_test_module, clippy::map_clone, clippy::multiple_unsafe_ops_per_block, clippy::unwrap_or_default, clippy::assign_op_pattern, clippy::needless_borrow, clippy::println_empty_string, clippy::unnecessary_cast, clippy::used_underscore_binding, clippy::create_dir, clippy::implicit_saturating_sub, clippy::exit, clippy::expect_fun_call, clippy::too_many_arguments, clippy::unnecessary_map_or, clippy::unwrap_used, dead_code, unused_imports, unused_variables, clippy::cloned_ref_to_slice_refs, clippy::neg_multiply, clippy::while_let_loop, clippy::bool_assert_comparison, clippy::excessive_precision, clippy::trivially_copy_pass_by_ref, clippy::op_ref, clippy::redundant_closure, clippy::unnecessary_lazy_evaluations, clippy::if_then_some_else_none, clippy::unnecessary_to_owned, clippy::single_component_path_imports, )] //! Walk-forward RL training binary for DQN and PPO models. //! //! Trains models using expanding walk-forward windows on real OHLCV data loaded //! from Databento DBN files. Supports early stopping, checkpoint saving, and //! normalization statistics export for reproducible evaluation. //! //! # Usage //! //! ```bash //! SQLX_OFFLINE=true cargo run -p ml --example train_baseline -- \ //! --model both --epochs 50 \ //! --data-dir test_data/futures-baseline \ //! --output-dir ml/trained_models //! ``` #![allow(unused_crate_dependencies)] #![deny( clippy::unwrap_used, clippy::expect_used, clippy::panic, clippy::indexing_slicing )] use std::path::{Path, PathBuf}; use anyhow::{Context, Result}; use clap::Parser; use serde_json::Value; use tracing::{error, info, warn}; use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer}; use ml::trainers::ppo::{PpoHyperparameters, PpoTrainer, PpoTrainingMetrics}; #[allow(unreachable_pub)] mod baseline_common; use baseline_common::completion::{write_failure_marker, write_success_marker, CompletionMetrics}; use baseline_common::{load_all_bars, spread_cost_bps}; use common::metrics::{server as metrics_server, training_metrics as metrics}; use ml::features::extraction::{extract_ml_features, FeatureVector}; use ml_core::gpu::profile::GpuProfile; use ml::types::OHLCVBar; use ml::walk_forward::{ generate_walk_forward_indices_from_timestamps, WalkForwardConfig, }; // --------------------------------------------------------------------------- // CLI Arguments // --------------------------------------------------------------------------- /// Walk-forward training binary for DQN and PPO baseline models. #[derive(Parser, Debug)] #[command(name = "train_baseline_rl", about = "Train DQN/PPO with walk-forward RL windows")] struct Args { /// Which model(s) to train: "dqn", "ppo", or "both" #[arg(long, default_value = "both")] model: String, /// Maximum training epochs per fold #[arg(long, default_value_t = 50)] epochs: usize, // batch_size always auto-scaled from VRAM (no CLI override). /// Path to directory containing .dbn.zst files (env: FOXHUNT_DATA_DIR) #[arg(long, env = "FOXHUNT_DATA_DIR")] data_dir: PathBuf, /// Output directory for trained model checkpoints #[arg(long, default_value = "ml/trained_models")] output_dir: PathBuf, /// Optional path to hyperopt results JSON -- overrides matching config fields #[arg(long)] hyperopt_params: Option, /// Feature dimension (42 market + 3 portfolio = 45, or 53 with OFI; must match trainer `state_dim`) #[arg(long, default_value_t = 43)] feature_dim: usize, /// Early stopping patience (epochs without improvement) #[arg(long, default_value_t = 10)] patience: usize, /// Number of actions (5 exposure levels for DQN, pass --num-actions 45 for PPO) #[arg(long, default_value_t = 5)] num_actions: usize, /// Walk-forward: initial training window in months #[arg(long, default_value_t = 12)] train_months: u32, /// Walk-forward: validation window in months #[arg(long, default_value_t = 3)] val_months: u32, /// Walk-forward: test window in months #[arg(long, default_value_t = 3)] test_months: u32, /// Walk-forward: step size in months between folds #[arg(long, default_value_t = 3)] step_months: u32, /// Learning rate for optimizer #[arg(long, default_value_t = 1e-4)] learning_rate: f64, /// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT") #[arg(long, default_value = "ES.FUT")] symbol: String, /// Maximum absolute per-bar return; larger moves are clamped (contract roll filter) #[arg(long, default_value_t = 0.01)] max_bar_return: f64, /// Round-trip commission cost in basis points (1 bps = 0.01%) /// Applied to BUY/SELL rewards; HOLD is free. /// Default 1.0 bps covers ~$4 exchange+broker for ES e-mini. #[arg(long, default_value_t = 1.0)] tx_cost_bps: f64, /// Instrument tick size in price units (ES=0.25, NQ=0.25, ZN=1/64) #[arg(long, default_value_t = 0.25)] tick_size: f64, /// Typical bid-ask spread in ticks (ES=1.0, ZN=1.0, 6E=2.0) /// Half-spread slippage is added to `tx_cost_bps` per trade. #[arg(long, default_value_t = 1.0)] spread_ticks: f64, /// Number of top hyperopt configs to train as ensemble (default 1 = best only). /// When > 1, loads `top_k_params` from the hyperopt JSON and trains a separate /// model for each param set, saving as `dqn_ensemble_{k}_fold_{fold}.safetensors`. #[arg(long, default_value_t = 1)] ensemble_top_k: usize, /// Optional path to MBP-10 order book data directory for OFI features. /// When set, enables 8 OFI features (OFI L1/L5, depth imbalance, VPIN, etc.) /// expanding state dimension from 43 to 51. #[arg(long)] mbp10_data_dir: Option, /// Optional path to trade data directory (.dbn.zst files with Schema::Trades). /// When set, real trade buy/sell classification feeds VPIN and Kyle's Lambda /// instead of the tick-rule proxy. Requires --mbp10-data-dir to take effect. #[arg(long)] trades_data_dir: Option, /// Enable offline RL mode: train exclusively from a pre-collected dataset, /// skipping online experience collection. Requires --dataset-path. #[arg(long, default_value_t = false)] offline: bool, /// Path to a pre-collected experience dataset (bincode format). /// Used with --offline to load a fixed dataset into the replay buffer. #[arg(long)] dataset_path: Option, /// Collect and save a dataset from the current policy, then exit. /// Use this to generate datasets for offline training. #[arg(long)] collect_dataset: Option, /// Disable Branching DQN (3-head: exposure, order, urgency). Enabled by default. #[arg(long)] no_branching: bool, /// Initial trading capital in dollars. Lower capital teaches conservative /// position sizing. Must match hyperopt --initial-capital for consistency. #[arg(long, default_value_t = 35_000.0)] initial_capital: f64, /// Minimum bars to hold a position before allowing exit (churn prevention). /// Lower = more trades, higher = fewer. Default from TOML (typically 5). #[arg(long)] min_hold_bars: Option, /// Named training profile to load from config/training/.toml. /// Profile values are applied after hyperopt JSON but before explicit CLI args. /// Known profiles: dqn-production, dqn-smoketest, dqn-hyperopt. #[arg(long, default_value = "dqn-production")] training_profile: String, /// Feature cache directory (overrides FOXHUNT_FEATURE_CACHE_DIR and auto-discovery) #[arg(long)] feature_cache_dir: Option, } // --------------------------------------------------------------------------- // Hyperopt parameter loading // --------------------------------------------------------------------------- /// Load `best_params` from a hyperopt results JSON file. /// /// Expected format: `{ "model_key": { "best_params": { ... }, ... } }` /// Returns `None` if the file doesn't exist or can't be parsed. #[allow(clippy::cognitive_complexity)] fn load_hyperopt_params(hp_path: &Option, model_key: &str) -> Option { let file_path = hp_path.as_ref()?; if !file_path.exists() { info!("Hyperopt params file not found: {}, using defaults", file_path.display()); return None; } let contents = match std::fs::read_to_string(file_path) { Ok(c) => c, Err(e) => { warn!("Failed to read hyperopt params {}: {}", file_path.display(), e); return None; } }; let json: Value = match serde_json::from_str(&contents) { Ok(v) => v, Err(e) => { warn!("Failed to parse hyperopt params JSON: {}", e); return None; } }; let params = json.get(model_key) .and_then(|m| m.get("best_params")) .cloned(); if params.is_some() { info!("Loaded hyperopt params for '{}' from {}", model_key, file_path.display()); } else { warn!("No best_params found for '{}' in {}", model_key, file_path.display()); } params } /// Load top-K param sets from a hyperopt results JSON file. /// /// Expected format: `{ "model_key": { "top_k_params": [{ "params": {...}, ... }, ...], "best_params": {...} } }` /// Falls back to `best_params` as a single entry when `top_k_params` is absent. /// Returns a vec of length `k` (or fewer if not enough entries), each entry `Some(params)` or `None`. fn load_top_k_params(hp_path: &Option, model_key: &str, k: usize) -> Vec> { let file_path = match hp_path.as_ref() { Some(p) if p.exists() => p, _ => return vec![None; k], }; let Ok(contents) = std::fs::read_to_string(file_path) else { return vec![None; k]; }; let Ok(json): Result = serde_json::from_str(&contents) else { return vec![None; k]; }; let top_k = json .get(model_key) .and_then(|m| m.get("top_k_params")) .and_then(|v| v.as_array()); if let Some(arr) = top_k { arr.iter() .take(k) .map(|entry| entry.get("params").cloned()) .collect() } else { // Fallback: just use best_params as the single entry let best = json .get(model_key) .and_then(|m| m.get("best_params")) .cloned(); vec![best] } } fn hp_f64(params: &Option, key: &str) -> Option { params.as_ref()?.get(key)?.as_f64() } fn hp_usize(params: &Option, key: &str) -> Option { params.as_ref()?.get(key)?.as_u64().map(|v| v as usize) } fn hp_bool(params: &Option, key: &str) -> Option { params.as_ref()?.get(key)?.as_bool() } // --------------------------------------------------------------------------- // Fold data helpers // --------------------------------------------------------------------------- // --------------------------------------------------------------------------- // DQN Training // --------------------------------------------------------------------------- /// Build DQN hyperparameters from args, hyperopt JSON, and GPU profile. /// /// Extracted from the old `train_dqn_fold` so it can be called ONCE before the fold /// loop. The returned hyperparams are ready for `DQNTrainer::new`. #[allow(clippy::cognitive_complexity)] fn build_dqn_hyperparams( args: &Args, hp: &Option, total_cost_bps: f64, ) -> DQNHyperparameters { let gpu_profile = GpuProfile::load(); let hp_hidden_base = hp_usize(hp, "hidden_dim_base"); let dqn_hidden_base = hp_hidden_base.or_else(|| { info!(" [DQN] hidden_dim_base: {} (from GPU profile)", gpu_profile.training.hidden_dim_base); Some(gpu_profile.training.hidden_dim_base) }); let epsilon_start = hp_f64(hp, "epsilon_start").unwrap_or(0.05); let mut hyperparams = DQNHyperparameters { learning_rate: hp_f64(hp, "learning_rate").unwrap_or(args.learning_rate), batch_size: gpu_profile.training.batch_size, gamma: hp_f64(hp, "gamma").unwrap_or(0.95), epsilon_start, epsilon_end: hp_f64(hp, "epsilon_end").unwrap_or(0.01), epsilon_decay: hp_f64(hp, "epsilon_decay").unwrap_or(0.995), buffer_size: hp_usize(hp, "buffer_size").unwrap_or(gpu_profile.training.buffer_size), min_replay_size: hp_usize(hp, "min_replay_size").unwrap_or(1000), epochs: args.epochs, checkpoint_frequency: 10, hidden_dim_base: dqn_hidden_base, warmup_steps: 0, early_stopping_enabled: true, transaction_cost_multiplier: total_cost_bps, per_alpha: hp_f64(hp, "per_alpha").unwrap_or(0.6), per_beta_start: hp_f64(hp, "per_beta_start").unwrap_or(0.4), dueling_hidden_dim: hp_usize(hp, "dueling_hidden_dim").unwrap_or(128), n_steps: hp_usize(hp, "n_steps").unwrap_or(3), tau: hp_f64(hp, "tau").unwrap_or(0.005), num_atoms: hp_usize(hp, "num_atoms") .unwrap_or(gpu_profile.training.num_atoms), v_min: hp_f64(hp, "v_min").unwrap_or_else(|| { let gamma = hp_f64(hp, "gamma").unwrap_or(0.95); -(10.0_f64 / (1.0 - gamma) * 1.2).clamp(20.0, 300.0) }), v_max: hp_f64(hp, "v_max").unwrap_or_else(|| { let gamma = hp_f64(hp, "gamma").unwrap_or(0.95); (10.0_f64 / (1.0 - gamma) * 1.2).clamp(20.0, 300.0) }), noisy_sigma_init: hp_f64(hp, "noisy_sigma_init").unwrap_or(0.5), num_quantiles: hp_usize(hp, "num_quantiles").unwrap_or(64), noisy_epsilon_floor: hp_f64(hp, "noisy_epsilon_floor").unwrap_or(0.05).into(), hold_penalty_weight: hp_f64(hp, "hold_penalty_weight").unwrap_or(0.01), max_position_absolute: hp_f64(hp, "max_position_absolute").unwrap_or(2.0), huber_delta: hp_f64(hp, "huber_delta").unwrap_or(10.0), entropy_coefficient: hp_f64(hp, "entropy_coefficient").unwrap_or(0.01), curiosity_weight: hp_f64(hp, "curiosity_weight").unwrap_or(0.1), weight_decay: hp_f64(hp, "weight_decay").unwrap_or(1e-4), kelly_fractional: hp_f64(hp, "kelly_fractional").unwrap_or(0.5), kelly_max_fraction: hp_f64(hp, "kelly_max_fraction").unwrap_or(0.25), mbp10_data_dir: args.mbp10_data_dir.as_ref().map(|p| p.to_string_lossy().into_owned()).unwrap_or_else(|| "test_data/futures-baseline-mbp10".to_string()), trades_data_dir: args.trades_data_dir.as_ref().map(|p| p.to_string_lossy().into_owned()).unwrap_or_else(|| "test_data/futures-baseline-trades".to_string()), offline_mode: args.offline, dataset_path: args.dataset_path.as_ref().map(|p| p.to_string_lossy().into_owned()), replay_buffer_vram_fraction: gpu_profile.training.replay_buffer_vram_fraction, gpu_timesteps_per_episode: gpu_profile.experience.gpu_timesteps_per_episode, ..DQNHyperparameters::default() }; // Load training profile (TOML) and apply to hyperparams. let profile = ml::training_profile::DqnTrainingProfile::load(&args.training_profile); profile.apply_to(&mut hyperparams); // CLI args override profile hyperparams.epochs = args.epochs; hyperparams.learning_rate = hp_f64(hp, "learning_rate").unwrap_or(args.learning_rate); hyperparams.initial_capital = args.initial_capital as f32; if let Some(mhb) = args.min_hold_bars { hyperparams.min_hold_bars = mhb; } hyperparams } /// Train a single DQN fold on a pre-initialized trainer. /// /// The trainer already has GPU data uploaded (via `init_from_fxcache`). /// This function sets per-fold ranges, resets state, and runs training. /// /// Returns the best validation loss achieved. #[allow(clippy::cognitive_complexity, clippy::too_many_arguments)] fn train_dqn_fold( rt: &tokio::runtime::Runtime, trainer: &mut DQNTrainer, fold: usize, train_features: &[[f64; 42]], val_features: &[[f64; 42]], train_targets: &[[f64; 4]], val_targets: &[[f64; 4]], range: &ml::walk_forward::FoldRange, output_dir: &Path, checkpoint_prefix: &str, ) -> Result { info!(" [DQN] Fold {} -- {} train, {} val features", fold, train_features.len(), val_features.len()); // Set fold range + val data + reset trainer.set_training_range(range.train_start, range.train_end, range.val_start, range.val_end); trainer.set_val_data_from_slices(val_features, val_targets, range.val_start); rt.block_on(trainer.reset_for_fold()) .context("reset_for_fold failed")?; // Checkpoint callback: save best model to output directory let output_dir_owned = output_dir.to_path_buf(); let prefix_owned = checkpoint_prefix.to_owned(); let checkpoint_callback = move |epoch: usize, data: Vec, is_best: bool| -> Result { let suffix = if is_best { "best" } else { &format!("epoch{}", epoch) }; let ckpt_path = output_dir_owned.join(format!("{}_fold{}_{}.safetensors", prefix_owned, fold, suffix)); let tmp_path = ckpt_path.with_extension("safetensors.tmp"); std::fs::write(&tmp_path, &data) .with_context(|| format!("Failed to write checkpoint tmp: {}", tmp_path.display()))?; if let Err(e) = std::fs::rename(&tmp_path, &ckpt_path) { drop(std::fs::remove_file(&tmp_path)); return Err(e).with_context(|| format!("Failed to rename checkpoint: {} -> {}", tmp_path.display(), ckpt_path.display())); } info!(" [DQN] Fold {} saved checkpoint: {} (prefix: {})", fold, ckpt_path.display(), prefix_owned); Ok(ckpt_path.to_string_lossy().into_owned()) }; let metrics = rt.block_on( trainer.train_fold_from_slices(train_features, train_targets, checkpoint_callback) ).map_err(|e| { error!(" [DQN] Fold {} training error chain: {:#}", fold, e); e }).context("DQNTrainer training failed")?; info!( " [DQN] Fold {} complete -- loss={:.6} epochs_trained={} converged={}", fold, metrics.loss, metrics.epochs_trained, metrics.convergence_achieved ); Ok(metrics.loss) } // --------------------------------------------------------------------------- // Training orchestration // --------------------------------------------------------------------------- /// Container for per-model RL results collected during training. struct RlTrainingResult { model_name: String, fold_results: Vec<(usize, f64)>, total_epochs: usize, } /// Run the full walk-forward RL training pipeline. /// /// Returns per-model results so `main()` can write completion markers. #[allow(clippy::cognitive_complexity, clippy::too_many_lines)] fn run_training(args: &Args) -> Result> { let train_dqn = args.model == "dqn" || args.model == "both"; let train_ppo = args.model == "ppo" || args.model == "both"; info!("=== Walk-Forward Baseline Training ==="); info!(" Model(s): {}", args.model); info!(" Symbol: {}", args.symbol); info!(" Epochs: {}", args.epochs); info!(" Batch size: auto (from VRAM)"); info!(" Data dir: {}", args.data_dir.display()); info!(" Output dir: {}", args.output_dir.display()); info!(" Feature dim: {}", args.feature_dim); info!(" Num actions: {}", args.num_actions); info!(" Learning rate: {:.1e}", args.learning_rate); info!(" Tx cost: {:.1} bps commission + {:.1} tick spread (tick_size={:.4})", args.tx_cost_bps, args.spread_ticks, args.tick_size); info!(" Patience: {}", args.patience); if let Some(ref hp_path) = args.hyperopt_params { info!(" Hyperopt params: {}", hp_path.display()); } if args.ensemble_top_k > 1 { info!(" Ensemble top-K: {} (training multiple models per fold)", args.ensemble_top_k); } // 1. Try fxcache first, fall back to DBN loading + feature extraction info!("Step 1/5: Loading data..."); let data_load_start = std::time::Instant::now(); let cache_dir_override = args.feature_cache_dir.as_ref().map(|s| std::path::PathBuf::from(s)); let mbp10 = args.mbp10_data_dir.as_ref().filter(|p| p.exists()); let trades = args.trades_data_dir.as_ref().filter(|p| p.exists()); let fxcache_data = ml::fxcache::discover_and_load( &args.data_dir, &args.symbol, mbp10.map(|p| p.as_path()), trades.map(|p| p.as_path()), "ohlcv", cache_dir_override.as_deref(), ); // Load data into fxcache-compatible arrays: features, targets, timestamps, ofi let fxcache = if let Some(cached) = fxcache_data { info!(" Loaded {} bars + features from fxcache in {:.1}s", cached.bar_count, data_load_start.elapsed().as_secs_f64()); cached } else { // Fall back to DBN loading — this is SLOW (148GB MBP-10 parsing) info!(" Loading OHLCV bars from DBN files..."); let bars = load_all_bars(&args.data_dir, &args.symbol)?; if bars.is_empty() { anyhow::bail!("No bars loaded from {}", args.data_dir.display()); } info!(" Loaded {} bars ({} to {})", bars.len(), bars.first().map(|b| b.timestamp.to_string()).unwrap_or_default(), bars.last().map(|b| b.timestamp.to_string()).unwrap_or_default(), ); info!(" Extracting {}-dimensional features...", args.feature_dim); let all_features = extract_ml_features(&bars) .context("Feature extraction failed")?; let warmup_offset = bars.len().saturating_sub(all_features.len()); info!(" Extracted {} feature vectors (warmup period consumed {} bars)", all_features.len(), warmup_offset); // Build FxCacheData from DBN results (features are already warmup-trimmed) let aligned_bars = &bars[warmup_offset..]; let n = all_features.len(); let timestamps: Vec = aligned_bars.iter() .map(|b| b.timestamp.timestamp_nanos_opt().unwrap_or(0)) .collect(); let targets: Vec<[f64; 4]> = aligned_bars.iter() .map(|b| [b.close, b.close, b.close, b.close]) .collect(); let ofi = vec![[0.0_f64; 8]; n]; ml::fxcache::FxCacheData { timestamps, features: all_features, targets, ofi, cache_key: [0u8; 32], bar_count: n, has_ofi: false, } }; // 2. Generate walk-forward fold ranges from timestamps (zero-copy) info!("Step 2/5: Generating walk-forward fold ranges..."); let wf_config = WalkForwardConfig { initial_train_months: args.train_months, val_months: args.val_months, test_months: args.test_months, step_months: args.step_months, }; let fold_ranges = generate_walk_forward_indices_from_timestamps(&fxcache.timestamps, &wf_config); if fold_ranges.is_empty() { anyhow::bail!( "No walk-forward folds generated. Need at least {} months of data.", wf_config.initial_train_months + wf_config.val_months + wf_config.test_months ); } info!(" Generated {} walk-forward folds (zero-copy index ranges)", fold_ranges.len()); // Record data loading + feature extraction time if train_dqn { metrics::record_data_load("dqn", data_load_start.elapsed().as_secs_f64()); } if train_ppo { metrics::record_data_load("ppo", data_load_start.elapsed().as_secs_f64()); } // Create output directory std::fs::create_dir_all(&args.output_dir) .with_context(|| format!("Failed to create output dir: {}", args.output_dir.display()))?; // 3. Create tokio runtime ONCE, create DQN trainer ONCE, upload fxcache to GPU ONCE let rt = tokio::runtime::Builder::new_current_thread() .enable_all() .build() .context("Failed to create tokio runtime")?; // Compute average spread slippage in bps from the full dataset let avg_price = { let sum: f64 = fxcache.targets.iter().map(|t| t[2]).sum(); // raw_close if fxcache.bar_count > 0 { sum / fxcache.bar_count as f64 } else { 0.0 } }; let avg_spread_bps = spread_cost_bps(avg_price, args.tick_size, args.spread_ticks); let total_cost_bps = args.tx_cost_bps + avg_spread_bps; info!(" Total tx cost: {:.2} bps (commission {:.1} + spread {:.2})", total_cost_bps, args.tx_cost_bps, avg_spread_bps); // Build DQN trainer ONCE (shared across folds) let mut dqn_trainer = if train_dqn { let hp = load_hyperopt_params(&args.hyperopt_params, "dqn"); let hyperparams = build_dqn_hyperparams(args, &hp, total_cost_bps); let mut trainer = DQNTrainer::new(hyperparams) .context("Failed to create DQNTrainer")?; // Upload full fxcache to GPU ONCE — all folds index into this data info!(" Uploading {} bars to GPU via init_from_fxcache...", fxcache.bar_count); rt.block_on(trainer.init_from_fxcache( &fxcache.features, &fxcache.targets, &fxcache.ofi, )).context("init_from_fxcache failed")?; info!(" GPU data uploaded — ready for fold loop"); Some(trainer) } else { None }; // Build ensemble trainers ONCE (shared across folds) — upload data once each. // Previously these were created inside the fold loop, re-uploading per fold. let mut ensemble_trainers: Vec = Vec::new(); if train_dqn && args.ensemble_top_k > 1 && args.hyperopt_params.is_some() { let param_sets = load_top_k_params( &args.hyperopt_params, "dqn", args.ensemble_top_k, ); // k=0 is the primary trainer, so build trainers for k=1.. for (k, hp) in param_sets.iter().enumerate().skip(1) { let ens_hyperparams = build_dqn_hyperparams(args, hp, total_cost_bps); match DQNTrainer::new(ens_hyperparams) { Ok(mut ens_trainer) => { info!(" Uploading fxcache to ensemble trainer {}...", k); if let Err(e) = rt.block_on(ens_trainer.init_from_fxcache( &fxcache.features, &fxcache.targets, &fxcache.ofi, )) { error!(" [DQN] Ensemble trainer {} init_from_fxcache failed: {}", k, e); continue; } ensemble_trainers.push(ens_trainer); } Err(e) => { error!(" [DQN] Failed to create ensemble trainer {}: {}", k, e); } } } info!(" Created {} ensemble trainers (data uploaded once each)", ensemble_trainers.len()); } // 4. Train each fold info!("Step 4/5: Training models on each fold..."); let mut dqn_results: Vec<(usize, f64)> = Vec::new(); let mut ppo_results: Vec<(usize, f64)> = Vec::new(); for (fold_idx, range) in fold_ranges.iter().enumerate() { info!("--- Fold {} ---", range.fold); info!( " Train: bars [{}..{}] ({} bars), Val: bars [{}..{}] ({} bars)", range.train_start, range.train_end, range.train_end - range.train_start, range.val_start, range.val_end, range.val_end - range.val_start, ); // Slice features/targets for this fold (zero-copy from fxcache arrays) let train_feat = &fxcache.features[range.train_start..range.train_end]; let val_feat = &fxcache.features[range.val_start..range.val_end]; let train_tgt = &fxcache.targets[range.train_start..range.train_end]; let val_tgt = &fxcache.targets[range.val_start..range.val_end]; if train_feat.is_empty() || val_feat.is_empty() { warn!(" Fold {} -- empty features, skipping", range.fold); continue; } // fxcache features are pre-normalized (z-score at precompute time). // NormStats saved alongside .fxcache file for inference denormalization. let train_norm = train_feat; let val_norm = val_feat; // Train DQN if let Some(ref mut trainer) = dqn_trainer { let fold_str = fold_idx.to_string(); let fold_start = std::time::Instant::now(); if args.ensemble_top_k > 1 && args.hyperopt_params.is_some() { // Ensemble mode: train one model per top-K hyperopt param set. // Trainers were created + data uploaded BEFORE the fold loop. // Here we just set_training_range + reset_for_fold on each. let total_members = 1 + ensemble_trainers.len(); // k=0 (primary) + secondaries // k=0: Primary ensemble member uses the shared trainer { let k = 0; info!( " [DQN] Training ensemble member {}/{} on fold {}", k + 1, total_members, range.fold ); let prefix = format!("dqn_ensemble_{}", k); match train_dqn_fold( &rt, trainer, range.fold, &train_norm, &val_norm, train_tgt, val_tgt, range, &args.output_dir, &prefix, ) { Ok(best_loss) => { info!(" [DQN] Ensemble member {} fold {} best_loss={:.6}", k, range.fold, best_loss); let elapsed = fold_start.elapsed().as_secs_f64(); metrics::set_epoch("dqn", &fold_str, fold_idx as f64); metrics::set_epoch_loss("dqn", &fold_str, best_loss); metrics::set_validation_loss("dqn", &fold_str, best_loss); metrics::set_iteration_seconds("dqn", &fold_str, elapsed); dqn_results.push((range.fold, best_loss)); } Err(e) => { error!(" [DQN] Ensemble member {} fold {} failed: {}", k, range.fold, e); } } } // k=1..: Secondary ensemble members reuse pre-created trainers for (ens_idx, ens_trainer) in ensemble_trainers.iter_mut().enumerate() { let k = ens_idx + 1; info!( " [DQN] Training ensemble member {}/{} on fold {}", k + 1, total_members, range.fold ); let prefix = format!("dqn_ensemble_{}", k); match train_dqn_fold( &rt, ens_trainer, range.fold, &train_norm, &val_norm, train_tgt, val_tgt, range, &args.output_dir, &prefix, ) { Ok(best_loss) => { info!(" [DQN] Ensemble member {} fold {} best_loss={:.6}", k, range.fold, best_loss); } Err(e) => { error!(" [DQN] Ensemble member {} fold {} failed: {}", k, range.fold, e); } } } } else { // Single-model mode (default) match train_dqn_fold( &rt, trainer, range.fold, &train_norm, &val_norm, train_tgt, val_tgt, range, &args.output_dir, "dqn", ) { Ok(best_loss) => { let elapsed = fold_start.elapsed().as_secs_f64(); metrics::set_epoch("dqn", &fold_str, fold_idx as f64); metrics::set_epoch_loss("dqn", &fold_str, best_loss); metrics::set_validation_loss("dqn", &fold_str, best_loss); metrics::set_iteration_seconds("dqn", &fold_str, elapsed); dqn_results.push((range.fold, best_loss)); } Err(e) => { error!(" [DQN] Fold {} failed: {:#}", range.fold, e); } } } } // Train PPO (zero-copy: pass fxcache feature slices directly) if train_ppo { let hp_ppo = load_hyperopt_params(&args.hyperopt_params, "ppo"); let fold_str = fold_idx.to_string(); let fold_start = std::time::Instant::now(); // Compute per-fold tx cost from fxcache targets (raw close at index 2) let train_closes: Vec = fxcache.targets[range.train_start..range.train_end] .iter().map(|t| t[2]).collect(); let avg_price = if train_closes.is_empty() { 0.0 } else { train_closes.iter().sum::() / train_closes.len() as f64 }; let avg_spread_bps_ppo = spread_cost_bps(avg_price, args.tick_size, args.spread_ticks); let total_cost_bps_ppo = args.tx_cost_bps + avg_spread_bps_ppo; // Build PpoHyperparameters inline (same as old train_ppo_fold) let hp_ppo_hidden_base = hp_usize(&hp_ppo, "hidden_dim_base"); let ppo_hidden_base = hp_ppo_hidden_base.or_else(|| { let profile = ml_core::gpu::profile::GpuProfile::load(); info!(" [PPO] hidden_dim_base: {} (from GPU profile)", profile.training.hidden_dim_base); Some(profile.training.hidden_dim_base) }); let ppo_hp = PpoHyperparameters { learning_rate: hp_f64(&hp_ppo, "learning_rate").unwrap_or(args.learning_rate), actor_learning_rate: Some(hp_f64(&hp_ppo, "policy_learning_rate").unwrap_or(args.learning_rate)), critic_learning_rate: Some(hp_f64(&hp_ppo, "value_learning_rate").unwrap_or(args.learning_rate * 3.0)), batch_size: ml_core::gpu::profile::GpuProfile::load().training.batch_size, gamma: hp_f64(&hp_ppo, "gamma").unwrap_or(0.99), clip_epsilon: hp_f64(&hp_ppo, "clip_epsilon").unwrap_or(0.2) as f32, vf_coef: hp_f64(&hp_ppo, "value_loss_coeff").unwrap_or(0.5) as f32, ent_coef: hp_f64(&hp_ppo, "entropy_coeff").unwrap_or(0.01) as f32, gae_lambda: hp_f64(&hp_ppo, "gae_lambda").unwrap_or(0.95) as f32, rollout_steps: hp_usize(&hp_ppo, "rollout_steps").unwrap_or(2048), minibatch_size: hp_usize(&hp_ppo, "minibatch_size").unwrap_or(64), epochs: args.epochs, early_stopping_enabled: true, transaction_cost_bps: total_cost_bps_ppo / 100.0, hidden_dim_base: ppo_hidden_base, ..PpoHyperparameters::conservative() }; let fold_ckpt_dir = args.output_dir.join(format!("ppo_fold{}", range.fold)); let ppo_trainer = PpoTrainer::new( ppo_hp, args.feature_dim, &fold_ckpt_dir, true, None, ); match ppo_trainer { Ok(trainer) => { let fold_idx_cap = range.fold; let epochs_cap = args.epochs; let progress_cb = move |m: PpoTrainingMetrics| { info!("[PPO] Fold {} Epoch {}/{} -- value_loss={:.6}", fold_idx_cap, m.epoch, epochs_cap, m.value_loss); }; match rt.block_on(trainer.train_from_slices(&train_norm, progress_cb)) { Ok(metrics) => { let best_loss = metrics.value_loss as f64; let elapsed = fold_start.elapsed().as_secs_f64(); metrics::set_epoch("ppo", &fold_str, fold_idx as f64); metrics::set_epoch_loss("ppo", &fold_str, best_loss); metrics::set_validation_loss("ppo", &fold_str, best_loss); metrics::set_iteration_seconds("ppo", &fold_str, elapsed); ppo_results.push((range.fold, best_loss)); } Err(e) => error!("PPO fold {} failed: {:#}", range.fold, e), } } Err(e) => error!(" [PPO] Fold {} trainer init failed: {:#}", range.fold, e), } } } // 5. Summary info!("Step 5/5: Training Summary"); info!(" ==================================="); let mut all_results = Vec::new(); let num_folds = fold_ranges.len(); if train_dqn { info!(" DQN Results ({} folds):", dqn_results.len()); for (fold, loss) in &dqn_results { info!(" Fold {}: best_val_metric = {:.6}", fold, loss); } if !dqn_results.is_empty() { let avg: f64 = dqn_results.iter().map(|(_, l)| l).sum::() / dqn_results.len() as f64; info!(" Average: {:.6}", avg); } all_results.push(RlTrainingResult { model_name: "dqn".to_owned(), fold_results: dqn_results, total_epochs: num_folds * args.epochs, }); } if train_ppo { info!(" PPO Results ({} folds):", ppo_results.len()); for (fold, loss) in &ppo_results { info!(" Fold {}: best_val_metric = {:.6}", fold, loss); } if !ppo_results.is_empty() { let avg: f64 = ppo_results.iter().map(|(_, l)| l).sum::() / ppo_results.len() as f64; info!(" Average: {:.6}", avg); } all_results.push(RlTrainingResult { model_name: "ppo".to_owned(), fold_results: ppo_results, total_epochs: num_folds * args.epochs, }); } info!(" Checkpoints saved to: {}", args.output_dir.display()); info!(" ==================================="); Ok(all_results) } // --------------------------------------------------------------------------- // Main // --------------------------------------------------------------------------- fn main() -> Result<()> { // Initialize tracing with optional OTLP export to Tempo let otlp_endpoint = std::env::var("OTEL_EXPORTER_OTLP_ENDPOINT").ok(); if let Err(e) = common::observability::init_observability( "train_baseline_rl", otlp_endpoint.as_deref(), ) { eprintln!("Observability init failed (non-fatal): {e}"); } // Pre-allocate CUBLAS workspace for deterministic + faster tensor core ops. // Enable TF32 for all FP32 matmuls — ~8x throughput on H100 tensor cores. // SAFETY: called once at startup before any multi-threading or CUDA work begins. #[allow(unsafe_code)] unsafe { std::env::set_var("CUBLAS_WORKSPACE_CONFIG", ":4096:8"); std::env::set_var("NVIDIA_TF32_OVERRIDE", "1"); } metrics::init(); metrics_server::start_metrics_server(9094); common::metrics::questdb_sink::init(None); metrics::set_active_workers(1.0); let args = Args::parse(); // Ensure output directory exists before training so markers can always be written. if let Err(e) = std::fs::create_dir_all(&args.output_dir) { error!("Failed to create output dir {}: {}", args.output_dir.display(), e); } let result = run_training(&args); metrics::set_active_workers(0.0); // Push final metrics to pushgateway so they persist after pod termination if let Err(e) = metrics_server::push_to_gateway(None, "train_baseline_rl") { tracing::warn!("Failed to push metrics to gateway (non-fatal): {e}"); } common::metrics::questdb_sink::flush(); match result { Ok(results) => { for training_result in &results { let best_val = training_result .fold_results .iter() .map(|(_, loss)| *loss) .fold(f64::MAX, f64::min); let metrics = CompletionMetrics { model: training_result.model_name.clone(), symbol: args.symbol.clone(), best_val_loss: (best_val < f64::MAX).then_some(best_val), sharpe_ratio: None, epochs_completed: training_result.total_epochs, folds_completed: training_result.fold_results.len(), }; write_success_marker(&args.output_dir, &metrics); } Ok(()) } Err(e) => { let msg = format!("{:#}", e); error!("Training failed: {}", msg); write_failure_marker(&args.output_dir, &msg); Err(e) } } }