Files
foxhunt/crates/ml/examples/train_baseline_rl.rs
jgrusewski 9c081ac7ce refactor: CLI binaries use shared fxcache::discover_and_load()
Replace 40-line inline fxcache discovery block in train_baseline_rl.rs
with a single call to ml::fxcache::discover_and_load(). precompute_features.rs
already uses correct symbol/data_source args — no changes needed there.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-04 10:05:26 +02:00

1010 lines
42 KiB
Rust

#![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<PathBuf>,
/// 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<PathBuf>,
/// 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<PathBuf>,
/// 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<PathBuf>,
/// Collect and save a dataset from the current policy, then exit.
/// Use this to generate datasets for offline training.
#[arg(long)]
collect_dataset: Option<PathBuf>,
/// 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<usize>,
/// Named training profile to load from config/training/<profile>.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<String>,
}
// ---------------------------------------------------------------------------
// 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<PathBuf>, model_key: &str) -> Option<Value> {
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<PathBuf>, model_key: &str, k: usize) -> Vec<Option<Value>> {
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<Value, _> = 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<Value>, key: &str) -> Option<f64> {
params.as_ref()?.get(key)?.as_f64()
}
fn hp_usize(params: &Option<Value>, key: &str) -> Option<usize> {
params.as_ref()?.get(key)?.as_u64().map(|v| v as usize)
}
fn hp_bool(params: &Option<Value>, key: &str) -> Option<bool> {
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<Value>,
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<f64> {
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<u8>, is_best: bool| -> Result<String> {
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<Vec<RlTrainingResult>> {
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<i64> = 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<DQNTrainer> = 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<f64> = 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::<f64>() / 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::<f64>()
/ 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::<f64>()
/ 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)
}
}
}