refactor(ml): consolidate 21 training binaries into 2 unified baselines
Replace 20 per-model training examples with: - train_baseline_rl: DQN + PPO (renamed from train_baseline) - train_baseline_supervised: TFT, Mamba2, Liquid, TGGN, TLOB, KAN, xLSTM, Diffusion via model factory + UnifiedTrainable generic training loop Update Dockerfile.training (16→7 binaries), train.sh MODEL_BINARY map, and job-template.yaml default. -12,759 lines. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -205,8 +205,12 @@ path = "examples/evaluate_ppo.rs"
|
||||
required-features = ["cuda"]
|
||||
|
||||
[[example]]
|
||||
name = "train_baseline"
|
||||
path = "examples/train_baseline.rs"
|
||||
name = "train_baseline_rl"
|
||||
path = "examples/train_baseline_rl.rs"
|
||||
|
||||
[[example]]
|
||||
name = "train_baseline_supervised"
|
||||
path = "examples/train_baseline_supervised.rs"
|
||||
|
||||
[[bench]]
|
||||
name = "microstructure_bench"
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
//! Walk-forward training binary for DQN and PPO models.
|
||||
//! 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
|
||||
@@ -46,7 +46,7 @@ use ml::walk_forward::{generate_walk_forward_windows, NormStats, WalkForwardConf
|
||||
|
||||
/// Walk-forward training binary for DQN and PPO baseline models.
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "train_baseline", about = "Train DQN/PPO with walk-forward windows")]
|
||||
#[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")]
|
||||
671
crates/ml/examples/train_baseline_supervised.rs
Normal file
671
crates/ml/examples/train_baseline_supervised.rs
Normal file
@@ -0,0 +1,671 @@
|
||||
//! Walk-forward supervised training binary for all non-RL models.
|
||||
//!
|
||||
//! Trains models using the `UnifiedTrainable` trait on real OHLCV data loaded
|
||||
//! from Databento DBN files. Supports walk-forward windows, early stopping,
|
||||
//! checkpoint saving, and z-score normalization.
|
||||
//!
|
||||
//! # Supported Models
|
||||
//!
|
||||
//! tft, mamba2, liquid, tggn, tlob, kan, xlstm, diffusion
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_baseline_supervised --release -- \
|
||||
//! --model kan --epochs 50 --batch-size 128 \
|
||||
//! --data-dir data/cache/futures-baseline \
|
||||
//! --output-dir ml/trained_models
|
||||
//! ```
|
||||
|
||||
#![allow(unused_crate_dependencies, clippy::cognitive_complexity)]
|
||||
#![deny(
|
||||
clippy::unwrap_used,
|
||||
clippy::expect_used,
|
||||
clippy::panic,
|
||||
clippy::indexing_slicing
|
||||
)]
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Instant;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::{Device, Tensor};
|
||||
use clap::Parser;
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use ml::features::extraction::extract_ml_features;
|
||||
use ml::training::unified_trainer::UnifiedTrainable;
|
||||
use ml::types::OHLCVBar;
|
||||
use ml::walk_forward::{generate_walk_forward_windows, NormStats, WalkForwardConfig};
|
||||
|
||||
#[allow(unreachable_pub)]
|
||||
mod baseline_common;
|
||||
use baseline_common::{load_all_bars, spread_cost_bps};
|
||||
|
||||
// Model adapter imports — using verified paths from existing examples
|
||||
use ml::diffusion::{DiffusionConfig, DiffusionTrainableAdapter};
|
||||
use ml::kan::{KANConfig, KANTrainableAdapter};
|
||||
use ml::liquid::{CfCTrainConfig, DeviceConfig, LiquidTrainableAdapter};
|
||||
use ml::mamba::{Mamba2Config, Mamba2SSM};
|
||||
use ml::tft::{TFTConfig, TrainableTFT};
|
||||
use ml::tgnn::trainable_adapter::TGGNTrainableAdapter;
|
||||
use ml::tgnn::TGGNConfig;
|
||||
use ml::tlob::{TLOBAdapterConfig, TLOBTrainableAdapter};
|
||||
use ml::xlstm::{XLSTMConfig, XLSTMTrainableAdapter};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CLI Arguments
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Walk-forward supervised training binary for non-RL models.
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "train_baseline_supervised", about = "Train supervised models with walk-forward windows")]
|
||||
struct Args {
|
||||
/// Model to train: tft, mamba2, liquid, tggn, tlob, kan, xlstm, diffusion, all
|
||||
#[arg(long)]
|
||||
model: String,
|
||||
|
||||
/// Path to directory containing .dbn.zst files (with symbol subdirectories)
|
||||
#[arg(long, default_value = "data/cache/futures-baseline")]
|
||||
data_dir: PathBuf,
|
||||
|
||||
/// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT")
|
||||
#[arg(long, default_value = "ES.FUT")]
|
||||
symbol: String,
|
||||
|
||||
/// Maximum training epochs per fold
|
||||
#[arg(long, default_value_t = 50)]
|
||||
epochs: usize,
|
||||
|
||||
/// Training batch size
|
||||
#[arg(long, default_value_t = 128)]
|
||||
batch_size: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value_t = 1e-3)]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Feature dimension (must match extract_ml_features output)
|
||||
#[arg(long, default_value_t = 51)]
|
||||
feature_dim: usize,
|
||||
|
||||
/// Max training steps per epoch (0 = use all bars)
|
||||
#[arg(long, default_value_t = 2000)]
|
||||
max_steps_per_epoch: usize,
|
||||
|
||||
/// Output directory for trained model checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: PathBuf,
|
||||
|
||||
/// Early stopping patience (epochs without improvement)
|
||||
#[arg(long, default_value_t = 10)]
|
||||
patience: 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,
|
||||
|
||||
/// Maximum absolute per-bar return; larger moves are clamped
|
||||
#[arg(long, default_value_t = 0.01)]
|
||||
max_bar_return: f64,
|
||||
|
||||
/// Round-trip commission cost in basis points
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
tx_cost_bps: f64,
|
||||
|
||||
/// Instrument tick size in price units (ES=0.25)
|
||||
#[arg(long, default_value_t = 0.25)]
|
||||
tick_size: f64,
|
||||
|
||||
/// Typical bid-ask spread in ticks
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
spread_ticks: f64,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Supported models
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
const ALL_MODELS: &[&str] = &[
|
||||
"tft", "mamba2", "liquid", "tggn", "tlob", "kan", "xlstm", "diffusion",
|
||||
];
|
||||
|
||||
fn validate_model(name: &str) -> Result<()> {
|
||||
if name == "all" || ALL_MODELS.contains(&name) {
|
||||
Ok(())
|
||||
} else {
|
||||
anyhow::bail!(
|
||||
"Unknown model '{}'. Valid: {} or 'all'",
|
||||
name,
|
||||
ALL_MODELS.join(", ")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Model Factory
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Create a model adapter by name with production-default configs.
|
||||
///
|
||||
/// All OHLCV models receive `feature_dim` as input dimension.
|
||||
/// The adapter's internal projection layers handle mapping to model-native dims.
|
||||
fn create_model(
|
||||
name: &str,
|
||||
feature_dim: usize,
|
||||
learning_rate: f64,
|
||||
device: &Device,
|
||||
) -> Result<Box<dyn UnifiedTrainable>> {
|
||||
match name {
|
||||
"tft" => {
|
||||
let config = TFTConfig {
|
||||
input_dim: feature_dim,
|
||||
hidden_dim: 128,
|
||||
num_heads: 4,
|
||||
num_layers: 2,
|
||||
num_quantiles: 3,
|
||||
num_static_features: 0,
|
||||
dropout_rate: 0.1,
|
||||
..TFTConfig::default()
|
||||
};
|
||||
let mut adapter = TrainableTFT::new(config)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create TFT: {}", e))?;
|
||||
adapter
|
||||
.set_learning_rate(learning_rate)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to set TFT learning rate: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"mamba2" => {
|
||||
let config = Mamba2Config {
|
||||
d_model: 128,
|
||||
num_layers: 4,
|
||||
d_state: 16,
|
||||
max_seq_len: 60,
|
||||
..Mamba2Config::default()
|
||||
};
|
||||
let mut adapter = Mamba2SSM::new(config, device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create Mamba2: {}", e))?;
|
||||
adapter
|
||||
.set_learning_rate(learning_rate)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to set Mamba2 learning rate: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"liquid" => {
|
||||
let config = CfCTrainConfig {
|
||||
input_size: feature_dim,
|
||||
hidden_size: 128,
|
||||
output_size: 1,
|
||||
backbone_hidden_sizes: vec![128, 64],
|
||||
learning_rate,
|
||||
device: DeviceConfig::Auto,
|
||||
..CfCTrainConfig::default()
|
||||
};
|
||||
let adapter = LiquidTrainableAdapter::new(config)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create Liquid: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"tggn" => {
|
||||
let config = TGGNConfig {
|
||||
node_dim: feature_dim,
|
||||
hidden_dim: 32,
|
||||
num_layers: 2,
|
||||
max_nodes: 64,
|
||||
max_edges: 128,
|
||||
edge_dim: 4,
|
||||
temporal_decay: 0.99,
|
||||
update_frequency_ns: 1_000_000,
|
||||
use_simd: false,
|
||||
};
|
||||
let mut adapter = TGGNTrainableAdapter::new(config, device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create TGGN: {}", e))?;
|
||||
adapter
|
||||
.set_learning_rate(learning_rate)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to set TGGN learning rate: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"tlob" => {
|
||||
let config = TLOBAdapterConfig {
|
||||
d_model: 128,
|
||||
num_heads: 4,
|
||||
num_layers: 2,
|
||||
seq_len: 1,
|
||||
feature_dim,
|
||||
};
|
||||
let mut adapter = TLOBTrainableAdapter::new(config, device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create TLOB: {}", e))?;
|
||||
adapter
|
||||
.set_learning_rate(learning_rate)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to set TLOB learning rate: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"kan" => {
|
||||
let config = KANConfig {
|
||||
grid_size: 5,
|
||||
spline_order: 4,
|
||||
layer_widths: vec![feature_dim, 32, 16, 1],
|
||||
learning_rate,
|
||||
weight_decay: 1e-4,
|
||||
grad_clip: 1.0,
|
||||
};
|
||||
let adapter = KANTrainableAdapter::new(config, device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create KAN: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"xlstm" => {
|
||||
let config = XLSTMConfig {
|
||||
input_dim: feature_dim,
|
||||
hidden_dim: 128,
|
||||
..XLSTMConfig::default()
|
||||
};
|
||||
let mut adapter = XLSTMTrainableAdapter::new(config, device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create xLSTM: {}", e))?;
|
||||
adapter
|
||||
.set_learning_rate(learning_rate)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to set xLSTM learning rate: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
"diffusion" => {
|
||||
let config = DiffusionConfig {
|
||||
feature_dim,
|
||||
hidden_dim: 128,
|
||||
..DiffusionConfig::default()
|
||||
};
|
||||
let adapter = DiffusionTrainableAdapter::new(config, device.clone())
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create Diffusion: {}", e))?;
|
||||
Ok(Box::new(adapter))
|
||||
}
|
||||
_ => anyhow::bail!("Unknown model: {}", name),
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Data preparation
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Build (input, target) tensor pairs from OHLCV bars for regression training.
|
||||
///
|
||||
/// Features are extracted via `extract_ml_features` (51-dim), z-score normalized,
|
||||
/// and the target is the next-bar return in basis points, clipped.
|
||||
fn prepare_fold_data(
|
||||
train_bars: &[OHLCVBar],
|
||||
val_bars: &[OHLCVBar],
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
) -> Result<(Vec<(Tensor, Tensor)>, Vec<(Tensor, Tensor)>)> {
|
||||
let train_features = extract_ml_features(train_bars)
|
||||
.context("Failed to extract training features")?;
|
||||
let val_features = extract_ml_features(val_bars)
|
||||
.context("Failed to extract validation features")?;
|
||||
|
||||
if train_features.len() < 2 {
|
||||
anyhow::bail!(
|
||||
"Insufficient training features: {}",
|
||||
train_features.len()
|
||||
);
|
||||
}
|
||||
if val_features.len() < 2 {
|
||||
anyhow::bail!(
|
||||
"Insufficient validation features: {}",
|
||||
val_features.len()
|
||||
);
|
||||
}
|
||||
|
||||
let norm_stats = NormStats::from_features(&train_features);
|
||||
let norm_train = norm_stats.normalize_batch(&train_features);
|
||||
let norm_val = norm_stats.normalize_batch(&val_features);
|
||||
|
||||
let train_bar_offset = train_bars.len().saturating_sub(train_features.len());
|
||||
let val_bar_offset = val_bars.len().saturating_sub(val_features.len());
|
||||
|
||||
let train_pairs = build_tensor_pairs(&norm_train, train_bars, train_bar_offset, args, device)?;
|
||||
let val_pairs = build_tensor_pairs(&norm_val, val_bars, val_bar_offset, args, device)?;
|
||||
|
||||
Ok((train_pairs, val_pairs))
|
||||
}
|
||||
|
||||
/// Convert normalized features + bars into (input, target) tensor pairs.
|
||||
fn build_tensor_pairs(
|
||||
norm_features: &[[f64; 51]],
|
||||
bars: &[OHLCVBar],
|
||||
bar_offset: usize,
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
) -> Result<Vec<(Tensor, Tensor)>> {
|
||||
let mut pairs = Vec::new();
|
||||
let n = norm_features.len();
|
||||
let limit = n.saturating_sub(1);
|
||||
let step_limit = if args.max_steps_per_epoch > 0 {
|
||||
args.max_steps_per_epoch.min(limit)
|
||||
} else {
|
||||
limit
|
||||
};
|
||||
|
||||
for i in 0..step_limit {
|
||||
let Some(feat) = norm_features.get(i) else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let bar_idx = i + bar_offset;
|
||||
let close_cur = bars.get(bar_idx).map(|b| b.close).unwrap_or(0.0);
|
||||
let close_next = bars
|
||||
.get(bar_idx + 1)
|
||||
.map(|b| b.close)
|
||||
.unwrap_or(close_cur);
|
||||
|
||||
let return_bps = if close_cur.abs() > 1e-10 {
|
||||
(close_next - close_cur) / close_cur * 10_000.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
let max_bps = args.max_bar_return * 10_000.0;
|
||||
let clipped = return_bps.clamp(-max_bps, max_bps);
|
||||
let spread = spread_cost_bps(close_cur, args.tick_size, args.spread_ticks);
|
||||
let net_return = clipped - (args.tx_cost_bps + spread);
|
||||
|
||||
let input_f32: Vec<f32> = feat.iter().map(|&v| v as f32).collect();
|
||||
let target_f32 = vec![net_return as f32];
|
||||
|
||||
let input_tensor = Tensor::from_vec(input_f32, &[1, args.feature_dim], device)
|
||||
.context("Failed to create input tensor")?;
|
||||
let target_tensor =
|
||||
Tensor::from_vec(target_f32, &[1, 1], device).context("Failed to create target tensor")?;
|
||||
|
||||
pairs.push((input_tensor, target_tensor));
|
||||
}
|
||||
|
||||
Ok(pairs)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Generic training loop
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Run one training epoch on mini-batches. Returns average loss.
|
||||
fn run_training_epoch(
|
||||
adapter: &mut dyn UnifiedTrainable,
|
||||
train_pairs: &[(Tensor, Tensor)],
|
||||
batch_size: usize,
|
||||
) -> Result<f64> {
|
||||
let mut epoch_loss_sum = 0.0_f64;
|
||||
let mut epoch_steps = 0_usize;
|
||||
let n_train = train_pairs.len();
|
||||
let mut batch_start = 0_usize;
|
||||
|
||||
while batch_start < n_train {
|
||||
let batch_end = (batch_start + batch_size).min(n_train);
|
||||
let Some(batch_slice) = train_pairs.get(batch_start..batch_end) else {
|
||||
break;
|
||||
};
|
||||
|
||||
let batch_inputs: Vec<&Tensor> = batch_slice.iter().map(|(inp, _)| inp).collect();
|
||||
let batch_targets: Vec<&Tensor> = batch_slice.iter().map(|(_, tgt)| tgt).collect();
|
||||
|
||||
if batch_inputs.is_empty() {
|
||||
batch_start = batch_end;
|
||||
continue;
|
||||
}
|
||||
|
||||
let input_cat =
|
||||
Tensor::cat(&batch_inputs, 0).context("Failed to concatenate batch inputs")?;
|
||||
let target_cat =
|
||||
Tensor::cat(&batch_targets, 0).context("Failed to concatenate batch targets")?;
|
||||
|
||||
adapter
|
||||
.zero_grad()
|
||||
.map_err(|e| anyhow::anyhow!("zero_grad failed: {}", e))?;
|
||||
let predictions = adapter
|
||||
.forward(&input_cat)
|
||||
.map_err(|e| anyhow::anyhow!("forward failed: {}", e))?;
|
||||
let loss = adapter
|
||||
.compute_loss(&predictions, &target_cat)
|
||||
.map_err(|e| anyhow::anyhow!("compute_loss failed: {}", e))?;
|
||||
let loss_val = loss
|
||||
.to_scalar::<f32>()
|
||||
.map_err(|e| anyhow::anyhow!("loss to_scalar failed: {}", e))?;
|
||||
|
||||
adapter
|
||||
.backward(&loss)
|
||||
.map_err(|e| anyhow::anyhow!("backward failed: {}", e))?;
|
||||
adapter
|
||||
.optimizer_step()
|
||||
.map_err(|e| anyhow::anyhow!("optimizer_step failed: {}", e))?;
|
||||
|
||||
epoch_loss_sum += loss_val as f64;
|
||||
epoch_steps += 1;
|
||||
batch_start = batch_end;
|
||||
}
|
||||
|
||||
Ok(if epoch_steps > 0 {
|
||||
epoch_loss_sum / epoch_steps as f64
|
||||
} else {
|
||||
0.0
|
||||
})
|
||||
}
|
||||
|
||||
/// Train a model on a single walk-forward fold. Returns best validation loss.
|
||||
fn train_fold(
|
||||
model_name: &str,
|
||||
fold: usize,
|
||||
train_pairs: &[(Tensor, Tensor)],
|
||||
val_pairs: &[(Tensor, Tensor)],
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
output_dir: &Path,
|
||||
) -> Result<f64> {
|
||||
info!(
|
||||
"[{}] Fold {} -- {} train, {} val pairs",
|
||||
model_name,
|
||||
fold,
|
||||
train_pairs.len(),
|
||||
val_pairs.len()
|
||||
);
|
||||
|
||||
let mut adapter = create_model(model_name, args.feature_dim, args.learning_rate, device)?;
|
||||
|
||||
let mut best_val_loss = f64::MAX;
|
||||
let mut epochs_without_improvement = 0_usize;
|
||||
|
||||
for epoch in 0..args.epochs {
|
||||
let epoch_start = Instant::now();
|
||||
|
||||
let avg_train_loss =
|
||||
run_training_epoch(adapter.as_mut(), train_pairs, args.batch_size)?;
|
||||
let val_loss = adapter
|
||||
.validate(val_pairs)
|
||||
.map_err(|e| anyhow::anyhow!("validation failed: {}", e))?;
|
||||
let elapsed = epoch_start.elapsed();
|
||||
|
||||
info!(
|
||||
" Fold {} Epoch {}/{}: train_loss={:.6}, val_loss={:.6}, lr={:.2e}, time={:.1}s",
|
||||
fold,
|
||||
epoch + 1,
|
||||
args.epochs,
|
||||
avg_train_loss,
|
||||
val_loss,
|
||||
adapter.get_learning_rate(),
|
||||
elapsed.as_secs_f64(),
|
||||
);
|
||||
|
||||
if val_loss < best_val_loss {
|
||||
best_val_loss = val_loss;
|
||||
epochs_without_improvement = 0;
|
||||
let ckpt_path = output_dir.join(format!("{}_fold{}_best", model_name, fold));
|
||||
let ckpt_str = ckpt_path.to_str().unwrap_or("checkpoint");
|
||||
if let Err(e) = adapter.save_checkpoint(ckpt_str) {
|
||||
warn!(" Failed to save checkpoint: {}", e);
|
||||
} else {
|
||||
info!(
|
||||
" [{}] New best val_loss={:.6}, checkpoint saved",
|
||||
model_name, val_loss
|
||||
);
|
||||
}
|
||||
} else {
|
||||
epochs_without_improvement += 1;
|
||||
}
|
||||
|
||||
if epochs_without_improvement >= args.patience {
|
||||
info!(
|
||||
" [{}] Early stopping at epoch {} (patience {} exhausted)",
|
||||
model_name,
|
||||
epoch + 1,
|
||||
args.patience
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Final checkpoint
|
||||
let final_path = output_dir.join(format!("{}_fold{}_final", model_name, fold));
|
||||
let final_str = final_path.to_str().unwrap_or("final");
|
||||
if let Err(e) = adapter.save_checkpoint(final_str) {
|
||||
warn!(" Failed to save final checkpoint: {}", e);
|
||||
}
|
||||
|
||||
Ok(best_val_loss)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Main
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
fn main() -> Result<()> {
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
|
||||
)
|
||||
.init();
|
||||
|
||||
let args = Args::parse();
|
||||
validate_model(&args.model)?;
|
||||
|
||||
let models_to_train: Vec<&str> = if args.model == "all" {
|
||||
ALL_MODELS.to_vec()
|
||||
} else {
|
||||
vec![args.model.as_str()]
|
||||
};
|
||||
|
||||
// Device selection
|
||||
let device = match Device::cuda_if_available(0) {
|
||||
Ok(d) => {
|
||||
info!("Using CUDA device");
|
||||
d
|
||||
}
|
||||
Err(_) => {
|
||||
info!("CUDA unavailable, using CPU");
|
||||
Device::Cpu
|
||||
}
|
||||
};
|
||||
|
||||
// Load OHLCV bars
|
||||
info!(
|
||||
"Loading OHLCV bars for {} from {}",
|
||||
args.symbol,
|
||||
args.data_dir.display()
|
||||
);
|
||||
let all_bars = load_all_bars(&args.data_dir, &args.symbol)?;
|
||||
info!("Loaded {} bars", all_bars.len());
|
||||
|
||||
if all_bars.len() < 100 {
|
||||
anyhow::bail!(
|
||||
"Insufficient data: {} bars (need >= 100)",
|
||||
all_bars.len()
|
||||
);
|
||||
}
|
||||
|
||||
// Walk-forward windows
|
||||
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 windows = generate_walk_forward_windows(&all_bars, &wf_config);
|
||||
if windows.is_empty() {
|
||||
anyhow::bail!(
|
||||
"No walk-forward windows generated. Data too short for configured window sizes."
|
||||
);
|
||||
}
|
||||
info!("Generated {} walk-forward folds", windows.len());
|
||||
|
||||
// Train each model
|
||||
for model_name in &models_to_train {
|
||||
info!("=== Training model: {} ===", model_name);
|
||||
let model_output = args.output_dir.join(model_name);
|
||||
std::fs::create_dir_all(&model_output)
|
||||
.with_context(|| format!("Failed to create {}", model_output.display()))?;
|
||||
|
||||
let mut fold_results: Vec<(usize, f64)> = Vec::new();
|
||||
|
||||
for window in &windows {
|
||||
let fold = window.fold;
|
||||
info!(
|
||||
"--- Fold {}: train={}, val={}, test={} bars ---",
|
||||
fold,
|
||||
window.train.len(),
|
||||
window.val.len(),
|
||||
window.test.len()
|
||||
);
|
||||
|
||||
let (train_pairs, val_pairs) =
|
||||
match prepare_fold_data(&window.train, &window.val, &args, &device) {
|
||||
Ok(data) => data,
|
||||
Err(e) => {
|
||||
warn!("Skipping fold {} -- {}", fold, e);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if train_pairs.is_empty() || val_pairs.is_empty() {
|
||||
warn!(
|
||||
"Skipping fold {} -- empty data after feature extraction",
|
||||
fold
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
match train_fold(
|
||||
model_name,
|
||||
fold,
|
||||
&train_pairs,
|
||||
&val_pairs,
|
||||
&args,
|
||||
&device,
|
||||
&model_output,
|
||||
) {
|
||||
Ok(best_val) => fold_results.push((fold, best_val)),
|
||||
Err(e) => error!("[{}] Fold {} failed: {}", model_name, fold, e),
|
||||
}
|
||||
}
|
||||
|
||||
// Summary
|
||||
info!(
|
||||
"=== {} Results ({} folds) ===",
|
||||
model_name,
|
||||
fold_results.len()
|
||||
);
|
||||
for (fold, loss) in &fold_results {
|
||||
info!(" Fold {}: best_val_loss = {:.6}", fold, loss);
|
||||
}
|
||||
if !fold_results.is_empty() {
|
||||
let avg: f64 =
|
||||
fold_results.iter().map(|(_, l)| l).sum::<f64>() / fold_results.len() as f64;
|
||||
info!(" Average: {:.6}", avg);
|
||||
}
|
||||
info!(" Checkpoints: {}", model_output.display());
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,920 +0,0 @@
|
||||
//! Continuous PPO Training Example with Parquet Data
|
||||
//!
|
||||
//! Trains a Continuous PPO model with Gaussian policies on market data from Parquet files:
|
||||
//! - Real OHLCV data + 51-dimensional features
|
||||
//! - Continuous position sizing in [-1.0, 1.0] range
|
||||
//! - PnL-based rewards with transaction costs
|
||||
//! - GAE advantages on real price trajectories
|
||||
//! - Dual learning rates (policy/value)
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (50 epochs, conservative exploration)
|
||||
//! cargo run -p ml --example train_continuous_ppo_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/ES_FUT_180d.parquet
|
||||
//!
|
||||
//! # Custom parameters with narrow action bounds
|
||||
//! cargo run -p ml --example train_continuous_ppo_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/ES_FUT_180d.parquet \
|
||||
//! --epochs 100 \
|
||||
//! --policy-lr 0.000001 \
|
||||
//! --value-lr 0.001 \
|
||||
//! --action-min -0.5 \
|
||||
//! --action-max 0.5
|
||||
//!
|
||||
//! # High exploration mode
|
||||
//! cargo run -p ml --example train_continuous_ppo_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/NQ_FUT_180d.parquet \
|
||||
//! --init-log-std 0.0
|
||||
//! ```
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::fs::File;
|
||||
use std::path::PathBuf;
|
||||
use tracing::{info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use arrow::array::{Array, Float64Array, PrimitiveArray, UInt64Array};
|
||||
use arrow::datatypes::TimestampNanosecondType;
|
||||
use arrow::record_batch::RecordBatch;
|
||||
use candle_core::Device;
|
||||
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
|
||||
|
||||
use ml::features::extraction::{extract_ml_features, OHLCVBar};
|
||||
use ml::ppo::continuous_ppo::{
|
||||
ContinuousPPO, ContinuousPPOConfig, ContinuousTrajectory,
|
||||
ContinuousTrajectoryBatch, ContinuousTrajectoryStep,
|
||||
};
|
||||
use ml::ppo::flow_policy::FlowPolicyConfig;
|
||||
use ml::ppo::gae::GAEConfig;
|
||||
use ml::evaluation::engine::{Action, EvaluationEngine};
|
||||
use ml::evaluation::metrics::{PerformanceMetrics, OHLCVBarF32 as MetricsOHLCVBar};
|
||||
|
||||
/// Results from dual-phase backtesting (exploration vs exploitation)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct DualPhaseBacktestResults {
|
||||
/// Metrics from exploration phase (epochs 0 to burn_in_epochs-1)
|
||||
pub exploration_metrics: PerformanceMetrics,
|
||||
/// Metrics from exploitation phase (epochs burn_in_epochs to total_epochs-1)
|
||||
pub exploitation_metrics: PerformanceMetrics,
|
||||
/// Number of burn-in epochs used
|
||||
pub burn_in_epochs: usize,
|
||||
/// Total number of epochs
|
||||
pub total_epochs: usize,
|
||||
}
|
||||
|
||||
/// Train Continuous PPO model on Parquet market data
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_continuous_ppo_parquet",
|
||||
about = "Train Continuous PPO model on Parquet market data"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Path to Parquet file with market data
|
||||
#[arg(long)]
|
||||
parquet_file: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "50")]
|
||||
epochs: usize,
|
||||
|
||||
/// Policy (actor) learning rate
|
||||
#[arg(long, default_value = "0.000001")]
|
||||
policy_lr: f64,
|
||||
|
||||
/// Value (critic) learning rate (reduced from 0.001 to prevent gradient explosion)
|
||||
#[arg(long, default_value = "0.0001")]
|
||||
value_lr: f64,
|
||||
|
||||
/// Minimum action bound (position size)
|
||||
#[arg(long, default_value = "-1.0")]
|
||||
action_min: f32,
|
||||
|
||||
/// Maximum action bound (position size)
|
||||
#[arg(long, default_value = "1.0")]
|
||||
action_max: f32,
|
||||
|
||||
/// Initial log standard deviation (exploration level)
|
||||
#[arg(long, default_value = "-1.0")]
|
||||
init_log_std: f32,
|
||||
|
||||
/// Checkpoint directory
|
||||
#[arg(long, default_value = "checkpoints/continuous_ppo")]
|
||||
checkpoint_dir: String,
|
||||
|
||||
/// Checkpoint save interval (epochs)
|
||||
#[arg(long, default_value = "10")]
|
||||
checkpoint_interval: usize,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
|
||||
/// Number of burn-in epochs for dual-phase backtesting
|
||||
/// Epochs 0 to burn_in_epochs-1 are "exploration" phase
|
||||
/// Epochs burn_in_epochs to total are "exploitation" phase
|
||||
#[arg(long, default_value = "50")]
|
||||
burn_in_epochs: usize,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 Starting Continuous PPO Training with Parquet Data");
|
||||
info!("Configuration:");
|
||||
info!(" • Parquet file: {}", opts.parquet_file);
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Policy learning rate: {}", opts.policy_lr);
|
||||
info!(" • Value learning rate: {}", opts.value_lr);
|
||||
info!(
|
||||
" • Action bounds: [{:.2}, {:.2}]",
|
||||
opts.action_min, opts.action_max
|
||||
);
|
||||
info!(" • Initial log std: {:.2}", opts.init_log_std);
|
||||
info!(" • GPU: CUDA if available (auto-fallback to CPU)");
|
||||
info!(" • Checkpoint directory: {}", opts.checkpoint_dir);
|
||||
info!(
|
||||
" • Checkpoint interval: {} epochs",
|
||||
opts.checkpoint_interval
|
||||
);
|
||||
|
||||
// Create checkpoint directory
|
||||
let checkpoint_path = PathBuf::from(&opts.checkpoint_dir);
|
||||
if !checkpoint_path.exists() {
|
||||
std::fs::create_dir_all(&checkpoint_path)
|
||||
.context("Failed to create checkpoint directory")?;
|
||||
info!("✅ Created checkpoint directory: {}", opts.checkpoint_dir);
|
||||
}
|
||||
|
||||
// Load market data from Parquet file
|
||||
info!("\n📊 Loading market data from Parquet file...");
|
||||
let bars = load_parquet_data(&opts.parquet_file)
|
||||
.await
|
||||
.context("Failed to load Parquet data")?;
|
||||
|
||||
info!("✅ Loaded {} OHLCV bars", bars.len());
|
||||
|
||||
// Extract 51-dimensional feature vectors (43 base + 8 OFI)
|
||||
info!("\n🏗️ Extracting 51-dimensional feature vectors...");
|
||||
let feature_vectors =
|
||||
extract_ml_features(&bars).context("Failed to extract 51-dimensional features")?;
|
||||
|
||||
info!(
|
||||
"✅ Extracted {} feature vectors (dim=51, warmup bars skipped=50)",
|
||||
feature_vectors.len()
|
||||
);
|
||||
|
||||
// Convert FeatureVector ([f64; 51]) to Vec<Vec<f32>> for PPO trainer
|
||||
let state_dim = 51;
|
||||
let market_data: Vec<Vec<f32>> = feature_vectors
|
||||
.iter()
|
||||
.map(|fv| fv.iter().map(|&v| v as f32).collect())
|
||||
.collect();
|
||||
|
||||
// Validate state dimensions
|
||||
if let Some(first_state) = market_data.first() {
|
||||
if first_state.len() != state_dim {
|
||||
return Err(anyhow::anyhow!(
|
||||
"State dimension mismatch: expected {}, got {}",
|
||||
state_dim,
|
||||
first_state.len()
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
info!(
|
||||
"✅ Feature extraction complete: {} samples",
|
||||
market_data.len()
|
||||
);
|
||||
|
||||
// Configure Flow-Based Policy for Continuous PPO
|
||||
let policy_config = FlowPolicyConfig {
|
||||
state_dim,
|
||||
action_dim: 1,
|
||||
context_dim: 128,
|
||||
num_layers: 4,
|
||||
scale_clamp: 5.0,
|
||||
};
|
||||
|
||||
let config = ContinuousPPOConfig {
|
||||
state_dim,
|
||||
policy_config,
|
||||
value_hidden_dims: vec![512, 384, 256, 128, 64],
|
||||
policy_learning_rate: opts.policy_lr,
|
||||
value_learning_rate: opts.value_lr,
|
||||
clip_epsilon: 0.2,
|
||||
value_loss_coeff: 0.5,
|
||||
entropy_coeff: 0.01,
|
||||
gae_config: GAEConfig {
|
||||
gamma: 0.99,
|
||||
lambda: 0.95,
|
||||
normalize_advantages: true,
|
||||
},
|
||||
batch_size: 2048,
|
||||
mini_batch_size: 64,
|
||||
num_epochs: 10,
|
||||
max_grad_norm: 0.5,
|
||||
};
|
||||
|
||||
// Create Continuous PPO agent
|
||||
let device = Device::cuda_if_available(0).unwrap_or(Device::Cpu);
|
||||
info!("Using device: {:?}", device);
|
||||
|
||||
let mut agent = ContinuousPPO::new(config.clone())
|
||||
.context("Failed to create Continuous PPO agent")?;
|
||||
|
||||
info!("✅ Continuous PPO agent initialized (state_dim={})", state_dim);
|
||||
|
||||
// Training loop
|
||||
info!("\n🏋️ Starting training...\n");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let transaction_cost_bps = 0.05; // 0.05% transaction cost
|
||||
let hold_penalty = 0.0001; // Small penalty for holding positions
|
||||
|
||||
for epoch in 0..opts.epochs {
|
||||
// Collect trajectories
|
||||
let trajectories = collect_trajectories(
|
||||
&agent,
|
||||
&market_data,
|
||||
transaction_cost_bps,
|
||||
hold_penalty,
|
||||
)?;
|
||||
|
||||
// Compute GAE advantages
|
||||
let mut batch = prepare_batch(trajectories, &config)?;
|
||||
|
||||
// Update agent
|
||||
let (policy_loss, value_loss) = agent
|
||||
.update(&mut batch)
|
||||
.context("Failed to update agent")?;
|
||||
|
||||
// Compute metrics
|
||||
let mean_reward = batch.advantages.iter().sum::<f32>() / batch.advantages.len() as f32;
|
||||
let sharpe_ratio = compute_sharpe_ratio(&batch.advantages);
|
||||
|
||||
info!(
|
||||
"📊 Epoch {}/{}: policy_loss={:.4}, value_loss={:.4}, mean_reward={:.4}, sharpe={:.4}",
|
||||
epoch + 1,
|
||||
opts.epochs,
|
||||
policy_loss,
|
||||
value_loss,
|
||||
mean_reward,
|
||||
sharpe_ratio
|
||||
);
|
||||
|
||||
// Save checkpoint
|
||||
if (epoch + 1) % opts.checkpoint_interval == 0 {
|
||||
let actor_path = checkpoint_path.join(format!("actor_epoch_{}.safetensors", epoch + 1));
|
||||
let critic_path =
|
||||
checkpoint_path.join(format!("critic_epoch_{}.safetensors", epoch + 1));
|
||||
|
||||
agent.actor.vars().save(&actor_path).with_context(|| {
|
||||
format!("Failed to save actor checkpoint: {:?}", actor_path)
|
||||
})?;
|
||||
agent.critic.vars().save(&critic_path).with_context(|| {
|
||||
format!("Failed to save critic checkpoint: {:?}", critic_path)
|
||||
})?;
|
||||
|
||||
info!(
|
||||
"💾 Checkpoint saved at epoch {} (actor: {:?}, critic: {:?})",
|
||||
epoch + 1,
|
||||
actor_path,
|
||||
critic_path
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Print final training metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
info!(" • Training steps: {}", agent.get_training_steps());
|
||||
|
||||
// Run dual-phase backtest to separate exploration from exploitation
|
||||
info!("\n🔍 Running dual-phase backtest (exploration vs exploitation)...");
|
||||
let backtest_start = std::time::Instant::now();
|
||||
|
||||
let dual_results = backtest_trained_agent_dual_phase(
|
||||
&agent,
|
||||
&market_data,
|
||||
&bars,
|
||||
opts.burn_in_epochs,
|
||||
opts.epochs,
|
||||
)?;
|
||||
|
||||
let backtest_duration = backtest_start.elapsed();
|
||||
|
||||
info!("\n📊 Dual-Phase Backtest Results:");
|
||||
info!(" • Backtest time: {:.2}s", backtest_duration.as_secs_f64());
|
||||
info!(" • Burn-in epochs: {} / {}", dual_results.burn_in_epochs, dual_results.total_epochs);
|
||||
|
||||
info!("\n--- Exploration Phase (Epochs 0-{}) ---", dual_results.burn_in_epochs.saturating_sub(1));
|
||||
info!(" • Total trades: {}", dual_results.exploration_metrics.total_trades);
|
||||
info!(" • Sharpe ratio: {:.4}", dual_results.exploration_metrics.sharpe_ratio);
|
||||
info!(" • Win rate: {:.2}%", dual_results.exploration_metrics.win_rate);
|
||||
info!(" • Max drawdown: {:.2}%", dual_results.exploration_metrics.max_drawdown_pct);
|
||||
info!(" • Total return: {:.2}%", dual_results.exploration_metrics.total_return_pct);
|
||||
info!(" • Average trade PnL: {:.4}", dual_results.exploration_metrics.avg_trade_pnl);
|
||||
info!(" • Final equity: ${:.2}", dual_results.exploration_metrics.final_equity);
|
||||
|
||||
info!("\n--- Exploitation Phase (Epochs {}-{}) ---",
|
||||
dual_results.burn_in_epochs,
|
||||
dual_results.total_epochs.saturating_sub(1));
|
||||
info!(" • Total trades: {}", dual_results.exploitation_metrics.total_trades);
|
||||
info!(" • Sharpe ratio: {:.4}", dual_results.exploitation_metrics.sharpe_ratio);
|
||||
info!(" • Win rate: {:.2}%", dual_results.exploitation_metrics.win_rate);
|
||||
info!(" • Max drawdown: {:.2}%", dual_results.exploitation_metrics.max_drawdown_pct);
|
||||
info!(" • Total return: {:.2}%", dual_results.exploitation_metrics.total_return_pct);
|
||||
info!(" • Average trade PnL: {:.4}", dual_results.exploitation_metrics.avg_trade_pnl);
|
||||
info!(" • Final equity: ${:.2}", dual_results.exploitation_metrics.final_equity);
|
||||
|
||||
// Calculate improvement (avoid division by zero)
|
||||
let sharpe_improvement = if dual_results.exploration_metrics.sharpe_ratio.abs() > 1e-6 {
|
||||
((dual_results.exploitation_metrics.sharpe_ratio - dual_results.exploration_metrics.sharpe_ratio)
|
||||
/ dual_results.exploration_metrics.sharpe_ratio.abs()) * 100.0
|
||||
} else if dual_results.exploitation_metrics.sharpe_ratio.abs() > 1e-6 {
|
||||
f64::INFINITY
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
info!("\n--- Performance Improvement ---");
|
||||
if sharpe_improvement.is_infinite() {
|
||||
info!(" • Sharpe improvement: +INF% (exploration Sharpe near zero)");
|
||||
} else {
|
||||
info!(" • Sharpe improvement: {:.2}%", sharpe_improvement);
|
||||
}
|
||||
|
||||
// Save final checkpoint
|
||||
let final_actor_path = checkpoint_path.join(format!("actor_epoch_{}.safetensors", opts.epochs));
|
||||
let final_critic_path =
|
||||
checkpoint_path.join(format!("critic_epoch_{}.safetensors", opts.epochs));
|
||||
|
||||
agent
|
||||
.actor
|
||||
.vars()
|
||||
.save(&final_actor_path)
|
||||
.with_context(|| format!("Failed to save final actor: {:?}", final_actor_path))?;
|
||||
agent
|
||||
.critic
|
||||
.vars()
|
||||
.save(&final_critic_path)
|
||||
.with_context(|| format!("Failed to save final critic: {:?}", final_critic_path))?;
|
||||
|
||||
info!(
|
||||
"\n💾 Final checkpoint saved to: {:?}, {:?}",
|
||||
final_actor_path, final_critic_path
|
||||
);
|
||||
info!("\n🎉 Continuous PPO training complete with Parquet data!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Collect continuous trajectories from market data
|
||||
fn collect_trajectories(
|
||||
agent: &ContinuousPPO,
|
||||
market_data: &[Vec<f32>],
|
||||
transaction_cost_bps: f32,
|
||||
hold_penalty: f32,
|
||||
) -> Result<Vec<ContinuousTrajectory>> {
|
||||
let mut trajectories = Vec::new();
|
||||
let mut current_trajectory = ContinuousTrajectory::new();
|
||||
|
||||
let mut position: f32 = 0.0; // Current position size
|
||||
let max_steps = market_data.len().min(2048);
|
||||
|
||||
// Diagnostic tracking
|
||||
let mut total_rewards = 0.0f32;
|
||||
let mut non_zero_rewards = 0usize;
|
||||
let mut position_samples = Vec::new();
|
||||
|
||||
for step_idx in 0..max_steps {
|
||||
let state = &market_data[step_idx];
|
||||
|
||||
// Get action from agent
|
||||
let (action, log_prob, value) = agent
|
||||
.act_with_log_prob(state)
|
||||
.context("Failed to select action")?;
|
||||
|
||||
let new_position = action.position_size();
|
||||
|
||||
// Compute reward using current position as old_position
|
||||
// (position holds the previous step's action, which is correct for PnL calculation)
|
||||
let log_return = state[state.len() - 1]; // Last feature is log return
|
||||
let reward = compute_reward(
|
||||
new_position,
|
||||
position, // old_position from previous step
|
||||
log_return,
|
||||
transaction_cost_bps,
|
||||
hold_penalty,
|
||||
);
|
||||
|
||||
let done = step_idx == max_steps - 1;
|
||||
|
||||
// Add step to trajectory
|
||||
let traj_step = ContinuousTrajectoryStep::new(
|
||||
state.clone(),
|
||||
action,
|
||||
log_prob,
|
||||
reward,
|
||||
value,
|
||||
done,
|
||||
);
|
||||
current_trajectory.add_step(traj_step);
|
||||
|
||||
// Update position for next step
|
||||
position = new_position;
|
||||
|
||||
// Track diagnostics (sample every 100 steps)
|
||||
if step_idx % 100 == 0 {
|
||||
position_samples.push(new_position);
|
||||
}
|
||||
total_rewards += reward;
|
||||
if reward.abs() > 1e-6 {
|
||||
non_zero_rewards += 1;
|
||||
}
|
||||
|
||||
// Start new trajectory every 1024 steps or at episode end
|
||||
if current_trajectory.len() >= 1024 || done {
|
||||
trajectories.push(current_trajectory);
|
||||
current_trajectory = ContinuousTrajectory::new();
|
||||
position = 0.0; // Reset position for new trajectory
|
||||
}
|
||||
}
|
||||
|
||||
// Add remaining trajectory if not empty
|
||||
if !current_trajectory.is_empty() {
|
||||
trajectories.push(current_trajectory);
|
||||
}
|
||||
|
||||
// Log diagnostic summary
|
||||
let avg_position = if !position_samples.is_empty() {
|
||||
position_samples.iter().sum::<f32>() / position_samples.len() as f32
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
info!(
|
||||
"Trajectory collection: {} steps, avg_reward={:.6}, non_zero_rewards={}/{}, avg_position={:.4}",
|
||||
max_steps, total_rewards / max_steps as f32, non_zero_rewards, max_steps, avg_position
|
||||
);
|
||||
|
||||
// Log first few position samples for debugging
|
||||
if !position_samples.is_empty() {
|
||||
let sample_slice = &position_samples[..position_samples.len().min(5)];
|
||||
info!("Position samples (first 5): {:?}", sample_slice);
|
||||
}
|
||||
|
||||
Ok(trajectories)
|
||||
}
|
||||
|
||||
/// Compute reward for continuous position sizing
|
||||
fn compute_reward(
|
||||
new_position: f32,
|
||||
old_position: f32,
|
||||
log_return: f32,
|
||||
transaction_cost_bps: f32,
|
||||
hold_penalty: f32,
|
||||
) -> f32 {
|
||||
// PnL from position and market movement (scaled to reasonable range)
|
||||
let pnl = old_position * log_return * 1000.0;
|
||||
|
||||
// Transaction cost (proportional to position change)
|
||||
let position_change = (new_position - old_position).abs();
|
||||
let transaction_cost = position_change * transaction_cost_bps;
|
||||
|
||||
// Hold penalty (small penalty for non-zero positions to encourage active trading)
|
||||
let hold_cost = new_position.abs() * hold_penalty;
|
||||
|
||||
// Total reward
|
||||
pnl - transaction_cost - hold_cost
|
||||
}
|
||||
|
||||
/// Prepare batch with GAE advantages
|
||||
fn prepare_batch(
|
||||
trajectories: Vec<ContinuousTrajectory>,
|
||||
config: &ContinuousPPOConfig,
|
||||
) -> Result<ContinuousTrajectoryBatch> {
|
||||
let gamma = config.gae_config.gamma;
|
||||
let lambda = config.gae_config.lambda;
|
||||
|
||||
// Compute GAE advantages for each trajectory
|
||||
let mut all_advantages = Vec::new();
|
||||
let mut all_returns = Vec::new();
|
||||
|
||||
for trajectory in &trajectories {
|
||||
let steps = trajectory.steps();
|
||||
|
||||
// Extract rewards, values, and dones
|
||||
let rewards: Vec<f32> = steps.iter().map(|s| s.reward).collect();
|
||||
let values: Vec<f32> = steps.iter().map(|s| s.value).collect();
|
||||
let dones: Vec<bool> = steps.iter().map(|s| s.done).collect();
|
||||
|
||||
// Compute GAE advantages
|
||||
let advantages = compute_gae_advantages(&rewards, &values, &dones, gamma, lambda);
|
||||
|
||||
// Compute returns
|
||||
let returns = compute_returns(&rewards, gamma);
|
||||
|
||||
all_advantages.extend(advantages);
|
||||
all_returns.extend(returns);
|
||||
}
|
||||
|
||||
// Create batch from trajectories
|
||||
let batch = ContinuousTrajectoryBatch::from_trajectories(
|
||||
trajectories,
|
||||
all_advantages,
|
||||
all_returns,
|
||||
);
|
||||
|
||||
Ok(batch)
|
||||
}
|
||||
|
||||
/// Compute GAE advantages
|
||||
fn compute_gae_advantages(
|
||||
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;
|
||||
|
||||
for t in (0..n).rev() {
|
||||
let reward = rewards[t];
|
||||
let value = values[t];
|
||||
let next_value = if t + 1 < n { values[t + 1] } else { 0.0 };
|
||||
let done = dones[t];
|
||||
|
||||
let mask = if done { 0.0 } else { 1.0 };
|
||||
let delta = reward + gamma * next_value * mask - value;
|
||||
gae = delta + gamma * lambda * mask * gae;
|
||||
|
||||
advantages[t] = gae;
|
||||
}
|
||||
|
||||
advantages
|
||||
}
|
||||
|
||||
/// Compute discounted returns
|
||||
fn compute_returns(rewards: &[f32], gamma: f32) -> Vec<f32> {
|
||||
let n = rewards.len();
|
||||
let mut returns = vec![0.0; n];
|
||||
let mut cumulative = 0.0;
|
||||
|
||||
for t in (0..n).rev() {
|
||||
cumulative = rewards[t] + gamma * cumulative;
|
||||
returns[t] = cumulative;
|
||||
}
|
||||
|
||||
returns
|
||||
}
|
||||
|
||||
/// Compute Sharpe ratio from rewards/advantages
|
||||
fn compute_sharpe_ratio(values: &[f32]) -> f32 {
|
||||
if values.is_empty() {
|
||||
return 0.0;
|
||||
}
|
||||
|
||||
let mean = values.iter().sum::<f32>() / values.len() as f32;
|
||||
let variance = values
|
||||
.iter()
|
||||
.map(|v| (v - mean).powi(2))
|
||||
.sum::<f32>()
|
||||
/ values.len() as f32;
|
||||
|
||||
let std = (variance + 1e-8).sqrt();
|
||||
|
||||
mean / std
|
||||
}
|
||||
|
||||
/// Load OHLCV data from Parquet file (Databento schema)
|
||||
async fn load_parquet_data(parquet_path: &str) -> Result<Vec<OHLCVBar>> {
|
||||
info!("Loading Parquet file: {}", parquet_path);
|
||||
|
||||
// Open Parquet file
|
||||
let file = File::open(parquet_path)
|
||||
.with_context(|| format!("Failed to open Parquet file: {}", parquet_path))?;
|
||||
|
||||
// Create Parquet reader
|
||||
let builder = ParquetRecordBatchReaderBuilder::try_new(file)
|
||||
.with_context(|| "Failed to create Parquet reader")?;
|
||||
|
||||
let reader = builder
|
||||
.build()
|
||||
.with_context(|| "Failed to build Parquet reader")?;
|
||||
|
||||
// Read all batches
|
||||
let mut all_ohlcv_bars = Vec::new();
|
||||
|
||||
for batch_result in reader {
|
||||
let batch: RecordBatch = batch_result.with_context(|| "Failed to read record batch")?;
|
||||
|
||||
// Extract columns from Databento Parquet schema:
|
||||
// Column 3: open, Column 4: high, Column 5: low, Column 6: close
|
||||
// Column 7: volume, Column 9: ts_event (Timestamp(Nanosecond, Some("UTC")))
|
||||
let timestamps = batch
|
||||
.column(9)
|
||||
.as_any()
|
||||
.downcast_ref::<PrimitiveArray<TimestampNanosecondType>>()
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!(
|
||||
"Failed to downcast timestamp column. Expected Timestamp(Nanosecond), got: {:?}",
|
||||
batch.column(9).data_type()
|
||||
)
|
||||
})?;
|
||||
|
||||
let opens = batch
|
||||
.column(3)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast open column"))?;
|
||||
|
||||
let highs = batch
|
||||
.column(4)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast high column"))?;
|
||||
|
||||
let lows = batch
|
||||
.column(5)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast low column"))?;
|
||||
|
||||
let closes = batch
|
||||
.column(6)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast close column"))?;
|
||||
|
||||
let volumes = batch
|
||||
.column(7)
|
||||
.as_any()
|
||||
.downcast_ref::<UInt64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast volume column"))?;
|
||||
|
||||
// Convert to OHLCVBar structs
|
||||
for i in 0..batch.num_rows() {
|
||||
let timestamp_ns = timestamps.value(i);
|
||||
|
||||
// Convert nanoseconds to DateTime<Utc>
|
||||
let timestamp = chrono::DateTime::from_timestamp(
|
||||
(timestamp_ns / 1_000_000_000) as i64,
|
||||
(timestamp_ns % 1_000_000_000) as u32,
|
||||
)
|
||||
.unwrap_or_else(|| chrono::Utc::now());
|
||||
|
||||
let bar = OHLCVBar {
|
||||
timestamp,
|
||||
open: opens.value(i),
|
||||
high: highs.value(i),
|
||||
low: lows.value(i),
|
||||
close: closes.value(i),
|
||||
volume: volumes.value(i) as f64,
|
||||
};
|
||||
|
||||
all_ohlcv_bars.push(bar);
|
||||
}
|
||||
}
|
||||
|
||||
info!("✅ Loaded {} OHLCV bars from Parquet", all_ohlcv_bars.len());
|
||||
|
||||
Ok(all_ohlcv_bars)
|
||||
}
|
||||
|
||||
/// Run backtest on trained PPO agent to compute actual trading performance
|
||||
fn backtest_trained_agent(
|
||||
agent: &ContinuousPPO,
|
||||
market_data: &[Vec<f32>],
|
||||
bars: &[OHLCVBar],
|
||||
) -> Result<PerformanceMetrics> {
|
||||
// Create evaluation engine with $10K initial capital
|
||||
let mut engine = EvaluationEngine::new(10000.0);
|
||||
|
||||
// Skip warmup bars (50 bars used for feature extraction)
|
||||
let warmup_bars = 50;
|
||||
let mut metrics_bars = Vec::new();
|
||||
|
||||
// Run backtest using trained agent (greedy, no exploration)
|
||||
for (step_idx, state) in market_data.iter().enumerate() {
|
||||
// Get corresponding OHLCV bar (accounting for warmup)
|
||||
let bar_idx = step_idx + warmup_bars;
|
||||
if bar_idx >= bars.len() {
|
||||
warn!("Bar index {} exceeds available bars ({})", bar_idx, bars.len());
|
||||
break;
|
||||
}
|
||||
let bar = &bars[bar_idx];
|
||||
|
||||
// Get action from agent (greedy policy - use mean without sampling)
|
||||
let (action, _value) = agent
|
||||
.act(state)
|
||||
.context("Failed to select action during backtest")?;
|
||||
|
||||
// Convert continuous position size to discrete trading action
|
||||
// Position size in [-1.0, 1.0]: negative = short, positive = long, near-zero = hold
|
||||
let trading_action = continuous_to_discrete_action(action.position_size());
|
||||
|
||||
// Convert to OHLCVBar for metrics (same timestamp, prices from original bar)
|
||||
let metrics_bar = MetricsOHLCVBar {
|
||||
timestamp: bar.timestamp.timestamp(),
|
||||
open: bar.open as f32,
|
||||
high: bar.high as f32,
|
||||
low: bar.low as f32,
|
||||
close: bar.close as f32,
|
||||
volume: bar.volume as f32,
|
||||
};
|
||||
|
||||
// Process bar in evaluation engine
|
||||
engine.process_bar(step_idx, &metrics_bar, trading_action);
|
||||
metrics_bars.push(metrics_bar);
|
||||
}
|
||||
|
||||
// Close any open position at end of backtest
|
||||
if let Some(last_bar) = metrics_bars.last() {
|
||||
engine.close_position(metrics_bars.len() - 1, last_bar);
|
||||
}
|
||||
|
||||
// Calculate performance metrics from actual trades
|
||||
let metrics = PerformanceMetrics::from_trades(
|
||||
&engine.trades,
|
||||
engine.initial_capital,
|
||||
&metrics_bars,
|
||||
);
|
||||
|
||||
Ok(metrics)
|
||||
}
|
||||
|
||||
/// Run dual-phase backtest to separate exploration from exploitation performance
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `agent` - Trained Continuous PPO agent
|
||||
/// * `market_data` - Feature vectors (51-dim)
|
||||
/// * `bars` - OHLCV bars (for timestamping)
|
||||
/// * `burn_in_epochs` - Number of epochs to treat as "exploration" phase
|
||||
/// * `total_epochs` - Total number of training epochs
|
||||
///
|
||||
/// # Returns
|
||||
/// Dual-phase backtest results with separate metrics for each phase
|
||||
fn backtest_trained_agent_dual_phase(
|
||||
agent: &ContinuousPPO,
|
||||
market_data: &[Vec<f32>],
|
||||
bars: &[OHLCVBar],
|
||||
burn_in_epochs: usize,
|
||||
total_epochs: usize,
|
||||
) -> Result<DualPhaseBacktestResults> {
|
||||
// Create separate engines for each phase
|
||||
let mut exploration_engine = EvaluationEngine::new(10000.0);
|
||||
let mut exploitation_engine = EvaluationEngine::new(10000.0);
|
||||
|
||||
// Skip warmup bars (50 bars used for feature extraction)
|
||||
let warmup_bars = 50;
|
||||
let mut exploration_bars = Vec::new();
|
||||
let mut exploitation_bars = Vec::new();
|
||||
|
||||
// Calculate steps per epoch (approximate)
|
||||
let total_steps = market_data.len();
|
||||
let steps_per_epoch = 1024; // From trajectory collection
|
||||
let actual_epochs = (total_steps + steps_per_epoch - 1) / steps_per_epoch;
|
||||
|
||||
info!(
|
||||
"Dual-phase backtest: burn_in={}, total_epochs={}, actual_epochs={}, total_steps={}",
|
||||
burn_in_epochs, total_epochs, actual_epochs, total_steps
|
||||
);
|
||||
|
||||
// Run backtest through all market data
|
||||
for (step_idx, state) in market_data.iter().enumerate() {
|
||||
// Calculate current epoch
|
||||
let current_epoch = step_idx / steps_per_epoch;
|
||||
|
||||
// Get corresponding OHLCV bar (accounting for warmup)
|
||||
let bar_idx = step_idx + warmup_bars;
|
||||
if bar_idx >= bars.len() {
|
||||
warn!("Bar index {} exceeds available bars ({})", bar_idx, bars.len());
|
||||
break;
|
||||
}
|
||||
let bar = &bars[bar_idx];
|
||||
|
||||
// Get action from agent (greedy policy - use mean without sampling)
|
||||
let (action, _value) = agent
|
||||
.act(state)
|
||||
.context("Failed to select action during backtest")?;
|
||||
|
||||
// Convert continuous position size to discrete trading action
|
||||
let trading_action = continuous_to_discrete_action(action.position_size());
|
||||
|
||||
// Convert to OHLCVBar for metrics
|
||||
let metrics_bar = MetricsOHLCVBar {
|
||||
timestamp: bar.timestamp.timestamp(),
|
||||
open: bar.open as f32,
|
||||
high: bar.high as f32,
|
||||
low: bar.low as f32,
|
||||
close: bar.close as f32,
|
||||
volume: bar.volume as f32,
|
||||
};
|
||||
|
||||
// Route to appropriate engine based on current epoch
|
||||
if current_epoch < burn_in_epochs {
|
||||
// Exploration phase
|
||||
exploration_engine.process_bar(step_idx, &metrics_bar, trading_action);
|
||||
exploration_bars.push(metrics_bar);
|
||||
} else {
|
||||
// Exploitation phase
|
||||
exploitation_engine.process_bar(step_idx, &metrics_bar, trading_action);
|
||||
exploitation_bars.push(metrics_bar);
|
||||
}
|
||||
}
|
||||
|
||||
// Close any open positions in both engines
|
||||
if let Some(last_bar) = exploration_bars.last() {
|
||||
let last_idx = exploration_bars.len() - 1;
|
||||
exploration_engine.close_position(last_idx, last_bar);
|
||||
}
|
||||
|
||||
if let Some(last_bar) = exploitation_bars.last() {
|
||||
let last_idx = exploitation_bars.len() - 1;
|
||||
exploitation_engine.close_position(last_idx, last_bar);
|
||||
}
|
||||
|
||||
// Calculate metrics for each phase
|
||||
let exploration_metrics = if burn_in_epochs > 0 && !exploration_engine.trades.is_empty() {
|
||||
PerformanceMetrics::from_trades(
|
||||
&exploration_engine.trades,
|
||||
exploration_engine.initial_capital,
|
||||
&exploration_bars,
|
||||
)
|
||||
} else {
|
||||
PerformanceMetrics::default()
|
||||
};
|
||||
|
||||
let exploitation_metrics = if actual_epochs > burn_in_epochs && !exploitation_engine.trades.is_empty() {
|
||||
PerformanceMetrics::from_trades(
|
||||
&exploitation_engine.trades,
|
||||
exploitation_engine.initial_capital,
|
||||
&exploitation_bars,
|
||||
)
|
||||
} else {
|
||||
PerformanceMetrics::default()
|
||||
};
|
||||
|
||||
info!(
|
||||
"Phase distribution: exploration_trades={}, exploitation_trades={}",
|
||||
exploration_engine.trades.len(),
|
||||
exploitation_engine.trades.len()
|
||||
);
|
||||
|
||||
Ok(DualPhaseBacktestResults {
|
||||
exploration_metrics,
|
||||
exploitation_metrics,
|
||||
burn_in_epochs,
|
||||
total_epochs,
|
||||
})
|
||||
}
|
||||
|
||||
/// Convert continuous position size to discrete trading action
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `position_size` - Continuous position size in [-1.0, 1.0]
|
||||
///
|
||||
/// # Returns
|
||||
/// Discrete action (Buy, Hold, Sell)
|
||||
///
|
||||
/// # Logic
|
||||
/// - position_size > 0.3: Buy (strong long signal)
|
||||
/// - position_size < -0.3: Sell (strong short signal)
|
||||
/// - otherwise: Hold (weak signal or neutral)
|
||||
fn continuous_to_discrete_action(position_size: f32) -> Action {
|
||||
const BUY_THRESHOLD: f32 = 0.3;
|
||||
const SELL_THRESHOLD: f32 = -0.3;
|
||||
|
||||
if position_size > BUY_THRESHOLD {
|
||||
Action::Buy
|
||||
} else if position_size < SELL_THRESHOLD {
|
||||
Action::Sell
|
||||
} else {
|
||||
Action::Hold
|
||||
}
|
||||
}
|
||||
@@ -1,703 +0,0 @@
|
||||
//! **Diffusion Model Training on Real OHLCV Market Data**
|
||||
//!
|
||||
//! Trains a DDPM/DDIM Diffusion model on real futures OHLCV data loaded from
|
||||
//! Databento DBN files. The model learns to denoise price path sequences,
|
||||
//! which can later be used for scenario generation and risk analysis.
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Quick training (10 epochs, CPU)
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_diffusion_dbn --release
|
||||
//!
|
||||
//! # Production training (100 epochs, GPU if available)
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_diffusion_dbn --release -- \
|
||||
//! --epochs 100 --batch-size 32 --learning-rate 1e-4 \
|
||||
//! --data-dir data/cache/futures-baseline --symbol ES.FUT
|
||||
//! ```
|
||||
//!
|
||||
//! # Output
|
||||
//!
|
||||
//! Checkpoints saved to `<output-dir>/diffusion_weights.safetensors` with
|
||||
//! accompanying `diffusion_meta.json` metadata.
|
||||
|
||||
#![allow(unused_crate_dependencies)]
|
||||
#![deny(
|
||||
clippy::unwrap_used,
|
||||
clippy::expect_used,
|
||||
clippy::panic,
|
||||
clippy::indexing_slicing
|
||||
)]
|
||||
#![allow(
|
||||
clippy::integer_division,
|
||||
clippy::doc_markdown,
|
||||
clippy::too_many_lines,
|
||||
clippy::missing_const_for_fn
|
||||
)]
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Instant;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::{Device, Tensor};
|
||||
use clap::Parser;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use ml::diffusion::config::{DiffusionConfig, NoiseSchedule};
|
||||
use ml::diffusion::trainable::DiffusionTrainableAdapter;
|
||||
use ml::training::unified_trainer::UnifiedTrainable;
|
||||
use ml::types::OHLCVBar;
|
||||
|
||||
#[allow(unreachable_pub)]
|
||||
mod baseline_common;
|
||||
use baseline_common::{load_all_bars, spread_cost_bps};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CLI Arguments
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Train a DDPM/DDIM Diffusion model on real OHLCV data from DBN files.
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "train_diffusion_dbn", about = "Train Diffusion model on DBN OHLCV data")]
|
||||
struct Args {
|
||||
/// Path to directory containing .dbn.zst files (with symbol subdirectories)
|
||||
#[arg(long, default_value = "data/cache/futures-baseline")]
|
||||
data_dir: PathBuf,
|
||||
|
||||
/// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT")
|
||||
#[arg(long, default_value = "ES.FUT")]
|
||||
symbol: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value_t = 20)]
|
||||
epochs: usize,
|
||||
|
||||
/// Batch size (keep <= 32 for RTX 3050 Ti 4GB VRAM)
|
||||
#[arg(long, default_value_t = 16)]
|
||||
batch_size: usize,
|
||||
|
||||
/// Learning rate for the optimizer
|
||||
#[arg(long, default_value_t = 1e-4)]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Maximum training steps per epoch (0 = use all available data)
|
||||
#[arg(long, default_value_t = 500)]
|
||||
max_steps_per_epoch: usize,
|
||||
|
||||
/// Output directory for checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models/diffusion")]
|
||||
output_dir: PathBuf,
|
||||
|
||||
/// Sequence length for diffusion model input
|
||||
#[arg(long, default_value_t = 64)]
|
||||
seq_len: usize,
|
||||
|
||||
/// Hidden dimension for the denoiser network
|
||||
#[arg(long, default_value_t = 128)]
|
||||
hidden_dim: usize,
|
||||
|
||||
/// Number of denoiser layers
|
||||
#[arg(long, default_value_t = 3)]
|
||||
num_layers: usize,
|
||||
|
||||
/// Number of diffusion timesteps (noise levels)
|
||||
#[arg(long, default_value_t = 1000)]
|
||||
num_timesteps: usize,
|
||||
|
||||
/// Number of DDIM sampling steps for inference
|
||||
#[arg(long, default_value_t = 10)]
|
||||
sampling_steps: usize,
|
||||
|
||||
/// Early stopping patience (epochs without improvement)
|
||||
#[arg(long, default_value_t = 10)]
|
||||
patience: usize,
|
||||
|
||||
/// Round-trip commission cost in basis points
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
tx_cost_bps: f64,
|
||||
|
||||
/// Instrument tick size in price units (ES=0.25)
|
||||
#[arg(long, default_value_t = 0.25)]
|
||||
tick_size: f64,
|
||||
|
||||
/// Typical bid-ask spread in ticks
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
spread_ticks: f64,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Data preparation
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Extract normalized net-of-cost return sequences from OHLCV bars.
|
||||
///
|
||||
/// Returns a vector of f32 sequences, each of length `seq_len`, created by
|
||||
/// sliding a window over bar-to-bar returns minus transaction and spread costs.
|
||||
/// Returns are z-score normalized within each window for stable training.
|
||||
fn prepare_sequences(
|
||||
bars: &[OHLCVBar],
|
||||
seq_len: usize,
|
||||
tx_cost_bps: f64,
|
||||
tick_size: f64,
|
||||
spread_ticks: f64,
|
||||
) -> Vec<Vec<f32>> {
|
||||
if bars.len() < seq_len + 1 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
// Compute net-of-cost returns for each consecutive pair of bars
|
||||
let mut returns: Vec<f64> = Vec::with_capacity(bars.len().saturating_sub(1));
|
||||
for i in 0..bars.len().saturating_sub(1) {
|
||||
let prev = match bars.get(i) {
|
||||
Some(b) => b,
|
||||
None => continue,
|
||||
};
|
||||
let cur = match bars.get(i + 1) {
|
||||
Some(b) => b,
|
||||
None => continue,
|
||||
};
|
||||
let raw_ret = if prev.close.abs() > 1e-10 {
|
||||
(cur.close - prev.close) / prev.close
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
let spread = spread_cost_bps(prev.close, tick_size, spread_ticks);
|
||||
let cost = (tx_cost_bps + spread) / 10_000.0;
|
||||
returns.push(raw_ret - cost);
|
||||
}
|
||||
|
||||
let n_sequences = returns.len().saturating_sub(seq_len);
|
||||
let mut sequences = Vec::with_capacity(n_sequences);
|
||||
|
||||
for start in 0..n_sequences {
|
||||
let end = start + seq_len;
|
||||
let window: Vec<f64> = returns
|
||||
.get(start..end)
|
||||
.map(|s| s.to_vec())
|
||||
.unwrap_or_default();
|
||||
|
||||
if window.len() != seq_len {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Z-score normalize within the window for stable training
|
||||
let mean: f64 = window.iter().sum::<f64>() / seq_len as f64;
|
||||
let var: f64 = window.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / seq_len as f64;
|
||||
let std_dev = var.sqrt().max(1e-8);
|
||||
|
||||
let normalized: Vec<f32> = window.iter().map(|&x| ((x - mean) / std_dev) as f32).collect();
|
||||
sequences.push(normalized);
|
||||
}
|
||||
|
||||
sequences
|
||||
}
|
||||
|
||||
/// Build a batch tensor from a slice of sequences.
|
||||
///
|
||||
/// Returns a tensor of shape `(batch_size, seq_len)` or `None` if the slice
|
||||
/// is too small.
|
||||
fn build_batch(
|
||||
sequences: &[Vec<f32>],
|
||||
start_idx: usize,
|
||||
batch_size: usize,
|
||||
seq_len: usize,
|
||||
device: &Device,
|
||||
) -> Result<Option<Tensor>> {
|
||||
let end_idx = (start_idx + batch_size).min(sequences.len());
|
||||
if end_idx <= start_idx {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let actual_batch = end_idx - start_idx;
|
||||
let mut flat = Vec::with_capacity(actual_batch * seq_len);
|
||||
for idx in start_idx..end_idx {
|
||||
if let Some(seq) = sequences.get(idx) {
|
||||
flat.extend_from_slice(seq);
|
||||
}
|
||||
}
|
||||
|
||||
if flat.len() != actual_batch * seq_len {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let tensor = Tensor::from_vec(flat, (actual_batch, seq_len), device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create batch tensor: {}", e))?;
|
||||
|
||||
Ok(Some(tensor))
|
||||
}
|
||||
|
||||
/// Try to build a batch, wrapping to start of data if current position is past the end.
|
||||
fn build_batch_with_wraparound(
|
||||
sequences: &[Vec<f32>],
|
||||
batch_start: &mut usize,
|
||||
batch_size: usize,
|
||||
seq_len: usize,
|
||||
device: &Device,
|
||||
) -> Result<Option<Tensor>> {
|
||||
if let Some(batch) = build_batch(sequences, *batch_start, batch_size, seq_len, device)? {
|
||||
return Ok(Some(batch));
|
||||
}
|
||||
// Wrap around to beginning of data
|
||||
*batch_start = 0;
|
||||
build_batch(sequences, *batch_start, batch_size, seq_len, device)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Device selection
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Select CUDA device if available, otherwise fall back to CPU.
|
||||
fn select_device() -> Device {
|
||||
match Device::cuda_if_available(0) {
|
||||
Ok(dev) => {
|
||||
if dev.is_cuda() {
|
||||
info!("Using CUDA device 0");
|
||||
} else {
|
||||
info!("CUDA not available, using CPU");
|
||||
}
|
||||
dev
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("CUDA init failed ({}), falling back to CPU", e);
|
||||
Device::Cpu
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Model construction
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Build a `DiffusionConfig` from CLI arguments.
|
||||
fn build_config(args: &Args) -> DiffusionConfig {
|
||||
DiffusionConfig {
|
||||
num_timesteps: args.num_timesteps,
|
||||
sampling_steps: args.sampling_steps,
|
||||
seq_len: args.seq_len,
|
||||
feature_dim: 1,
|
||||
hidden_dim: args.hidden_dim,
|
||||
num_layers: args.num_layers,
|
||||
time_embed_dim: 32,
|
||||
schedule: NoiseSchedule::Cosine,
|
||||
learning_rate: args.learning_rate,
|
||||
weight_decay: 1e-4,
|
||||
grad_clip: 1.0,
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Training
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// State tracked across the training loop.
|
||||
struct TrainState {
|
||||
best_val_loss: f64,
|
||||
epochs_without_improvement: usize,
|
||||
loss_history: Vec<f64>,
|
||||
}
|
||||
|
||||
/// Execute a single training step (forward, loss, backward, optimizer).
|
||||
///
|
||||
/// Returns the loss value if the step succeeded, or `None` if it should be skipped.
|
||||
#[allow(clippy::cognitive_complexity)]
|
||||
fn execute_train_step(
|
||||
adapter: &mut DiffusionTrainableAdapter,
|
||||
batch: &Tensor,
|
||||
epoch: usize,
|
||||
step: usize,
|
||||
) -> Result<Option<f64>> {
|
||||
adapter.zero_grad()?;
|
||||
|
||||
let predictions = match adapter.forward(batch) {
|
||||
Ok(p) => p,
|
||||
Err(e) => {
|
||||
warn!(" Forward pass error at epoch {} step {}: {}", epoch + 1, step, e);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
// Compute loss: MSE between predicted noise and input.
|
||||
// The diffusion adapter generates noise targets internally during forward;
|
||||
// using the input as pseudo-target exercises the denoiser gradient path.
|
||||
let loss = match adapter.compute_loss(&predictions, batch) {
|
||||
Ok(l) => l,
|
||||
Err(e) => {
|
||||
warn!(" compute_loss error at epoch {} step {}: {}", epoch + 1, step, e);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
let loss_val = loss
|
||||
.to_scalar::<f32>()
|
||||
.map(|v| v as f64)
|
||||
.unwrap_or(f64::NAN);
|
||||
|
||||
if !loss_val.is_finite() {
|
||||
warn!(" NaN/Inf loss at epoch {} step {}, skipping", epoch + 1, step);
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Err(e) = adapter.backward(&loss) {
|
||||
warn!(" Backward error at epoch {} step {}: {}", epoch + 1, step, e);
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Err(e) = adapter.optimizer_step() {
|
||||
warn!(" Optimizer step error: {}", e);
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some(loss_val))
|
||||
}
|
||||
|
||||
/// Run one training epoch and return the average training loss.
|
||||
fn run_train_epoch(
|
||||
adapter: &mut DiffusionTrainableAdapter,
|
||||
train_sequences: &[Vec<f32>],
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
steps_per_epoch: usize,
|
||||
epoch: usize,
|
||||
) -> Result<f64> {
|
||||
let mut epoch_loss = 0.0_f64;
|
||||
let mut epoch_steps = 0_usize;
|
||||
let mut batch_start = 0_usize;
|
||||
|
||||
for step in 0..steps_per_epoch {
|
||||
let Some(batch) = build_batch_with_wraparound(
|
||||
train_sequences,
|
||||
&mut batch_start,
|
||||
args.batch_size,
|
||||
args.seq_len,
|
||||
device,
|
||||
)? else {
|
||||
break;
|
||||
};
|
||||
|
||||
if let Some(loss_val) = execute_train_step(adapter, &batch, epoch, step)? {
|
||||
epoch_loss += loss_val;
|
||||
epoch_steps += 1;
|
||||
}
|
||||
|
||||
// Advance batch position with wraparound
|
||||
batch_start += args.batch_size;
|
||||
if batch_start >= train_sequences.len() {
|
||||
batch_start = 0;
|
||||
}
|
||||
}
|
||||
|
||||
if epoch_steps > 0 {
|
||||
Ok(epoch_loss / epoch_steps as f64)
|
||||
} else {
|
||||
Ok(f64::NAN)
|
||||
}
|
||||
}
|
||||
|
||||
/// Save a checkpoint to the given subdirectory under the output dir.
|
||||
fn save_checkpoint(
|
||||
adapter: &DiffusionTrainableAdapter,
|
||||
output_dir: &Path,
|
||||
subdir: &str,
|
||||
label: &str,
|
||||
) -> Result<()> {
|
||||
let dir = output_dir.join(subdir);
|
||||
std::fs::create_dir_all(&dir)
|
||||
.with_context(|| format!("Failed to create {} dir: {}", label, dir.display()))?;
|
||||
match adapter.save_checkpoint(dir.to_str().unwrap_or(subdir)) {
|
||||
Ok(path) => info!(" {} checkpoint saved: {}", label, path),
|
||||
Err(e) => warn!(" Failed to save {} checkpoint: {}", label, e),
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Print the final training summary.
|
||||
fn print_summary(
|
||||
state: &TrainState,
|
||||
adapter: &DiffusionTrainableAdapter,
|
||||
training_time: std::time::Duration,
|
||||
total_time: std::time::Duration,
|
||||
output_dir: &Path,
|
||||
) {
|
||||
let final_metrics = adapter.collect_metrics();
|
||||
println!();
|
||||
println!("{}", "=".repeat(80));
|
||||
println!(" Training Complete");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
println!(" Results:");
|
||||
println!(" Epochs trained: {}", state.loss_history.len());
|
||||
println!(" Final train loss: {:.6}", state.loss_history.last().copied().unwrap_or(f64::NAN));
|
||||
println!(" Best val loss: {:.6}", state.best_val_loss);
|
||||
println!(" Final LR: {:.1e}", final_metrics.learning_rate);
|
||||
println!(" Total steps: {}", adapter.get_step());
|
||||
println!();
|
||||
println!(" Performance:");
|
||||
println!(
|
||||
" Training time: {:.1}s ({:.1} min)",
|
||||
training_time.as_secs_f64(),
|
||||
training_time.as_secs_f64() / 60.0,
|
||||
);
|
||||
if !state.loss_history.is_empty() {
|
||||
println!(
|
||||
" Avg epoch time: {:.2}s",
|
||||
training_time.as_secs_f64() / state.loss_history.len() as f64,
|
||||
);
|
||||
}
|
||||
println!();
|
||||
println!(" Checkpoints:");
|
||||
println!(" Best: {}", output_dir.join("best").display());
|
||||
println!(" Final: {}", output_dir.join("final").display());
|
||||
println!();
|
||||
println!(
|
||||
" Total time: {:.1}s ({:.1} min)",
|
||||
total_time.as_secs_f64(),
|
||||
total_time.as_secs_f64() / 60.0,
|
||||
);
|
||||
println!();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Validation
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Run validation on held-out sequences and return average loss.
|
||||
fn run_validation(
|
||||
adapter: &mut DiffusionTrainableAdapter,
|
||||
val_sequences: &[Vec<f32>],
|
||||
batch_size: usize,
|
||||
seq_len: usize,
|
||||
device: &Device,
|
||||
) -> Result<f64> {
|
||||
let mut total_loss = 0.0_f64;
|
||||
let mut total_batches = 0_usize;
|
||||
let mut batch_start = 0_usize;
|
||||
|
||||
while batch_start < val_sequences.len() {
|
||||
let Some(batch) = build_batch(val_sequences, batch_start, batch_size, seq_len, device)? else {
|
||||
break;
|
||||
};
|
||||
|
||||
let Ok(predictions) = adapter.forward(&batch) else {
|
||||
break;
|
||||
};
|
||||
|
||||
let Ok(loss) = adapter.compute_loss(&predictions, &batch) else {
|
||||
break;
|
||||
};
|
||||
|
||||
let loss_val = loss
|
||||
.to_scalar::<f32>()
|
||||
.map(|v| v as f64)
|
||||
.unwrap_or(f64::NAN);
|
||||
|
||||
if loss_val.is_finite() {
|
||||
total_loss += loss_val;
|
||||
total_batches += 1;
|
||||
}
|
||||
|
||||
batch_start += batch_size;
|
||||
}
|
||||
|
||||
if total_batches > 0 {
|
||||
Ok(total_loss / total_batches as f64)
|
||||
} else {
|
||||
Ok(f64::MAX)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Data loading and model init
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Load bars and prepare train/val sequences. Returns (sequences, train_size).
|
||||
#[allow(clippy::cognitive_complexity)]
|
||||
fn load_and_prepare_data(args: &Args) -> Result<(Vec<Vec<f32>>, usize)> {
|
||||
info!("Step 1/4: 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(), args.symbol);
|
||||
}
|
||||
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!("Step 2/4: Preparing training sequences (seq_len={})...", args.seq_len);
|
||||
let sequences = prepare_sequences(&bars, args.seq_len, args.tx_cost_bps, args.tick_size, args.spread_ticks);
|
||||
if sequences.is_empty() {
|
||||
anyhow::bail!(
|
||||
"No sequences generated. Need at least {} bars, got {}.",
|
||||
args.seq_len,
|
||||
bars.len()
|
||||
);
|
||||
}
|
||||
|
||||
// Split 90/10 into train/val
|
||||
let val_size = (sequences.len() / 10).max(1);
|
||||
let train_size = sequences.len().saturating_sub(val_size);
|
||||
|
||||
info!(" Total sequences: {}", sequences.len());
|
||||
info!(" Train sequences: {}", train_size);
|
||||
info!(" Val sequences: {}", val_size);
|
||||
|
||||
Ok((sequences, train_size))
|
||||
}
|
||||
|
||||
/// Initialize the diffusion model adapter from config.
|
||||
fn init_model(args: &Args, device: &Device) -> Result<(DiffusionTrainableAdapter, DiffusionConfig)> {
|
||||
info!("Step 3/4: Initializing Diffusion model...");
|
||||
let config = build_config(args);
|
||||
|
||||
let mut adapter = DiffusionTrainableAdapter::new(config.clone(), device.clone())
|
||||
.context("Failed to create DiffusionTrainableAdapter")?;
|
||||
|
||||
info!(
|
||||
" Model: {} (data_dim={}, hidden={}, layers={}, timesteps={})",
|
||||
adapter.model_type(),
|
||||
config.data_dim(),
|
||||
config.hidden_dim,
|
||||
config.num_layers,
|
||||
config.num_timesteps,
|
||||
);
|
||||
|
||||
adapter
|
||||
.set_learning_rate(args.learning_rate)
|
||||
.context("Failed to set learning rate")?;
|
||||
|
||||
std::fs::create_dir_all(&args.output_dir)
|
||||
.with_context(|| format!("Failed to create output dir: {}", args.output_dir.display()))?;
|
||||
|
||||
Ok((adapter, config))
|
||||
}
|
||||
|
||||
/// Run the main training loop over all epochs.
|
||||
fn run_training_loop(
|
||||
adapter: &mut DiffusionTrainableAdapter,
|
||||
train_sequences: &[Vec<f32>],
|
||||
val_sequences: &[Vec<f32>],
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
) -> Result<(TrainState, std::time::Duration)> {
|
||||
let training_start = Instant::now();
|
||||
let mut state = TrainState {
|
||||
best_val_loss: f64::MAX,
|
||||
epochs_without_improvement: 0,
|
||||
loss_history: Vec::new(),
|
||||
};
|
||||
|
||||
let steps_per_epoch = if args.max_steps_per_epoch > 0 {
|
||||
args.max_steps_per_epoch
|
||||
} else {
|
||||
train_sequences.len() / args.batch_size.max(1)
|
||||
};
|
||||
|
||||
for epoch in 0..args.epochs {
|
||||
let epoch_start = Instant::now();
|
||||
|
||||
let avg_train_loss = run_train_epoch(
|
||||
adapter, train_sequences, args, device, steps_per_epoch, epoch,
|
||||
)?;
|
||||
state.loss_history.push(avg_train_loss);
|
||||
|
||||
let val_loss = if val_sequences.is_empty() {
|
||||
avg_train_loss
|
||||
} else {
|
||||
run_validation(adapter, val_sequences, args.batch_size, args.seq_len, device)?
|
||||
};
|
||||
|
||||
let epoch_time = epoch_start.elapsed();
|
||||
let metrics = adapter.collect_metrics();
|
||||
|
||||
info!(
|
||||
" Epoch {}/{} -- train_loss={:.6} val_loss={:.6} lr={:.1e} step={} ({:.1}s)",
|
||||
epoch + 1, args.epochs, avg_train_loss, val_loss,
|
||||
metrics.learning_rate, adapter.get_step(), epoch_time.as_secs_f64(),
|
||||
);
|
||||
|
||||
if (epoch + 1) % 5 == 0 {
|
||||
save_checkpoint(adapter, &args.output_dir, &format!("epoch_{}", epoch + 1), "Periodic")?;
|
||||
}
|
||||
|
||||
if val_loss < state.best_val_loss {
|
||||
state.best_val_loss = val_loss;
|
||||
state.epochs_without_improvement = 0;
|
||||
save_checkpoint(adapter, &args.output_dir, "best", "Best")?;
|
||||
} else {
|
||||
state.epochs_without_improvement += 1;
|
||||
if state.epochs_without_improvement >= args.patience {
|
||||
info!(
|
||||
" Early stopping at epoch {} (patience {} exhausted)",
|
||||
epoch + 1, args.patience,
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok((state, training_start.elapsed()))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Main
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
#[allow(clippy::cognitive_complexity)]
|
||||
fn main() -> Result<()> {
|
||||
// Initialize tracing
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
|
||||
)
|
||||
.init();
|
||||
|
||||
let args = Args::parse();
|
||||
|
||||
println!();
|
||||
println!("{}", "=".repeat(80));
|
||||
println!(" Diffusion Model Training on Real OHLCV Data (DDPM/DDIM)");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
println!(" Configuration:");
|
||||
println!(" Symbol: {}", args.symbol);
|
||||
println!(" Data dir: {}", args.data_dir.display());
|
||||
println!(" Epochs: {}", args.epochs);
|
||||
println!(" Batch size: {}", args.batch_size);
|
||||
println!(" Learning rate: {:.1e}", args.learning_rate);
|
||||
println!(" Seq len: {}", args.seq_len);
|
||||
println!(" Hidden dim: {}", args.hidden_dim);
|
||||
println!(" Num layers: {}", args.num_layers);
|
||||
println!(" Num timesteps: {}", args.num_timesteps);
|
||||
println!(" Sampling steps: {}", args.sampling_steps);
|
||||
println!(" Max steps/epoch: {}", args.max_steps_per_epoch);
|
||||
println!(" Patience: {}", args.patience);
|
||||
println!(" Output dir: {}", args.output_dir.display());
|
||||
println!();
|
||||
|
||||
let total_start = Instant::now();
|
||||
let device = select_device();
|
||||
|
||||
let (sequences, train_size) = load_and_prepare_data(&args)?;
|
||||
let train_sequences = sequences.get(..train_size).unwrap_or(&sequences);
|
||||
let val_sequences = sequences.get(train_size..).unwrap_or(&[]);
|
||||
println!();
|
||||
|
||||
let (mut adapter, _config) = init_model(&args, &device)?;
|
||||
|
||||
info!("Step 4/4: Starting training...");
|
||||
println!();
|
||||
println!("{}", "=".repeat(80));
|
||||
println!(" Training Loop");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
|
||||
let (state, training_time) =
|
||||
run_training_loop(&mut adapter, train_sequences, val_sequences, &args, &device)?;
|
||||
|
||||
save_checkpoint(&adapter, &args.output_dir, "final", "Final")?;
|
||||
print_summary(&state, &adapter, training_time, total_start.elapsed(), &args.output_dir);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,324 +0,0 @@
|
||||
//! **DQN Training on ES.FUT Real Market Data**
|
||||
//!
|
||||
//! Production training script for DQN model on ES.FUT futures data.
|
||||
//!
|
||||
//! ## Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Fast training (10 epochs, ~5 seconds)
|
||||
//! cargo run -p ml --example train_dqn_es_fut --release
|
||||
//!
|
||||
//! # Production training (50 epochs, ~20 seconds)
|
||||
//! cargo run -p ml --example train_dqn_es_fut --release -- --epochs 50
|
||||
//!
|
||||
//! # Full training (200 epochs, ~80 seconds)
|
||||
//! cargo run -p ml --example train_dqn_es_fut --release -- --epochs 200
|
||||
//! ```
|
||||
//!
|
||||
//! ## Expected Results
|
||||
//!
|
||||
//! - **10 epochs**: Loss ~0.15, Q-value ~3.0
|
||||
//! - **50 epochs**: Loss ~0.04, Q-value ~0.9 (production checkpoint)
|
||||
//! - **200 epochs**: Loss ~0.01, Q-value ~0.5 (maximum convergence)
|
||||
//!
|
||||
//! ## Output
|
||||
//!
|
||||
//! Checkpoint saved to: `ml/checkpoints/dqn_es_fut_v1.safetensors`
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer};
|
||||
use std::path::PathBuf;
|
||||
use std::time::Instant;
|
||||
use tracing::{info, Level};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(author, version, about, long_about = None)]
|
||||
struct Args {
|
||||
/// Number of training epochs
|
||||
#[arg(short, long, default_value_t = 10)]
|
||||
epochs: usize,
|
||||
|
||||
/// Batch size (max 230 for RTX 3050 Ti 4GB VRAM)
|
||||
#[arg(short, long, default_value_t = 128)]
|
||||
batch_size: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(short, long, default_value_t = 0.0001)]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Data directory
|
||||
#[arg(
|
||||
short,
|
||||
long,
|
||||
default_value = "../test_data/real/databento/ml_training_small"
|
||||
)]
|
||||
data_dir: String,
|
||||
|
||||
/// Output checkpoint path
|
||||
#[arg(short, long, default_value = "checkpoints/dqn_es_fut_v1.safetensors")]
|
||||
output: String,
|
||||
|
||||
/// Enable early stopping
|
||||
#[arg(long, default_value_t = true)]
|
||||
early_stopping: bool,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
let args = Args::parse();
|
||||
|
||||
// Setup logging
|
||||
let log_level = if args.verbose {
|
||||
Level::DEBUG
|
||||
} else {
|
||||
Level::INFO
|
||||
};
|
||||
let subscriber = FmtSubscriber::builder()
|
||||
.with_max_level(log_level)
|
||||
.with_target(false)
|
||||
.with_thread_ids(false)
|
||||
.with_file(false)
|
||||
.with_line_number(false)
|
||||
.finish();
|
||||
tracing::subscriber::set_global_default(subscriber)?;
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("🚀 DQN Training on ES.FUT Real Market Data");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
println!("⚙️ Configuration:");
|
||||
println!(" Epochs: {}", args.epochs);
|
||||
println!(" Batch Size: {}", args.batch_size);
|
||||
println!(" Learning Rate: {}", args.learning_rate);
|
||||
println!(" Data Dir: {}", args.data_dir);
|
||||
println!(" Output: {}", args.output);
|
||||
println!(" Early Stopping: {}", args.early_stopping);
|
||||
println!();
|
||||
|
||||
let start_time = Instant::now();
|
||||
|
||||
// ========================================================================
|
||||
// Step 1: Verify data directory exists
|
||||
// ========================================================================
|
||||
info!("Verifying data directory...");
|
||||
|
||||
let data_path = PathBuf::from(&args.data_dir);
|
||||
if !data_path.exists() {
|
||||
eprintln!("❌ Error: Data directory not found: {}", args.data_dir);
|
||||
eprintln!(" Run data acquisition first or check path.");
|
||||
std::process::exit(1);
|
||||
}
|
||||
|
||||
// Count DBN files
|
||||
let dbn_files: Vec<_> = std::fs::read_dir(&data_path)?
|
||||
.filter_map(|entry| entry.ok())
|
||||
.filter(|entry| entry.path().extension().and_then(|s| s.to_str()) == Some("dbn"))
|
||||
.collect();
|
||||
|
||||
if dbn_files.is_empty() {
|
||||
eprintln!("❌ Error: No DBN files found in: {}", args.data_dir);
|
||||
std::process::exit(1);
|
||||
}
|
||||
|
||||
info!("Found {} DBN files", dbn_files.len());
|
||||
println!(
|
||||
"✅ Data directory validated ({} DBN files)\n",
|
||||
dbn_files.len()
|
||||
);
|
||||
|
||||
// ========================================================================
|
||||
// Step 2: Configure DQN hyperparameters
|
||||
// ========================================================================
|
||||
info!("Configuring DQN hyperparameters...");
|
||||
|
||||
let mut hyperparams = DQNHyperparameters::conservative();
|
||||
hyperparams.epochs = args.epochs;
|
||||
hyperparams.batch_size = args.batch_size;
|
||||
hyperparams.learning_rate = args.learning_rate;
|
||||
hyperparams.gamma = 0.99;
|
||||
hyperparams.epsilon_start = 1.0;
|
||||
hyperparams.epsilon_end = 0.01;
|
||||
hyperparams.epsilon_decay = 0.995;
|
||||
hyperparams.buffer_size = 100_000;
|
||||
hyperparams.checkpoint_frequency = args.epochs / 5; // Save 5 checkpoints
|
||||
hyperparams.early_stopping_enabled = args.early_stopping;
|
||||
hyperparams.q_value_floor = 0.5;
|
||||
hyperparams.min_loss_improvement_pct = 2.0;
|
||||
hyperparams.plateau_window = 30;
|
||||
hyperparams.min_epochs_before_stopping = args.epochs / 2;
|
||||
|
||||
// Validate batch size
|
||||
if hyperparams.batch_size > 230 {
|
||||
eprintln!(
|
||||
"❌ Error: Batch size {} exceeds GPU limit (230)",
|
||||
hyperparams.batch_size
|
||||
);
|
||||
eprintln!(" Reduce batch size to fit in 4GB VRAM.");
|
||||
std::process::exit(1);
|
||||
}
|
||||
|
||||
println!("✅ Hyperparameters configured\n");
|
||||
|
||||
// ========================================================================
|
||||
// Step 3: Create DQN trainer
|
||||
// ========================================================================
|
||||
info!("Initializing DQN trainer...");
|
||||
|
||||
let mut trainer =
|
||||
DQNTrainer::new(hyperparams.clone()).context("Failed to create DQN trainer")?;
|
||||
|
||||
println!("✅ DQN trainer initialized\n");
|
||||
|
||||
// ========================================================================
|
||||
// Step 4: Setup checkpoint directory
|
||||
// ========================================================================
|
||||
info!("Setting up checkpoint directory...");
|
||||
|
||||
let output_path = PathBuf::from(&args.output);
|
||||
let checkpoint_dir = output_path.parent().context("Invalid output path")?;
|
||||
|
||||
std::fs::create_dir_all(checkpoint_dir)?;
|
||||
|
||||
println!(
|
||||
"✅ Checkpoint directory ready: {}\n",
|
||||
checkpoint_dir.display()
|
||||
);
|
||||
|
||||
// ========================================================================
|
||||
// Step 5: Run training
|
||||
// ========================================================================
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("🏋️ Starting DQN Training");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
|
||||
let training_start = Instant::now();
|
||||
let mut checkpoint_count = 0;
|
||||
|
||||
let metrics = trainer
|
||||
.train(&args.data_dir, |epoch, checkpoint_data, is_final| {
|
||||
checkpoint_count += 1;
|
||||
|
||||
let checkpoint_path = if is_final {
|
||||
// Final checkpoint
|
||||
output_path.clone()
|
||||
} else {
|
||||
// Intermediate checkpoint
|
||||
checkpoint_dir.join(format!("dqn_es_fut_epoch_{}.safetensors", epoch))
|
||||
};
|
||||
|
||||
std::fs::write(&checkpoint_path, checkpoint_data)
|
||||
.context("Failed to write checkpoint")?;
|
||||
|
||||
let size_kb = std::fs::metadata(&checkpoint_path)?.len() / 1024;
|
||||
info!(
|
||||
"Checkpoint saved: epoch {} ({} KB) -> {}",
|
||||
epoch,
|
||||
size_kb,
|
||||
checkpoint_path.display()
|
||||
);
|
||||
|
||||
Ok(checkpoint_path.to_string_lossy().to_string())
|
||||
})
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_time = training_start.elapsed();
|
||||
|
||||
println!();
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("✅ Training Complete");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
|
||||
// ========================================================================
|
||||
// Step 6: Report results
|
||||
// ========================================================================
|
||||
println!("📊 Training Metrics:");
|
||||
println!();
|
||||
println!(" Epochs Completed: {}", metrics.epochs_trained);
|
||||
println!(" Final Loss: {:.6}", metrics.loss);
|
||||
println!(" Convergence: {}", metrics.convergence_achieved);
|
||||
println!();
|
||||
|
||||
if let Some(avg_q_value) = metrics.additional_metrics.get("avg_q_value") {
|
||||
println!(" Avg Q-value: {:.4}", avg_q_value);
|
||||
}
|
||||
|
||||
if let Some(avg_grad_norm) = metrics.additional_metrics.get("avg_gradient_norm") {
|
||||
println!(" Avg Gradient Norm: {:.6}", avg_grad_norm);
|
||||
}
|
||||
|
||||
if let Some(final_epsilon) = metrics.additional_metrics.get("final_epsilon") {
|
||||
println!(" Final Epsilon: {:.4}", final_epsilon);
|
||||
}
|
||||
|
||||
println!();
|
||||
println!("⏱️ Performance:");
|
||||
println!();
|
||||
println!(
|
||||
" Training Time: {:.2}s ({:.1} min)",
|
||||
training_time.as_secs_f64(),
|
||||
training_time.as_secs_f64() / 60.0
|
||||
);
|
||||
println!(
|
||||
" Avg Epoch Time: {:.3}s",
|
||||
training_time.as_secs_f64() / metrics.epochs_trained as f64
|
||||
);
|
||||
println!(" Checkpoints Saved: {}", checkpoint_count);
|
||||
println!();
|
||||
|
||||
// ========================================================================
|
||||
// Step 7: Verify final checkpoint
|
||||
// ========================================================================
|
||||
if output_path.exists() {
|
||||
let checkpoint_size = std::fs::metadata(&output_path)?.len();
|
||||
println!("💾 Final Checkpoint:");
|
||||
println!();
|
||||
println!(" Path: {}", output_path.display());
|
||||
println!(
|
||||
" Size: {} KB ({} bytes)",
|
||||
checkpoint_size / 1024,
|
||||
checkpoint_size
|
||||
);
|
||||
println!();
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Summary
|
||||
// ========================================================================
|
||||
let total_time = start_time.elapsed();
|
||||
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("🎉 DQN Training Successful");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
println!("✅ Model trained and saved to: {}", args.output);
|
||||
println!(
|
||||
"⏱️ Total time: {:.2}s ({:.1} min)",
|
||||
total_time.as_secs_f64(),
|
||||
total_time.as_secs_f64() / 60.0
|
||||
);
|
||||
println!();
|
||||
|
||||
// Next steps
|
||||
println!("📌 Next Steps:");
|
||||
println!();
|
||||
println!(" 1. Run inference test:");
|
||||
println!(" cargo test -p ml dqn_training_pipeline_test");
|
||||
println!();
|
||||
println!(" 2. Integrate with paper trading:");
|
||||
println!(" See services/trading_service/src/paper_trading_executor.rs");
|
||||
println!();
|
||||
println!(" 3. Monitor performance:");
|
||||
println!(" Check Grafana dashboard for ML metrics");
|
||||
println!();
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,290 +0,0 @@
|
||||
//! Production DQN Training Script
|
||||
//!
|
||||
//! Trains a DQN model for 50 epochs using production hyperparameters.
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::Tensor;
|
||||
use ml::dqn::{DQNConfig, DQN};
|
||||
use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer};
|
||||
use std::path::PathBuf;
|
||||
use std::time::Instant;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Setup logging
|
||||
tracing_subscriber::fmt()
|
||||
.with_max_level(tracing::Level::INFO)
|
||||
.init();
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("🚀 DQN Production Training - 50 Epochs");
|
||||
println!("{}", "=".repeat(80));
|
||||
|
||||
let start_time = Instant::now();
|
||||
|
||||
// Get data directory
|
||||
let workspace_root = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
||||
.parent()
|
||||
.context("Failed to get workspace root")?
|
||||
.to_path_buf();
|
||||
|
||||
let data_dir = workspace_root.join("test_data/real/databento/ml_training_small");
|
||||
|
||||
if !data_dir.exists() {
|
||||
anyhow::bail!(
|
||||
"Data directory not found: {}. Please check the path.",
|
||||
data_dir.display()
|
||||
);
|
||||
}
|
||||
|
||||
// Create checkpoint directory
|
||||
let checkpoint_dir = PathBuf::from("/tmp");
|
||||
std::fs::create_dir_all(&checkpoint_dir)?;
|
||||
|
||||
println!("\n📋 Configuration:");
|
||||
println!(" Data Directory: {}", data_dir.display());
|
||||
println!(" Checkpoint Directory: {}", checkpoint_dir.display());
|
||||
|
||||
// Configure production hyperparameters (conservative baseline)
|
||||
let mut hyperparams = DQNHyperparameters::conservative();
|
||||
hyperparams.epochs = 50;
|
||||
hyperparams.batch_size = 64;
|
||||
hyperparams.learning_rate = 0.0001;
|
||||
hyperparams.gamma = 0.99;
|
||||
hyperparams.epsilon_start = 0.3;
|
||||
hyperparams.epsilon_end = 0.05;
|
||||
hyperparams.epsilon_decay = 0.995;
|
||||
hyperparams.checkpoint_frequency = 10;
|
||||
hyperparams.early_stopping_enabled = true;
|
||||
hyperparams.min_epochs_before_stopping = 50; // Allow all 50 epochs
|
||||
|
||||
println!("\n⚙️ Hyperparameters:");
|
||||
println!(" Epochs: {}", hyperparams.epochs);
|
||||
println!(" Batch Size: {}", hyperparams.batch_size);
|
||||
println!(" Learning Rate: {}", hyperparams.learning_rate);
|
||||
println!(" Gamma: {}", hyperparams.gamma);
|
||||
println!(
|
||||
" Epsilon: {} → {} (decay: {})",
|
||||
hyperparams.epsilon_start, hyperparams.epsilon_end, hyperparams.epsilon_decay
|
||||
);
|
||||
|
||||
// Create trainer
|
||||
println!("\n🏗️ Initializing DQN trainer...");
|
||||
let mut trainer = DQNTrainer::new(hyperparams.clone())?;
|
||||
|
||||
// Train the model
|
||||
println!("\n🚀 Starting training...\n");
|
||||
|
||||
let mut best_checkpoint_path = PathBuf::new();
|
||||
|
||||
let metrics = trainer
|
||||
.train(
|
||||
&data_dir.to_string_lossy().to_string(),
|
||||
|epoch, checkpoint_data, is_best| {
|
||||
let filename = if is_best {
|
||||
"dqn_prod_best.safetensors".to_string()
|
||||
} else {
|
||||
format!("dqn_prod_epoch_{}.safetensors", epoch)
|
||||
};
|
||||
let path = checkpoint_dir.join(filename);
|
||||
std::fs::write(&path, checkpoint_data)?;
|
||||
if is_best {
|
||||
best_checkpoint_path = path.clone();
|
||||
println!(
|
||||
" 💾 ⭐ BEST checkpoint saved: epoch {} -> {}",
|
||||
epoch,
|
||||
path.display()
|
||||
);
|
||||
} else {
|
||||
println!(
|
||||
" 💾 Checkpoint saved: epoch {} -> {}",
|
||||
epoch,
|
||||
path.display()
|
||||
);
|
||||
}
|
||||
Ok(path.to_string_lossy().to_string())
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
|
||||
let training_time = start_time.elapsed();
|
||||
|
||||
// Report results
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("✅ TRAINING COMPLETE");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("\n📊 Results:");
|
||||
println!(" Epochs Completed: {}", metrics.epochs_trained);
|
||||
println!(" Final Loss: {:.6}", metrics.loss);
|
||||
println!(
|
||||
" Training Time: {:.2}s ({:.1} min)",
|
||||
training_time.as_secs_f64(),
|
||||
training_time.as_secs_f64() / 60.0
|
||||
);
|
||||
println!(" Convergence: {}", metrics.convergence_achieved);
|
||||
|
||||
if let Some(avg_q_value) = metrics.additional_metrics.get("avg_q_value") {
|
||||
println!(" Avg Q-value: {:.4}", avg_q_value);
|
||||
}
|
||||
|
||||
if let Some(final_epsilon) = metrics.additional_metrics.get("final_epsilon") {
|
||||
println!(" Final Epsilon: {:.4}", final_epsilon);
|
||||
}
|
||||
|
||||
println!("\n💾 Best Checkpoint: {}", best_checkpoint_path.display());
|
||||
|
||||
let checkpoint_size = std::fs::metadata(&best_checkpoint_path)?.len();
|
||||
println!(" Size: {} KB", checkpoint_size / 1024);
|
||||
|
||||
// =====================================================================
|
||||
// Inference Demo: load best checkpoint and run on synthetic states
|
||||
// =====================================================================
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("INFERENCE DEMO");
|
||||
println!("{}", "=".repeat(80));
|
||||
|
||||
let inference_result: Result<()> = (|| -> Result<()> {
|
||||
let checkpoint_str = best_checkpoint_path
|
||||
.to_str()
|
||||
.context("Non-UTF8 checkpoint path")?;
|
||||
|
||||
// Load checkpoint tensors to discover architecture params
|
||||
let checkpoint_tensors =
|
||||
candle_core::safetensors::load(checkpoint_str, &candle_core::Device::Cpu)?;
|
||||
|
||||
// Detect noisy nets from key names
|
||||
let uses_noisy = checkpoint_tensors.keys().any(|k| k.starts_with("noisy_"));
|
||||
|
||||
// Discover state_dim from first layer weight tensor
|
||||
let state_dim = if uses_noisy {
|
||||
checkpoint_tensors
|
||||
.iter()
|
||||
.find(|(name, _)| name.contains("noisy_hidden_0") && name.contains("mu_w"))
|
||||
.map(|(_, t)| {
|
||||
let d = t.dims();
|
||||
if d.len() == 2 { d[1] } else { 54 }
|
||||
})
|
||||
.unwrap_or(54)
|
||||
} else {
|
||||
checkpoint_tensors
|
||||
.iter()
|
||||
.find(|(name, _)| name.contains("hidden_0") && name.contains("weight"))
|
||||
.map(|(_, t)| {
|
||||
let d = t.dims();
|
||||
if d.len() == 2 { d[1] } else { 54 }
|
||||
})
|
||||
.unwrap_or(54)
|
||||
};
|
||||
|
||||
// Build matching config
|
||||
let mut config = DQNConfig::conservative();
|
||||
config.state_dim = state_dim;
|
||||
config.num_actions = 45;
|
||||
config.hidden_dims = vec![256, 128, 64];
|
||||
config.use_noisy_nets = uses_noisy;
|
||||
config.noisy_sigma_init = 0.5;
|
||||
config.use_iqn = true;
|
||||
config.use_cql = true;
|
||||
|
||||
let num_actions = config.num_actions;
|
||||
|
||||
println!(
|
||||
" Config: state_dim={}, num_actions={}, noisy={}",
|
||||
state_dim, num_actions, uses_noisy
|
||||
);
|
||||
|
||||
// Load into fresh DQN
|
||||
let mut fresh_dqn = DQN::new(config)?;
|
||||
fresh_dqn.load_from_safetensors(checkpoint_str)?;
|
||||
let device = fresh_dqn.device().clone();
|
||||
|
||||
println!(" Loaded checkpoint into fresh DQN\n");
|
||||
|
||||
// Run inference on 5 synthetic state vectors
|
||||
let num_samples: usize = 5;
|
||||
for i in 0..num_samples {
|
||||
// Deterministic synthetic state in [-0.5, 0.5]
|
||||
let state_vec: Vec<f32> = (0..state_dim)
|
||||
.map(|j| ((i * state_dim + j) as f32 * 0.037).sin() * 0.5)
|
||||
.collect();
|
||||
|
||||
let state_tensor = match Tensor::from_vec(state_vec, (1, state_dim), &device) {
|
||||
Ok(t) => t,
|
||||
Err(e) => {
|
||||
eprintln!(" [WARN] Sample {}: failed to create tensor: {}", i + 1, e);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
match fresh_dqn.forward(&state_tensor) {
|
||||
Ok(q_values) => {
|
||||
let q_vec: Vec<f32> = q_values
|
||||
.to_vec2::<f32>()?
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.collect();
|
||||
|
||||
// Find best action (argmax)
|
||||
let (best_action, max_q) = q_vec
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|(_, a), (_, b)| {
|
||||
a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)
|
||||
})
|
||||
.map(|(idx, &val)| (idx, val))
|
||||
.unwrap_or((0, 0.0));
|
||||
|
||||
let min_q = q_vec
|
||||
.iter()
|
||||
.copied()
|
||||
.min_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
|
||||
.unwrap_or(0.0);
|
||||
|
||||
let q_spread = max_q - min_q;
|
||||
|
||||
println!(
|
||||
" Sample {}/{}: action={:>2} max_Q={:>+10.4} Q-spread={:.4}",
|
||||
i + 1,
|
||||
num_samples,
|
||||
best_action,
|
||||
max_q,
|
||||
q_spread,
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!(" [WARN] Sample {}: inference failed: {}", i + 1, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
println!("\n Inference demo complete.");
|
||||
Ok(())
|
||||
})();
|
||||
|
||||
if let Err(e) = inference_result {
|
||||
eprintln!(
|
||||
"\n [WARN] Inference demo failed (training results above are still valid): {}",
|
||||
e
|
||||
);
|
||||
}
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
|
||||
// Save metrics to JSON
|
||||
let metrics_json = serde_json::json!({
|
||||
"epochs_trained": metrics.epochs_trained,
|
||||
"final_loss": metrics.loss,
|
||||
"training_time_seconds": training_time.as_secs_f64(),
|
||||
"convergence_achieved": metrics.convergence_achieved,
|
||||
"avg_q_value": metrics.additional_metrics.get("avg_q_value"),
|
||||
"final_epsilon": metrics.additional_metrics.get("final_epsilon"),
|
||||
"checkpoint_path": best_checkpoint_path.to_string_lossy().to_string(),
|
||||
"checkpoint_size_kb": checkpoint_size / 1024,
|
||||
});
|
||||
|
||||
let metrics_path = PathBuf::from("/tmp/dqn_production_test_training.json");
|
||||
std::fs::write(&metrics_path, serde_json::to_string_pretty(&metrics_json)?)?;
|
||||
println!("📄 Training metrics saved to: {}", metrics_path.display());
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,615 +0,0 @@
|
||||
//! KAN (Kolmogorov-Arnold Network) training on real OHLCV data from DBN files.
|
||||
//!
|
||||
//! Trains a KAN model using the `UnifiedTrainable` adapter on Databento OHLCV bars
|
||||
//! with walk-forward evaluation windows, z-score normalization, and checkpoint
|
||||
//! saving.
|
||||
//!
|
||||
//! KAN replaces fixed activation functions with learnable B-spline
|
||||
//! activations on each edge, enabling the network to discover arbitrary
|
||||
//! non-linear relationships in price data.
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Quick test (5 epochs on ES data)
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_kan_dbn --release -- \
|
||||
//! --data-dir data/cache/futures-baseline --symbol ES.FUT --epochs 5
|
||||
//!
|
||||
//! # Production training (50 epochs, custom grid)
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_kan_dbn --release -- \
|
||||
//! --data-dir data/cache/futures-baseline --symbol ES.FUT --epochs 50 \
|
||||
//! --grid-size 8 --spline-order 4 --learning-rate 0.0005
|
||||
//! ```
|
||||
//!
|
||||
//! # Output
|
||||
//!
|
||||
//! Checkpoints saved to `<output-dir>/kan_fold<N>_best.safetensors` (one per fold).
|
||||
|
||||
#![allow(unused_crate_dependencies, clippy::cognitive_complexity)]
|
||||
#![deny(
|
||||
clippy::unwrap_used,
|
||||
clippy::expect_used,
|
||||
clippy::panic,
|
||||
clippy::indexing_slicing
|
||||
)]
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Instant;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::{Device, Tensor};
|
||||
use clap::Parser;
|
||||
use tracing::{info, warn, Level};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::features::extraction::extract_ml_features;
|
||||
use ml::kan::{KANConfig, KANTrainableAdapter};
|
||||
use ml::training::unified_trainer::UnifiedTrainable;
|
||||
use ml::types::OHLCVBar;
|
||||
use ml::walk_forward::{generate_walk_forward_windows, NormStats, WalkForwardConfig};
|
||||
|
||||
#[allow(unreachable_pub)]
|
||||
mod baseline_common;
|
||||
use baseline_common::{load_all_bars, spread_cost_bps};
|
||||
|
||||
/// Type alias for a batch of (input, target) tensor pairs.
|
||||
type TensorPairs = Vec<(Tensor, Tensor)>;
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CLI Arguments
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Train a KAN (Kolmogorov-Arnold Network) on real OHLCV data from DBN files.
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "train_kan_dbn", about = "Train KAN with walk-forward windows on DBN data")]
|
||||
struct Args {
|
||||
/// Path to directory containing .dbn.zst files (with symbol subdirectories)
|
||||
#[arg(long, default_value = "data/cache/futures-baseline")]
|
||||
data_dir: PathBuf,
|
||||
|
||||
/// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT")
|
||||
#[arg(long, default_value = "ES.FUT")]
|
||||
symbol: String,
|
||||
|
||||
/// Maximum training epochs per fold
|
||||
#[arg(long, default_value_t = 20)]
|
||||
epochs: usize,
|
||||
|
||||
/// Training batch size
|
||||
#[arg(long, default_value_t = 128)]
|
||||
batch_size: usize,
|
||||
|
||||
/// Learning rate for `AdamW` optimizer
|
||||
#[arg(long, default_value_t = 1e-3)]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Max environment steps per epoch (0 = use all bars)
|
||||
#[arg(long, default_value_t = 2000)]
|
||||
max_steps_per_epoch: usize,
|
||||
|
||||
/// Output directory for trained model checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: PathBuf,
|
||||
|
||||
/// B-spline grid size (more points = finer approximation)
|
||||
#[arg(long, default_value_t = 5)]
|
||||
grid_size: usize,
|
||||
|
||||
/// B-spline order (4 = cubic splines)
|
||||
#[arg(long, default_value_t = 4)]
|
||||
spline_order: usize,
|
||||
|
||||
/// Feature dimension (must match feature extraction output)
|
||||
#[arg(long, default_value_t = 51)]
|
||||
feature_dim: 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,
|
||||
|
||||
/// L2 weight decay for regularization
|
||||
#[arg(long, default_value_t = 1e-4)]
|
||||
weight_decay: f64,
|
||||
|
||||
/// Maximum gradient norm for clipping
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
grad_clip: f64,
|
||||
|
||||
/// Early stopping patience (epochs without improvement)
|
||||
#[arg(long, default_value_t = 10)]
|
||||
patience: usize,
|
||||
|
||||
/// 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
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
tx_cost_bps: f64,
|
||||
|
||||
/// Instrument tick size in price units (ES=0.25)
|
||||
#[arg(long, default_value_t = 0.25)]
|
||||
tick_size: f64,
|
||||
|
||||
/// Typical bid-ask spread in ticks
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
spread_ticks: f64,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Feature preparation
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Build (input, target) tensor pairs from OHLCV bars for regression training.
|
||||
///
|
||||
/// Features are extracted via `extract_ml_features` (51-dim), z-score normalized,
|
||||
/// and the target is the next-bar return in basis points, clipped to `[-10, 10]`.
|
||||
///
|
||||
/// Returns `(train_pairs, val_pairs)` ready for the `UnifiedTrainable` API.
|
||||
fn prepare_fold_data(
|
||||
train_bars: &[OHLCVBar],
|
||||
val_bars: &[OHLCVBar],
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
) -> Result<(TensorPairs, TensorPairs)> {
|
||||
// Extract features
|
||||
let train_features = extract_ml_features(train_bars)
|
||||
.context("Failed to extract training features")?;
|
||||
let val_features = extract_ml_features(val_bars)
|
||||
.context("Failed to extract validation features")?;
|
||||
|
||||
if train_features.len() < 2 {
|
||||
anyhow::bail!("Insufficient training features: {}", train_features.len());
|
||||
}
|
||||
if val_features.len() < 2 {
|
||||
anyhow::bail!("Insufficient validation features: {}", val_features.len());
|
||||
}
|
||||
|
||||
// Compute normalization stats from training data only
|
||||
let norm_stats = NormStats::from_features(&train_features);
|
||||
let norm_train = norm_stats.normalize_batch(&train_features);
|
||||
let norm_val = norm_stats.normalize_batch(&val_features);
|
||||
|
||||
// The feature extractor skips a warmup period, so the number of features
|
||||
// is less than the number of bars. Align bars to features by offsetting
|
||||
// from the end.
|
||||
let train_bar_offset = train_bars.len().saturating_sub(train_features.len());
|
||||
let val_bar_offset = val_bars.len().saturating_sub(val_features.len());
|
||||
|
||||
// Build (input, target) pairs for training
|
||||
let train_pairs = build_tensor_pairs(
|
||||
&norm_train, train_bars, train_bar_offset, args, device,
|
||||
)?;
|
||||
let val_pairs = build_tensor_pairs(
|
||||
&norm_val, val_bars, val_bar_offset, args, device,
|
||||
)?;
|
||||
|
||||
Ok((train_pairs, val_pairs))
|
||||
}
|
||||
|
||||
/// Convert normalized feature vectors and bars into tensor pairs.
|
||||
///
|
||||
/// Each pair: input = feature vector at time t, target = clipped next-bar return (bps).
|
||||
fn build_tensor_pairs(
|
||||
norm_features: &[[f64; 51]],
|
||||
bars: &[OHLCVBar],
|
||||
bar_offset: usize,
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
) -> Result<TensorPairs> {
|
||||
let mut pairs = Vec::new();
|
||||
let n = norm_features.len();
|
||||
|
||||
// We need at least 2 features to form (input_t, target from bar_t+1)
|
||||
let limit = n.saturating_sub(1);
|
||||
let step_limit = if args.max_steps_per_epoch > 0 {
|
||||
args.max_steps_per_epoch.min(limit)
|
||||
} else {
|
||||
limit
|
||||
};
|
||||
|
||||
for i in 0..step_limit {
|
||||
let Some(feat) = norm_features.get(i) else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let bar_idx = i + bar_offset;
|
||||
let close_cur = bars.get(bar_idx).map(|b| b.close).unwrap_or(0.0);
|
||||
let close_next = bars.get(bar_idx + 1).map(|b| b.close).unwrap_or(close_cur);
|
||||
|
||||
// Target: next-bar return in basis points, clipped
|
||||
let return_bps = if close_cur.abs() > 1e-10 {
|
||||
(close_next - close_cur) / close_cur * 10_000.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
// Clamp large moves (contract rolls)
|
||||
let max_bps = args.max_bar_return * 10_000.0;
|
||||
let clipped = return_bps.clamp(-max_bps, max_bps);
|
||||
|
||||
// Subtract transaction cost for a round-trip trade
|
||||
let spread = spread_cost_bps(close_cur, args.tick_size, args.spread_ticks);
|
||||
let net_return = clipped - (args.tx_cost_bps + spread);
|
||||
|
||||
let input_f32: Vec<f32> = feat.iter().map(|&v| v as f32).collect();
|
||||
let target_f32 = vec![net_return as f32];
|
||||
|
||||
let input_tensor = Tensor::from_vec(input_f32, &[1, args.feature_dim], device)
|
||||
.context("Failed to create input tensor")?;
|
||||
let target_tensor = Tensor::from_vec(target_f32, &[1, 1], device)
|
||||
.context("Failed to create target tensor")?;
|
||||
|
||||
pairs.push((input_tensor, target_tensor));
|
||||
}
|
||||
|
||||
Ok(pairs)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Training loop
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Run one training epoch (all mini-batches) and return average loss.
|
||||
fn run_training_epoch(
|
||||
adapter: &mut KANTrainableAdapter,
|
||||
train_pairs: &[(Tensor, Tensor)],
|
||||
batch_size: usize,
|
||||
) -> Result<f64> {
|
||||
let mut epoch_loss_sum = 0.0_f64;
|
||||
let mut epoch_steps = 0_usize;
|
||||
let n_train = train_pairs.len();
|
||||
let mut batch_start = 0_usize;
|
||||
|
||||
while batch_start < n_train {
|
||||
let batch_end = (batch_start + batch_size).min(n_train);
|
||||
let Some(batch_slice) = train_pairs.get(batch_start..batch_end) else {
|
||||
break;
|
||||
};
|
||||
|
||||
let batch_inputs: Vec<&Tensor> = batch_slice.iter().map(|(inp, _)| inp).collect();
|
||||
let batch_targets: Vec<&Tensor> = batch_slice.iter().map(|(_, tgt)| tgt).collect();
|
||||
|
||||
if batch_inputs.is_empty() {
|
||||
batch_start = batch_end;
|
||||
continue;
|
||||
}
|
||||
|
||||
let input_cat = Tensor::cat(&batch_inputs, 0)
|
||||
.context("Failed to concatenate batch inputs")?;
|
||||
let target_cat = Tensor::cat(&batch_targets, 0)
|
||||
.context("Failed to concatenate batch targets")?;
|
||||
|
||||
adapter.zero_grad()?;
|
||||
let predictions = adapter.forward(&input_cat)?;
|
||||
let loss = adapter.compute_loss(&predictions, &target_cat)?;
|
||||
let loss_val = loss
|
||||
.to_scalar::<f32>()
|
||||
.map_err(|e| anyhow::anyhow!("Failed to extract loss scalar: {e}"))?;
|
||||
|
||||
adapter.backward(&loss)?;
|
||||
adapter.optimizer_step()?;
|
||||
|
||||
epoch_loss_sum += loss_val as f64;
|
||||
epoch_steps += 1;
|
||||
batch_start = batch_end;
|
||||
}
|
||||
|
||||
Ok(if epoch_steps > 0 {
|
||||
epoch_loss_sum / epoch_steps as f64
|
||||
} else {
|
||||
0.0
|
||||
})
|
||||
}
|
||||
|
||||
/// Save a checkpoint at the given path, returning the path string used.
|
||||
fn save_checkpoint(adapter: &KANTrainableAdapter, path: &Path) -> Result<()> {
|
||||
let path_str = path.to_str().unwrap_or("kan_checkpoint");
|
||||
adapter.save_checkpoint(path_str)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Train a KAN model on a single walk-forward fold.
|
||||
///
|
||||
/// Returns the best validation loss achieved.
|
||||
fn train_kan_fold(
|
||||
fold: usize,
|
||||
train_pairs: &[(Tensor, Tensor)],
|
||||
val_pairs: &[(Tensor, Tensor)],
|
||||
args: &Args,
|
||||
device: &Device,
|
||||
output_dir: &Path,
|
||||
) -> Result<f64> {
|
||||
info!(
|
||||
"[KAN] Fold {} -- {} train pairs, {} val pairs",
|
||||
fold,
|
||||
train_pairs.len(),
|
||||
val_pairs.len(),
|
||||
);
|
||||
|
||||
let config = KANConfig {
|
||||
grid_size: args.grid_size,
|
||||
spline_order: args.spline_order,
|
||||
layer_widths: vec![args.feature_dim, 32, 16, 1],
|
||||
learning_rate: args.learning_rate,
|
||||
weight_decay: args.weight_decay,
|
||||
grad_clip: args.grad_clip,
|
||||
};
|
||||
|
||||
let mut adapter = KANTrainableAdapter::new(config, device)
|
||||
.context("Failed to create KANTrainableAdapter")?;
|
||||
|
||||
let mut best_val_loss = f64::MAX;
|
||||
let mut epochs_without_improvement = 0_usize;
|
||||
|
||||
for epoch in 0..args.epochs {
|
||||
let epoch_start = Instant::now();
|
||||
|
||||
let avg_train_loss = run_training_epoch(&mut adapter, train_pairs, args.batch_size)?;
|
||||
let val_loss = adapter.validate(val_pairs)?;
|
||||
let elapsed = epoch_start.elapsed();
|
||||
|
||||
info!(
|
||||
" Fold {} Epoch {}/{}: train_loss={:.6}, val_loss={:.6}, lr={:.2e}, step={}, time={:.1}s",
|
||||
fold,
|
||||
epoch + 1,
|
||||
args.epochs,
|
||||
avg_train_loss,
|
||||
val_loss,
|
||||
adapter.get_learning_rate(),
|
||||
adapter.get_step(),
|
||||
elapsed.as_secs_f64(),
|
||||
);
|
||||
|
||||
// Checkpoint on improvement
|
||||
if val_loss < best_val_loss {
|
||||
best_val_loss = val_loss;
|
||||
epochs_without_improvement = 0;
|
||||
let ckpt_path = output_dir.join(format!("kan_fold{fold}_best"));
|
||||
save_checkpoint(&adapter, &ckpt_path)?;
|
||||
info!(" [KAN] New best val_loss={:.6}, checkpoint saved", val_loss);
|
||||
} else {
|
||||
epochs_without_improvement += 1;
|
||||
}
|
||||
|
||||
// Periodic checkpoint every 5 epochs
|
||||
if (epoch + 1) % 5 == 0 {
|
||||
let periodic_path = output_dir.join(format!("kan_fold{fold}_epoch{}", epoch + 1));
|
||||
save_checkpoint(&adapter, &periodic_path)?;
|
||||
info!(" [KAN] Periodic checkpoint at epoch {}", epoch + 1);
|
||||
}
|
||||
|
||||
// Early stopping
|
||||
if epochs_without_improvement >= args.patience {
|
||||
info!(
|
||||
" [KAN] Early stopping at epoch {} (no improvement for {} epochs)",
|
||||
epoch + 1,
|
||||
args.patience,
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Final checkpoint
|
||||
let final_path = output_dir.join(format!("kan_fold{fold}_final"));
|
||||
save_checkpoint(&adapter, &final_path)?;
|
||||
|
||||
let metrics = adapter.collect_metrics();
|
||||
info!(
|
||||
" [KAN] Fold {} complete: best_val_loss={:.6}, total_steps={}, grid_size={}, spline_order={}",
|
||||
fold,
|
||||
best_val_loss,
|
||||
metrics.custom_metrics.get("training_steps").copied().unwrap_or(0.0),
|
||||
metrics.custom_metrics.get("grid_size").copied().unwrap_or(0.0),
|
||||
metrics.custom_metrics.get("spline_order").copied().unwrap_or(0.0),
|
||||
);
|
||||
|
||||
Ok(best_val_loss)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Main helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Select compute device (CUDA if available, otherwise CPU).
|
||||
fn select_device() -> Device {
|
||||
match Device::cuda_if_available(0) {
|
||||
Ok(dev) => {
|
||||
info!("Using CUDA device");
|
||||
dev
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("CUDA not available ({}), falling back to CPU", e);
|
||||
Device::Cpu
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Print configuration summary to stdout.
|
||||
fn print_config(args: &Args) {
|
||||
println!();
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("KAN Training on DBN OHLCV Data");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
println!("Configuration:");
|
||||
println!(" Symbol: {}", args.symbol);
|
||||
println!(" Data dir: {}", args.data_dir.display());
|
||||
println!(" Epochs: {}", args.epochs);
|
||||
println!(" Batch size: {}", args.batch_size);
|
||||
println!(" Learning rate: {}", args.learning_rate);
|
||||
println!(" Grid size: {}", args.grid_size);
|
||||
println!(" Spline order: {}", args.spline_order);
|
||||
println!(" Feature dim: {}", args.feature_dim);
|
||||
println!(" Weight decay: {}", args.weight_decay);
|
||||
println!(" Grad clip: {}", args.grad_clip);
|
||||
println!(" Max steps/ep: {}", args.max_steps_per_epoch);
|
||||
println!(" Output dir: {}", args.output_dir.display());
|
||||
println!(" Patience: {}", args.patience);
|
||||
println!();
|
||||
}
|
||||
|
||||
/// Print final summary of all fold results.
|
||||
fn print_summary(fold_results: &[(usize, f64)], total_elapsed: std::time::Duration, output_dir: &Path) {
|
||||
println!();
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("KAN Training Complete");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!();
|
||||
|
||||
if fold_results.is_empty() {
|
||||
println!("WARNING: No folds completed successfully.");
|
||||
} else {
|
||||
println!("Results by fold:");
|
||||
let mut total_val_loss = 0.0_f64;
|
||||
for (fold, val_loss) in fold_results {
|
||||
println!(" Fold {}: val_loss = {:.6}", fold, val_loss);
|
||||
total_val_loss += val_loss;
|
||||
}
|
||||
let avg_val = total_val_loss / fold_results.len() as f64;
|
||||
println!();
|
||||
println!("Average validation loss: {:.6}", avg_val);
|
||||
println!("Number of folds completed: {}", fold_results.len());
|
||||
}
|
||||
|
||||
println!();
|
||||
println!(
|
||||
"Total time: {:.1}s ({:.1} min)",
|
||||
total_elapsed.as_secs_f64(),
|
||||
total_elapsed.as_secs_f64() / 60.0,
|
||||
);
|
||||
println!("Checkpoints saved to: {}", output_dir.display());
|
||||
println!();
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Main
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
fn main() -> Result<()> {
|
||||
let args = Args::parse();
|
||||
|
||||
// Setup logging
|
||||
let log_level = if args.verbose { Level::DEBUG } else { Level::INFO };
|
||||
let subscriber = FmtSubscriber::builder()
|
||||
.with_max_level(log_level)
|
||||
.with_target(false)
|
||||
.with_thread_ids(false)
|
||||
.with_file(false)
|
||||
.with_line_number(false)
|
||||
.finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
print_config(&args);
|
||||
|
||||
let total_start = Instant::now();
|
||||
let device = select_device();
|
||||
|
||||
// Load OHLCV bars
|
||||
info!("Loading OHLCV bars for {} from {}", args.symbol, args.data_dir.display());
|
||||
let all_bars = load_all_bars(&args.data_dir, &args.symbol)?;
|
||||
println!("Loaded {} bars for {}", all_bars.len(), args.symbol);
|
||||
|
||||
if all_bars.len() < 100 {
|
||||
anyhow::bail!(
|
||||
"Insufficient data: {} bars loaded (need at least 100)",
|
||||
all_bars.len()
|
||||
);
|
||||
}
|
||||
|
||||
// Generate walk-forward windows
|
||||
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 windows = generate_walk_forward_windows(&all_bars, &wf_config);
|
||||
|
||||
if windows.is_empty() {
|
||||
anyhow::bail!(
|
||||
"No walk-forward windows generated. Data may be too short for the configured window sizes."
|
||||
);
|
||||
}
|
||||
println!("Generated {} walk-forward folds", windows.len());
|
||||
println!();
|
||||
|
||||
// Create output directory
|
||||
std::fs::create_dir_all(&args.output_dir)
|
||||
.with_context(|| format!("Failed to create output directory: {}", args.output_dir.display()))?;
|
||||
|
||||
// Train on each fold
|
||||
let mut fold_results: Vec<(usize, f64)> = Vec::new();
|
||||
|
||||
for window in &windows {
|
||||
let fold = window.fold;
|
||||
println!("{}", "-".repeat(60));
|
||||
println!(
|
||||
"Fold {}: train={} bars, val={} bars, test={} bars",
|
||||
fold,
|
||||
window.train.len(),
|
||||
window.val.len(),
|
||||
window.test.len(),
|
||||
);
|
||||
println!(
|
||||
" Train end: {}, Val end: {}, Test end: {}",
|
||||
window.train_end, window.val_end, window.test_end,
|
||||
);
|
||||
|
||||
let fold_start = Instant::now();
|
||||
|
||||
let (train_pairs, val_pairs) = match prepare_fold_data(
|
||||
&window.train,
|
||||
&window.val,
|
||||
&args,
|
||||
&device,
|
||||
) {
|
||||
Ok(data) => data,
|
||||
Err(e) => {
|
||||
warn!("Skipping fold {} -- {}", fold, e);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if train_pairs.is_empty() || val_pairs.is_empty() {
|
||||
warn!("Skipping fold {} -- empty train or val data after feature extraction", fold);
|
||||
continue;
|
||||
}
|
||||
|
||||
let best_val_loss = train_kan_fold(
|
||||
fold,
|
||||
&train_pairs,
|
||||
&val_pairs,
|
||||
&args,
|
||||
&device,
|
||||
&args.output_dir,
|
||||
)?;
|
||||
|
||||
let fold_elapsed = fold_start.elapsed();
|
||||
println!(
|
||||
"Fold {} complete: best_val_loss={:.6}, time={:.1}s",
|
||||
fold, best_val_loss, fold_elapsed.as_secs_f64(),
|
||||
);
|
||||
fold_results.push((fold, best_val_loss));
|
||||
}
|
||||
|
||||
print_summary(&fold_results, total_start.elapsed(), &args.output_dir);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,258 +0,0 @@
|
||||
// Pilot Training Example for Liquid Neural Network with DBN Real Market Data
|
||||
//
|
||||
// This example demonstrates training a Liquid Time-Constant Neural Network (LTC)
|
||||
// on real market data from Databento (DBN format) for HFT price prediction.
|
||||
//
|
||||
// Architecture:
|
||||
// - Input: 16 features (5 OHLCV + 10 technical indicators + 1 volume)
|
||||
// - Hidden: 128 LTC neurons with adaptive time constants
|
||||
// - Output: 3 classes (buy=0, hold=1, sell=2)
|
||||
//
|
||||
// Expected Performance:
|
||||
// - Accuracy: 55-65% (better than random 33.3%)
|
||||
// - Convergence: 20-30 epochs with early stopping
|
||||
// - Training time: ~5 minutes (CPU) or ~30 seconds (GPU)
|
||||
// - Inference latency: <100μs (fixed-point arithmetic)
|
||||
|
||||
use anyhow::Result;
|
||||
use ml::data_loaders::dbn_sequence_loader::DbnSequenceLoader;
|
||||
use ml::liquid::{
|
||||
ActivationType, FixedPoint, LTCConfig, LayerConfig, LiquidNetwork, LiquidNetworkConfig,
|
||||
LiquidTrainer, LiquidTrainingConfig, NetworkType, OutputLayerConfig, SolverType,
|
||||
TrainingSample, TrainingUtils, PRECISION,
|
||||
};
|
||||
use std::time::Instant;
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
println!("========================================");
|
||||
println!("Liquid Neural Network Pilot Training");
|
||||
println!("========================================");
|
||||
println!();
|
||||
println!("Architecture:");
|
||||
println!(" Input: 16 features (normalized OHLCV sequence)");
|
||||
println!(" Hidden: 128 LTC neurons (τ=0.01-1.0)");
|
||||
println!(" Output: 3 classes (buy/hold/sell)");
|
||||
println!(" Solver: RK4 (4th order accuracy)");
|
||||
println!();
|
||||
|
||||
// Step 1: Load DBN market data
|
||||
println!("[1/6] Loading DBN market data (6E.FUT)...");
|
||||
|
||||
// Create DbnSequenceLoader with sequence length 60 and feature dimension 16
|
||||
let mut loader = DbnSequenceLoader::new(60, 16).await?;
|
||||
|
||||
// Load OHLCV sequences from the production data directory
|
||||
let data_dir = "test_data/real/databento/ml_training";
|
||||
let (train_sequences, _val_sequences) = loader.load_sequences(data_dir, 0.8).await?;
|
||||
|
||||
println!(" ✓ Loaded {} training sequences", train_sequences.len());
|
||||
|
||||
// Step 2: Convert sequences to training samples
|
||||
println!();
|
||||
println!("[2/6] Converting sequences to training samples...");
|
||||
|
||||
let mut training_samples = Vec::new();
|
||||
|
||||
for (input_tensor, _target_tensor) in train_sequences.iter() {
|
||||
// Extract the last timestep from the input sequence for feature extraction
|
||||
// Input shape: [seq_len, d_model] = [60, 16]
|
||||
// Target shape: [d_model] = [16] (next timestep prediction)
|
||||
let seq_data = input_tensor.to_vec2::<f64>()?;
|
||||
|
||||
// Use the last timestep as features (16 features)
|
||||
if let Some(last_step) = seq_data.last() {
|
||||
// Convert input to FixedPoint
|
||||
let features: Vec<FixedPoint> =
|
||||
last_step.iter().map(|&f| FixedPoint::from_f64(f)).collect();
|
||||
|
||||
// For this pilot, we'll create synthetic labels based on the trend in the sequence
|
||||
// In production, you'd use actual price change labels from target_data
|
||||
let label = if seq_data.len() >= 2 {
|
||||
// Compare last few prices to determine trend
|
||||
let recent_prices: Vec<f64> =
|
||||
seq_data.iter().rev().take(5).map(|step| step[3]).collect(); // Close price at index 3
|
||||
let first = recent_prices.last().unwrap_or(&0.0);
|
||||
let last = recent_prices.first().unwrap_or(&0.0);
|
||||
let price_change = (last - first) / first.abs().max(1e-6);
|
||||
|
||||
// Thresholds for buy/hold/sell (0.1% = 10 basis points)
|
||||
if price_change > 0.001 {
|
||||
0 // Buy signal
|
||||
} else if price_change < -0.001 {
|
||||
2 // Sell signal
|
||||
} else {
|
||||
1 // Hold signal
|
||||
}
|
||||
} else {
|
||||
1 // Hold for insufficient data
|
||||
};
|
||||
|
||||
// One-hot encode label [buy, hold, sell]
|
||||
let target = match label {
|
||||
0 => vec![FixedPoint::one(), FixedPoint::zero(), FixedPoint::zero()], // Buy
|
||||
1 => vec![FixedPoint::zero(), FixedPoint::one(), FixedPoint::zero()], // Hold
|
||||
2 => vec![FixedPoint::zero(), FixedPoint::zero(), FixedPoint::one()], // Sell
|
||||
_ => vec![FixedPoint::zero(), FixedPoint::one(), FixedPoint::zero()], // Default: Hold
|
||||
};
|
||||
|
||||
training_samples.push(TrainingSample {
|
||||
input: features,
|
||||
target,
|
||||
timestamp: None,
|
||||
market_regime: None,
|
||||
volatility: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
println!(
|
||||
" ✓ Created {} training samples from sequences",
|
||||
training_samples.len()
|
||||
);
|
||||
|
||||
// Step 3: Normalize features (Z-score normalization)
|
||||
println!();
|
||||
println!("[3/6] Normalizing features...");
|
||||
let (means, _stds) = TrainingUtils::normalize_features(&mut training_samples)?;
|
||||
println!(" ✓ Normalized {} features (mean=0, std=1)", means.len());
|
||||
|
||||
// Step 4: Split into training and validation sets
|
||||
println!();
|
||||
println!("[4/6] Splitting data (80% train, 20% validation)...");
|
||||
let (train_samples, val_samples) = TrainingUtils::train_validation_split(
|
||||
training_samples,
|
||||
0.2, // 20% validation
|
||||
);
|
||||
println!(" ✓ Training samples: {}", train_samples.len());
|
||||
println!(" ✓ Validation samples: {}", val_samples.len());
|
||||
|
||||
// Create batches
|
||||
let batch_size = 32;
|
||||
let train_batches = TrainingUtils::create_batches(train_samples, batch_size);
|
||||
let val_batches = TrainingUtils::create_batches(val_samples, batch_size);
|
||||
println!(" ✓ Training batches: {}", train_batches.len());
|
||||
println!(" ✓ Validation batches: {}", val_batches.len());
|
||||
|
||||
// Step 5: Create Liquid Neural Network
|
||||
println!();
|
||||
println!("[5/6] Creating Liquid Neural Network...");
|
||||
|
||||
// Create LTC layer configuration
|
||||
let ltc_config = LTCConfig {
|
||||
input_size: 16, // 5 OHLCV + 10 technical indicators + 1 volume
|
||||
hidden_size: 128,
|
||||
tau_min: FixedPoint(PRECISION / 100), // 0.01
|
||||
tau_max: FixedPoint(PRECISION), // 1.0
|
||||
use_bias: true,
|
||||
solver_type: SolverType::RK4, // 4th order accuracy
|
||||
activation: ActivationType::Tanh,
|
||||
};
|
||||
|
||||
let network_config = LiquidNetworkConfig {
|
||||
network_type: NetworkType::LTC,
|
||||
input_size: 16, // 5 OHLCV + 10 technical indicators + 1 volume
|
||||
output_size: 3, // buy/hold/sell
|
||||
layer_configs: vec![LayerConfig::LTC(ltc_config)],
|
||||
output_layer: OutputLayerConfig {
|
||||
use_linear_output: false,
|
||||
output_activation: Some(ActivationType::Sigmoid),
|
||||
dropout_rate: None,
|
||||
},
|
||||
default_dt: FixedPoint(PRECISION / 100), // 0.01 time step
|
||||
market_regime_adaptation: true,
|
||||
};
|
||||
|
||||
let mut network = LiquidNetwork::new(network_config)?;
|
||||
println!(
|
||||
" ✓ Network created with {} parameters",
|
||||
network.parameter_count()
|
||||
);
|
||||
println!(
|
||||
" ✓ Memory footprint: ~{} KB",
|
||||
(network.parameter_count() * 8) / 1024
|
||||
);
|
||||
|
||||
// Step 6: Train the network
|
||||
println!();
|
||||
println!("[6/6] Training Liquid Neural Network (50 epochs)...");
|
||||
println!();
|
||||
|
||||
let training_config = LiquidTrainingConfig {
|
||||
learning_rate: FixedPoint(PRECISION / 1000), // 0.001
|
||||
batch_size,
|
||||
max_epochs: 50, // Pilot training
|
||||
early_stopping_patience: 10,
|
||||
gradient_clip_threshold: FixedPoint(PRECISION), // 1.0
|
||||
l2_regularization: FixedPoint(PRECISION / 10000), // 0.0001
|
||||
adaptive_learning_rate: true,
|
||||
market_regime_adaptation: false, // No regime data in pilot
|
||||
validation_split: 0.2,
|
||||
};
|
||||
|
||||
let mut trainer = LiquidTrainer::new(training_config);
|
||||
|
||||
let training_start = Instant::now();
|
||||
trainer.train(&mut network, &train_batches, Some(&val_batches))?;
|
||||
let training_duration = training_start.elapsed();
|
||||
|
||||
// Training results
|
||||
println!();
|
||||
println!("========================================");
|
||||
println!("Training Complete!");
|
||||
println!("========================================");
|
||||
println!();
|
||||
println!("Training Metrics:");
|
||||
println!(" Total time: {:.2}s", training_duration.as_secs_f64());
|
||||
println!(" Epochs trained: {}", trainer.training_history.len());
|
||||
|
||||
if let Some(final_metrics) = trainer.training_history.last() {
|
||||
println!(" Final loss: {:.6}", final_metrics.training_loss);
|
||||
if let Some(val_loss) = final_metrics.validation_loss {
|
||||
println!(" Val loss: {:.6}", val_loss);
|
||||
}
|
||||
println!(" Learning rate: {:.6}", final_metrics.learning_rate);
|
||||
println!(" Gradient norm: {:.4}", final_metrics.gradient_norm);
|
||||
println!(" Samples/sec: {:.1}", final_metrics.samples_per_second);
|
||||
}
|
||||
|
||||
// Test inference latency
|
||||
println!();
|
||||
println!("Inference Performance:");
|
||||
let test_input: Vec<FixedPoint> = (0..16)
|
||||
.map(|i| FixedPoint::from_f64((i as f64) / 16.0))
|
||||
.collect();
|
||||
|
||||
let inference_start = Instant::now();
|
||||
for _ in 0..1000 {
|
||||
let _ = network.forward(&test_input)?;
|
||||
}
|
||||
let avg_inference_time = inference_start.elapsed().as_micros() / 1000;
|
||||
|
||||
println!(" Average latency: {}μs (1000 runs)", avg_inference_time);
|
||||
println!(" Target latency: <100μs");
|
||||
|
||||
if avg_inference_time < 100 {
|
||||
println!(" ✓ Latency target MET");
|
||||
} else {
|
||||
println!(" ⚠ Latency target EXCEEDED (consider GPU optimization)");
|
||||
}
|
||||
|
||||
// Save network state (optional - future work)
|
||||
println!();
|
||||
println!("Checkpoint Status:");
|
||||
println!(" ⚠ Checkpoint saving not implemented (future: MinIO/S3)");
|
||||
println!(" ✓ Network state can be serialized via serde");
|
||||
|
||||
println!();
|
||||
println!("========================================");
|
||||
println!("Next Steps:");
|
||||
println!("========================================");
|
||||
println!("1. Run full training (100 epochs, 90 days data)");
|
||||
println!("2. Integrate with ML Training Service (gRPC)");
|
||||
println!("3. Add checkpoint saving (MinIO)");
|
||||
println!("4. GPU acceleration (if latency >100μs)");
|
||||
println!("5. Hyperparameter tuning (Optuna)");
|
||||
println!();
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,946 +0,0 @@
|
||||
//! MAMBA-2 Production Training with Real DBN Market Data
|
||||
//!
|
||||
//! **Complete end-to-end MAMBA-2 training pipeline with real DataBento market data**
|
||||
//!
|
||||
//! This script implements production-ready MAMBA-2 training using:
|
||||
//! - Real DBN data (ES.FUT, NQ.FUT, ZN.FUT, 6E.FUT)
|
||||
//! - DbnSequenceLoader for data loading
|
||||
//! - GPU acceleration (CUDA) with 4GB VRAM optimization
|
||||
//! - Comprehensive checkpointing every 10 epochs
|
||||
//! - Early stopping with patience=20
|
||||
//! - Training metrics and loss curves
|
||||
//! - SSM state stability monitoring
|
||||
//!
|
||||
//! ## Configuration
|
||||
//! ```yaml
|
||||
//! Model: MAMBA-2 State Space Model
|
||||
//! Default Epochs: 200 (configurable)
|
||||
//! Batch Size: 32 (MAMBA-2 optimized)
|
||||
//! Learning Rate: 0.0001
|
||||
//! Hidden Dim: 256
|
||||
//! State Size: 16
|
||||
//! Layers: 6
|
||||
//! Sequence Length: 60
|
||||
//! Device: CUDA (GPU) with CPU fallback
|
||||
//! Data: Real DBN files from test_data/
|
||||
//! Checkpoints: ml/checkpoints/mamba2_dbn/
|
||||
//! ```
|
||||
//!
|
||||
//! ## Features
|
||||
//! - **Real Market Data**: Loads OHLCV bars from DBN files
|
||||
//! - **Feature Engineering**: 16 features + 10 technical indicators per timestep
|
||||
//! - **GPU Training**: RTX 3050 Ti optimized (~2GB VRAM usage)
|
||||
//! - **Checkpointing**: Saves best model based on validation loss
|
||||
//! - **Early Stopping**: Stops if no improvement for 20 epochs
|
||||
//! - **Monitoring**: Loss curves, perplexity, SSM state statistics
|
||||
//! - **Production Ready**: Follows Agent 78 fixes and best practices
|
||||
//!
|
||||
//! ## Usage
|
||||
//! ```bash
|
||||
//! # Default: 200 epochs, all available DBN data
|
||||
//! cargo run -p ml --example train_mamba2_dbn --release
|
||||
//!
|
||||
//! # Custom epochs:
|
||||
//! cargo run -p ml --example train_mamba2_dbn --release -- --epochs 50
|
||||
//!
|
||||
//! # Pilot run (50 epochs):
|
||||
//! cargo run -p ml --example train_mamba2_dbn --release -- --epochs 50
|
||||
//! ```
|
||||
//!
|
||||
//! ## Expected Training Time
|
||||
//! - 50 epochs: ~30-45 minutes (pilot)
|
||||
//! - 200 epochs: ~2-3 hours (full training)
|
||||
//! - GPU utilization: ~60-70% (memory-bound)
|
||||
//!
|
||||
//! ## Output
|
||||
//! - Checkpoints: ml/checkpoints/mamba2_dbn/checkpoint_epoch_*.safetensors
|
||||
//! - Best model: ml/checkpoints/mamba2_dbn/best_model.safetensors
|
||||
//! - Loss curves: ml/checkpoints/mamba2_dbn/training_losses.csv
|
||||
//! - Metrics: ml/checkpoints/mamba2_dbn/training_metrics.json
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::{Device, Tensor};
|
||||
use std::path::PathBuf;
|
||||
use std::time::Instant;
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use ml::data_loaders::DbnSequenceLoader;
|
||||
use ml::mamba::{Mamba2Config, Mamba2SSM};
|
||||
|
||||
/// Training configuration
|
||||
#[derive(Debug, Clone)]
|
||||
struct TrainingConfig {
|
||||
/// Number of training epochs
|
||||
pub epochs: usize,
|
||||
/// Batch size (MAMBA-2 is memory-intensive)
|
||||
pub batch_size: usize,
|
||||
/// Learning rate
|
||||
pub learning_rate: f64,
|
||||
/// Model dimension
|
||||
pub d_model: usize,
|
||||
/// Number of layers
|
||||
pub n_layers: usize,
|
||||
/// SSM state size
|
||||
pub state_size: usize,
|
||||
/// Sequence length for training
|
||||
pub seq_len: usize,
|
||||
/// Dropout rate
|
||||
pub dropout: f64,
|
||||
/// Gradient clipping
|
||||
pub grad_clip: f64,
|
||||
/// Weight decay
|
||||
pub weight_decay: f64,
|
||||
/// Warmup steps
|
||||
pub warmup_steps: usize,
|
||||
/// DBN data directory
|
||||
pub data_dir: PathBuf,
|
||||
/// Output directory for checkpoints
|
||||
pub checkpoint_dir: PathBuf,
|
||||
/// Early stopping patience
|
||||
pub early_stopping_patience: usize,
|
||||
}
|
||||
|
||||
impl Default for TrainingConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
epochs: 200,
|
||||
batch_size: 32, // Conservative for 4GB VRAM
|
||||
learning_rate: 0.0001,
|
||||
d_model: 54, // Wave D: 201 Wave C + 24 Wave D features (auto-adjusted from feature_config)
|
||||
n_layers: 6,
|
||||
state_size: 16, // SSM state dimension
|
||||
seq_len: 60, // 60 timesteps per sequence
|
||||
dropout: 0.1,
|
||||
grad_clip: 1.0,
|
||||
weight_decay: 1e-4,
|
||||
warmup_steps: 1000,
|
||||
data_dir: PathBuf::from("test_data/real/databento/ml_training_small"),
|
||||
checkpoint_dir: PathBuf::from("ml/checkpoints/mamba2_dbn"),
|
||||
early_stopping_patience: 20,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Training monitor for metrics tracking
|
||||
struct TrainingMonitor {
|
||||
pub start_time: Instant,
|
||||
pub best_val_loss: f64,
|
||||
pub best_epoch: usize,
|
||||
pub patience_counter: usize,
|
||||
pub epoch_losses: Vec<f64>,
|
||||
pub val_losses: Vec<f64>,
|
||||
pub learning_rates: Vec<f64>,
|
||||
}
|
||||
|
||||
impl TrainingMonitor {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
start_time: Instant::now(),
|
||||
best_val_loss: f64::INFINITY,
|
||||
best_epoch: 0,
|
||||
patience_counter: 0,
|
||||
epoch_losses: Vec::new(),
|
||||
val_losses: Vec::new(),
|
||||
learning_rates: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn update(
|
||||
&mut self,
|
||||
epoch: usize,
|
||||
train_loss: f64,
|
||||
val_loss: f64,
|
||||
lr: f64,
|
||||
patience: usize,
|
||||
) -> bool {
|
||||
self.epoch_losses.push(train_loss);
|
||||
self.val_losses.push(val_loss);
|
||||
self.learning_rates.push(lr);
|
||||
|
||||
if val_loss < self.best_val_loss {
|
||||
self.best_val_loss = val_loss;
|
||||
self.best_epoch = epoch;
|
||||
self.patience_counter = 0;
|
||||
true // Save checkpoint
|
||||
} else {
|
||||
self.patience_counter += 1;
|
||||
if self.patience_counter >= patience {
|
||||
info!(
|
||||
"Early stopping triggered: no improvement for {} epochs",
|
||||
patience
|
||||
);
|
||||
return false;
|
||||
}
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
fn should_stop(&self, patience: usize) -> bool {
|
||||
self.patience_counter >= patience
|
||||
}
|
||||
|
||||
fn get_summary(&self) -> String {
|
||||
let elapsed = self.start_time.elapsed();
|
||||
let avg_train_loss = if !self.epoch_losses.is_empty() {
|
||||
self.epoch_losses.iter().sum::<f64>() / self.epoch_losses.len() as f64
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
format!(
|
||||
"Training Summary:\n\
|
||||
- Duration: {:.2}h\n\
|
||||
- Best Val Loss: {:.6} (epoch {})\n\
|
||||
- Avg Train Loss: {:.6}\n\
|
||||
- Total Epochs: {}\n\
|
||||
- Perplexity: {:.4}",
|
||||
elapsed.as_secs_f64() / 3600.0,
|
||||
self.best_val_loss,
|
||||
self.best_epoch,
|
||||
avg_train_loss,
|
||||
self.epoch_losses.len(),
|
||||
self.best_val_loss.exp()
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Main training function
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Initialize tracing
|
||||
tracing_subscriber::fmt()
|
||||
.with_max_level(tracing::Level::INFO)
|
||||
.with_target(false)
|
||||
.with_thread_ids(false)
|
||||
.init();
|
||||
|
||||
info!("╔═══════════════════════════════════════════════════════════╗");
|
||||
info!("║ MAMBA-2 Production Training with Real DBN Data ║");
|
||||
info!("╚═══════════════════════════════════════════════════════════╝");
|
||||
|
||||
// Parse command-line arguments
|
||||
let args: Vec<String> = std::env::args().collect();
|
||||
let mut config = TrainingConfig::default();
|
||||
|
||||
// Wave B: Alternative bar sampling configuration
|
||||
#[allow(unused_assignments)]
|
||||
let mut bar_method: Option<String> = None;
|
||||
#[allow(unused_assignments)]
|
||||
let mut bar_threshold: Option<f64> = None;
|
||||
|
||||
// Parse all command-line arguments
|
||||
for i in 0..args.len() {
|
||||
match args[i].as_str() {
|
||||
"--epochs" if i + 1 < args.len() => {
|
||||
if let Ok(epochs) = args[i + 1].parse::<usize>() {
|
||||
config.epochs = epochs;
|
||||
info!("Custom epochs: {}", epochs);
|
||||
}
|
||||
},
|
||||
"--batch-size" if i + 1 < args.len() => {
|
||||
if let Ok(batch_size) = args[i + 1].parse::<usize>() {
|
||||
config.batch_size = batch_size;
|
||||
info!("Custom batch size: {}", batch_size);
|
||||
}
|
||||
},
|
||||
"--learning-rate" if i + 1 < args.len() => {
|
||||
if let Ok(lr) = args[i + 1].parse::<f64>() {
|
||||
config.learning_rate = lr;
|
||||
info!("Custom learning rate: {}", lr);
|
||||
}
|
||||
},
|
||||
"--sequence-length" if i + 1 < args.len() => {
|
||||
if let Ok(seq_len) = args[i + 1].parse::<usize>() {
|
||||
config.seq_len = seq_len;
|
||||
info!("Custom sequence length: {}", seq_len);
|
||||
}
|
||||
},
|
||||
"--hidden-dim" if i + 1 < args.len() => {
|
||||
if let Ok(d_model) = args[i + 1].parse::<usize>() {
|
||||
config.d_model = d_model;
|
||||
info!("Custom hidden dimension: {}", d_model);
|
||||
}
|
||||
},
|
||||
"--bar-method" if i + 1 < args.len() => {
|
||||
bar_method = Some(args[i + 1].clone());
|
||||
info!("Alternative bar method: {}", args[i + 1]);
|
||||
},
|
||||
"--bar-threshold" if i + 1 < args.len() => {
|
||||
if let Ok(threshold) = args[i + 1].parse::<f64>() {
|
||||
bar_threshold = Some(threshold);
|
||||
info!("Bar threshold: {}", threshold);
|
||||
}
|
||||
},
|
||||
"--state-dim" if i + 1 < args.len() => {
|
||||
if let Ok(state_size) = args[i + 1].parse::<usize>() {
|
||||
config.state_size = state_size;
|
||||
info!("Custom state dimension: {}", state_size);
|
||||
}
|
||||
},
|
||||
"--data-dir" if i + 1 < args.len() => {
|
||||
config.data_dir = PathBuf::from(&args[i + 1]);
|
||||
info!("Custom data directory: {:?}", config.data_dir);
|
||||
},
|
||||
"--output-dir" if i + 1 < args.len() => {
|
||||
config.checkpoint_dir = PathBuf::from(&args[i + 1]);
|
||||
info!("Custom output directory: {:?}", config.checkpoint_dir);
|
||||
},
|
||||
"--use-gpu" => {
|
||||
info!("GPU acceleration requested");
|
||||
},
|
||||
_ => {},
|
||||
}
|
||||
}
|
||||
|
||||
info!("Configuration:");
|
||||
info!(" Epochs: {}", config.epochs);
|
||||
info!(" Batch Size: {}", config.batch_size);
|
||||
info!(" Learning Rate: {}", config.learning_rate);
|
||||
info!(" Model Dimension: {}", config.d_model);
|
||||
info!(" State Size: {}", config.state_size);
|
||||
info!(" Sequence Length: {}", config.seq_len);
|
||||
info!(" Layers: {}", config.n_layers);
|
||||
info!(
|
||||
" Early Stopping Patience: {}",
|
||||
config.early_stopping_patience
|
||||
);
|
||||
|
||||
// Create checkpoint directory
|
||||
std::fs::create_dir_all(&config.checkpoint_dir)
|
||||
.context("Failed to create checkpoint directory")?;
|
||||
info!("Checkpoint directory: {:?}", config.checkpoint_dir);
|
||||
|
||||
// Initialize device (FORCE CUDA - no CPU fallback)
|
||||
info!("Initializing CUDA device (GPU-only mode)...");
|
||||
let device = Device::new_cuda(0).context(
|
||||
"CUDA GPU required for MAMBA-2 training. Ensure CUDA is installed and GPU is available.",
|
||||
)?;
|
||||
info!("✓ Using CUDA GPU (RTX 3050 Ti) - Device confirmed");
|
||||
|
||||
// Load DBN sequences with Wave D configuration (54 features)
|
||||
info!("Loading DBN sequences from: {:?}", config.data_dir);
|
||||
info!("Using Wave D feature configuration (54 features)");
|
||||
use ml::features::config::FeatureConfig;
|
||||
let feature_config = FeatureConfig::wave_d();
|
||||
info!(
|
||||
"Feature config phase: {:?}, feature_count: {}",
|
||||
feature_config.phase,
|
||||
feature_config.feature_count()
|
||||
);
|
||||
|
||||
// Override d_model to match Wave D feature count
|
||||
config.d_model = feature_config.feature_count();
|
||||
info!(
|
||||
"Adjusted d_model to {} to match Wave D feature count",
|
||||
config.d_model
|
||||
);
|
||||
|
||||
let mut loader = DbnSequenceLoader::with_feature_config(config.seq_len, feature_config)
|
||||
.await
|
||||
.context("Failed to create DBN sequence loader")?;
|
||||
|
||||
// Wave B: Configure alternative bar sampling if specified
|
||||
use ml::data_loaders::BarSamplingMethod;
|
||||
if let Some(method) = bar_method {
|
||||
let threshold = bar_threshold.unwrap_or_else(|| {
|
||||
// Default thresholds if not specified
|
||||
match method.as_str() {
|
||||
"tick" => 100.0,
|
||||
"volume" => 10000.0,
|
||||
"dollar" => 2_000_000.0, // $2M for ES.FUT
|
||||
"imbalance" => 1000.0,
|
||||
"run" => 50.0,
|
||||
_ => 100.0,
|
||||
}
|
||||
});
|
||||
|
||||
let bar_sampling = match method.as_str() {
|
||||
"tick" => BarSamplingMethod::TickBars(threshold as usize),
|
||||
"volume" => BarSamplingMethod::VolumeBars(threshold),
|
||||
"dollar" => BarSamplingMethod::DollarBars(threshold),
|
||||
"imbalance" => BarSamplingMethod::ImbalanceBars(threshold),
|
||||
"run" => BarSamplingMethod::RunBars(threshold as usize),
|
||||
_ => {
|
||||
warn!("Unknown bar method '{}', using time bars (default)", method);
|
||||
BarSamplingMethod::TimeBars
|
||||
},
|
||||
};
|
||||
|
||||
info!("✓ Alternative bar sampling configured: {:?}", bar_sampling);
|
||||
loader.set_bar_sampling_method(bar_sampling);
|
||||
}
|
||||
|
||||
let (train_data, val_data) = loader
|
||||
.load_sequences(&config.data_dir, 0.8) // 80% train, 20% validation
|
||||
.await
|
||||
.context("Failed to load DBN sequences")?;
|
||||
|
||||
info!("✓ Loaded {} training sequences", train_data.len());
|
||||
info!("✓ Loaded {} validation sequences", val_data.len());
|
||||
|
||||
if train_data.is_empty() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"No training data loaded! Check DBN files in {:?}",
|
||||
config.data_dir
|
||||
));
|
||||
}
|
||||
|
||||
// ===== SHAPE VALIDATION (Agent 200) =====
|
||||
// Verify that loader output matches expected dimensions [batch, seq_len, d_model]
|
||||
info!("╔═══════════════════════════════════════════════════════════╗");
|
||||
info!("║ Shape Validation (Agent 200) ║");
|
||||
info!("╚═══════════════════════════════════════════════════════════╝");
|
||||
|
||||
if !train_data.is_empty() {
|
||||
let (first_input, first_target) = &train_data[0];
|
||||
let input_shape = first_input.dims();
|
||||
let target_shape = first_target.dims();
|
||||
|
||||
info!("First training sequence shape validation:");
|
||||
info!(" Input shape: {:?}", input_shape);
|
||||
info!(" Target shape: {:?}", target_shape);
|
||||
info!(
|
||||
" Expected input: [1, {}, {}]",
|
||||
config.seq_len, config.d_model
|
||||
);
|
||||
info!(" Expected target: [1, 1, 1] (regression: next close price)");
|
||||
|
||||
// Validate input dimensions
|
||||
if input_shape.len() != 3 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Invalid input tensor rank! Expected 3D [batch, seq_len, d_model], got {}D: {:?}",
|
||||
input_shape.len(),
|
||||
input_shape
|
||||
));
|
||||
}
|
||||
|
||||
if input_shape[0] != 1 {
|
||||
warn!(
|
||||
"⚠️ Input batch dimension is {}, expected 1 (will be batched during training)",
|
||||
input_shape[0]
|
||||
);
|
||||
}
|
||||
|
||||
if input_shape[1] != config.seq_len {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input sequence length mismatch! Expected seq_len={}, got {}",
|
||||
config.seq_len,
|
||||
input_shape[1]
|
||||
));
|
||||
}
|
||||
|
||||
if input_shape[2] != config.d_model {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input feature dimension mismatch! Expected d_model={}, got {}",
|
||||
config.d_model,
|
||||
input_shape[2]
|
||||
));
|
||||
}
|
||||
|
||||
// FIXED (Agent 254): Validate target dimensions for regression
|
||||
// Agent 246 changed model output_dim to 1 for price prediction (regression)
|
||||
// Target shape should be [batch, 1, 1] not [batch, 1, d_model]
|
||||
if target_shape.len() != 3 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Invalid target tensor rank! Expected 3D [batch, 1, 1], got {}D: {:?}",
|
||||
target_shape.len(),
|
||||
target_shape
|
||||
));
|
||||
}
|
||||
|
||||
if target_shape[2] != 1 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Target dimension mismatch! Expected output_dim=1 (regression), got {}",
|
||||
target_shape[2]
|
||||
));
|
||||
}
|
||||
|
||||
info!("✓ Shape validation PASSED");
|
||||
info!(
|
||||
" Input: [batch={}, seq_len={}, d_model={}]",
|
||||
input_shape[0], input_shape[1], input_shape[2]
|
||||
);
|
||||
info!(
|
||||
" Target: [batch={}, steps={}, output_dim={}] (regression: next close price)",
|
||||
target_shape[0], target_shape[1], target_shape[2]
|
||||
);
|
||||
}
|
||||
// ===== END SHAPE VALIDATION =====
|
||||
|
||||
// Estimate memory usage
|
||||
let params_per_layer = config.d_model * config.state_size * 3; // A, B, C matrices
|
||||
let total_params = params_per_layer * config.n_layers;
|
||||
let memory_mb = (total_params * 4 * 3) / (1024 * 1024); // params + gradients + optimizer (f32)
|
||||
info!("Estimated VRAM usage: ~{}MB (model parameters)", memory_mb);
|
||||
|
||||
if memory_mb > 3500 {
|
||||
warn!("⚠ Memory usage may exceed 4GB VRAM constraint!");
|
||||
}
|
||||
|
||||
// Create MAMBA-2 model
|
||||
info!("Initializing MAMBA-2 model...");
|
||||
let mamba_config = Mamba2Config {
|
||||
d_model: config.d_model,
|
||||
d_state: config.state_size,
|
||||
d_head: config.d_model / 8,
|
||||
num_heads: 8,
|
||||
expand: 2,
|
||||
num_layers: config.n_layers,
|
||||
dropout: config.dropout,
|
||||
use_ssd: true, // Structured State Duality
|
||||
use_selective_state: true, // Selective state mechanism
|
||||
hardware_aware: true,
|
||||
target_latency_us: 5,
|
||||
max_seq_len: config.seq_len * 2,
|
||||
learning_rate: config.learning_rate,
|
||||
weight_decay: config.weight_decay,
|
||||
grad_clip: config.grad_clip,
|
||||
warmup_steps: config.warmup_steps,
|
||||
adam_beta1: 0.9, // P1: Adam beta1 parameter
|
||||
adam_beta2: 0.999, // P1: Adam beta2 parameter
|
||||
adam_epsilon: 1e-8, // P1: Adam epsilon
|
||||
total_decay_steps: config.epochs * 100, // P1: Total decay steps
|
||||
optimizer_type: ml::mamba::OptimizerType::Adam, // P1: Optimizer type
|
||||
sgd_momentum: 0.9, // P1: SGD momentum (unused for Adam)
|
||||
batch_size: config.batch_size,
|
||||
seq_len: config.seq_len,
|
||||
shuffle_batches: false, // Reproducibility
|
||||
sequence_stride: config.seq_len / 2, // P2: Overlapping windows
|
||||
norm_eps: 1e-5, // P2: Layer norm epsilon
|
||||
early_stopping_enabled: true,
|
||||
early_stopping_patience: config.early_stopping_patience,
|
||||
early_stopping_min_delta: 1e-6,
|
||||
early_stopping_min_epochs: 10,
|
||||
};
|
||||
|
||||
let mut model =
|
||||
Mamba2SSM::new(mamba_config.clone(), &device).context("Failed to create MAMBA-2 model")?;
|
||||
|
||||
let param_count = model.metadata.num_parameters;
|
||||
info!("✓ Model initialized: {} parameters", param_count);
|
||||
|
||||
// Initialize training monitor
|
||||
let mut monitor = TrainingMonitor::new();
|
||||
|
||||
// Training loop
|
||||
info!("╔═══════════════════════════════════════════════════════════╗");
|
||||
info!("║ Starting Training Loop ║");
|
||||
info!("╚═══════════════════════════════════════════════════════════╝");
|
||||
|
||||
// Debug logging: show first batch shapes (Agent 200)
|
||||
info!("Debug: First batch tensor shapes (Agent 200):");
|
||||
for (idx, (input, target)) in train_data.iter().take(3).enumerate() {
|
||||
info!(
|
||||
" Sequence {}: input={:?}, target={:?}",
|
||||
idx,
|
||||
input.dims(),
|
||||
target.dims()
|
||||
);
|
||||
|
||||
// Verify shape consistency
|
||||
if input.dims().len() != 3 || input.dims()[2] != config.d_model {
|
||||
error!(
|
||||
"⚠️ SHAPE MISMATCH: Sequence {} has invalid input shape: {:?}",
|
||||
idx,
|
||||
input.dims()
|
||||
);
|
||||
return Err(anyhow::anyhow!(
|
||||
"Training data shape mismatch at sequence {}: expected [1, {}, {}], got {:?}",
|
||||
idx,
|
||||
config.seq_len,
|
||||
config.d_model,
|
||||
input.dims()
|
||||
));
|
||||
}
|
||||
}
|
||||
info!(
|
||||
"✓ First batch shapes verified: all sequences match [1, {}, {}]",
|
||||
config.seq_len, config.d_model
|
||||
);
|
||||
|
||||
// Validate ALL training and validation sequences before expensive training loop
|
||||
validate_training_batch(&train_data, config.seq_len, config.d_model)
|
||||
.context("Training data validation failed")?;
|
||||
validate_training_batch(&val_data, config.seq_len, config.d_model)
|
||||
.context("Validation data validation failed")?;
|
||||
|
||||
let training_history = model
|
||||
.train(&train_data, &val_data, config.epochs, None)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
// Process training history with early stopping
|
||||
for (epoch_idx, epoch) in training_history.iter().enumerate() {
|
||||
let train_loss = epoch.loss;
|
||||
let should_save = monitor.update(
|
||||
epoch_idx,
|
||||
train_loss,
|
||||
train_loss, // Using train loss as val loss for now
|
||||
epoch.learning_rate,
|
||||
config.early_stopping_patience,
|
||||
);
|
||||
|
||||
// Save checkpoint if best model
|
||||
if should_save {
|
||||
let checkpoint_path = config
|
||||
.checkpoint_dir
|
||||
.join(format!("best_model_epoch_{}.ckpt", epoch_idx));
|
||||
|
||||
let path_str = checkpoint_path
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid checkpoint path"))?;
|
||||
model
|
||||
.save_checkpoint(path_str)
|
||||
.await
|
||||
.context("Failed to save checkpoint")?;
|
||||
|
||||
info!(
|
||||
"✓ Saved best model at epoch {} (loss: {:.6})",
|
||||
epoch_idx, epoch.loss
|
||||
);
|
||||
}
|
||||
|
||||
// Save periodic checkpoints every 10 epochs
|
||||
if epoch_idx % 10 == 0 && epoch_idx > 0 {
|
||||
let checkpoint_path = config
|
||||
.checkpoint_dir
|
||||
.join(format!("checkpoint_epoch_{}.ckpt", epoch_idx));
|
||||
|
||||
let path_str = checkpoint_path
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid checkpoint path"))?;
|
||||
model
|
||||
.save_checkpoint(path_str)
|
||||
.await
|
||||
.context("Failed to save checkpoint")?;
|
||||
|
||||
info!("✓ Checkpoint saved: epoch {}", epoch_idx);
|
||||
}
|
||||
|
||||
// Log progress every 5 epochs
|
||||
if epoch_idx % 5 == 0 {
|
||||
let perplexity = epoch.loss.exp();
|
||||
let elapsed = monitor.start_time.elapsed();
|
||||
let epochs_per_min = (epoch_idx + 1) as f64 / elapsed.as_secs_f64() * 60.0;
|
||||
|
||||
info!(
|
||||
"Epoch {:3}/{}: Loss={:.6}, Perplexity={:.4}, LR={:.2e}, Time={:.1}s, Speed={:.1} ep/min",
|
||||
epoch_idx + 1,
|
||||
config.epochs,
|
||||
train_loss,
|
||||
perplexity,
|
||||
epoch.learning_rate,
|
||||
epoch.duration_seconds,
|
||||
epochs_per_min
|
||||
);
|
||||
}
|
||||
|
||||
// Check for early stopping
|
||||
if monitor.should_stop(config.early_stopping_patience) {
|
||||
info!("Early stopping at epoch {}", epoch_idx);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Training completed
|
||||
info!("╔═══════════════════════════════════════════════════════════╗");
|
||||
info!("║ Training Completed ║");
|
||||
info!("╚═══════════════════════════════════════════════════════════╝");
|
||||
info!("{}", monitor.get_summary());
|
||||
|
||||
// Save final model
|
||||
let final_model_path = config.checkpoint_dir.join("final_model.ckpt");
|
||||
let final_path_str = final_model_path
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid final model path"))?;
|
||||
model
|
||||
.save_checkpoint(final_path_str)
|
||||
.await
|
||||
.context("Failed to save final model")?;
|
||||
info!("✓ Final model saved: {:?}", final_model_path);
|
||||
|
||||
// Export training curves
|
||||
export_training_metrics(&monitor, &config)?;
|
||||
|
||||
// Final analysis
|
||||
info!("╔═══════════════════════════════════════════════════════════╗");
|
||||
info!("║ Final Model Analysis ║");
|
||||
info!("╚═══════════════════════════════════════════════════════════╝");
|
||||
|
||||
let model_metrics = model.get_performance_metrics();
|
||||
info!("Model Performance Metrics:");
|
||||
info!(
|
||||
" Total Inferences: {}",
|
||||
model_metrics.get("total_inferences").unwrap_or(&0.0)
|
||||
);
|
||||
info!(
|
||||
" Total Training Steps: {}",
|
||||
model_metrics.get("total_training_steps").unwrap_or(&0.0)
|
||||
);
|
||||
info!(
|
||||
" Model Parameters: {}",
|
||||
model_metrics.get("model_parameters").unwrap_or(&0.0)
|
||||
);
|
||||
|
||||
if let Some(compression_ratio) = model_metrics.get("compression_ratio") {
|
||||
info!(" State Compression Ratio: {:.4}", compression_ratio);
|
||||
}
|
||||
|
||||
// Convergence analysis
|
||||
if monitor.epoch_losses.len() >= 10 {
|
||||
let recent_losses: Vec<f64> = monitor
|
||||
.epoch_losses
|
||||
.iter()
|
||||
.rev()
|
||||
.take(10)
|
||||
.copied()
|
||||
.collect();
|
||||
let avg_recent = recent_losses.iter().sum::<f64>() / recent_losses.len() as f64;
|
||||
let std_dev = {
|
||||
let variance = recent_losses
|
||||
.iter()
|
||||
.map(|l| (l - avg_recent).powi(2))
|
||||
.sum::<f64>()
|
||||
/ recent_losses.len() as f64;
|
||||
variance.sqrt()
|
||||
};
|
||||
|
||||
info!("Convergence Analysis (last 10 epochs):");
|
||||
info!(" Avg Loss: {:.6}", avg_recent);
|
||||
info!(" Std Dev: {:.6}", std_dev);
|
||||
|
||||
if std_dev < 0.01 {
|
||||
info!("✓ Model has CONVERGED (low variance in recent losses)");
|
||||
} else if std_dev < 0.05 {
|
||||
info!("⚠ Model is CONVERGING (moderate variance)");
|
||||
} else {
|
||||
info!("⚠ Model still LEARNING (high variance - may need more epochs)");
|
||||
}
|
||||
}
|
||||
|
||||
// Loss reduction
|
||||
if !monitor.epoch_losses.is_empty() {
|
||||
let initial_loss = monitor.epoch_losses[0];
|
||||
let final_loss = monitor
|
||||
.epoch_losses
|
||||
.last()
|
||||
.copied()
|
||||
.unwrap_or(0.0);
|
||||
let reduction = ((initial_loss - final_loss) / initial_loss) * 100.0;
|
||||
|
||||
info!("Loss Reduction:");
|
||||
info!(" Initial: {:.6}", initial_loss);
|
||||
info!(" Final: {:.6}", final_loss);
|
||||
info!(" Reduction: {:.2}%", reduction);
|
||||
|
||||
if reduction > 30.0 {
|
||||
info!("✓ EXCELLENT: >30% loss reduction");
|
||||
} else if reduction > 10.0 {
|
||||
info!("✓ GOOD: 10-30% loss reduction");
|
||||
} else {
|
||||
info!("⚠ LOW: <10% loss reduction (may need more epochs or hyperparameter tuning)");
|
||||
}
|
||||
}
|
||||
|
||||
info!("╔═══════════════════════════════════════════════════════════╗");
|
||||
info!("║ MAMBA-2 Training Successfully Completed ║");
|
||||
info!("╚═══════════════════════════════════════════════════════════╝");
|
||||
info!(
|
||||
"Best model: {:?}/best_model_epoch_{}.ckpt",
|
||||
config.checkpoint_dir, monitor.best_epoch
|
||||
);
|
||||
info!(
|
||||
"Training metrics: {:?}/training_metrics.json",
|
||||
config.checkpoint_dir
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Export training metrics to CSV and JSON
|
||||
fn export_training_metrics(monitor: &TrainingMonitor, config: &TrainingConfig) -> Result<()> {
|
||||
use std::io::Write;
|
||||
|
||||
// Export training losses to CSV
|
||||
let loss_csv_path = config.checkpoint_dir.join("training_losses.csv");
|
||||
let mut loss_file = std::fs::File::create(&loss_csv_path)?;
|
||||
writeln!(loss_file, "epoch,train_loss,val_loss,learning_rate")?;
|
||||
|
||||
for (i, ((train_loss, val_loss), lr)) in monitor
|
||||
.epoch_losses
|
||||
.iter()
|
||||
.zip(monitor.val_losses.iter())
|
||||
.zip(monitor.learning_rates.iter())
|
||||
.enumerate()
|
||||
{
|
||||
writeln!(loss_file, "{},{},{},{}", i, train_loss, val_loss, lr)?;
|
||||
}
|
||||
info!("✓ Training losses exported: {:?}", loss_csv_path);
|
||||
|
||||
// Export summary metrics to JSON
|
||||
let metrics_json_path = config.checkpoint_dir.join("training_metrics.json");
|
||||
let summary = serde_json::json!({
|
||||
"total_epochs": monitor.epoch_losses.len(),
|
||||
"best_val_loss": monitor.best_val_loss,
|
||||
"best_epoch": monitor.best_epoch,
|
||||
"training_duration_hours": monitor.start_time.elapsed().as_secs_f64() / 3600.0,
|
||||
"final_perplexity": monitor.best_val_loss.exp(),
|
||||
"config": {
|
||||
"d_model": config.d_model,
|
||||
"n_layers": config.n_layers,
|
||||
"state_size": config.state_size,
|
||||
"seq_len": config.seq_len,
|
||||
"batch_size": config.batch_size,
|
||||
"learning_rate": config.learning_rate,
|
||||
"dropout": config.dropout,
|
||||
}
|
||||
});
|
||||
|
||||
let mut metrics_file = std::fs::File::create(&metrics_json_path)?;
|
||||
metrics_file.write_all(serde_json::to_string_pretty(&summary)?.as_bytes())?;
|
||||
info!("✓ Training metrics exported: {:?}", metrics_json_path);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Validate tensor shapes for training (Agent 201, updated by Agent 254)
|
||||
///
|
||||
/// FIXED (Agent 254): Ensures that input and target tensors have correct shapes for MAMBA-2 regression:
|
||||
/// - Input: [batch_size, seq_len, d_model]
|
||||
/// - Target: [batch_size, 1, 1] (regression: next close price)
|
||||
///
|
||||
/// Also validates that tensors are contiguous in memory for efficient GPU operations.
|
||||
fn validate_tensor_shapes(
|
||||
input: &Tensor,
|
||||
target: &Tensor,
|
||||
expected_batch_size: usize,
|
||||
expected_seq_len: usize,
|
||||
expected_d_model: usize,
|
||||
) -> Result<()> {
|
||||
// Validate input tensor shape
|
||||
let input_dims = input.dims();
|
||||
|
||||
if input_dims.len() != 3 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input tensor must be 3D [batch, seq, features], got {} dimensions: {:?}",
|
||||
input_dims.len(),
|
||||
input_dims
|
||||
));
|
||||
}
|
||||
|
||||
if input_dims[0] != expected_batch_size {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input batch size mismatch: expected {}, got {}",
|
||||
expected_batch_size,
|
||||
input_dims[0]
|
||||
));
|
||||
}
|
||||
|
||||
if input_dims[1] != expected_seq_len {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input sequence length mismatch: expected {}, got {}",
|
||||
expected_seq_len,
|
||||
input_dims[1]
|
||||
));
|
||||
}
|
||||
|
||||
if input_dims[2] != expected_d_model {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input feature dimension mismatch: expected {}, got {}",
|
||||
expected_d_model,
|
||||
input_dims[2]
|
||||
));
|
||||
}
|
||||
|
||||
// FIXED (Agent 254): Validate target tensor shape for regression
|
||||
// Target should be [batch, 1, 1] for price prediction (regression)
|
||||
let target_dims = target.dims();
|
||||
|
||||
if target_dims.len() != 3 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Target tensor must be 3D [batch, 1, 1] for regression, got {} dimensions: {:?}",
|
||||
target_dims.len(),
|
||||
target_dims
|
||||
));
|
||||
}
|
||||
|
||||
let target_batch_size = target_dims[0];
|
||||
if target_batch_size != expected_batch_size {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Target batch size mismatch: expected {}, got {}",
|
||||
expected_batch_size,
|
||||
target_batch_size
|
||||
));
|
||||
}
|
||||
|
||||
// Validate target shape: [batch, 1, 1] for regression
|
||||
if target_dims[1] != 1 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Target tensor middle dimension must be 1 for [batch, 1, 1], got {}",
|
||||
target_dims[1]
|
||||
));
|
||||
}
|
||||
|
||||
if target_dims[2] != 1 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Target output dimension must be 1 for regression, got {}",
|
||||
target_dims[2]
|
||||
));
|
||||
}
|
||||
|
||||
// Validate tensors are contiguous for GPU efficiency
|
||||
if !input.is_contiguous() {
|
||||
warn!("⚠ Input tensor is not contiguous - may impact GPU performance");
|
||||
}
|
||||
|
||||
if !target.is_contiguous() {
|
||||
warn!("⚠ Target tensor is not contiguous - may impact GPU performance");
|
||||
}
|
||||
|
||||
// Check for empty tensors
|
||||
if input_dims.iter().any(|&d| d == 0) {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Input tensor has zero dimension: {:?}",
|
||||
input_dims
|
||||
));
|
||||
}
|
||||
|
||||
if target_dims.iter().any(|&d| d == 0) {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Target tensor has zero dimension: {:?}",
|
||||
target_dims
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Validate a batch of training sequences (Agent 201)
|
||||
///
|
||||
/// Performs shape validation on all training sequences to catch issues early
|
||||
/// before starting the expensive training loop. This helps prevent CUDA errors
|
||||
/// and ensures data integrity.
|
||||
fn validate_training_batch(
|
||||
batch: &[(Tensor, Tensor)],
|
||||
expected_seq_len: usize,
|
||||
expected_d_model: usize,
|
||||
) -> Result<()> {
|
||||
if batch.is_empty() {
|
||||
return Err(anyhow::anyhow!("Training batch is empty"));
|
||||
}
|
||||
|
||||
info!("Validating {} training sequences...", batch.len());
|
||||
|
||||
for (idx, (input, target)) in batch.iter().enumerate() {
|
||||
// Each sequence has batch_size=1 in the loader
|
||||
validate_tensor_shapes(input, target, 1, expected_seq_len, expected_d_model)
|
||||
.context(format!("Validation failed for sequence {}", idx))?;
|
||||
}
|
||||
|
||||
info!(
|
||||
"✓ All {} training sequences validated successfully",
|
||||
batch.len()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,242 +0,0 @@
|
||||
//! Train PPO on ES.FUT Market Data
|
||||
//!
|
||||
//! This example demonstrates training a PPO (Proximal Policy Optimization) model
|
||||
//! on real ES.FUT (E-mini S&P 500) futures data for trading strategy development.
|
||||
//!
|
||||
//! ## Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! cargo run -p ml --example train_ppo_es_fut --release
|
||||
//! ```
|
||||
//!
|
||||
//! ## Configuration
|
||||
//!
|
||||
//! - State dimension: 26 (OHLCV + technical indicators)
|
||||
//! - Actions: 3 (Buy, Sell, Hold)
|
||||
//! - Training epochs: 50 (configurable)
|
||||
//! - GPU: Automatic (RTX 3050 Ti if available, else CPU)
|
||||
//! - Checkpoints: Saved every 10 epochs to `ml/checkpoints/`
|
||||
//!
|
||||
//! ## Expected Results
|
||||
//!
|
||||
//! - Policy improvement > 20% over 50 epochs
|
||||
//! - Value loss decreasing trend
|
||||
//! - Checkpoint file ~10-20 KB per epoch
|
||||
//! - Training time: ~5-10 minutes on CPU, ~2-3 minutes on GPU
|
||||
//!
|
||||
//! ## Output
|
||||
//!
|
||||
//! - Checkpoint: `ml/checkpoints/ppo_es_fut_v1_actor_epoch_50.safetensors`
|
||||
//! - Checkpoint: `ml/checkpoints/ppo_es_fut_v1_critic_epoch_50.safetensors`
|
||||
//! - Metrics: Epoch-by-epoch training progress
|
||||
|
||||
use anyhow::Result;
|
||||
use ml::trainers::ppo::{PpoHyperparameters, PpoTrainer, PpoTrainingMetrics};
|
||||
use std::f32::consts::PI;
|
||||
|
||||
/// Generate synthetic ES.FUT market data
|
||||
///
|
||||
/// In production, this would load from Parquet files or database.
|
||||
/// For now, we generate realistic synthetic data with:
|
||||
/// - OHLCV patterns (sine wave price movements)
|
||||
/// - Technical indicators (RSI, MACD, Bollinger Bands, etc.)
|
||||
/// - Realistic price ranges (~4000-4200 for ES.FUT)
|
||||
fn generate_market_data(num_bars: usize, state_dim: usize) -> Vec<Vec<f32>> {
|
||||
println!(
|
||||
"🔄 Generating {} bars of synthetic ES.FUT data...",
|
||||
num_bars
|
||||
);
|
||||
|
||||
let mut data: Vec<Vec<f32>> = Vec::with_capacity(num_bars);
|
||||
|
||||
for i in 0..num_bars {
|
||||
let t = i as f32 / num_bars as f32;
|
||||
|
||||
// Base price with trend and volatility
|
||||
let base_price = 4100.0 + 50.0 * (t * 2.0 * PI).sin() + 20.0 * (t * 10.0 * PI).sin();
|
||||
|
||||
// OHLCV features
|
||||
let close = base_price;
|
||||
let high = close * 1.005; // 0.5% above close
|
||||
let low = close * 0.995; // 0.5% below close
|
||||
let open = close * (1.0 + 0.002 * (t * 5.0 * PI).sin());
|
||||
let volume = 1000.0 + 200.0 * (t * 4.0 * PI).cos();
|
||||
|
||||
// Technical indicators
|
||||
let rsi = 50.0 + 20.0 * (t * PI).sin(); // RSI oscillating around 50
|
||||
let macd = (t * 2.0 * PI).sin(); // MACD signal
|
||||
let signal = (t * 2.0 * PI - 0.5).sin(); // Signal line
|
||||
let atr = 15.0 + 5.0 * (t * 3.0 * PI).cos(); // ATR
|
||||
let bb_lower = close * 0.98; // Bollinger lower
|
||||
let bb_upper = close * 1.02; // Bollinger upper
|
||||
let ema = close * (1.0 + 0.001 * (t * PI).cos()); // EMA
|
||||
|
||||
// Build state vector
|
||||
let mut state = vec![
|
||||
close, high, low, open, volume, rsi, macd, signal, atr, bb_lower, bb_upper, ema,
|
||||
];
|
||||
|
||||
// Add log return (used for reward calculation)
|
||||
let log_return = if i > 0 {
|
||||
let prev_close = data[i - 1][0]; // Previous close
|
||||
(close / prev_close).ln()
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
state.push(log_return);
|
||||
|
||||
// Pad to state_dim with zeros
|
||||
while state.len() < state_dim {
|
||||
state.push(0.0);
|
||||
}
|
||||
|
||||
data.push(state);
|
||||
}
|
||||
|
||||
println!("✓ Generated {} bars (state_dim={})", data.len(), state_dim);
|
||||
data
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Initialize logging
|
||||
tracing_subscriber::fmt()
|
||||
.with_max_level(tracing::Level::INFO)
|
||||
.init();
|
||||
|
||||
println!("\n🚀 PPO Training on ES.FUT Market Data");
|
||||
println!("=====================================\n");
|
||||
|
||||
// Configuration
|
||||
let state_dim = 26;
|
||||
let num_bars = 5000; // 5000 bars for more robust training
|
||||
let num_epochs = 50;
|
||||
let checkpoint_dir = "ml/checkpoints";
|
||||
|
||||
// Generate synthetic market data (in production, load from Parquet)
|
||||
let market_data = generate_market_data(num_bars, state_dim);
|
||||
|
||||
// Configure PPO hyperparameters
|
||||
let mut hyperparams = PpoHyperparameters::conservative();
|
||||
hyperparams.epochs = num_epochs;
|
||||
hyperparams.learning_rate = 3e-4; // Standard PPO learning rate
|
||||
hyperparams.batch_size = 128; // Larger batch for stability
|
||||
hyperparams.rollout_steps = 2048; // Standard rollout length
|
||||
hyperparams.minibatch_size = 64; // Mini-batch size
|
||||
hyperparams.gamma = 0.99; // Discount factor
|
||||
hyperparams.gae_lambda = 0.95; // GAE parameter
|
||||
hyperparams.clip_epsilon = 0.2; // PPO clip range
|
||||
hyperparams.vf_coef = 0.5; // Value loss coefficient
|
||||
hyperparams.ent_coef = 0.01; // Entropy coefficient
|
||||
hyperparams.early_stopping_enabled = true;
|
||||
hyperparams.min_value_loss_improvement_pct = 2.0;
|
||||
hyperparams.min_explained_variance = 0.4;
|
||||
hyperparams.plateau_window = 30;
|
||||
hyperparams.min_epochs_before_stopping = 50; // No early stopping for full training
|
||||
|
||||
println!("📋 Training Configuration:");
|
||||
println!(" • State dimension: {}", state_dim);
|
||||
println!(" • Market data: {} bars", num_bars);
|
||||
println!(" • Training epochs: {}", num_epochs);
|
||||
println!(" • Learning rate: {}", hyperparams.learning_rate);
|
||||
println!(" • Batch size: {}", hyperparams.batch_size);
|
||||
println!(" • Rollout steps: {}", hyperparams.rollout_steps);
|
||||
println!(" • Checkpoint dir: {}", checkpoint_dir);
|
||||
|
||||
// Detect GPU availability
|
||||
let use_gpu = candle_core::Device::cuda_if_available(0).is_ok();
|
||||
println!(
|
||||
" • Device: {}\n",
|
||||
if use_gpu { "GPU (CUDA)" } else { "CPU" }
|
||||
);
|
||||
|
||||
// Create PPO trainer
|
||||
let trainer = PpoTrainer::new(hyperparams, state_dim, checkpoint_dir, use_gpu, None)?;
|
||||
|
||||
println!("✓ PPO trainer initialized\n");
|
||||
println!("🏋️ Starting training...\n");
|
||||
println!(
|
||||
"{:<8} {:<12} {:<12} {:<12} {:<12}",
|
||||
"Epoch", "Policy Loss", "Value Loss", "Expl. Var.", "Mean Reward"
|
||||
);
|
||||
println!("{}", "-".repeat(64));
|
||||
|
||||
// Track metrics for summary
|
||||
let mut all_metrics = Vec::new();
|
||||
|
||||
// Train PPO model
|
||||
let final_metrics = trainer
|
||||
.train(market_data, |metrics: PpoTrainingMetrics| {
|
||||
println!(
|
||||
"{:<8} {:<12.4} {:<12.4} {:<12.4} {:<12.4}",
|
||||
metrics.epoch,
|
||||
metrics.policy_loss,
|
||||
metrics.value_loss,
|
||||
metrics.explained_variance,
|
||||
metrics.mean_reward
|
||||
);
|
||||
all_metrics.push(metrics);
|
||||
})
|
||||
.await?;
|
||||
|
||||
println!("{}", "-".repeat(64));
|
||||
println!("\n✅ Training complete!\n");
|
||||
|
||||
// Print summary statistics
|
||||
println!("📊 Training Summary:");
|
||||
println!(" • Final epoch: {}", final_metrics.epoch);
|
||||
println!(" • Policy loss: {:.4}", final_metrics.policy_loss);
|
||||
println!(" • Value loss: {:.4}", final_metrics.value_loss);
|
||||
println!(" • KL divergence: {:.4}", final_metrics.kl_divergence);
|
||||
println!(
|
||||
" • Explained variance: {:.4}",
|
||||
final_metrics.explained_variance
|
||||
);
|
||||
println!(" • Mean reward: {:.4}", final_metrics.mean_reward);
|
||||
println!(" • Std reward: {:.4}", final_metrics.std_reward);
|
||||
println!(" • Entropy: {:.4}\n", final_metrics.entropy);
|
||||
|
||||
// Compute improvement metrics
|
||||
if let (Some(first), Some(last)) = (all_metrics.first(), all_metrics.last()) {
|
||||
let policy_improvement =
|
||||
((first.policy_loss - last.policy_loss) / first.policy_loss.abs()) * 100.0;
|
||||
let value_improvement = ((first.value_loss - last.value_loss) / first.value_loss) * 100.0;
|
||||
|
||||
println!("📈 Improvement Over Training:");
|
||||
println!(" • Policy loss: {:.2}%", policy_improvement);
|
||||
println!(" • Value loss: {:.2}%\n", value_improvement);
|
||||
|
||||
// Check if target achieved
|
||||
if policy_improvement > 20.0 {
|
||||
println!(
|
||||
"🎯 Target achieved: Policy improved by {:.2}% (target: >20%)",
|
||||
policy_improvement
|
||||
);
|
||||
} else {
|
||||
println!(
|
||||
"⚠️ Target not met: Policy improved by {:.2}% (target: >20%)",
|
||||
policy_improvement
|
||||
);
|
||||
println!(" Consider training for more epochs or tuning hyperparameters");
|
||||
}
|
||||
}
|
||||
|
||||
// Print checkpoint locations
|
||||
println!("\n💾 Model Checkpoints:");
|
||||
println!(
|
||||
" • Actor: {}/ppo_es_fut_v1_actor_epoch_{}.safetensors",
|
||||
checkpoint_dir, final_metrics.epoch
|
||||
);
|
||||
println!(
|
||||
" • Critic: {}/ppo_es_fut_v1_critic_epoch_{}.safetensors",
|
||||
checkpoint_dir, final_metrics.epoch
|
||||
);
|
||||
|
||||
println!("\n🎉 PPO training pipeline complete!");
|
||||
println!("\nNext steps:");
|
||||
println!("1. Validate checkpoint loading: cargo test -p ml test_checkpoint_loading");
|
||||
println!("2. Backtest strategy with trained model");
|
||||
println!("3. Deploy to paper trading environment\n");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,532 +0,0 @@
|
||||
//! PPO Extended Training with Hyperparameter Tuning (Agent F6)
|
||||
//!
|
||||
//! This script implements the Agent F6 task: extended PPO training with:
|
||||
//! - 100 epochs (increased from 20)
|
||||
//! - Hyperparameter tuning (learning rate, clip_ratio, entropy_coef)
|
||||
//! - Comprehensive training curve monitoring
|
||||
//! - Policy improvement validation
|
||||
//! - Inference latency benchmarking
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with 100 epochs and tuned hyperparameters
|
||||
//! cargo run -p ml --example train_ppo_extended --release --features cuda
|
||||
//!
|
||||
//! # Custom configuration
|
||||
//! cargo run -p ml --example train_ppo_extended --release --features cuda -- \
|
||||
//! --epochs 100 \
|
||||
//! --learning-rate 0.0001 \
|
||||
//! --clip-epsilon 0.2 \
|
||||
//! --entropy-coef 0.05 \
|
||||
//! --value-coef 1.0 \
|
||||
//! --output-dir ml/trained_models/ppo_extended
|
||||
//! ```
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::path::PathBuf;
|
||||
use std::time::Instant;
|
||||
use tracing::{info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::real_data_loader::RealDataLoader;
|
||||
use ml::trainers::ppo::{PpoHyperparameters, PpoTrainer, PpoTrainingMetrics};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_ppo_extended",
|
||||
about = "PPO Extended Training with Hyperparameter Tuning (Agent F6)"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Number of training epochs (Agent F6: 100 epochs)
|
||||
#[arg(long, default_value = "100")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate (tuned for value network convergence)
|
||||
#[arg(long, default_value = "0.0001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Clip epsilon (PPO clip range, 0.1-0.3)
|
||||
#[arg(long, default_value = "0.2")]
|
||||
clip_epsilon: f32,
|
||||
|
||||
/// Value function coefficient (increased for value learning)
|
||||
#[arg(long, default_value = "1.0")]
|
||||
value_coef: f32,
|
||||
|
||||
/// Entropy coefficient (exploration vs exploitation)
|
||||
#[arg(long, default_value = "0.05")]
|
||||
entropy_coef: f32,
|
||||
|
||||
/// Batch size (max 230 for RTX 3050 Ti 4GB)
|
||||
#[arg(long, default_value = "64")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Output directory for trained model
|
||||
#[arg(long, default_value = "ml/trained_models/ppo_extended")]
|
||||
output_dir: String,
|
||||
|
||||
/// Data directory containing DBN files
|
||||
#[arg(long, default_value = "test_data/real/databento")]
|
||||
data_dir: String,
|
||||
|
||||
/// Symbol to train on (ZN.FUT has ~29K bars)
|
||||
#[arg(long, default_value = "ZN.FUT")]
|
||||
symbol: String,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
|
||||
/// Disable early stopping (run all 100 epochs)
|
||||
#[arg(long)]
|
||||
no_early_stopping: bool,
|
||||
|
||||
/// Minimum value loss improvement percentage for plateau detection
|
||||
#[arg(long, default_value = "2.0")]
|
||||
min_value_loss_improvement: f64,
|
||||
|
||||
/// Minimum explained variance threshold
|
||||
#[arg(long, default_value = "0.4")]
|
||||
min_explained_variance: f64,
|
||||
|
||||
/// Plateau detection window size (epochs)
|
||||
#[arg(long, default_value = "30")]
|
||||
plateau_window: usize,
|
||||
|
||||
/// Run inference latency benchmark after training
|
||||
#[arg(long)]
|
||||
benchmark_inference: bool,
|
||||
|
||||
/// Number of inference iterations for benchmarking
|
||||
#[arg(long, default_value = "1000")]
|
||||
benchmark_iterations: usize,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 Agent F6: PPO Extended Training & Hyperparameter Tuning");
|
||||
info!("Objective: Improve PPO production readiness from 75% to 100%");
|
||||
info!("\n📋 Configuration:");
|
||||
info!(" • Epochs: {} (increased from 20 baseline)", opts.epochs);
|
||||
info!(
|
||||
" • Learning rate: {} (tuned for value network)",
|
||||
opts.learning_rate
|
||||
);
|
||||
info!(" • Clip epsilon: {} (PPO clip range)", opts.clip_epsilon);
|
||||
info!(
|
||||
" • Value coefficient: {} (prioritize value learning)",
|
||||
opts.value_coef
|
||||
);
|
||||
info!(
|
||||
" • Entropy coefficient: {} (exploration boost)",
|
||||
opts.entropy_coef
|
||||
);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(" • GPU: CUDA MANDATORY (no CPU fallback)");
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
info!(" • Data directory: {}", opts.data_dir);
|
||||
info!(" • Symbol: {}", opts.symbol);
|
||||
|
||||
// Early stopping configuration
|
||||
let early_stopping_enabled = !opts.no_early_stopping;
|
||||
info!(
|
||||
" • Early stopping: {}",
|
||||
if early_stopping_enabled {
|
||||
"enabled"
|
||||
} else {
|
||||
"disabled"
|
||||
}
|
||||
);
|
||||
if early_stopping_enabled {
|
||||
info!(
|
||||
" - Min value loss improvement: {}%",
|
||||
opts.min_value_loss_improvement
|
||||
);
|
||||
info!(
|
||||
" - Min explained variance: {}",
|
||||
opts.min_explained_variance
|
||||
);
|
||||
info!(" - Plateau window: {} epochs", opts.plateau_window);
|
||||
}
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Load real market data from DBN files
|
||||
info!("\n📊 Loading real market data from DBN files...");
|
||||
let mut loader = RealDataLoader::new(&opts.data_dir);
|
||||
let bars = loader
|
||||
.load_symbol_data(&opts.symbol)
|
||||
.await
|
||||
.context(format!("Failed to load data for symbol: {}", opts.symbol))?;
|
||||
|
||||
info!("✅ Loaded {} OHLCV bars for {}", bars.len(), opts.symbol);
|
||||
|
||||
// Extract features and indicators
|
||||
info!("\n🔧 Extracting features and technical indicators...");
|
||||
let features = loader
|
||||
.extract_features(&bars)
|
||||
.context("Failed to extract features")?;
|
||||
let indicators = loader
|
||||
.calculate_indicators(&bars)
|
||||
.context("Failed to calculate indicators")?;
|
||||
|
||||
info!("✅ Feature extraction complete:");
|
||||
info!(" • OHLCV bars: {}", features.prices.len());
|
||||
info!(" • Returns: {}", features.returns.len());
|
||||
info!(" • Volume: {}", features.volume.len());
|
||||
info!(" • Indicators: 10 technical indicators");
|
||||
|
||||
// Build PPO state vectors (16-feature baseline)
|
||||
info!("\n🏗️ Building PPO state vectors (16-feature baseline)...");
|
||||
let state_dim = 16; // 5 (OHLCV) + 10 (indicators) + 1 (return)
|
||||
let mut market_data = Vec::with_capacity(bars.len());
|
||||
|
||||
for i in 0..bars.len() {
|
||||
let mut state = Vec::with_capacity(state_dim);
|
||||
|
||||
// OHLCV (normalized 0-1)
|
||||
state.extend_from_slice(&features.prices[i]);
|
||||
|
||||
// Technical indicators (10 values)
|
||||
state.push(indicators.rsi[i]);
|
||||
state.push(indicators.macd[i]);
|
||||
state.push(indicators.macd_signal[i]);
|
||||
state.push(indicators.bb_upper[i]);
|
||||
state.push(indicators.bb_middle[i]);
|
||||
state.push(indicators.bb_lower[i]);
|
||||
state.push(indicators.atr[i]);
|
||||
state.push(indicators.ema_fast[i]);
|
||||
state.push(indicators.ema_slow[i]);
|
||||
state.push(indicators.volume_ma[i]);
|
||||
|
||||
// Log return
|
||||
state.push(features.returns[i]);
|
||||
|
||||
market_data.push(state);
|
||||
}
|
||||
|
||||
info!(
|
||||
"✅ Built {} state vectors (dim={})",
|
||||
market_data.len(),
|
||||
state_dim
|
||||
);
|
||||
|
||||
// Configure PPO hyperparameters with Agent F6 tuning
|
||||
let hyperparams = PpoHyperparameters {
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
gamma: 0.99,
|
||||
clip_epsilon: opts.clip_epsilon,
|
||||
vf_coef: opts.value_coef,
|
||||
ent_coef: opts.entropy_coef,
|
||||
gae_lambda: 0.95,
|
||||
rollout_steps: 2048,
|
||||
minibatch_size: opts.batch_size,
|
||||
epochs: opts.epochs,
|
||||
early_stopping_enabled,
|
||||
min_value_loss_improvement_pct: opts.min_value_loss_improvement,
|
||||
min_explained_variance: opts.min_explained_variance,
|
||||
plateau_window: opts.plateau_window,
|
||||
min_epochs_before_stopping: 50,
|
||||
};
|
||||
|
||||
info!("\n🎛️ Hyperparameter Tuning (Agent F6):");
|
||||
info!(
|
||||
" • Learning rate: {} (baseline: 0.0003)",
|
||||
hyperparams.learning_rate
|
||||
);
|
||||
info!(
|
||||
" • Clip epsilon: {} (baseline: 0.2)",
|
||||
hyperparams.clip_epsilon
|
||||
);
|
||||
info!(
|
||||
" • Value coef: {} (baseline: 0.5, +100% increase)",
|
||||
hyperparams.vf_coef
|
||||
);
|
||||
info!(
|
||||
" • Entropy coef: {} (baseline: 0.01, +400% increase)",
|
||||
hyperparams.ent_coef
|
||||
);
|
||||
info!(
|
||||
" • Epochs: {} (baseline: 20, +400% increase)",
|
||||
hyperparams.epochs
|
||||
);
|
||||
|
||||
// Create PPO trainer with real data state dimension
|
||||
let trainer = PpoTrainer::new(
|
||||
hyperparams.clone(),
|
||||
state_dim,
|
||||
&opts.output_dir,
|
||||
true, // CUDA always required
|
||||
)
|
||||
.context("Failed to create PPO trainer")?;
|
||||
|
||||
info!("✅ PPO trainer initialized (state_dim={})", state_dim);
|
||||
|
||||
// Training curve tracking
|
||||
let mut policy_losses = Vec::new();
|
||||
let mut value_losses = Vec::new();
|
||||
let mut kl_divergences = Vec::new();
|
||||
let mut explained_variances = Vec::new();
|
||||
let mut mean_rewards = Vec::new();
|
||||
let mut entropies = Vec::new();
|
||||
let mut policy_updates = 0;
|
||||
|
||||
let progress_callback = |metrics: PpoTrainingMetrics| {
|
||||
// Track policy updates (KL divergence > 0 indicates policy changed)
|
||||
if metrics.kl_divergence > 0.0 {
|
||||
policy_updates += 1;
|
||||
}
|
||||
|
||||
// Store training curves
|
||||
policy_losses.push(metrics.policy_loss);
|
||||
value_losses.push(metrics.value_loss);
|
||||
kl_divergences.push(metrics.kl_divergence);
|
||||
explained_variances.push(metrics.explained_variance);
|
||||
mean_rewards.push(metrics.mean_reward);
|
||||
entropies.push(metrics.entropy);
|
||||
|
||||
// Log progress every 10 epochs
|
||||
if metrics.epoch % 10 == 0 || metrics.epoch == 1 {
|
||||
info!(
|
||||
"📊 Epoch {}/{}: policy_loss={:.6}, value_loss={:.4}, kl_div={:.6}, expl_var={:.4}, reward={:.4}, entropy={:.4}",
|
||||
metrics.epoch,
|
||||
hyperparams.epochs,
|
||||
metrics.policy_loss,
|
||||
metrics.value_loss,
|
||||
metrics.kl_divergence,
|
||||
metrics.explained_variance,
|
||||
metrics.mean_reward,
|
||||
metrics.entropy
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
// Train the model
|
||||
info!("\n🏋️ Starting training (Agent F6 Extended Training)...\n");
|
||||
let start_time = Instant::now();
|
||||
|
||||
let final_metrics = trainer
|
||||
.train(market_data.clone(), progress_callback)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Print final metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!("\n📊 Final Metrics:");
|
||||
info!(" • Policy loss: {:.6}", final_metrics.policy_loss);
|
||||
info!(" • Value loss: {:.6}", final_metrics.value_loss);
|
||||
info!(" • KL divergence: {:.6}", final_metrics.kl_divergence);
|
||||
info!(
|
||||
" • Explained variance: {:.4}",
|
||||
final_metrics.explained_variance
|
||||
);
|
||||
info!(" • Mean reward: {:.4}", final_metrics.mean_reward);
|
||||
info!(" • Std reward: {:.4}", final_metrics.std_reward);
|
||||
info!(" • Entropy: {:.4}", final_metrics.entropy);
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
|
||||
// Analyze training curves
|
||||
info!("\n📈 Training Curve Analysis:");
|
||||
|
||||
// Policy loss trend
|
||||
let policy_loss_improvement = if policy_losses.len() > 1 {
|
||||
let initial = policy_losses.first().unwrap();
|
||||
let final_loss = policy_losses.last().unwrap();
|
||||
((initial - final_loss) / initial.abs()) * 100.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
info!(
|
||||
" • Policy loss improvement: {:.2}%",
|
||||
policy_loss_improvement
|
||||
);
|
||||
|
||||
// Value loss trend
|
||||
let value_loss_improvement = if value_losses.len() > 1 {
|
||||
let initial = value_losses.first().unwrap();
|
||||
let final_loss = value_losses.last().unwrap();
|
||||
((initial - final_loss) / initial.abs()) * 100.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
info!(" • Value loss improvement: {:.2}%", value_loss_improvement);
|
||||
|
||||
// Explained variance trend
|
||||
let expl_var_mean = explained_variances.iter().sum::<f32>() / explained_variances.len() as f32;
|
||||
let expl_var_max = explained_variances
|
||||
.iter()
|
||||
.copied()
|
||||
.fold(f32::NEG_INFINITY, f32::max);
|
||||
info!(" • Explained variance (mean): {:.4}", expl_var_mean);
|
||||
info!(" • Explained variance (max): {:.4}", expl_var_max);
|
||||
|
||||
// Reward trend
|
||||
let reward_mean = mean_rewards.iter().sum::<f32>() / mean_rewards.len() as f32;
|
||||
let reward_max = mean_rewards
|
||||
.iter()
|
||||
.copied()
|
||||
.fold(f32::NEG_INFINITY, f32::max);
|
||||
info!(" • Mean reward (avg): {:.4}", reward_mean);
|
||||
info!(" • Mean reward (max): {:.4}", reward_max);
|
||||
|
||||
// Policy convergence analysis
|
||||
info!("\n🔍 Policy Convergence Analysis:");
|
||||
info!(" • Total epochs: {}", hyperparams.epochs);
|
||||
info!(" • Policy updates (KL > 0): {}", policy_updates);
|
||||
info!(
|
||||
" • Policy update rate: {:.1}%",
|
||||
(policy_updates as f64 / hyperparams.epochs as f64) * 100.0
|
||||
);
|
||||
|
||||
// KL divergence statistics
|
||||
let kl_mean = kl_divergences.iter().sum::<f32>() / kl_divergences.len() as f32;
|
||||
let kl_max = kl_divergences
|
||||
.iter()
|
||||
.copied()
|
||||
.fold(f32::NEG_INFINITY, f32::max);
|
||||
info!(" • KL divergence (mean): {:.6}", kl_mean);
|
||||
info!(" • KL divergence (max): {:.6}", kl_max);
|
||||
|
||||
// Validation checks
|
||||
let mut passed_checks = 0;
|
||||
let mut total_checks = 0;
|
||||
|
||||
total_checks += 1;
|
||||
if final_metrics.kl_divergence > 0.0 || policy_updates > 0 {
|
||||
info!(" ✅ PASS: Policy updates detected");
|
||||
passed_checks += 1;
|
||||
} else {
|
||||
warn!(" ⚠️ WARN: No policy updates (may indicate convergence)");
|
||||
}
|
||||
|
||||
total_checks += 1;
|
||||
if final_metrics.explained_variance > 0.5 {
|
||||
info!(" ✅ PASS: Value network learning (explained variance > 0.5)");
|
||||
passed_checks += 1;
|
||||
} else if final_metrics.explained_variance > 0.0 {
|
||||
warn!(
|
||||
" ⚠️ WARN: Value network below target (explained variance = {:.4})",
|
||||
final_metrics.explained_variance
|
||||
);
|
||||
} else {
|
||||
warn!(
|
||||
" ❌ FAIL: Value network not learning (explained variance = {:.4})",
|
||||
final_metrics.explained_variance
|
||||
);
|
||||
}
|
||||
|
||||
total_checks += 1;
|
||||
if value_loss_improvement > 0.0 {
|
||||
info!(
|
||||
" ✅ PASS: Value loss improved by {:.2}%",
|
||||
value_loss_improvement
|
||||
);
|
||||
passed_checks += 1;
|
||||
} else {
|
||||
warn!(" ⚠️ WARN: Value loss did not improve");
|
||||
}
|
||||
|
||||
// Note: Inference latency benchmarking is not available in PpoTrainer
|
||||
// PPO inference latency is estimated at ~320μs based on previous benchmarks
|
||||
info!("\n⏱️ Inference Latency Estimate:");
|
||||
info!(" • Estimated latency: ~320μs (from previous benchmarks)");
|
||||
info!(" • Target: <500μs");
|
||||
info!(" ✅ PASS: Estimated within target");
|
||||
|
||||
// Final checkpoint
|
||||
let final_checkpoint = output_path.join(format!(
|
||||
"ppo_checkpoint_epoch_{}.safetensors",
|
||||
hyperparams.epochs
|
||||
));
|
||||
info!(
|
||||
"\n💾 Final checkpoint saved to: {}",
|
||||
final_checkpoint.display()
|
||||
);
|
||||
|
||||
// Agent F6 Summary
|
||||
info!("\n🎉 Agent F6: PPO Extended Training Complete!");
|
||||
info!("\n📋 Summary:");
|
||||
info!(
|
||||
" • Training epochs: {} (vs. 20 baseline, +400%)",
|
||||
hyperparams.epochs
|
||||
);
|
||||
info!(
|
||||
" • Training time: {:.1} min (vs. 3.0 min baseline)",
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
info!(
|
||||
" • Policy loss improvement: {:.2}%",
|
||||
policy_loss_improvement
|
||||
);
|
||||
info!(" • Value loss improvement: {:.2}%", value_loss_improvement);
|
||||
info!(
|
||||
" • Explained variance: {:.4} (baseline: -0.69)",
|
||||
final_metrics.explained_variance
|
||||
);
|
||||
info!(
|
||||
" • Mean reward: {:.4} (baseline: -0.0002)",
|
||||
final_metrics.mean_reward
|
||||
);
|
||||
info!(
|
||||
" • Validation checks: {}/{} passed",
|
||||
passed_checks, total_checks
|
||||
);
|
||||
|
||||
info!("\n📁 Model files saved to: {}", opts.output_dir);
|
||||
info!("\n🎯 Production Readiness Assessment:");
|
||||
|
||||
let production_ready_pct = (passed_checks as f64 / total_checks as f64) * 100.0;
|
||||
if production_ready_pct >= 75.0 {
|
||||
info!(
|
||||
" ✅ READY: {:.0}% of validation checks passed",
|
||||
production_ready_pct
|
||||
);
|
||||
} else {
|
||||
warn!(
|
||||
" ⚠️ NOT READY: {:.0}% of validation checks passed",
|
||||
production_ready_pct
|
||||
);
|
||||
}
|
||||
|
||||
info!("\n📝 Recommendations:");
|
||||
if final_metrics.explained_variance < 0.5 {
|
||||
info!(
|
||||
" • Consider further tuning value coefficient (current: {})",
|
||||
hyperparams.vf_coef
|
||||
);
|
||||
}
|
||||
if final_metrics.mean_reward < 0.0 {
|
||||
info!(" • Negative rewards suggest 54-feature retraining is critical");
|
||||
}
|
||||
info!(" • Next step: Retrain with 54-feature set for +25-50% Sharpe improvement");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,456 +0,0 @@
|
||||
//! PPO Training Example with Parquet Data
|
||||
//!
|
||||
//! Trains a PPO model on market data from Parquet files with:
|
||||
//! - Real OHLCV data + 51-dimensional features
|
||||
//! - Actual PnL-based rewards
|
||||
//! - GAE advantages on real price trajectories
|
||||
//! - Policy convergence validation (KL divergence > 0)
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (30 epochs, hyperopt-optimized learning rates)
|
||||
//! cargo run -p ml --example train_ppo_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/ZN_FUT_90d_clean.parquet
|
||||
//!
|
||||
//! # Custom epochs, batch size, and learning rates
|
||||
//! cargo run -p ml --example train_ppo_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/ZN_FUT_90d_clean.parquet \
|
||||
//! --epochs 50 \
|
||||
//! --batch-size 128 \
|
||||
//! --policy-lr 0.000001 \
|
||||
//! --value-lr 0.001
|
||||
//!
|
||||
//! # With early stopping disabled
|
||||
//! cargo run -p ml --example train_ppo_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/NQ_FUT_180d.parquet \
|
||||
//! --no-early-stopping
|
||||
//! ```
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::fs::File;
|
||||
use std::path::PathBuf;
|
||||
use tracing::{info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use arrow::array::{Array, Float64Array, PrimitiveArray, UInt64Array};
|
||||
use arrow::datatypes::TimestampNanosecondType;
|
||||
use arrow::record_batch::RecordBatch;
|
||||
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
|
||||
|
||||
use ml::features::extraction::{extract_ml_features, OHLCVBar};
|
||||
use ml::trainers::ppo::{PpoHyperparameters, PpoTrainer, PpoTrainingMetrics};
|
||||
|
||||
/// Train PPO model on Parquet market data
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_ppo_parquet",
|
||||
about = "Train PPO model on Parquet market data"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Path to Parquet file with market data
|
||||
#[arg(long)]
|
||||
parquet_file: String,
|
||||
|
||||
/// Number of training epochs (default: 30 for policy convergence)
|
||||
#[arg(long, default_value = "30")]
|
||||
epochs: usize,
|
||||
|
||||
/// Policy (actor) learning rate (default: 1e-6, ultra-conservative for stability)
|
||||
#[arg(long, default_value = "0.000001")]
|
||||
policy_lr: f64,
|
||||
|
||||
/// Value (critic) learning rate (default: 0.001, aggressive for faster convergence)
|
||||
#[arg(long, default_value = "0.001")]
|
||||
value_lr: f64,
|
||||
|
||||
/// Batch size (max 230 for RTX 3050 Ti 4GB)
|
||||
#[arg(long, default_value = "64")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Output directory for trained model
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
|
||||
/// Enable early stopping (recommended, use --no-early-stopping to disable)
|
||||
#[arg(long)]
|
||||
early_stopping: bool,
|
||||
|
||||
/// Disable early stopping
|
||||
#[arg(long)]
|
||||
no_early_stopping: bool,
|
||||
|
||||
/// Minimum value loss improvement percentage for plateau detection
|
||||
#[arg(long, default_value = "2.0")]
|
||||
min_value_loss_improvement: f64,
|
||||
|
||||
/// Minimum explained variance threshold
|
||||
#[arg(long, default_value = "0.4")]
|
||||
min_explained_variance: f64,
|
||||
|
||||
/// Plateau detection window size (epochs)
|
||||
#[arg(long, default_value = "30")]
|
||||
plateau_window: usize,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 Starting PPO Training with Parquet Data");
|
||||
info!("Configuration:");
|
||||
info!(" • Parquet file: {}", opts.parquet_file);
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Policy learning rate: {}", opts.policy_lr);
|
||||
info!(" • Value learning rate: {}", opts.value_lr);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(" • GPU: CUDA if available (auto-fallback to CPU)");
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
|
||||
// Determine early stopping (enabled by default, unless --no-early-stopping is specified)
|
||||
let early_stopping_enabled = !opts.no_early_stopping;
|
||||
info!(
|
||||
" • Early stopping: {}",
|
||||
if early_stopping_enabled {
|
||||
"enabled"
|
||||
} else {
|
||||
"disabled"
|
||||
}
|
||||
);
|
||||
if early_stopping_enabled {
|
||||
info!(
|
||||
" - Min value loss improvement: {}%",
|
||||
opts.min_value_loss_improvement
|
||||
);
|
||||
info!(
|
||||
" - Min explained variance: {}",
|
||||
opts.min_explained_variance
|
||||
);
|
||||
info!(" - Plateau window: {} epochs", opts.plateau_window);
|
||||
}
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Load market data from Parquet file
|
||||
info!("\n📊 Loading market data from Parquet file...");
|
||||
let bars = load_parquet_data(&opts.parquet_file)
|
||||
.await
|
||||
.context("Failed to load Parquet data")?;
|
||||
|
||||
info!("✅ Loaded {} OHLCV bars", bars.len());
|
||||
|
||||
// Extract 51-dimensional feature vectors (43 base + 8 OFI)
|
||||
info!("\n🏗️ Extracting 51-dimensional feature vectors...");
|
||||
let feature_vectors =
|
||||
extract_ml_features(&bars).context("Failed to extract 51-dimensional features")?;
|
||||
|
||||
info!(
|
||||
"✅ Extracted {} feature vectors (dim=51, warmup bars skipped=50)",
|
||||
feature_vectors.len()
|
||||
);
|
||||
|
||||
// Convert FeatureVector ([f64; 51]) to Vec<Vec<f32>> for PPO trainer
|
||||
let state_dim = 51;
|
||||
let market_data: Vec<Vec<f32>> = feature_vectors
|
||||
.iter()
|
||||
.map(|fv| fv.iter().map(|&v| v as f32).collect())
|
||||
.collect();
|
||||
|
||||
// Validate state dimensions
|
||||
if let Some(first_state) = market_data.first() {
|
||||
if first_state.len() != state_dim {
|
||||
return Err(anyhow::anyhow!(
|
||||
"State dimension mismatch: expected {}, got {}",
|
||||
state_dim,
|
||||
first_state.len()
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
info!(
|
||||
"✅ Feature extraction complete: {} samples",
|
||||
market_data.len()
|
||||
);
|
||||
|
||||
// Configure PPO hyperparameters
|
||||
let hyperparams = PpoHyperparameters {
|
||||
learning_rate: 1e-4, // Deprecated field, kept for backward compatibility
|
||||
actor_learning_rate: Some(opts.policy_lr),
|
||||
critic_learning_rate: Some(opts.value_lr),
|
||||
batch_size: opts.batch_size,
|
||||
gamma: 0.99,
|
||||
clip_epsilon: 0.2,
|
||||
vf_coef: 0.5,
|
||||
ent_coef: 0.01,
|
||||
gae_lambda: 0.95,
|
||||
rollout_steps: 2048,
|
||||
minibatch_size: opts.batch_size,
|
||||
epochs: opts.epochs,
|
||||
early_stopping_enabled,
|
||||
min_value_loss_improvement_pct: opts.min_value_loss_improvement,
|
||||
min_explained_variance: opts.min_explained_variance,
|
||||
plateau_window: opts.plateau_window,
|
||||
min_epochs_before_stopping: 50,
|
||||
..PpoHyperparameters::conservative()
|
||||
};
|
||||
|
||||
// Create PPO trainer
|
||||
let trainer = PpoTrainer::new(
|
||||
hyperparams.clone(),
|
||||
state_dim,
|
||||
&opts.output_dir,
|
||||
true, // Use GPU if available
|
||||
None, // Single environment (standard mode)
|
||||
)
|
||||
.context("Failed to create PPO trainer")?;
|
||||
|
||||
info!("✅ PPO trainer initialized (state_dim={})", state_dim);
|
||||
|
||||
// Create progress callback with convergence tracking
|
||||
let mut policy_updates = 0;
|
||||
let mut kl_divergence_history = Vec::new();
|
||||
|
||||
let progress_callback = |metrics: PpoTrainingMetrics| {
|
||||
// Track policy updates (KL divergence > 0 indicates policy changed)
|
||||
if metrics.kl_divergence > 0.0 {
|
||||
policy_updates += 1;
|
||||
}
|
||||
kl_divergence_history.push(metrics.kl_divergence);
|
||||
|
||||
info!(
|
||||
"📊 Epoch {}/{}: policy_loss={:.4}, value_loss={:.4}, kl_div={:.6}, expl_var={:.4}, mean_reward={:.4}",
|
||||
metrics.epoch,
|
||||
hyperparams.epochs,
|
||||
metrics.policy_loss,
|
||||
metrics.value_loss,
|
||||
metrics.kl_divergence,
|
||||
metrics.explained_variance,
|
||||
metrics.mean_reward
|
||||
);
|
||||
};
|
||||
|
||||
// Train the model
|
||||
info!("\n🏋️ Starting training...\n");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let final_metrics = trainer
|
||||
.train(market_data, progress_callback)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Print final metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!("\n📊 Final Metrics:");
|
||||
info!(" • Policy loss: {:.6}", final_metrics.policy_loss);
|
||||
info!(" • Value loss: {:.6}", final_metrics.value_loss);
|
||||
info!(" • KL divergence: {:.6}", final_metrics.kl_divergence);
|
||||
info!(
|
||||
" • Explained variance: {:.4}",
|
||||
final_metrics.explained_variance
|
||||
);
|
||||
info!(" • Mean reward: {:.4}", final_metrics.mean_reward);
|
||||
info!(" • Std reward: {:.4}", final_metrics.std_reward);
|
||||
info!(" • Entropy: {:.4}", final_metrics.entropy);
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
|
||||
// Validate policy convergence
|
||||
info!("\n🔍 Policy Convergence Analysis:");
|
||||
info!(" • Total epochs: {}", hyperparams.epochs);
|
||||
info!(" • Policy updates (KL > 0): {}", policy_updates);
|
||||
info!(
|
||||
" • Policy update rate: {:.1}%",
|
||||
(policy_updates as f64 / hyperparams.epochs as f64) * 100.0
|
||||
);
|
||||
|
||||
// Calculate KL divergence statistics
|
||||
let kl_mean = kl_divergence_history.iter().sum::<f32>() / kl_divergence_history.len() as f32;
|
||||
let kl_max = kl_divergence_history
|
||||
.iter()
|
||||
.copied()
|
||||
.fold(f32::NEG_INFINITY, f32::max);
|
||||
let kl_min = kl_divergence_history
|
||||
.iter()
|
||||
.copied()
|
||||
.fold(f32::INFINITY, f32::min);
|
||||
|
||||
info!(" • KL divergence (mean): {:.6}", kl_mean);
|
||||
info!(" • KL divergence (max): {:.6}", kl_max);
|
||||
info!(" • KL divergence (min): {:.6}", kl_min);
|
||||
|
||||
// Convergence validation
|
||||
if final_metrics.kl_divergence > 0.0 {
|
||||
info!(" ✅ PASS: Policy updates detected (KL divergence > 0)");
|
||||
} else {
|
||||
warn!(" ⚠️ WARN: No policy updates in final epoch (KL divergence = 0)");
|
||||
warn!(" This may indicate learning rate too low or convergence");
|
||||
}
|
||||
|
||||
// Value function validation
|
||||
if final_metrics.explained_variance > 0.5 {
|
||||
info!(" ✅ PASS: Value network learning (explained variance > 0.5)");
|
||||
} else {
|
||||
warn!(" ⚠️ WARN: Value network may need tuning (explained variance < 0.5)");
|
||||
}
|
||||
|
||||
// Checkpoint is already saved by trainer (every 10 epochs)
|
||||
let final_checkpoint = output_path.join(format!(
|
||||
"ppo_checkpoint_epoch_{}.safetensors",
|
||||
hyperparams.epochs
|
||||
));
|
||||
|
||||
info!(
|
||||
"\n💾 Final checkpoint saved to: {}",
|
||||
final_checkpoint.display()
|
||||
);
|
||||
info!("\n🎉 PPO training complete with Parquet data!");
|
||||
info!("📁 Model files saved to: {}", opts.output_dir);
|
||||
|
||||
info!("\n📈 Training Summary:");
|
||||
info!(" • Data source: Parquet file ({})", opts.parquet_file);
|
||||
info!(" • Training samples: {}", bars.len());
|
||||
info!(
|
||||
" • Feature samples: {} (after warmup)",
|
||||
feature_vectors.len()
|
||||
);
|
||||
info!(" • State dimension: {}", state_dim);
|
||||
info!(" • Features: 51-dimensional");
|
||||
info!(
|
||||
" • Policy updates: {}/{} epochs ({:.1}%)",
|
||||
policy_updates,
|
||||
hyperparams.epochs,
|
||||
(policy_updates as f64 / hyperparams.epochs as f64) * 100.0
|
||||
);
|
||||
info!(
|
||||
" • Convergence: {}",
|
||||
if final_metrics.kl_divergence > 0.0 {
|
||||
"✅ Achieved"
|
||||
} else {
|
||||
"⚠️ Check logs"
|
||||
}
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Load OHLCV data from Parquet file (Databento schema)
|
||||
async fn load_parquet_data(parquet_path: &str) -> Result<Vec<OHLCVBar>> {
|
||||
info!("Loading Parquet file: {}", parquet_path);
|
||||
|
||||
// Open Parquet file
|
||||
let file = File::open(parquet_path)
|
||||
.with_context(|| format!("Failed to open Parquet file: {}", parquet_path))?;
|
||||
|
||||
// Create Parquet reader
|
||||
let builder = ParquetRecordBatchReaderBuilder::try_new(file)
|
||||
.with_context(|| "Failed to create Parquet reader")?;
|
||||
|
||||
let reader = builder
|
||||
.build()
|
||||
.with_context(|| "Failed to build Parquet reader")?;
|
||||
|
||||
// Read all batches
|
||||
let mut all_ohlcv_bars = Vec::new();
|
||||
|
||||
for batch_result in reader {
|
||||
let batch: RecordBatch = batch_result.with_context(|| "Failed to read record batch")?;
|
||||
|
||||
// Extract columns from Databento Parquet schema:
|
||||
// Column 3: open, Column 4: high, Column 5: low, Column 6: close
|
||||
// Column 7: volume, Column 9: ts_event (Timestamp(Nanosecond, Some("UTC")))
|
||||
let timestamps = batch
|
||||
.column(9)
|
||||
.as_any()
|
||||
.downcast_ref::<PrimitiveArray<TimestampNanosecondType>>()
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!(
|
||||
"Failed to downcast timestamp column. Expected Timestamp(Nanosecond), got: {:?}",
|
||||
batch.column(9).data_type()
|
||||
)
|
||||
})?;
|
||||
|
||||
let opens = batch
|
||||
.column(3)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast open column"))?;
|
||||
|
||||
let highs = batch
|
||||
.column(4)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast high column"))?;
|
||||
|
||||
let lows = batch
|
||||
.column(5)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast low column"))?;
|
||||
|
||||
let closes = batch
|
||||
.column(6)
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast close column"))?;
|
||||
|
||||
let volumes = batch
|
||||
.column(7)
|
||||
.as_any()
|
||||
.downcast_ref::<UInt64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast volume column"))?;
|
||||
|
||||
// Convert to OHLCVBar structs
|
||||
for i in 0..batch.num_rows() {
|
||||
let timestamp_ns = timestamps.value(i);
|
||||
|
||||
// Convert nanoseconds to DateTime<Utc>
|
||||
let timestamp = chrono::DateTime::from_timestamp(
|
||||
(timestamp_ns / 1_000_000_000) as i64,
|
||||
(timestamp_ns % 1_000_000_000) as u32,
|
||||
)
|
||||
.unwrap_or_else(|| chrono::Utc::now());
|
||||
|
||||
let bar = OHLCVBar {
|
||||
timestamp,
|
||||
open: opens.value(i),
|
||||
high: highs.value(i),
|
||||
low: lows.value(i),
|
||||
close: closes.value(i),
|
||||
volume: volumes.value(i) as f64,
|
||||
};
|
||||
|
||||
all_ohlcv_bars.push(bar);
|
||||
}
|
||||
}
|
||||
|
||||
info!("✅ Loaded {} OHLCV bars from Parquet", all_ohlcv_bars.len());
|
||||
|
||||
Ok(all_ohlcv_bars)
|
||||
}
|
||||
@@ -1,898 +0,0 @@
|
||||
//! Rainbow DQN Training Example
|
||||
//!
|
||||
//! Trains a Rainbow DQN model on market data using all 6 components:
|
||||
//! 1. Double Q-learning, 2. Dueling Networks, 3. Prioritized Experience Replay,
|
||||
//! 4. Multi-step Learning, 5. Distributional RL (C51), 6. Noisy Networks
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (100 epochs)
|
||||
//! cargo run -p ml --example train_rainbow --release --features cuda
|
||||
//!
|
||||
//! # Custom epochs and output path
|
||||
//! cargo run -p ml --example train_rainbow --release --features cuda -- \
|
||||
//! --epochs 500 \
|
||||
//! --output ml/trained_models/rainbow_model.safetensors
|
||||
//!
|
||||
//! # Custom parameters (C51 distributional)
|
||||
//! cargo run -p ml --example train_rainbow --release --features cuda -- \
|
||||
//! --num-atoms 51 \
|
||||
//! --v-min -10.0 \
|
||||
//! --v-max 10.0 \
|
||||
//! --n-step 3 \
|
||||
//! --priority-alpha 0.6 \
|
||||
//! --priority-beta 0.4
|
||||
//! ```
|
||||
|
||||
// Use mimalloc allocator for 10-25% performance improvement
|
||||
use mimalloc::MiMalloc;
|
||||
#[global_allocator]
|
||||
static GLOBAL: MiMalloc = MiMalloc;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use tokio::signal;
|
||||
use tracing::{info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
// Use full Rainbow DQN implementation (not stub)
|
||||
use ml::checkpoint::{CheckpointConfig, CheckpointManager};
|
||||
use ml::data_loaders::BarSamplingMethod;
|
||||
use ml::dqn::distributional::DistributionalConfig;
|
||||
use ml::dqn::multi_step::MultiStepConfig;
|
||||
use ml::dqn::rainbow_agent_impl::RainbowAgent;
|
||||
use ml::dqn::rainbow_config::RainbowAgentConfig;
|
||||
use ml::dqn::rainbow_network::RainbowNetworkConfig;
|
||||
use ml::features::extraction::OHLCVBar;
|
||||
|
||||
// Feature vector type: 128 features (125 market + 3 portfolio placeholders)
|
||||
type FeatureVector128 = [f64; 128];
|
||||
|
||||
/// Train Rainbow DQN model on market data
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_rainbow",
|
||||
about = "Train Rainbow DQN model on market data"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "100")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate (conservative for Rainbow)
|
||||
#[arg(long, default_value = "0.0001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Batch size (max 230 for RTX 3050 Ti 4GB)
|
||||
#[arg(long, default_value = "32")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Discount factor (gamma)
|
||||
#[arg(long, default_value = "0.99")]
|
||||
gamma: f64,
|
||||
|
||||
/// Checkpoint save frequency (epochs)
|
||||
#[arg(long, default_value = "10")]
|
||||
checkpoint_frequency: usize,
|
||||
|
||||
/// Output directory for trained model
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Data directory containing DBN files
|
||||
#[arg(long, default_value = "test_data/real/databento/ml_training")]
|
||||
data_dir: String,
|
||||
|
||||
/// Parquet file path (overrides data_dir if specified)
|
||||
#[arg(long)]
|
||||
parquet_file: Option<String>,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
|
||||
/// Replay buffer capacity
|
||||
#[arg(long, default_value = "100000")]
|
||||
buffer_size: usize,
|
||||
|
||||
/// Minimum replay buffer size before training starts
|
||||
#[arg(long, default_value = "10000")]
|
||||
min_replay_size: usize,
|
||||
|
||||
/// Checkpoint directory (overrides output_dir for checkpoints)
|
||||
#[arg(long)]
|
||||
checkpoint_dir: Option<String>,
|
||||
|
||||
/// Alternative bar sampling method (time, tick, volume, dollar, imbalance, run)
|
||||
#[arg(long, default_value = "time")]
|
||||
bar_method: String,
|
||||
|
||||
/// Bar threshold for alternative sampling methods
|
||||
#[arg(long)]
|
||||
bar_threshold: Option<f64>,
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// RAINBOW-SPECIFIC PARAMETERS (No epsilon - uses noisy networks instead)
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
/// Number of atoms for C51 distributional RL (default: 51)
|
||||
/// Higher = more accurate distribution approximation but more memory
|
||||
#[arg(long, default_value = "51")]
|
||||
num_atoms: usize,
|
||||
|
||||
/// Minimum value of support for C51 distribution (default: -10.0)
|
||||
/// Should be lower than expected minimum return
|
||||
#[arg(long, default_value = "-10.0")]
|
||||
v_min: f64,
|
||||
|
||||
/// Maximum value of support for C51 distribution (default: 10.0)
|
||||
/// Should be higher than expected maximum return
|
||||
#[arg(long, default_value = "10.0")]
|
||||
v_max: f64,
|
||||
|
||||
/// N-step for multi-step learning (default: 3)
|
||||
/// Higher = faster credit assignment but more bias
|
||||
#[arg(long, default_value = "3")]
|
||||
n_step: usize,
|
||||
|
||||
/// Priority replay alpha (default: 0.6)
|
||||
/// 0 = uniform sampling, 1 = full prioritization
|
||||
#[arg(long, default_value = "0.6")]
|
||||
priority_alpha: f64,
|
||||
|
||||
/// Priority replay beta (default: 0.4, anneals to 1.0)
|
||||
/// Importance sampling correction strength
|
||||
#[arg(long, default_value = "0.4")]
|
||||
priority_beta: f64,
|
||||
|
||||
/// Priority replay beta increment per step (default: 0.00025)
|
||||
#[arg(long, default_value = "0.00025")]
|
||||
priority_beta_increment: f64,
|
||||
|
||||
/// Noisy network sigma (default: 0.5)
|
||||
/// Controls exploration via parameter noise
|
||||
#[arg(long, default_value = "0.5")]
|
||||
noisy_sigma: f64,
|
||||
|
||||
/// Target network update frequency (steps)
|
||||
#[arg(long, default_value = "1000")]
|
||||
target_update_freq: usize,
|
||||
|
||||
/// Training frequency (steps between training)
|
||||
#[arg(long, default_value = "4")]
|
||||
train_freq: usize,
|
||||
|
||||
/// Noisy network noise reset frequency (steps)
|
||||
#[arg(long, default_value = "100")]
|
||||
noise_reset_freq: usize,
|
||||
|
||||
/// Input state dimension (default: 128 for DQN features)
|
||||
#[arg(long, default_value = "128")]
|
||||
state_dim: usize,
|
||||
|
||||
/// Number of actions (default: 3 for BUY/SELL/HOLD)
|
||||
#[arg(long, default_value = "3")]
|
||||
num_actions: usize,
|
||||
|
||||
/// Hidden layer sizes (comma-separated, default: 512,512)
|
||||
#[arg(long, default_value = "512,512")]
|
||||
hidden_sizes: String,
|
||||
}
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// HELPER FUNCTIONS
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
/// Load training data from Parquet file
|
||||
/// Returns vector of (features, [current_close, next_close]) tuples
|
||||
async fn load_training_data_from_parquet(
|
||||
parquet_path: &str,
|
||||
) -> Result<Vec<(FeatureVector128, Vec<f64>)>> {
|
||||
use arrow::array::{Array, Float64Array, PrimitiveArray, UInt64Array};
|
||||
use arrow::datatypes::TimestampNanosecondType;
|
||||
use arrow::record_batch::RecordBatch;
|
||||
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
|
||||
use std::fs::File;
|
||||
|
||||
info!("Loading Parquet file: {}", parquet_path);
|
||||
|
||||
// Open Parquet file
|
||||
let file = File::open(parquet_path)
|
||||
.with_context(|| format!("Failed to open Parquet file: {}", parquet_path))?;
|
||||
|
||||
// Create Parquet reader
|
||||
let builder = ParquetRecordBatchReaderBuilder::try_new(file)
|
||||
.with_context(|| "Failed to create Parquet reader")?;
|
||||
|
||||
let reader = builder
|
||||
.build()
|
||||
.with_context(|| "Failed to build Parquet reader")?;
|
||||
|
||||
// Read all batches
|
||||
let mut all_ohlcv_bars = Vec::new();
|
||||
|
||||
for batch_result in reader {
|
||||
let batch: RecordBatch = batch_result.with_context(|| "Failed to read record batch")?;
|
||||
|
||||
// Extract timestamp column
|
||||
let timestamp_col = batch
|
||||
.column_by_name("timestamp_ns")
|
||||
.or_else(|| batch.column_by_name("ts_event"))
|
||||
.ok_or_else(|| {
|
||||
anyhow::anyhow!("Missing timestamp column. Expected 'timestamp_ns' or 'ts_event'")
|
||||
})?;
|
||||
|
||||
let timestamps = timestamp_col
|
||||
.as_any()
|
||||
.downcast_ref::<PrimitiveArray<TimestampNanosecondType>>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Failed to downcast timestamp column"))?;
|
||||
|
||||
// Extract OHLCV columns
|
||||
let opens = batch
|
||||
.column_by_name("open")
|
||||
.ok_or_else(|| anyhow::anyhow!("Missing 'open' column"))?
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid 'open' column type"))?;
|
||||
|
||||
let highs = batch
|
||||
.column_by_name("high")
|
||||
.ok_or_else(|| anyhow::anyhow!("Missing 'high' column"))?
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid 'high' column type"))?;
|
||||
|
||||
let lows = batch
|
||||
.column_by_name("low")
|
||||
.ok_or_else(|| anyhow::anyhow!("Missing 'low' column"))?
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid 'low' column type"))?;
|
||||
|
||||
let closes = batch
|
||||
.column_by_name("close")
|
||||
.ok_or_else(|| anyhow::anyhow!("Missing 'close' column"))?
|
||||
.as_any()
|
||||
.downcast_ref::<Float64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid 'close' column type"))?;
|
||||
|
||||
let volumes = batch
|
||||
.column_by_name("volume")
|
||||
.ok_or_else(|| anyhow::anyhow!("Missing 'volume' column"))?
|
||||
.as_any()
|
||||
.downcast_ref::<UInt64Array>()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid 'volume' column type"))?;
|
||||
|
||||
// Convert to OHLCVBar structs
|
||||
for i in 0..batch.num_rows() {
|
||||
let timestamp_ns = timestamps.value(i);
|
||||
let timestamp = chrono::DateTime::from_timestamp_nanos(timestamp_ns);
|
||||
|
||||
let bar = OHLCVBar {
|
||||
timestamp,
|
||||
open: opens.value(i),
|
||||
high: highs.value(i),
|
||||
low: lows.value(i),
|
||||
close: closes.value(i),
|
||||
volume: volumes.value(i) as f64,
|
||||
};
|
||||
all_ohlcv_bars.push(bar);
|
||||
}
|
||||
}
|
||||
|
||||
info!(
|
||||
"Successfully loaded {} OHLCV bars from Parquet",
|
||||
all_ohlcv_bars.len()
|
||||
);
|
||||
|
||||
// Sort bars chronologically
|
||||
all_ohlcv_bars.sort_by_key(|bar| bar.timestamp);
|
||||
|
||||
// Extract features
|
||||
let feature_vectors = extract_features(&all_ohlcv_bars)?;
|
||||
info!(
|
||||
"Extracted {} feature vectors (128 dimensions)",
|
||||
feature_vectors.len()
|
||||
);
|
||||
|
||||
// Create training data pairs (features, [current_close, next_close])
|
||||
let mut training_data = Vec::new();
|
||||
for i in 0..feature_vectors.len().saturating_sub(1) {
|
||||
let current_close = all_ohlcv_bars[i + 50].close; // +50 for warmup
|
||||
let next_close = all_ohlcv_bars[i + 51].close;
|
||||
training_data.push((feature_vectors[i], vec![current_close, next_close]));
|
||||
}
|
||||
|
||||
// Last sample targets itself
|
||||
if !feature_vectors.is_empty() {
|
||||
let idx = all_ohlcv_bars.len() - 1;
|
||||
let current_close = all_ohlcv_bars[idx].close;
|
||||
training_data.push((
|
||||
feature_vectors[feature_vectors.len() - 1],
|
||||
vec![current_close, current_close],
|
||||
));
|
||||
}
|
||||
|
||||
info!("Created {} training samples", training_data.len());
|
||||
|
||||
Ok(training_data)
|
||||
}
|
||||
|
||||
/// Extract 128-dim features from OHLCV bars
|
||||
fn extract_features(bars: &[OHLCVBar]) -> Result<Vec<FeatureVector128>> {
|
||||
use ml::features::extraction::FeatureExtractor;
|
||||
|
||||
if bars.is_empty() {
|
||||
anyhow::bail!("Cannot extract features from empty bar sequence");
|
||||
}
|
||||
|
||||
const WARMUP_PERIOD: usize = 50;
|
||||
if bars.len() < WARMUP_PERIOD {
|
||||
anyhow::bail!(
|
||||
"Insufficient data: {} bars provided, {} required for warmup",
|
||||
bars.len(),
|
||||
WARMUP_PERIOD
|
||||
);
|
||||
}
|
||||
|
||||
let mut extractor = FeatureExtractor::new();
|
||||
let mut feature_vectors = Vec::with_capacity(bars.len() - WARMUP_PERIOD);
|
||||
|
||||
// Feed bars sequentially to build rolling windows
|
||||
for (i, bar) in bars.iter().enumerate() {
|
||||
extractor.update(bar)?;
|
||||
|
||||
// Start extracting features after warmup
|
||||
if i >= WARMUP_PERIOD {
|
||||
// Extract 54 features and reduce to 125 market features
|
||||
let features_225 = extractor.extract_current_features()?;
|
||||
|
||||
// Take first 125 features
|
||||
let mut features_125 = [0.0; 125];
|
||||
features_125.copy_from_slice(&features_225[0..125]);
|
||||
|
||||
// Convert to 128-dim (125 market + 3 portfolio placeholder zeros)
|
||||
let mut features_128 = [0.0; 128];
|
||||
features_128[0..125].copy_from_slice(&features_125);
|
||||
// features_128[125..128] remain as zeros (portfolio placeholders)
|
||||
|
||||
feature_vectors.push(features_128);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(feature_vectors)
|
||||
}
|
||||
|
||||
/// Convert feature vector to state representation (Vec<f32> for Rainbow agent)
|
||||
fn feature_vector_to_state(feature_vec: &FeatureVector128) -> Vec<f32> {
|
||||
feature_vec.iter().map(|&v| v as f32).collect()
|
||||
}
|
||||
|
||||
/// Simple trading environment for Rainbow DQN
|
||||
struct TradingEnvironment {
|
||||
position: f32, // Current position (-1.0 to +1.0)
|
||||
portfolio_value: f32, // Current portfolio value
|
||||
last_price: f32, // Last observed price
|
||||
initial_value: f32, // Initial portfolio value
|
||||
}
|
||||
|
||||
impl TradingEnvironment {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
position: 0.0,
|
||||
portfolio_value: 10000.0, // Start with $10,000
|
||||
last_price: 0.0,
|
||||
initial_value: 10000.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute action and return reward
|
||||
/// action: 0=BUY, 1=SELL, 2=HOLD
|
||||
fn step(&mut self, action: usize, current_price: f32, next_price: f32) -> f32 {
|
||||
// Update last price
|
||||
if self.last_price == 0.0 {
|
||||
self.last_price = current_price;
|
||||
}
|
||||
|
||||
// Calculate price change
|
||||
let price_change = next_price - current_price;
|
||||
let price_change_pct = price_change / current_price;
|
||||
|
||||
// Execute action and calculate reward
|
||||
let reward = match action {
|
||||
0 => {
|
||||
// BUY: Go long (or add to long position)
|
||||
let _old_position = self.position;
|
||||
self.position = (self.position + 0.5).min(1.0); // Add 0.5, cap at 1.0
|
||||
|
||||
// Reward is P&L from position
|
||||
let pnl = self.position * price_change_pct * self.portfolio_value;
|
||||
self.portfolio_value += pnl;
|
||||
|
||||
// Return normalized reward
|
||||
pnl / 100.0 // Scale to reasonable range
|
||||
},
|
||||
1 => {
|
||||
// SELL: Go short (or add to short position)
|
||||
let _old_position = self.position;
|
||||
self.position = (self.position - 0.5).max(-1.0); // Subtract 0.5, floor at -1.0
|
||||
|
||||
// Reward is P&L from position
|
||||
let pnl = self.position * price_change_pct * self.portfolio_value;
|
||||
self.portfolio_value += pnl;
|
||||
|
||||
// Return normalized reward
|
||||
pnl / 100.0
|
||||
},
|
||||
2 => {
|
||||
// HOLD: Maintain current position
|
||||
let pnl = self.position * price_change_pct * self.portfolio_value;
|
||||
self.portfolio_value += pnl;
|
||||
|
||||
// Small penalty for holding to encourage action
|
||||
let hold_penalty = -0.01;
|
||||
(pnl / 100.0) + hold_penalty
|
||||
},
|
||||
_ => 0.0,
|
||||
};
|
||||
|
||||
self.last_price = next_price;
|
||||
reward
|
||||
}
|
||||
|
||||
fn reset(&mut self) {
|
||||
self.position = 0.0;
|
||||
self.portfolio_value = self.initial_value;
|
||||
self.last_price = 0.0;
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("Using mimalloc allocator for improved performance");
|
||||
info!("Starting Rainbow DQN Training");
|
||||
info!("╔══════════════════════════════════════════════════════════════════════════╗");
|
||||
info!("║ Rainbow DQN: No epsilon-greedy! Uses noisy networks for exploration ║");
|
||||
info!("║ Components: Double-Q + Dueling + Priority Replay + Multi-step + C51 ║");
|
||||
info!("╚══════════════════════════════════════════════════════════════════════════╝");
|
||||
info!("\nConfiguration:");
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Learning rate: {}", opts.learning_rate);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(" • Gamma: {}", opts.gamma);
|
||||
info!(
|
||||
" • Checkpoint frequency: {} epochs",
|
||||
opts.checkpoint_frequency
|
||||
);
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
info!(" • Data directory: {}", opts.data_dir);
|
||||
info!(" • Bar sampling method: {}", opts.bar_method);
|
||||
if let Some(threshold) = opts.bar_threshold {
|
||||
info!(" • Bar threshold: {}", threshold);
|
||||
}
|
||||
info!(" • Buffer size: {}", opts.buffer_size);
|
||||
info!(" • Min replay size: {}", opts.min_replay_size);
|
||||
|
||||
info!("\n📊 Rainbow DQN Parameters:");
|
||||
info!(" • C51 Distributional:");
|
||||
info!(" - Num atoms: {}", opts.num_atoms);
|
||||
info!(" - V-min: {}", opts.v_min);
|
||||
info!(" - V-max: {}", opts.v_max);
|
||||
info!(" • Multi-step learning:");
|
||||
info!(" - N-step: {}", opts.n_step);
|
||||
info!(" • Priority Replay:");
|
||||
info!(
|
||||
" - Alpha: {} (prioritization strength)",
|
||||
opts.priority_alpha
|
||||
);
|
||||
info!(
|
||||
" - Beta: {} → 1.0 (importance sampling)",
|
||||
opts.priority_beta
|
||||
);
|
||||
info!(" - Beta increment: {}", opts.priority_beta_increment);
|
||||
info!(" • Noisy Networks:");
|
||||
info!(" - Sigma: {} (parameter noise)", opts.noisy_sigma);
|
||||
info!(" - Noise reset freq: {} steps", opts.noise_reset_freq);
|
||||
info!(" • Network Updates:");
|
||||
info!(
|
||||
" - Target update freq: {} steps",
|
||||
opts.target_update_freq
|
||||
);
|
||||
info!(" - Train freq: {} steps", opts.train_freq);
|
||||
|
||||
// Setup graceful shutdown handler
|
||||
let shutdown_flag = Arc::new(AtomicBool::new(false));
|
||||
let shutdown_clone = shutdown_flag.clone();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let ctrl_c = signal::ctrl_c();
|
||||
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use tokio::signal::unix::{signal, SignalKind};
|
||||
let mut sigterm =
|
||||
signal(SignalKind::terminate()).expect("Failed to setup SIGTERM handler");
|
||||
|
||||
tokio::select! {
|
||||
_ = ctrl_c => {
|
||||
info!("🛑 Received Ctrl+C, initiating graceful shutdown...");
|
||||
}
|
||||
_ = sigterm.recv() => {
|
||||
info!("🛑 Received SIGTERM, initiating graceful shutdown...");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(unix))]
|
||||
{
|
||||
ctrl_c.await.expect("Failed to listen for Ctrl+C");
|
||||
info!("🛑 Received Ctrl+C, initiating graceful shutdown...");
|
||||
}
|
||||
|
||||
shutdown_clone.store(true, Ordering::Relaxed);
|
||||
});
|
||||
|
||||
info!("✅ Graceful shutdown handler registered (Ctrl+C / SIGTERM)");
|
||||
|
||||
// Create output and checkpoint directories
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
let checkpoint_path = if let Some(ref dir) = opts.checkpoint_dir {
|
||||
PathBuf::from(dir)
|
||||
} else {
|
||||
output_path.clone()
|
||||
};
|
||||
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
if !checkpoint_path.exists() && checkpoint_path != output_path {
|
||||
std::fs::create_dir_all(&checkpoint_path)
|
||||
.context("Failed to create checkpoint directory")?;
|
||||
info!(
|
||||
"✅ Created checkpoint directory: {}",
|
||||
checkpoint_path.display()
|
||||
);
|
||||
}
|
||||
|
||||
if opts.checkpoint_dir.is_some() {
|
||||
info!(" • Checkpoint directory: {}", checkpoint_path.display());
|
||||
}
|
||||
|
||||
// Parse hidden layer sizes
|
||||
let hidden_sizes: Vec<usize> = opts
|
||||
.hidden_sizes
|
||||
.split(',')
|
||||
.map(|s| s.trim().parse::<usize>())
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.context("Failed to parse hidden_sizes")?;
|
||||
|
||||
// Configure Rainbow DQN
|
||||
let config = RainbowAgentConfig {
|
||||
device: if cfg!(feature = "cuda") {
|
||||
"cuda".to_string()
|
||||
} else {
|
||||
"cpu".to_string()
|
||||
},
|
||||
network_config: RainbowNetworkConfig {
|
||||
input_size: opts.state_dim,
|
||||
hidden_sizes,
|
||||
num_actions: opts.num_actions,
|
||||
activation: ml::dqn::rainbow_network::ActivationType::ReLU,
|
||||
dropout_rate: 0.1,
|
||||
distributional: DistributionalConfig {
|
||||
num_atoms: opts.num_atoms,
|
||||
v_min: opts.v_min,
|
||||
v_max: opts.v_max,
|
||||
},
|
||||
use_noisy_layers: true,
|
||||
dueling: true,
|
||||
},
|
||||
min_replay_size: opts.min_replay_size,
|
||||
replay_buffer_size: opts.buffer_size,
|
||||
batch_size: opts.batch_size,
|
||||
learning_rate: opts.learning_rate,
|
||||
gamma: opts.gamma,
|
||||
target_update_freq: opts.target_update_freq,
|
||||
train_freq: opts.train_freq,
|
||||
multi_step: MultiStepConfig {
|
||||
enabled: true,
|
||||
n_steps: opts.n_step,
|
||||
gamma: opts.gamma,
|
||||
},
|
||||
priority_alpha: opts.priority_alpha,
|
||||
priority_beta: opts.priority_beta,
|
||||
priority_beta_increment: opts.priority_beta_increment,
|
||||
noise_reset_freq: opts.noise_reset_freq,
|
||||
};
|
||||
|
||||
// Create Rainbow agent
|
||||
let agent = RainbowAgent::new(config).context("Failed to create Rainbow agent")?;
|
||||
|
||||
info!("✅ Rainbow DQN agent initialized");
|
||||
|
||||
// Configure alternative bar sampling
|
||||
let bar_sampling = match opts.bar_method.as_str() {
|
||||
"tick" => BarSamplingMethod::TickBars(opts.bar_threshold.unwrap_or(100.0) as usize),
|
||||
"volume" => BarSamplingMethod::VolumeBars(opts.bar_threshold.unwrap_or(10000.0)),
|
||||
"dollar" => BarSamplingMethod::DollarBars(opts.bar_threshold.unwrap_or(2_000_000.0)),
|
||||
"imbalance" => BarSamplingMethod::ImbalanceBars(opts.bar_threshold.unwrap_or(1000.0)),
|
||||
"run" => BarSamplingMethod::RunBars(opts.bar_threshold.unwrap_or(50.0) as usize),
|
||||
_ => BarSamplingMethod::TimeBars,
|
||||
};
|
||||
|
||||
info!("✅ Bar sampling configured: {:?}", bar_sampling);
|
||||
|
||||
// Setup checkpoint manager
|
||||
let checkpoint_config = CheckpointConfig {
|
||||
base_dir: output_path.clone(),
|
||||
max_checkpoints_per_model: 10,
|
||||
auto_cleanup: true,
|
||||
validate_checksums: true,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let _checkpoint_manager =
|
||||
CheckpointManager::new(checkpoint_config).context("Failed to create checkpoint manager")?;
|
||||
|
||||
info!("✅ Checkpoint manager initialized (max 10 checkpoints, auto-cleanup enabled)");
|
||||
|
||||
// Create checkpoint callback with interruption handling
|
||||
let checkpoint_dir_for_callback = opts
|
||||
.checkpoint_dir
|
||||
.clone()
|
||||
.unwrap_or_else(|| opts.output_dir.clone());
|
||||
let shutdown_check = shutdown_flag.clone();
|
||||
|
||||
let checkpoint_callback =
|
||||
move |epoch: usize, model_data: Vec<u8>, is_best: bool| -> Result<String> {
|
||||
// Check if shutdown was requested
|
||||
let interrupted = shutdown_check.load(Ordering::Relaxed);
|
||||
|
||||
let filename = if is_best {
|
||||
"rainbow_best_model.safetensors".to_string()
|
||||
} else if interrupted {
|
||||
format!("rainbow_interrupted_epoch{}.safetensors", epoch)
|
||||
} else {
|
||||
format!("rainbow_epoch_{}.safetensors", epoch)
|
||||
};
|
||||
|
||||
let checkpoint_path = PathBuf::from(&checkpoint_dir_for_callback).join(filename);
|
||||
|
||||
// Save checkpoint to disk
|
||||
std::fs::write(&checkpoint_path, &model_data)
|
||||
.context(format!("Failed to save checkpoint: {:?}", checkpoint_path))?;
|
||||
|
||||
let checkpoint_type = if is_best {
|
||||
"🎉 BEST"
|
||||
} else if interrupted {
|
||||
"⚠️ INTERRUPTED"
|
||||
} else {
|
||||
"💾 PERIODIC"
|
||||
};
|
||||
|
||||
info!(
|
||||
"{} Checkpoint saved: {} ({} bytes)",
|
||||
checkpoint_type,
|
||||
checkpoint_path.display(),
|
||||
model_data.len()
|
||||
);
|
||||
|
||||
Ok(checkpoint_path.to_string_lossy().to_string())
|
||||
};
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// DATA LOADING - Load ES futures data from parquet
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
info!("\n📊 Loading training data from parquet...");
|
||||
|
||||
let parquet_path = if let Some(ref path) = opts.parquet_file {
|
||||
path.clone()
|
||||
} else {
|
||||
// Default to ES futures test data
|
||||
"test_data/ES_FUT_180d.parquet".to_string()
|
||||
};
|
||||
|
||||
let training_data = load_training_data_from_parquet(&parquet_path)
|
||||
.await
|
||||
.context("Failed to load training data")?;
|
||||
|
||||
info!("✅ Loaded {} samples from parquet", training_data.len());
|
||||
|
||||
if training_data.is_empty() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"No training data loaded! Check parquet file path"
|
||||
));
|
||||
}
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// TRAINING ENVIRONMENT SETUP
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
let mut env = TradingEnvironment::new();
|
||||
let mut best_episode_reward = f32::NEG_INFINITY;
|
||||
let mut total_training_steps = 0_usize;
|
||||
|
||||
info!("\n🏋️ Starting Rainbow DQN training loop...\n");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
// TRAINING LOOP - Full implementation with real market data
|
||||
// ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
for epoch in 0..opts.epochs {
|
||||
// Check for shutdown
|
||||
if shutdown_flag.load(Ordering::Relaxed) {
|
||||
warn!("\n⚠️ Training interrupted at epoch {}", epoch);
|
||||
break;
|
||||
}
|
||||
|
||||
let mut episode_reward = 0.0_f32;
|
||||
let mut episode_steps = 0_usize;
|
||||
let mut action_counts = [0_usize; 3]; // [BUY, SELL, HOLD]
|
||||
let mut cumulative_reward = 0.0_f64; // Track cumulative reward as Q-value proxy
|
||||
|
||||
env.reset();
|
||||
|
||||
// Episode loop - iterate through all training samples
|
||||
for (step, (feature_vec, targets)) in training_data.iter().enumerate() {
|
||||
// Convert feature vector to state representation
|
||||
let state = feature_vector_to_state(feature_vec);
|
||||
|
||||
// Extract current and next close prices
|
||||
let current_price = targets[0] as f32;
|
||||
let next_price = if step + 1 < training_data.len() {
|
||||
training_data[step + 1].1[0] as f32
|
||||
} else {
|
||||
targets[1] as f32 // Terminal state, use self
|
||||
};
|
||||
|
||||
// Select action using Rainbow agent (noisy networks provide exploration)
|
||||
let action = agent
|
||||
.select_action(&state)
|
||||
.context("Failed to select action")?;
|
||||
let action_usize = action as usize;
|
||||
|
||||
// Track action distribution
|
||||
if action_usize < 3 {
|
||||
action_counts[action_usize] += 1;
|
||||
}
|
||||
|
||||
// Execute action in environment and get reward
|
||||
let reward = env.step(action_usize, current_price, next_price);
|
||||
episode_reward += reward;
|
||||
episode_steps += 1;
|
||||
total_training_steps += 1;
|
||||
cumulative_reward += reward as f64; // Accumulate for Q-value estimation
|
||||
|
||||
// Get next state
|
||||
let next_state = if step + 1 < training_data.len() {
|
||||
feature_vector_to_state(&training_data[step + 1].0)
|
||||
} else {
|
||||
state.clone() // Terminal state
|
||||
};
|
||||
|
||||
// Check if episode is done
|
||||
let done = step + 1 >= training_data.len();
|
||||
|
||||
// Add experience to Rainbow replay buffer
|
||||
let experience =
|
||||
ml::dqn::Experience::new(state, action as u8, reward, next_state, done);
|
||||
|
||||
agent
|
||||
.add_experience(experience)
|
||||
.context("Failed to add experience")?;
|
||||
|
||||
// Train Rainbow agent (after replay buffer has enough samples)
|
||||
let metrics = agent.metrics();
|
||||
if metrics.replay_buffer_size >= opts.min_replay_size
|
||||
&& total_training_steps % opts.train_freq == 0
|
||||
{
|
||||
if let Some(training_result) = agent.train()? {
|
||||
// Log metrics every 100 training steps
|
||||
if total_training_steps % 100 == 0 {
|
||||
// Compute average Q-value estimate from cumulative rewards
|
||||
// Note: This is a proxy since TrainingResult.q_values is empty
|
||||
// In C51 distributional RL, Q-values typically range from -10 to +10
|
||||
let avg_q_estimate = cumulative_reward / (episode_steps.max(1) as f64);
|
||||
|
||||
info!(
|
||||
"Epoch {}/{}, Step {}: Loss={:.4}, AvgQ≈{:.3}, Buffer={}, Steps={}",
|
||||
epoch + 1,
|
||||
opts.epochs,
|
||||
step,
|
||||
training_result.loss,
|
||||
avg_q_estimate,
|
||||
metrics.replay_buffer_size,
|
||||
metrics.total_steps
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Epoch summary
|
||||
let buy_pct = (action_counts[0] as f32 / episode_steps as f32) * 100.0;
|
||||
let sell_pct = (action_counts[1] as f32 / episode_steps as f32) * 100.0;
|
||||
let hold_pct = (action_counts[2] as f32 / episode_steps as f32) * 100.0;
|
||||
|
||||
info!(
|
||||
"Epoch {}/{} completed: Reward={:.2}, Steps={}, Actions=[BUY:{:.1}%, SELL:{:.1}%, HOLD:{:.1}%]",
|
||||
epoch + 1, opts.epochs, episode_reward, episode_steps,
|
||||
buy_pct, sell_pct, hold_pct
|
||||
);
|
||||
|
||||
// Track best episode
|
||||
if episode_reward > best_episode_reward {
|
||||
best_episode_reward = episode_reward;
|
||||
info!("🎉 New best episode reward: {:.2}", best_episode_reward);
|
||||
}
|
||||
|
||||
// Periodic checkpoint
|
||||
if (epoch + 1) % opts.checkpoint_frequency == 0 {
|
||||
let checkpoint_data = vec![0u8; 1024]; // Placeholder - would serialize agent state
|
||||
checkpoint_callback(
|
||||
epoch + 1,
|
||||
checkpoint_data,
|
||||
episode_reward >= best_episode_reward,
|
||||
)?;
|
||||
}
|
||||
}
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Check if training was interrupted
|
||||
if shutdown_flag.load(Ordering::Relaxed) {
|
||||
info!("\n⚠️ Training was interrupted by shutdown signal");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Print final metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!("\n📊 Final Metrics:");
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
|
||||
// Save final model
|
||||
let final_model_path =
|
||||
output_path.join(format!("rainbow_final_epoch{}.safetensors", opts.epochs));
|
||||
info!("\n💾 Saving final model to: {}", final_model_path.display());
|
||||
|
||||
// Placeholder - full implementation would serialize agent state
|
||||
let final_checkpoint_data = vec![0u8; 1024];
|
||||
std::fs::write(&final_model_path, &final_checkpoint_data)
|
||||
.context("Failed to save final model")?;
|
||||
|
||||
info!(
|
||||
"✅ Final model saved: {} ({} bytes)",
|
||||
final_model_path.display(),
|
||||
final_checkpoint_data.len()
|
||||
);
|
||||
|
||||
info!("\n🎉 Rainbow DQN training complete!");
|
||||
info!("📁 Model files saved to: {}", opts.output_dir);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,277 +0,0 @@
|
||||
//! TFT (Temporal Fusion Transformer) Training Example
|
||||
//!
|
||||
//! Trains a TFT model for time series forecasting and saves checkpoints.
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (100 epochs)
|
||||
//! cargo run -p ml --example train_tft --release --features cuda
|
||||
//!
|
||||
//! # Custom configuration
|
||||
//! cargo run -p ml --example train_tft --release --features cuda -- \
|
||||
//! --epochs 500 \
|
||||
//! --batch-size 32 \
|
||||
//! --hidden-dim 256
|
||||
//! ```
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use ndarray::Array2;
|
||||
use std::path::PathBuf;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::info;
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::checkpoint::FileSystemStorage;
|
||||
use ml::tft::training::TFTDataLoader;
|
||||
use ml::trainers::tft::{TFTTrainer, TFTTrainerConfig};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(name = "train_tft", about = "Train TFT model on time series data")]
|
||||
struct Opts {
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "100")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value = "0.001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Batch size (max 32 for 4GB VRAM)
|
||||
#[arg(long, default_value = "32")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Hidden dimension
|
||||
#[arg(long, default_value = "256")]
|
||||
hidden_dim: usize,
|
||||
|
||||
/// Number of attention heads
|
||||
#[arg(long, default_value = "8")]
|
||||
num_attention_heads: usize,
|
||||
|
||||
/// Lookback window
|
||||
#[arg(long, default_value = "60")]
|
||||
lookback_window: usize,
|
||||
|
||||
/// Forecast horizon
|
||||
#[arg(long, default_value = "10")]
|
||||
forecast_horizon: usize,
|
||||
|
||||
/// Output directory for trained model
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Use GPU
|
||||
#[arg(long)]
|
||||
use_gpu: bool,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 Starting TFT Training");
|
||||
info!("Configuration:");
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Learning rate: {}", opts.learning_rate);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(" • Hidden dimension: {}", opts.hidden_dim);
|
||||
info!(" • Attention heads: {}", opts.num_attention_heads);
|
||||
info!(" • Lookback window: {}", opts.lookback_window);
|
||||
info!(" • Forecast horizon: {}", opts.forecast_horizon);
|
||||
info!(" • GPU enabled: {}", opts.use_gpu);
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Configure TFT trainer
|
||||
let trainer_config = TFTTrainerConfig {
|
||||
epochs: opts.epochs,
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
auto_batch_size: false,
|
||||
validation_batch_size: opts.batch_size,
|
||||
hidden_dim: opts.hidden_dim,
|
||||
num_attention_heads: opts.num_attention_heads,
|
||||
dropout_rate: 0.1,
|
||||
lstm_layers: 2,
|
||||
quantiles: vec![0.1, 0.5, 0.9],
|
||||
lookback_window: opts.lookback_window,
|
||||
forecast_horizon: opts.forecast_horizon,
|
||||
use_gpu: opts.use_gpu,
|
||||
use_int8_quantization: false,
|
||||
use_qat: false,
|
||||
qat_calibration_batches: 100,
|
||||
qat_warmup_epochs: 10,
|
||||
qat_cooldown_factor: 0.1,
|
||||
qat_min_batch_size: 2,
|
||||
use_gradient_checkpointing: false,
|
||||
max_validation_batches: None,
|
||||
validation_frequency: 1,
|
||||
checkpoint_dir: opts.output_dir.clone(),
|
||||
};
|
||||
|
||||
// Create checkpoint storage
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(output_path.clone()));
|
||||
|
||||
// Create TFT trainer
|
||||
let mut trainer =
|
||||
TFTTrainer::new(trainer_config.clone(), storage).context("Failed to create TFT trainer")?;
|
||||
|
||||
info!("✅ TFT trainer initialized");
|
||||
|
||||
// Generate synthetic time series data
|
||||
info!("\n📊 Generating training data...");
|
||||
let num_train_samples = 3200; // 100 batches of size 32
|
||||
let num_val_samples = 320; // 10 batches of size 32
|
||||
|
||||
let train_loader = generate_data_loader(
|
||||
num_train_samples,
|
||||
opts.batch_size,
|
||||
opts.lookback_window,
|
||||
opts.forecast_horizon,
|
||||
true, // shuffle training data
|
||||
)?;
|
||||
|
||||
let val_loader = generate_data_loader(
|
||||
num_val_samples,
|
||||
opts.batch_size,
|
||||
opts.lookback_window,
|
||||
opts.forecast_horizon,
|
||||
false, // don't shuffle validation data
|
||||
)?;
|
||||
|
||||
info!(
|
||||
"✅ Generated {} training samples, {} validation samples",
|
||||
num_train_samples, num_val_samples
|
||||
);
|
||||
|
||||
// Setup progress callback
|
||||
let (progress_tx, mut progress_rx) = mpsc::unbounded_channel();
|
||||
trainer.set_progress_callback(progress_tx);
|
||||
|
||||
// Spawn progress monitor task
|
||||
let monitor_task = tokio::spawn(async move {
|
||||
while let Some(progress) = progress_rx.recv().await {
|
||||
if progress.current_epoch % 10 == 0 {
|
||||
info!("{}", progress.message);
|
||||
if let Some(loss) = progress.metrics.get("train_loss") {
|
||||
info!(" • Train loss: {:.6}", loss);
|
||||
}
|
||||
if let Some(val_loss) = progress.metrics.get("val_loss") {
|
||||
info!(" • Val loss: {:.6}", val_loss);
|
||||
}
|
||||
if let Some(rmse) = progress.metrics.get("rmse") {
|
||||
info!(" • RMSE: {:.6}", rmse);
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Train the model
|
||||
info!("\n🏋️ Starting training...\n");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let final_metrics = trainer
|
||||
.train(train_loader, val_loader)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Wait for progress monitor to finish
|
||||
drop(trainer); // Drop trainer to close progress channel
|
||||
let _ = monitor_task.await;
|
||||
|
||||
// Print final metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!("\n📊 Final Metrics:");
|
||||
info!(" • Training loss: {:.6}", final_metrics.train_loss);
|
||||
info!(" • Validation loss: {:.6}", final_metrics.val_loss);
|
||||
info!(" • Quantile loss: {:.6}", final_metrics.quantile_loss);
|
||||
info!(" • RMSE: {:.6}", final_metrics.rmse);
|
||||
info!(
|
||||
" • Attention entropy: {:.4}",
|
||||
final_metrics.attention_entropy
|
||||
);
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min)",
|
||||
final_metrics.training_time_seconds,
|
||||
final_metrics.training_time_seconds / 60.0
|
||||
);
|
||||
info!(
|
||||
" • Wall-clock duration: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
|
||||
info!("\n💾 Model checkpoints saved to: {}", opts.output_dir);
|
||||
info!("\n🎉 TFT training complete!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Generate synthetic data loader for training/validation
|
||||
fn generate_data_loader(
|
||||
num_samples: usize,
|
||||
batch_size: usize,
|
||||
lookback_window: usize,
|
||||
forecast_horizon: usize,
|
||||
shuffle: bool,
|
||||
) -> Result<TFTDataLoader> {
|
||||
use ndarray::Array1;
|
||||
|
||||
// Generate synthetic data samples
|
||||
let mut data = Vec::with_capacity(num_samples);
|
||||
|
||||
for i in 0..num_samples {
|
||||
// Static features: [num_static_features] = [10]
|
||||
let static_features = Array1::from_shape_fn(10, |j| (i as f64 * 0.1 + j as f64 * 0.01));
|
||||
|
||||
// Historical features: [lookback_window, num_hist_features] = [60, 54] (Wave E)
|
||||
let historical_features = Array2::from_shape_fn((lookback_window, 54), |(t, f)| {
|
||||
(i as f64 * 0.1 + t as f64 * 0.01 + f as f64 * 0.001).sin()
|
||||
});
|
||||
|
||||
// Future features: [forecast_horizon, num_fut_features] = [10, 10]
|
||||
let future_features = Array2::from_shape_fn((forecast_horizon, 10), |(t, f)| {
|
||||
(i as f64 * 0.1 + (lookback_window + t) as f64 * 0.01 + f as f64 * 0.001).cos()
|
||||
});
|
||||
|
||||
// Targets: [forecast_horizon] = [10]
|
||||
let targets = Array1::from_shape_fn(forecast_horizon, |t| {
|
||||
(i as f64 * 0.1 + (lookback_window + t) as f64 * 0.01).sin() * 100.0
|
||||
});
|
||||
|
||||
data.push((
|
||||
static_features,
|
||||
historical_features,
|
||||
future_features,
|
||||
targets,
|
||||
));
|
||||
}
|
||||
|
||||
Ok(TFTDataLoader::new(data, batch_size, shuffle))
|
||||
}
|
||||
@@ -1,728 +0,0 @@
|
||||
//! TFT (Temporal Fusion Transformer) Training with Real DataBento Data
|
||||
//!
|
||||
//! Trains a TFT model using real market data from DataBento DBN files with proper
|
||||
//! static covariates, time-varying features, and multi-horizon forecasting.
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (20 epochs)
|
||||
//! cargo run -p ml --example train_tft_dbn --release --features cuda
|
||||
//!
|
||||
//! # Custom configuration
|
||||
//! cargo run -p ml --example train_tft_dbn --release --features cuda -- \
|
||||
//! --epochs 50 \
|
||||
//! --batch-size 32 \
|
||||
//! --lookback 60 \
|
||||
//! --horizon 10
|
||||
//! ```
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use chrono::{DateTime, Datelike, TimeZone, Timelike, Utc};
|
||||
use clap::Parser;
|
||||
use dbn::decode::{DbnDecoder, DecodeRecordRef};
|
||||
use dbn::OhlcvMsg;
|
||||
use ndarray::{Array1, Array2};
|
||||
use std::path::PathBuf;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::checkpoint::FileSystemStorage;
|
||||
use ml::data_loaders::BarSamplingMethod;
|
||||
use ml::features::config::FeatureConfig;
|
||||
use ml::features::extraction::{extract_ml_features, OHLCVBar as ExtractorBar};
|
||||
use ml::tft::training::TFTDataLoader;
|
||||
use ml::trainers::tft::{TFTTrainer, TFTTrainerConfig};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_tft_dbn",
|
||||
about = "Train TFT model on real DataBento data"
|
||||
)]
|
||||
struct Opts {
|
||||
/// DBN file path (or directory containing multiple DBN files)
|
||||
#[arg(
|
||||
long,
|
||||
default_value = "test_data/real/databento/ES.FUT_ohlcv-1m_2024-01-02.dbn"
|
||||
)]
|
||||
data_path: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "20")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value = "0.001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Batch size (max 32 for 4GB VRAM)
|
||||
#[arg(long, default_value = "32")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Hidden dimension
|
||||
#[arg(long, default_value = "256")]
|
||||
hidden_dim: usize,
|
||||
|
||||
/// Number of attention heads
|
||||
#[arg(long, default_value = "8")]
|
||||
num_attention_heads: usize,
|
||||
|
||||
/// Lookback window
|
||||
#[arg(long, default_value = "60")]
|
||||
lookback_window: usize,
|
||||
|
||||
/// Forecast horizon
|
||||
#[arg(long, default_value = "10")]
|
||||
forecast_horizon: usize,
|
||||
|
||||
/// Training/validation split (0.0-1.0)
|
||||
#[arg(long, default_value = "0.8")]
|
||||
train_split: f64,
|
||||
|
||||
/// Output directory for trained model
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Early stopping patience (epochs without improvement)
|
||||
#[arg(long, default_value = "20")]
|
||||
early_stopping_patience: usize,
|
||||
|
||||
/// Early stopping threshold (minimum improvement)
|
||||
#[arg(long, default_value = "0.0001")]
|
||||
early_stopping_threshold: f64,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
|
||||
/// Alternative bar sampling method (time, tick, volume, dollar, imbalance, run)
|
||||
#[arg(long, default_value = "time")]
|
||||
bar_method: String,
|
||||
|
||||
/// Bar sampling threshold (tick count, volume, dollar value, imbalance, or run length)
|
||||
#[arg(long)]
|
||||
bar_threshold: Option<f64>,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 Starting TFT Training with Real DataBento Data");
|
||||
// Initialize Wave D feature configuration (225 features)
|
||||
// TODO: Update to wave_e() when 54-feature config is implemented
|
||||
let feature_config = FeatureConfig::wave_d();
|
||||
|
||||
info!("Configuration:");
|
||||
// Validate feature count matches expected 225-dimensional input
|
||||
let total_features = feature_config.feature_count();
|
||||
assert_eq!(
|
||||
total_features, 225,
|
||||
"Feature config must provide exactly 225 features for Wave D (201 Wave C + 24 Wave D)"
|
||||
);
|
||||
info!(" • Data path: {}", opts.data_path);
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Learning rate: {}", opts.learning_rate);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(" • Hidden dimension: {}", opts.hidden_dim);
|
||||
info!(" • Attention heads: {}", opts.num_attention_heads);
|
||||
info!(" • Lookback window: {}", opts.lookback_window);
|
||||
info!(" • Forecast horizon: {}", opts.forecast_horizon);
|
||||
info!(
|
||||
" • Feature count: {} (Wave D: Wave C 201 + Wave D 24)",
|
||||
total_features
|
||||
);
|
||||
info!(
|
||||
" • Train/val split: {:.1}%/{:.1}%",
|
||||
opts.train_split * 100.0,
|
||||
(1.0 - opts.train_split) * 100.0
|
||||
);
|
||||
info!(" • GPU: CUDA MANDATORY (no CPU fallback)");
|
||||
info!(
|
||||
" • Early stopping patience: {} epochs",
|
||||
opts.early_stopping_patience
|
||||
);
|
||||
info!(
|
||||
" • Early stopping threshold: {:.2e}",
|
||||
opts.early_stopping_threshold
|
||||
);
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
info!(" • Bar sampling method: {}", opts.bar_method);
|
||||
if let Some(threshold) = opts.bar_threshold {
|
||||
info!(" • Bar threshold: {}", threshold);
|
||||
}
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Configure alternative bar sampling (Wave B)
|
||||
let bar_sampling = match opts.bar_method.as_str() {
|
||||
"tick" => BarSamplingMethod::TickBars(opts.bar_threshold.unwrap_or(100.0) as usize),
|
||||
"volume" => BarSamplingMethod::VolumeBars(opts.bar_threshold.unwrap_or(10000.0)),
|
||||
"dollar" => BarSamplingMethod::DollarBars(opts.bar_threshold.unwrap_or(2_000_000.0)),
|
||||
"imbalance" => BarSamplingMethod::ImbalanceBars(opts.bar_threshold.unwrap_or(1000.0)),
|
||||
"run" => BarSamplingMethod::RunBars(opts.bar_threshold.unwrap_or(50.0) as usize),
|
||||
_ => BarSamplingMethod::TimeBars,
|
||||
};
|
||||
|
||||
info!("✅ Bar sampling configured: {:?}", bar_sampling);
|
||||
|
||||
// Load real market data from DBN files
|
||||
info!("\n📊 Loading real market data from DataBento...");
|
||||
|
||||
// Note: load_dbn_ohlcv_bars will need to support alternative bar sampling
|
||||
// This requires extending the function to accept bar_sampling parameter
|
||||
// For now, it loads time-based bars
|
||||
|
||||
// Check if path is a file or directory
|
||||
let path = std::path::Path::new(&opts.data_path);
|
||||
let bars = if path.is_dir() {
|
||||
// Load all .dbn files from directory
|
||||
info!("Loading DBN files from directory: {}", opts.data_path);
|
||||
let mut all_bars = Vec::new();
|
||||
|
||||
for entry in std::fs::read_dir(path)? {
|
||||
let entry = entry?;
|
||||
let file_path = entry.path();
|
||||
|
||||
if file_path.extension().and_then(|s| s.to_str()) == Some("dbn") {
|
||||
info!(" • Loading: {:?}", file_path.file_name().unwrap());
|
||||
let file_bars = load_dbn_ohlcv_bars(file_path.to_str().unwrap())
|
||||
.await
|
||||
.context(format!("Failed to load DBN file: {:?}", file_path))?;
|
||||
all_bars.extend(file_bars);
|
||||
}
|
||||
}
|
||||
|
||||
// Sort by timestamp
|
||||
all_bars.sort_by_key(|b| b.timestamp);
|
||||
all_bars
|
||||
} else {
|
||||
// Load single file
|
||||
load_dbn_ohlcv_bars(&opts.data_path)
|
||||
.await
|
||||
.context("Failed to load DBN data")?
|
||||
};
|
||||
|
||||
info!("✅ Loaded {} OHLCV bars from DataBento", bars.len());
|
||||
|
||||
// Convert to TFT data structure with Wave D features (225)
|
||||
info!(
|
||||
"\n🔄 Converting to TFT data format with {} features...",
|
||||
total_features
|
||||
);
|
||||
let tft_data = convert_to_tft_data(
|
||||
&bars,
|
||||
opts.lookback_window,
|
||||
opts.forecast_horizon,
|
||||
&feature_config,
|
||||
)
|
||||
.context("Failed to convert to TFT format")?;
|
||||
|
||||
info!("✅ Created {} TFT samples", tft_data.len());
|
||||
|
||||
// Split into train/validation
|
||||
let split_idx = (tft_data.len() as f64 * opts.train_split) as usize;
|
||||
let train_data = tft_data[..split_idx].to_vec();
|
||||
let val_data = tft_data[split_idx..].to_vec();
|
||||
|
||||
info!(
|
||||
"✅ Split: {} training, {} validation samples",
|
||||
train_data.len(),
|
||||
val_data.len()
|
||||
);
|
||||
|
||||
// Create data loaders
|
||||
let train_loader = TFTDataLoader::new(train_data, opts.batch_size, true);
|
||||
let val_loader = TFTDataLoader::new(val_data, opts.batch_size, false);
|
||||
|
||||
// Configure TFT trainer
|
||||
// Configure TFT trainer with 225 features (Wave D)
|
||||
// Static features: 10 (symbol metadata, volatility, liquidity)
|
||||
// Historical features: 225 (Wave C 201 + Wave D 24)
|
||||
// Future features: 10 (calendar features)
|
||||
let trainer_config = TFTTrainerConfig {
|
||||
epochs: opts.epochs,
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
auto_batch_size: false,
|
||||
validation_batch_size: opts.batch_size,
|
||||
hidden_dim: opts.hidden_dim,
|
||||
num_attention_heads: opts.num_attention_heads,
|
||||
dropout_rate: 0.1,
|
||||
lstm_layers: 2,
|
||||
quantiles: vec![0.1, 0.5, 0.9],
|
||||
lookback_window: opts.lookback_window,
|
||||
forecast_horizon: opts.forecast_horizon,
|
||||
use_gpu: true, // CUDA always required
|
||||
use_int8_quantization: false,
|
||||
use_qat: false,
|
||||
qat_calibration_batches: 100,
|
||||
qat_warmup_epochs: 10,
|
||||
qat_cooldown_factor: 0.1,
|
||||
qat_min_batch_size: 2,
|
||||
use_gradient_checkpointing: false,
|
||||
max_validation_batches: None,
|
||||
validation_frequency: 1,
|
||||
checkpoint_dir: opts.output_dir.clone(),
|
||||
};
|
||||
|
||||
// Create checkpoint storage
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(output_path.clone()));
|
||||
|
||||
// Create TFT trainer
|
||||
let mut trainer =
|
||||
TFTTrainer::new(trainer_config.clone(), storage).context("Failed to create TFT trainer")?;
|
||||
|
||||
info!("✅ TFT trainer initialized");
|
||||
|
||||
// Setup progress callback
|
||||
let (progress_tx, mut progress_rx) = mpsc::unbounded_channel();
|
||||
trainer.set_progress_callback(progress_tx);
|
||||
|
||||
// Spawn progress monitor task
|
||||
let monitor_task = tokio::spawn(async move {
|
||||
while let Some(progress) = progress_rx.recv().await {
|
||||
info!("{}", progress.message);
|
||||
if let Some(loss) = progress.metrics.get("train_loss") {
|
||||
info!(" • Train loss: {:.6}", loss);
|
||||
}
|
||||
if let Some(val_loss) = progress.metrics.get("val_loss") {
|
||||
info!(" • Val loss: {:.6}", val_loss);
|
||||
}
|
||||
if let Some(quantile_loss) = progress.metrics.get("quantile_loss") {
|
||||
info!(" • Quantile loss: {:.6}", quantile_loss);
|
||||
}
|
||||
if let Some(rmse) = progress.metrics.get("rmse") {
|
||||
info!(" • RMSE: {:.6}", rmse);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Train the model
|
||||
info!("\n🏋️ Starting training...\n");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let final_metrics = trainer
|
||||
.train(train_loader, val_loader)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Wait for progress monitor to finish
|
||||
drop(trainer); // Drop trainer to close progress channel
|
||||
let _ = monitor_task.await;
|
||||
|
||||
// Print final metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!("\n📊 Final Metrics:");
|
||||
info!(" • Training loss: {:.6}", final_metrics.train_loss);
|
||||
info!(" • Validation loss: {:.6}", final_metrics.val_loss);
|
||||
info!(" • Quantile loss: {:.6}", final_metrics.quantile_loss);
|
||||
info!(" • RMSE: {:.6}", final_metrics.rmse);
|
||||
info!(
|
||||
" • Attention entropy: {:.4}",
|
||||
final_metrics.attention_entropy
|
||||
);
|
||||
info!(
|
||||
" • Training duration: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min)",
|
||||
final_metrics.training_time_seconds,
|
||||
final_metrics.training_time_seconds / 60.0
|
||||
);
|
||||
|
||||
info!("\n💾 Model checkpoints saved to: {}", opts.output_dir);
|
||||
info!("\n🎉 TFT training with real DataBento data complete!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// OHLCV bar structure (intermediate format)
|
||||
#[derive(Debug, Clone)]
|
||||
struct OhlcvBar {
|
||||
timestamp: DateTime<Utc>,
|
||||
open: f64,
|
||||
high: f64,
|
||||
low: f64,
|
||||
close: f64,
|
||||
volume: f64,
|
||||
}
|
||||
|
||||
/// Load OHLCV bars from DBN file with price anomaly correction
|
||||
async fn load_dbn_ohlcv_bars(file_path: &str) -> Result<Vec<OhlcvBar>> {
|
||||
debug!("Loading DBN file: {}", file_path);
|
||||
|
||||
let mut decoder = DbnDecoder::from_file(file_path).context(format!(
|
||||
"Failed to create DBN decoder for file: {}",
|
||||
file_path
|
||||
))?;
|
||||
|
||||
let mut bars = Vec::new();
|
||||
let mut prev_close: Option<f64> = None;
|
||||
let mut corrections_applied = 0;
|
||||
|
||||
while let Some(record_ref) = decoder
|
||||
.decode_record_ref()
|
||||
.context("Failed to decode DBN record")?
|
||||
{
|
||||
if let Some(ohlcv) = record_ref.get::<OhlcvMsg>() {
|
||||
// Convert timestamp
|
||||
let ts_nanos = ohlcv.hd.ts_event as i64;
|
||||
let secs = ts_nanos / 1_000_000_000;
|
||||
let nanos = (ts_nanos % 1_000_000_000) as u32;
|
||||
let timestamp = Utc
|
||||
.timestamp_opt(secs, nanos)
|
||||
.single()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid timestamp: {}", ts_nanos))?;
|
||||
|
||||
// Convert prices (DBN uses 9 decimal places)
|
||||
let mut open_f64 = ohlcv.open as f64 / 1_000_000_000.0;
|
||||
let mut high_f64 = ohlcv.high as f64 / 1_000_000_000.0;
|
||||
let mut low_f64 = ohlcv.low as f64 / 1_000_000_000.0;
|
||||
let mut close_f64 = ohlcv.close as f64 / 1_000_000_000.0;
|
||||
|
||||
// Price anomaly detection (same logic as backtesting service)
|
||||
if let Some(prev) = prev_close {
|
||||
let pct_change = ((close_f64 - prev) / prev).abs();
|
||||
|
||||
if pct_change > 0.5 && close_f64 < 1000.0 {
|
||||
let corrected_close = close_f64 * 100.0;
|
||||
|
||||
if corrected_close >= 3000.0 && corrected_close <= 6000.0 {
|
||||
open_f64 *= 100.0;
|
||||
high_f64 *= 100.0;
|
||||
low_f64 *= 100.0;
|
||||
close_f64 = corrected_close;
|
||||
corrections_applied += 1;
|
||||
|
||||
if corrections_applied <= 5 {
|
||||
debug!(
|
||||
"Applied 100x price correction at bar {} ({}% change)",
|
||||
bars.len() + 1,
|
||||
pct_change * 100.0
|
||||
);
|
||||
}
|
||||
} else {
|
||||
warn!(
|
||||
"Skipping corrupted bar at index {} (timestamp: {})",
|
||||
bars.len() + 1,
|
||||
timestamp
|
||||
);
|
||||
prev_close = Some(prev);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
prev_close = Some(close_f64);
|
||||
|
||||
let bar = OhlcvBar {
|
||||
timestamp,
|
||||
open: open_f64,
|
||||
high: high_f64,
|
||||
low: low_f64,
|
||||
close: close_f64,
|
||||
volume: ohlcv.volume as f64,
|
||||
};
|
||||
|
||||
bars.push(bar);
|
||||
}
|
||||
}
|
||||
|
||||
if corrections_applied > 0 {
|
||||
info!(
|
||||
"Applied {} automatic price corrections for encoding inconsistencies",
|
||||
corrections_applied
|
||||
);
|
||||
}
|
||||
|
||||
Ok(bars)
|
||||
}
|
||||
|
||||
/// Convert OHLCV bars to TFT data structure using production feature extraction
|
||||
///
|
||||
/// ## Wave D Feature Set (225 total):
|
||||
/// Uses `ml::features::extraction::extract_ml_features()` for all 225 features:
|
||||
/// ### Wave C Base (indices 0-200): 201 features
|
||||
/// - OHLCV: 5 features
|
||||
/// - Technical Indicators: 21 features (SMA, EMA, RSI, MACD, Bollinger, ATR)
|
||||
/// - Microstructure: 3 features (Roll, Amihud, Corwin-Schultz)
|
||||
/// - Statistical Features: 172 features (quantiles, correlations, rolling stats)
|
||||
///
|
||||
/// ### Wave D Regime Detection (indices 201-224): 24 features
|
||||
/// - CUSUM Statistics (201-210): 10 features
|
||||
/// - ADX & Directional Indicators (211-215): 5 features
|
||||
/// - Regime Transition Probabilities (216-220): 5 features
|
||||
/// - Adaptive Strategy Metrics (221-224): 4 features
|
||||
fn convert_to_tft_data(
|
||||
bars: &[OhlcvBar],
|
||||
lookback_window: usize,
|
||||
forecast_horizon: usize,
|
||||
_feature_config: &FeatureConfig,
|
||||
) -> Result<Vec<(Array1<f64>, Array2<f64>, Array2<f64>, Array1<f64>)>> {
|
||||
if bars.len() < lookback_window + forecast_horizon {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Not enough data: need {} bars, got {}",
|
||||
lookback_window + forecast_horizon,
|
||||
bars.len()
|
||||
));
|
||||
}
|
||||
|
||||
// Convert OhlcvBar to ExtractorBar for feature extraction
|
||||
let extractor_bars: Vec<ExtractorBar> = bars
|
||||
.iter()
|
||||
.map(|b| ExtractorBar {
|
||||
timestamp: b.timestamp,
|
||||
open: b.open,
|
||||
high: b.high,
|
||||
low: b.low,
|
||||
close: b.close,
|
||||
volume: b.volume,
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Extract 225-dimensional features using production pipeline
|
||||
info!("🔍 Extracting 225-dim features via production pipeline...");
|
||||
let feature_vectors =
|
||||
extract_ml_features(&extractor_bars).context("Failed to extract ML features")?;
|
||||
|
||||
info!(
|
||||
"✅ Extracted {} feature vectors (225-dim each)",
|
||||
feature_vectors.len()
|
||||
);
|
||||
|
||||
// Calculate statistics for static features and normalization
|
||||
let prices: Vec<f64> = bars.iter().map(|b| b.close).collect();
|
||||
let mean_price = prices.iter().sum::<f64>() / prices.len() as f64;
|
||||
let price_std =
|
||||
(prices.iter().map(|p| (p - mean_price).powi(2)).sum::<f64>() / prices.len() as f64).sqrt();
|
||||
|
||||
let volumes: Vec<f64> = bars.iter().map(|b| b.volume).collect();
|
||||
let mean_volume = volumes.iter().sum::<f64>() / volumes.len() as f64;
|
||||
let volume_std = (volumes
|
||||
.iter()
|
||||
.map(|v| (v - mean_volume).powi(2))
|
||||
.sum::<f64>()
|
||||
/ volumes.len() as f64)
|
||||
.sqrt();
|
||||
|
||||
let mut tft_samples = Vec::new();
|
||||
|
||||
// Note: feature_vectors starts AFTER warmup period (50 bars)
|
||||
// So feature_vectors[0] corresponds to bars[50]
|
||||
const WARMUP_PERIOD: usize = 50;
|
||||
|
||||
// Create sliding windows from feature vectors
|
||||
for i in 0..feature_vectors.len() {
|
||||
// Check if we have enough data for this window
|
||||
if i + lookback_window + forecast_horizon > feature_vectors.len() {
|
||||
break;
|
||||
}
|
||||
|
||||
// Calculate the bar index in the original bars array
|
||||
let bar_idx = WARMUP_PERIOD + i;
|
||||
let first_bar = &bars[bar_idx];
|
||||
|
||||
// Static features: Symbol metadata (10 features)
|
||||
let hour = first_bar.timestamp.hour() as f64;
|
||||
let day_of_week = first_bar.timestamp.weekday().num_days_from_monday() as f64;
|
||||
let is_morning = if hour < 12.0 { 1.0 } else { 0.0 };
|
||||
let is_afternoon = if hour >= 12.0 && hour < 17.0 {
|
||||
1.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
// Calculate volatility over lookback window
|
||||
let lookback_slice = &bars[bar_idx..bar_idx + lookback_window];
|
||||
let returns: Vec<f64> = lookback_slice
|
||||
.windows(2)
|
||||
.map(|w| (w[1].close / w[0].close).ln())
|
||||
.collect();
|
||||
let volatility = if returns.len() > 1 {
|
||||
let mean_return = returns.iter().sum::<f64>() / returns.len() as f64;
|
||||
(returns
|
||||
.iter()
|
||||
.map(|r| (r - mean_return).powi(2))
|
||||
.sum::<f64>()
|
||||
/ returns.len() as f64)
|
||||
.sqrt()
|
||||
} else {
|
||||
0.01
|
||||
};
|
||||
|
||||
let liquidity = mean_volume / mean_price;
|
||||
|
||||
let static_features = Array1::from_vec(vec![
|
||||
mean_price / 5000.0,
|
||||
price_std / 100.0,
|
||||
mean_volume / 1000.0,
|
||||
volume_std / 1000.0,
|
||||
hour / 24.0,
|
||||
day_of_week / 7.0,
|
||||
is_morning,
|
||||
is_afternoon,
|
||||
volatility * 100.0,
|
||||
liquidity / 100.0,
|
||||
]);
|
||||
|
||||
// Historical features: Use production-extracted 225-dim feature vectors
|
||||
let mut hist_features = Vec::with_capacity(lookback_window * 225);
|
||||
for t in 0..lookback_window {
|
||||
let feature_vec = &feature_vectors[i + t];
|
||||
hist_features.extend_from_slice(feature_vec);
|
||||
}
|
||||
|
||||
let historical_features = Array2::from_shape_vec((lookback_window, 225), hist_features)?;
|
||||
|
||||
// Future features: Known future events (10 features per timestep)
|
||||
let mut fut_features = Vec::with_capacity(forecast_horizon * 10);
|
||||
|
||||
for t in 0..forecast_horizon {
|
||||
let future_bar = &bars[bar_idx + lookback_window + t];
|
||||
let fut_hour = future_bar.timestamp.hour() as f64;
|
||||
let fut_day = future_bar.timestamp.weekday().num_days_from_monday() as f64;
|
||||
let is_weekend = if fut_day >= 5.0 { 1.0 } else { 0.0 };
|
||||
let fut_is_morning = if fut_hour < 12.0 { 1.0 } else { 0.0 };
|
||||
let fut_is_afternoon = if fut_hour >= 12.0 && fut_hour < 17.0 {
|
||||
1.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
let week_of_month = ((future_bar.timestamp.day() - 1) / 7) as f64;
|
||||
let month = future_bar.timestamp.month() as f64;
|
||||
let quarter = ((month - 1.0) / 3.0).floor();
|
||||
let is_month_start = if future_bar.timestamp.day() <= 5 {
|
||||
1.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
let is_month_end = if future_bar.timestamp.day() >= 25 {
|
||||
1.0
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
fut_features.extend(vec![
|
||||
fut_hour / 24.0,
|
||||
fut_day / 7.0,
|
||||
is_weekend,
|
||||
fut_is_morning,
|
||||
fut_is_afternoon,
|
||||
week_of_month / 4.0,
|
||||
month / 12.0,
|
||||
quarter / 4.0,
|
||||
is_month_start,
|
||||
is_month_end,
|
||||
]);
|
||||
}
|
||||
|
||||
let future_features = Array2::from_shape_vec((forecast_horizon, 10), fut_features)?;
|
||||
|
||||
// Targets: Multi-horizon price forecast (normalized)
|
||||
let targets: Vec<f64> = (0..forecast_horizon)
|
||||
.map(|t| bars[bar_idx + lookback_window + t].close / mean_price)
|
||||
.collect();
|
||||
|
||||
let target_array = Array1::from_vec(targets);
|
||||
|
||||
tft_samples.push((
|
||||
static_features,
|
||||
historical_features,
|
||||
future_features,
|
||||
target_array,
|
||||
));
|
||||
}
|
||||
|
||||
Ok(tft_samples)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_load_dbn_bars() {
|
||||
let dbn_file = "test_data/real/databento/ES.FUT_ohlcv-1m_2024-01-02.dbn";
|
||||
|
||||
if !std::path::Path::new(dbn_file).exists() {
|
||||
eprintln!("Skipping test: DBN file not found");
|
||||
return;
|
||||
}
|
||||
|
||||
let result = load_dbn_ohlcv_bars(dbn_file).await;
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"Failed to load DBN bars: {:?}",
|
||||
result.err()
|
||||
);
|
||||
|
||||
let bars = result.unwrap();
|
||||
assert!(!bars.is_empty(), "Should load bars");
|
||||
println!("✅ Loaded {} bars", bars.len());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_convert_to_tft_format() {
|
||||
let dbn_file = "test_data/real/databento/ES.FUT_ohlcv-1m_2024-01-02.dbn";
|
||||
|
||||
if !std::path::Path::new(dbn_file).exists() {
|
||||
eprintln!("Skipping test: DBN file not found");
|
||||
return;
|
||||
}
|
||||
|
||||
let bars = load_dbn_ohlcv_bars(dbn_file).await.unwrap();
|
||||
let feature_config = FeatureConfig::wave_e();
|
||||
let tft_data = convert_to_tft_data(&bars, 60, 10, &feature_config).unwrap();
|
||||
|
||||
assert!(!tft_data.is_empty(), "Should create TFT samples");
|
||||
|
||||
let (static_feat, hist_feat, fut_feat, targets) = &tft_data[0];
|
||||
|
||||
// Verify shapes
|
||||
assert_eq!(
|
||||
static_feat.len(),
|
||||
10,
|
||||
"Static features should have 10 dimensions"
|
||||
);
|
||||
assert_eq!(
|
||||
hist_feat.shape(),
|
||||
&[60, 225],
|
||||
"Historical features should be [60, 225] (Wave D)"
|
||||
);
|
||||
assert_eq!(
|
||||
fut_feat.shape(),
|
||||
&[10, 10],
|
||||
"Future features should be [10, 10]"
|
||||
);
|
||||
assert_eq!(targets.len(), 10, "Targets should have 10 timesteps");
|
||||
|
||||
println!("✅ TFT data structure validated:");
|
||||
println!(" • Static features: {:?}", static_feat.shape());
|
||||
println!(" • Historical features: {:?}", hist_feat.shape());
|
||||
println!(" • Future features: {:?}", fut_feat.shape());
|
||||
println!(" • Targets: {:?}", targets.shape());
|
||||
}
|
||||
}
|
||||
@@ -1,486 +0,0 @@
|
||||
//! TFT (Temporal Fusion Transformer) Training with Parquet Data
|
||||
//!
|
||||
//! Trains a TFT model using market data from Parquet files with lazy batch loading
|
||||
//! to avoid OOM issues on large datasets. Uses the TFTParquetExt trait for efficient
|
||||
//! memory management and 54-feature extraction (Wave E).
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (20 epochs)
|
||||
//! cargo run -p ml --example train_tft_parquet --release --features cuda
|
||||
//!
|
||||
//! # Custom configuration
|
||||
//! cargo run -p ml --example train_tft_parquet --release --features cuda -- \
|
||||
//! --parquet-file test_data/ES_FUT_180d.parquet \
|
||||
//! --epochs 50 \
|
||||
//! --batch-size 32 \
|
||||
//! --lookback-window 60 \
|
||||
//! --forecast-horizon 10
|
||||
//! ```
|
||||
//!
|
||||
//! # Features
|
||||
//!
|
||||
//! - Lazy batch loading (10,000 rows at a time) to avoid OOM crashes
|
||||
//! - 54-feature extraction (Wave E) from OHLCV bars
|
||||
//! - Sliding window creation (configurable lookback/horizon)
|
||||
//! - GPU-accelerated training (RTX 3050 Ti, 4GB VRAM)
|
||||
//! - Automatic train/validation split (80/20)
|
||||
//! - Model checkpointing and early stopping
|
||||
//!
|
||||
//! # Parquet Schema Requirements
|
||||
//!
|
||||
//! The Parquet file must follow Databento schema:
|
||||
//! - Column 3: open (Float64)
|
||||
//! - Column 4: high (Float64)
|
||||
//! - Column 5: low (Float64)
|
||||
//! - Column 6: close (Float64)
|
||||
//! - Column 7: volume (UInt64)
|
||||
//! - Column 9: ts_event (Timestamp[ns, UTC])
|
||||
|
||||
// Suppress warnings for unused dependencies in this example
|
||||
// (examples have access to all crate dependencies but typically only use a subset)
|
||||
#![allow(unused_crate_dependencies)]
|
||||
|
||||
// Use mimalloc allocator for 10-25% performance improvement
|
||||
#[cfg(feature = "mimalloc-allocator")]
|
||||
use mimalloc::MiMalloc;
|
||||
#[cfg(feature = "mimalloc-allocator")]
|
||||
#[global_allocator]
|
||||
static GLOBAL: MiMalloc = MiMalloc;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::path::PathBuf;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::checkpoint::FileSystemStorage;
|
||||
use ml::trainers::tft::{TFTTrainer, TFTTrainerConfig};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_tft_parquet",
|
||||
about = "Train TFT model on Parquet market data with lazy loading"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Parquet file path containing OHLCV bars (Databento schema)
|
||||
#[arg(long, default_value = "test_data/ES_FUT_small.parquet")]
|
||||
parquet_file: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "3")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value = "0.001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Batch size (max 32 for 4GB VRAM)
|
||||
#[arg(long, default_value = "32")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Validation batch size (defaults to match training batch_size)
|
||||
#[arg(long)]
|
||||
validation_batch_size: Option<usize>,
|
||||
|
||||
/// Hidden dimension
|
||||
#[arg(long, default_value = "256")]
|
||||
hidden_dim: usize,
|
||||
|
||||
/// Number of attention heads
|
||||
#[arg(long, default_value = "8")]
|
||||
num_attention_heads: usize,
|
||||
|
||||
/// Lookback window (historical bars)
|
||||
#[arg(long, default_value = "60")]
|
||||
lookback_window: usize,
|
||||
|
||||
/// Forecast horizon (future bars)
|
||||
#[arg(long, default_value = "10")]
|
||||
forecast_horizon: usize,
|
||||
|
||||
/// Dropout rate for regularization
|
||||
#[arg(long, default_value = "0.1")]
|
||||
dropout_rate: f64,
|
||||
|
||||
/// Number of LSTM layers
|
||||
#[arg(long, default_value = "2")]
|
||||
lstm_layers: usize,
|
||||
|
||||
/// Quantiles for probabilistic forecasting (comma-separated)
|
||||
#[arg(long, default_value = "0.1,0.5,0.9")]
|
||||
quantiles: String,
|
||||
|
||||
/// Output directory for trained model checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Use GPU for training (CUDA required)
|
||||
#[arg(long)]
|
||||
use_gpu: bool,
|
||||
|
||||
/// Use INT8 quantization for memory efficiency (reduces VRAM usage by 3-8x)
|
||||
#[arg(long)]
|
||||
use_int8: bool,
|
||||
|
||||
/// Use Quantization-Aware Training (1-2% better accuracy than PTQ)
|
||||
/// Trains with fake quantization, converts to INT8 at the end
|
||||
#[arg(long)]
|
||||
use_qat: bool,
|
||||
|
||||
/// Number of batches for QAT calibration (default: 100)
|
||||
/// Higher values improve accuracy but increase training time
|
||||
#[arg(long, default_value = "100")]
|
||||
qat_calibration_batches: usize,
|
||||
|
||||
/// Minimum batch size for QAT calibration OOM recovery (default: 2)
|
||||
/// If OOM occurs, batch size is halved automatically. Training aborts if below this threshold.
|
||||
#[arg(long, default_value = "2")]
|
||||
qat_min_batch_size: usize,
|
||||
|
||||
/// ⚠️ WARNING: Gradient checkpointing NOT IMPLEMENTED for QAT (see QAT_GRADIENT_CHECKPOINTING_WORKAROUND.md)
|
||||
///
|
||||
/// This flag is IGNORED when --use-qat is enabled. For QAT memory reduction:
|
||||
/// 1. Use 2-phase workaround (calibrate → freeze observers → train with checkpointing)
|
||||
/// 2. See ml/docs/QAT_GRADIENT_CHECKPOINTING_WORKAROUND.md for details
|
||||
///
|
||||
/// For non-QAT training: Reduces GPU memory usage by 30-40% but increases training time by ~20%
|
||||
#[arg(long)]
|
||||
use_gradient_checkpointing: bool,
|
||||
|
||||
/// Auto-detect optimal batch size based on available GPU memory
|
||||
/// Overrides --batch-size if enabled. Prevents OOM errors and maximizes GPU utilization.
|
||||
#[arg(long)]
|
||||
auto_batch_size: bool,
|
||||
|
||||
/// Verbose logging (debug level)
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
|
||||
/// Maximum validation batches to run (default: unlimited, use 50 for 4GB GPUs)
|
||||
/// Limits validation to N batches to reduce memory usage. Each batch uses ~10MB,
|
||||
/// so 50 batches = ~500MB vs 1760MB for full validation (176 batches).
|
||||
#[arg(long)]
|
||||
max_validation_batches: Option<usize>,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
#[cfg(feature = "mimalloc-allocator")]
|
||||
info!("🚀 Using mimalloc allocator for improved performance");
|
||||
#[cfg(not(feature = "mimalloc-allocator"))]
|
||||
info!("ℹ️ Using system allocator (consider --features mimalloc-allocator for 10-25% speedup)");
|
||||
info!("🚀 Starting TFT Training with Parquet Data (Lazy Loading)");
|
||||
info!("");
|
||||
info!("Configuration:");
|
||||
info!(" • Parquet file: {}", opts.parquet_file);
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Learning rate: {}", opts.learning_rate);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(
|
||||
" • Validation batch size: {}",
|
||||
opts.validation_batch_size.unwrap_or(opts.batch_size)
|
||||
);
|
||||
if let Some(max_val_batches) = opts.max_validation_batches {
|
||||
info!(
|
||||
" • Max validation batches: {} (memory optimization)",
|
||||
max_val_batches
|
||||
);
|
||||
} else {
|
||||
info!(" • Max validation batches: unlimited");
|
||||
}
|
||||
info!(" • Hidden dimension: {}", opts.hidden_dim);
|
||||
info!(" • Attention heads: {}", opts.num_attention_heads);
|
||||
info!(" • Lookback window: {}", opts.lookback_window);
|
||||
info!(" • Forecast horizon: {}", opts.forecast_horizon);
|
||||
info!(" • Dropout rate: {}", opts.dropout_rate);
|
||||
info!(" • LSTM layers: {}", opts.lstm_layers);
|
||||
info!(" • Quantiles: {}", opts.quantiles);
|
||||
info!(" • Feature count: 54 (Wave E)");
|
||||
info!(" • GPU enabled: {}", opts.use_gpu);
|
||||
info!(" • INT8 quantization: {}", opts.use_int8);
|
||||
info!(" • Quantization-Aware Training: {}", opts.use_qat);
|
||||
if opts.use_qat {
|
||||
info!(
|
||||
" • QAT calibration batches: {}",
|
||||
opts.qat_calibration_batches
|
||||
);
|
||||
}
|
||||
info!(
|
||||
" • Gradient checkpointing: {}",
|
||||
opts.use_gradient_checkpointing
|
||||
);
|
||||
if opts.use_gradient_checkpointing {
|
||||
if opts.use_qat {
|
||||
warn!("⚠️ WARNING: --use-gradient-checkpointing is IGNORED with --use-qat (not implemented)");
|
||||
warn!(" → For QAT memory reduction, use 2-phase workaround:");
|
||||
warn!(" → See ml/docs/QAT_GRADIENT_CHECKPOINTING_WORKAROUND.md");
|
||||
} else {
|
||||
info!(" → Expected: 30-40% memory reduction, ~20% slower training");
|
||||
}
|
||||
}
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
info!("");
|
||||
|
||||
// Verify Parquet file exists
|
||||
let parquet_path = PathBuf::from(&opts.parquet_file);
|
||||
if !parquet_path.exists() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Parquet file not found: {}",
|
||||
opts.parquet_file
|
||||
));
|
||||
}
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Parse quantiles
|
||||
let quantiles: Vec<f64> = opts
|
||||
.quantiles
|
||||
.split(',')
|
||||
.map(|s| s.trim().parse::<f64>())
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.context("Failed to parse quantiles")?;
|
||||
|
||||
if quantiles.is_empty() || quantiles.len() > 10 {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Invalid number of quantiles: {} (must be 1-10)",
|
||||
quantiles.len()
|
||||
));
|
||||
}
|
||||
|
||||
info!(
|
||||
"📊 Quantiles for probabilistic forecasting: {:?}",
|
||||
quantiles
|
||||
);
|
||||
|
||||
// Configure TFT trainer
|
||||
// Static features: 5 (symbol metadata)
|
||||
// Historical features: 39 (Wave E 54 features - 5 static - 10 known)
|
||||
// Future features: 10 (calendar features, time-based)
|
||||
// Total input features: 54 (5 + 10 + 39)
|
||||
let trainer_config = TFTTrainerConfig {
|
||||
epochs: opts.epochs,
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
auto_batch_size: opts.auto_batch_size,
|
||||
validation_batch_size: opts.validation_batch_size.unwrap_or(opts.batch_size),
|
||||
hidden_dim: opts.hidden_dim,
|
||||
num_attention_heads: opts.num_attention_heads,
|
||||
dropout_rate: opts.dropout_rate,
|
||||
lstm_layers: opts.lstm_layers,
|
||||
quantiles,
|
||||
lookback_window: opts.lookback_window,
|
||||
forecast_horizon: opts.forecast_horizon,
|
||||
use_gpu: opts.use_gpu,
|
||||
use_int8_quantization: opts.use_int8,
|
||||
use_qat: opts.use_qat,
|
||||
qat_calibration_batches: opts.qat_calibration_batches,
|
||||
qat_min_batch_size: opts.qat_min_batch_size,
|
||||
qat_warmup_epochs: 10, // Default: 10 epochs LR warmup after calibration
|
||||
qat_cooldown_factor: 0.1, // Default: 10x LR reduction in final 10% of training
|
||||
use_gradient_checkpointing: opts.use_gradient_checkpointing,
|
||||
max_validation_batches: opts.max_validation_batches,
|
||||
validation_frequency: 1, // Validate every epoch
|
||||
checkpoint_dir: opts.output_dir.clone(),
|
||||
};
|
||||
|
||||
// Create checkpoint storage
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(output_path.clone()));
|
||||
|
||||
// Create TFT trainer
|
||||
let mut trainer =
|
||||
TFTTrainer::new(trainer_config.clone(), storage).context("Failed to create TFT trainer")?;
|
||||
|
||||
info!(
|
||||
"✅ TFT trainer initialized with {} quantiles",
|
||||
trainer_config.quantiles.len()
|
||||
);
|
||||
|
||||
if opts.use_qat {
|
||||
info!("🧠 Quantization-Aware Training (QAT) enabled");
|
||||
info!(
|
||||
" Phase 1: Calibration ({} batches) - collecting activation statistics",
|
||||
opts.qat_calibration_batches
|
||||
);
|
||||
info!(" Phase 2: Training with fake quantization - simulating INT8 ops");
|
||||
info!(" Phase 3: Conversion to true INT8 model");
|
||||
info!(" Expected: 1-2% better accuracy than post-training quantization");
|
||||
} else if opts.use_int8 {
|
||||
info!("⚡ INT8 quantization enabled (PTQ mode) - expect 3-8x memory reduction");
|
||||
info!(" Memory usage: ~125MB (vs ~1GB FP32)");
|
||||
}
|
||||
|
||||
// Setup progress callback
|
||||
let (progress_tx, mut progress_rx) = mpsc::unbounded_channel();
|
||||
trainer.set_progress_callback(progress_tx);
|
||||
|
||||
// Spawn progress monitor task
|
||||
let monitor_task = tokio::spawn(async move {
|
||||
while let Some(progress) = progress_rx.recv().await {
|
||||
info!("{}", progress.message);
|
||||
if let Some(loss) = progress.metrics.get("train_loss") {
|
||||
info!(" • Train loss: {:.6}", loss);
|
||||
}
|
||||
if let Some(val_loss) = progress.metrics.get("val_loss") {
|
||||
info!(" • Val loss: {:.6}", val_loss);
|
||||
}
|
||||
if let Some(quantile_loss) = progress.metrics.get("quantile_loss") {
|
||||
info!(" • Quantile loss: {:.6}", quantile_loss);
|
||||
}
|
||||
if let Some(rmse) = progress.metrics.get("rmse") {
|
||||
info!(" • RMSE: {:.6}", rmse);
|
||||
}
|
||||
if let Some(attention_entropy) = progress.metrics.get("attention_entropy") {
|
||||
info!(" • Attention entropy: {:.4}", attention_entropy);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Train the model using lazy-loading Parquet pipeline
|
||||
info!("");
|
||||
info!("🏋️ Starting training with lazy-loading Parquet pipeline...");
|
||||
info!(" (Loading 10,000 rows at a time to avoid OOM)");
|
||||
info!("");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let final_metrics = trainer
|
||||
.train_from_parquet(&opts.parquet_file)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Wait for progress monitor to finish
|
||||
drop(trainer); // Drop trainer to close progress channel
|
||||
let _ = monitor_task.await;
|
||||
|
||||
// Print final metrics
|
||||
info!("");
|
||||
info!("✅ Training completed successfully!");
|
||||
info!("");
|
||||
info!("📊 Final Metrics:");
|
||||
info!(" • Training loss: {:.6}", final_metrics.train_loss);
|
||||
info!(" • Validation loss: {:.6}", final_metrics.val_loss);
|
||||
info!(" • Quantile loss: {:.6}", final_metrics.quantile_loss);
|
||||
info!(" • RMSE: {:.6}", final_metrics.rmse);
|
||||
info!(
|
||||
" • Attention entropy: {:.4}",
|
||||
final_metrics.attention_entropy
|
||||
);
|
||||
info!(
|
||||
" • Training duration: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
info!(
|
||||
" • Reported training time: {:.1}s ({:.1} min)",
|
||||
final_metrics.training_time_seconds,
|
||||
final_metrics.training_time_seconds / 60.0
|
||||
);
|
||||
info!("");
|
||||
info!("💾 Model checkpoints saved to: {}", opts.output_dir);
|
||||
info!("");
|
||||
info!("🎉 TFT training with Parquet data complete!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_cli_parsing() {
|
||||
// Test default parameters
|
||||
let args = vec!["train_tft_parquet"];
|
||||
let opts = Opts::try_parse_from(args).expect("Failed to parse default args");
|
||||
|
||||
assert_eq!(opts.parquet_file, "test_data/ES_FUT_180d.parquet");
|
||||
assert_eq!(opts.epochs, 20);
|
||||
assert_eq!(opts.batch_size, 32);
|
||||
assert_eq!(opts.lookback_window, 60);
|
||||
assert_eq!(opts.forecast_horizon, 10);
|
||||
assert_eq!(opts.learning_rate, 0.001);
|
||||
assert!(!opts.use_gpu); // Default is false (flag not set)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cli_custom_parameters() {
|
||||
// Test custom parameters with GPU enabled
|
||||
let args = vec![
|
||||
"train_tft_parquet",
|
||||
"--parquet-file",
|
||||
"test_data/NQ_FUT_180d.parquet",
|
||||
"--epochs",
|
||||
"50",
|
||||
"--batch-size",
|
||||
"16",
|
||||
"--lookback-window",
|
||||
"120",
|
||||
"--forecast-horizon",
|
||||
"20",
|
||||
"--learning-rate",
|
||||
"0.0005",
|
||||
"--use-gpu", // Flag to enable GPU
|
||||
"--verbose",
|
||||
];
|
||||
|
||||
let opts = Opts::try_parse_from(args).expect("Failed to parse custom args");
|
||||
|
||||
assert_eq!(opts.parquet_file, "test_data/NQ_FUT_180d.parquet");
|
||||
assert_eq!(opts.epochs, 50);
|
||||
assert_eq!(opts.batch_size, 16);
|
||||
assert_eq!(opts.lookback_window, 120);
|
||||
assert_eq!(opts.forecast_horizon, 20);
|
||||
assert_eq!(opts.learning_rate, 0.0005);
|
||||
assert!(opts.use_gpu); // GPU should be enabled
|
||||
assert!(opts.verbose);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quantile_parsing() {
|
||||
let quantiles_str = "0.1,0.5,0.9";
|
||||
let quantiles: Vec<f64> = quantiles_str
|
||||
.split(',')
|
||||
.map(|s| s.trim().parse::<f64>())
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.expect("Failed to parse quantiles");
|
||||
|
||||
assert_eq!(quantiles.len(), 3);
|
||||
assert_eq!(quantiles[0], 0.1);
|
||||
assert_eq!(quantiles[1], 0.5);
|
||||
assert_eq!(quantiles[2], 0.9);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_invalid_quantiles() {
|
||||
let quantiles_str = "invalid,0.5,0.9";
|
||||
let result: Result<Vec<f64>, _> = quantiles_str
|
||||
.split(',')
|
||||
.map(|s| s.trim().parse::<f64>())
|
||||
.collect();
|
||||
|
||||
assert!(result.is_err(), "Should fail on invalid quantile");
|
||||
}
|
||||
}
|
||||
@@ -1,529 +0,0 @@
|
||||
//! TFT Quantization-Aware Training (QAT) Example
|
||||
//!
|
||||
//! Demonstrates how to use QAT to train a TFT model with INT8 quantization
|
||||
//! for better accuracy compared to post-training quantization (PTQ).
|
||||
//!
|
||||
//! # QAT vs PTQ
|
||||
//!
|
||||
//! - **PTQ (Post-Training Quantization)**: Quantize weights after training
|
||||
//! - Pros: Fast, no retraining required
|
||||
//! - Cons: Can lose 2-5% accuracy on complex models
|
||||
//!
|
||||
//! - **QAT (Quantization-Aware Training)**: Train with simulated quantization
|
||||
//! - Pros: 1-2% better accuracy than PTQ, model learns to compensate
|
||||
//! - Cons: Slower training (adds fake quantization ops)
|
||||
//!
|
||||
//! # Three-Phase QAT Process
|
||||
//!
|
||||
//! 1. **Calibration Phase**: Collect activation statistics (100 batches)
|
||||
//! 2. **Training Phase**: Train with fake quantization (simulates INT8 ops)
|
||||
//! 3. **Conversion Phase**: Convert FP32 model to true INT8 model
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Basic QAT training
|
||||
//! cargo run -p ml --example train_tft_qat --release --features cuda
|
||||
//!
|
||||
//! # Custom calibration batches (higher = better accuracy, slower training)
|
||||
//! cargo run -p ml --example train_tft_qat --release --features cuda -- \
|
||||
//! --parquet-file test_data/ES_FUT_180d.parquet \
|
||||
//! --epochs 50 \
|
||||
//! --qat-calibration-batches 200
|
||||
//!
|
||||
//! # Compare FP32 vs PTQ vs QAT accuracy
|
||||
//! cargo run -p ml --example train_tft_qat --release --features cuda -- \
|
||||
//! --compare-accuracy
|
||||
//! ```
|
||||
//!
|
||||
//! # Expected Results
|
||||
//!
|
||||
//! - Training time: 1.2-1.5x slower than FP32 (due to fake quantization)
|
||||
//! - Memory usage: Same as FP32 during training, 3-8x reduction after conversion
|
||||
//! - Accuracy: 1-2% better than PTQ, within 0.5% of FP32
|
||||
//! - Final model: INT8 quantized (~125MB vs ~1GB FP32)
|
||||
|
||||
// Suppress warnings for unused dependencies in this example
|
||||
#![allow(unused_crate_dependencies)]
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::path::PathBuf;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::info;
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::checkpoint::FileSystemStorage;
|
||||
use ml::trainers::tft::{TFTTrainer, TFTTrainerConfig};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_tft_qat",
|
||||
about = "Train TFT with Quantization-Aware Training (QAT) for better INT8 accuracy"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Parquet file path containing OHLCV bars (Databento schema)
|
||||
#[arg(long, default_value = "test_data/ES_FUT_small.parquet")]
|
||||
parquet_file: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "20")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value = "0.001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Batch size (max 32 for 4GB VRAM)
|
||||
#[arg(long, default_value = "32")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Number of batches for QAT calibration (default: 100)
|
||||
/// Higher values improve accuracy but increase training time
|
||||
/// Recommended range: 50-500 batches
|
||||
#[arg(long, default_value = "100")]
|
||||
qat_calibration_batches: usize,
|
||||
|
||||
/// Output directory for trained model checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Use GPU for training (CUDA required)
|
||||
#[arg(long)]
|
||||
use_gpu: bool,
|
||||
|
||||
/// Compare FP32 vs PTQ vs QAT accuracy (trains 3 models)
|
||||
#[arg(long)]
|
||||
compare_accuracy: bool,
|
||||
|
||||
/// Verbose logging (debug level)
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 TFT Quantization-Aware Training (QAT) Example");
|
||||
info!("");
|
||||
info!("This example demonstrates the three-phase QAT process:");
|
||||
info!(
|
||||
" 1. Calibration: Collect activation statistics ({} batches)",
|
||||
opts.qat_calibration_batches
|
||||
);
|
||||
info!(" 2. Training: Train with fake quantization (simulates INT8)");
|
||||
info!(" 3. Conversion: Convert FP32 model to true INT8 model");
|
||||
info!("");
|
||||
|
||||
if opts.compare_accuracy {
|
||||
// Train 3 models and compare accuracy
|
||||
info!("📊 Running accuracy comparison: FP32 vs PTQ vs QAT");
|
||||
info!("");
|
||||
run_accuracy_comparison(&opts).await?;
|
||||
} else {
|
||||
// Train single QAT model
|
||||
info!("🧠 Training QAT model...");
|
||||
info!("");
|
||||
run_qat_training(&opts).await?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Train a single QAT model and show detailed phase logging
|
||||
async fn run_qat_training(opts: &Opts) -> Result<()> {
|
||||
info!("Configuration:");
|
||||
info!(" • Parquet file: {}", opts.parquet_file);
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Learning rate: {}", opts.learning_rate);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(
|
||||
" • QAT calibration batches: {}",
|
||||
opts.qat_calibration_batches
|
||||
);
|
||||
info!(" • GPU enabled: {}", opts.use_gpu);
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
info!("");
|
||||
|
||||
// Verify Parquet file exists
|
||||
let parquet_path = PathBuf::from(&opts.parquet_file);
|
||||
if !parquet_path.exists() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Parquet file not found: {}",
|
||||
opts.parquet_file
|
||||
));
|
||||
}
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Configure TFT trainer with QAT enabled
|
||||
let trainer_config = TFTTrainerConfig {
|
||||
epochs: opts.epochs,
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
validation_batch_size: opts.batch_size,
|
||||
hidden_dim: 256,
|
||||
num_attention_heads: 8,
|
||||
dropout_rate: 0.1,
|
||||
lstm_layers: 2,
|
||||
quantiles: vec![0.1, 0.5, 0.9],
|
||||
lookback_window: 60,
|
||||
forecast_horizon: 10,
|
||||
use_gpu: opts.use_gpu,
|
||||
use_int8_quantization: true, // Enable INT8 quantization
|
||||
use_qat: true, // Enable QAT (the key difference!)
|
||||
qat_calibration_batches: opts.qat_calibration_batches,
|
||||
validation_frequency: 1,
|
||||
checkpoint_dir: opts.output_dir.clone(),
|
||||
};
|
||||
|
||||
// Create checkpoint storage
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(output_path.clone()));
|
||||
|
||||
// Create TFT trainer
|
||||
let mut trainer =
|
||||
TFTTrainer::new(trainer_config.clone(), storage).context("Failed to create TFT trainer")?;
|
||||
|
||||
info!("✅ TFT trainer initialized with QAT enabled");
|
||||
info!("");
|
||||
info!("📋 QAT Training Process:");
|
||||
info!("");
|
||||
info!(
|
||||
"Phase 1: Calibration ({} batches)",
|
||||
opts.qat_calibration_batches
|
||||
);
|
||||
info!(" • Insert fake quantization nodes in model graph");
|
||||
info!(" • Run forward passes to collect activation statistics");
|
||||
info!(" • Compute optimal scale/zero-point for each layer");
|
||||
info!(" • No gradient updates (calibration only)");
|
||||
info!("");
|
||||
info!("Phase 2: Training with Fake Quantization");
|
||||
info!(" • Forward pass: Simulate INT8 operations (FP32→INT8→FP32)");
|
||||
info!(" • Backward pass: Standard FP32 gradients");
|
||||
info!(" • Model learns to compensate for quantization errors");
|
||||
info!(" • Training time: ~1.2-1.5x slower than FP32");
|
||||
info!("");
|
||||
info!("Phase 3: Conversion to True INT8");
|
||||
info!(" • Extract FP32 weights from trained model");
|
||||
info!(" • Quantize weights using calibrated scales");
|
||||
info!(" • Create INT8 model (3-8x memory reduction)");
|
||||
info!(" • Expect 1-2% better accuracy than PTQ");
|
||||
info!("");
|
||||
|
||||
// Setup progress callback
|
||||
let (progress_tx, mut progress_rx) = mpsc::unbounded_channel();
|
||||
trainer.set_progress_callback(progress_tx);
|
||||
|
||||
// Spawn progress monitor task
|
||||
let monitor_task = tokio::spawn(async move {
|
||||
while let Some(progress) = progress_rx.recv().await {
|
||||
info!("{}", progress.message);
|
||||
if let Some(loss) = progress.metrics.get("train_loss") {
|
||||
info!(" • Train loss: {:.6}", loss);
|
||||
}
|
||||
if let Some(val_loss) = progress.metrics.get("val_loss") {
|
||||
info!(" • Val loss: {:.6}", val_loss);
|
||||
}
|
||||
if let Some(quantile_loss) = progress.metrics.get("quantile_loss") {
|
||||
info!(" • Quantile loss: {:.6}", quantile_loss);
|
||||
}
|
||||
if let Some(rmse) = progress.metrics.get("rmse") {
|
||||
info!(" • RMSE: {:.6}", rmse);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Train the model with QAT
|
||||
info!("🏋️ Starting QAT training...");
|
||||
info!("");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let final_metrics = trainer
|
||||
.train_from_parquet(&opts.parquet_file)
|
||||
.await
|
||||
.context("QAT training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Wait for progress monitor to finish
|
||||
drop(trainer);
|
||||
let _ = monitor_task.await;
|
||||
|
||||
// Print final metrics
|
||||
info!("");
|
||||
info!("✅ QAT Training completed successfully!");
|
||||
info!("");
|
||||
info!("📊 Final Metrics:");
|
||||
info!(" • Training loss: {:.6}", final_metrics.train_loss);
|
||||
info!(" • Validation loss: {:.6}", final_metrics.val_loss);
|
||||
info!(" • Quantile loss: {:.6}", final_metrics.quantile_loss);
|
||||
info!(" • RMSE: {:.6}", final_metrics.rmse);
|
||||
info!(
|
||||
" • Attention entropy: {:.4}",
|
||||
final_metrics.attention_entropy
|
||||
);
|
||||
info!(
|
||||
" • Training duration: {:.1}s ({:.1} min)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0
|
||||
);
|
||||
info!("");
|
||||
info!("💾 Quantized model saved to: {}", opts.output_dir);
|
||||
info!(" Memory footprint: ~125MB (vs ~1GB FP32)");
|
||||
info!(" Expected accuracy: Within 0.5% of FP32 model");
|
||||
info!("");
|
||||
info!("🎉 QAT training complete!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Train 3 models (FP32, PTQ, QAT) and compare their accuracy
|
||||
async fn run_accuracy_comparison(opts: &Opts) -> Result<()> {
|
||||
info!("Training 3 models for accuracy comparison:");
|
||||
info!(" 1. FP32 Baseline (no quantization)");
|
||||
info!(" 2. PTQ (Post-Training Quantization)");
|
||||
info!(" 3. QAT (Quantization-Aware Training)");
|
||||
info!("");
|
||||
|
||||
// Verify Parquet file exists
|
||||
let parquet_path = PathBuf::from(&opts.parquet_file);
|
||||
if !parquet_path.exists() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"Parquet file not found: {}",
|
||||
opts.parquet_file
|
||||
));
|
||||
}
|
||||
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path)?;
|
||||
}
|
||||
|
||||
// 1. Train FP32 baseline
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("1️⃣ Training FP32 Baseline Model");
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("");
|
||||
|
||||
let fp32_config = TFTTrainerConfig {
|
||||
epochs: opts.epochs,
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
validation_batch_size: opts.batch_size,
|
||||
hidden_dim: 256,
|
||||
num_attention_heads: 8,
|
||||
dropout_rate: 0.1,
|
||||
lstm_layers: 2,
|
||||
quantiles: vec![0.1, 0.5, 0.9],
|
||||
lookback_window: 60,
|
||||
forecast_horizon: 10,
|
||||
use_gpu: opts.use_gpu,
|
||||
use_int8_quantization: false, // FP32 only
|
||||
use_qat: false,
|
||||
qat_calibration_batches: 0,
|
||||
validation_frequency: 1,
|
||||
checkpoint_dir: format!("{}/fp32", opts.output_dir),
|
||||
};
|
||||
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(PathBuf::from(
|
||||
fp32_config.checkpoint_dir.clone(),
|
||||
)));
|
||||
let mut fp32_trainer = TFTTrainer::new(fp32_config.clone(), storage)?;
|
||||
|
||||
let fp32_start = std::time::Instant::now();
|
||||
let fp32_metrics = fp32_trainer.train_from_parquet(&opts.parquet_file).await?;
|
||||
let fp32_duration = fp32_start.elapsed();
|
||||
|
||||
info!("");
|
||||
info!("✅ FP32 training complete!");
|
||||
info!(" Val loss: {:.6}", fp32_metrics.val_loss);
|
||||
info!(" RMSE: {:.6}", fp32_metrics.rmse);
|
||||
info!(" Time: {:.1}s", fp32_duration.as_secs_f64());
|
||||
info!("");
|
||||
|
||||
// 2. Train PTQ model
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("2️⃣ Training PTQ Model (Post-Training Quantization)");
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("");
|
||||
|
||||
let ptq_config = TFTTrainerConfig {
|
||||
use_int8_quantization: true, // PTQ enabled
|
||||
use_qat: false, // No QAT
|
||||
qat_calibration_batches: 0,
|
||||
checkpoint_dir: format!("{}/ptq", opts.output_dir),
|
||||
..fp32_config.clone()
|
||||
};
|
||||
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(PathBuf::from(
|
||||
ptq_config.checkpoint_dir.clone(),
|
||||
)));
|
||||
let mut ptq_trainer = TFTTrainer::new(ptq_config.clone(), storage)?;
|
||||
|
||||
let ptq_start = std::time::Instant::now();
|
||||
let ptq_metrics = ptq_trainer.train_from_parquet(&opts.parquet_file).await?;
|
||||
let ptq_duration = ptq_start.elapsed();
|
||||
|
||||
info!("");
|
||||
info!("✅ PTQ training complete!");
|
||||
info!(" Val loss: {:.6}", ptq_metrics.val_loss);
|
||||
info!(" RMSE: {:.6}", ptq_metrics.rmse);
|
||||
info!(" Time: {:.1}s", ptq_duration.as_secs_f64());
|
||||
info!("");
|
||||
|
||||
// 3. Train QAT model
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("3️⃣ Training QAT Model (Quantization-Aware Training)");
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("");
|
||||
|
||||
let qat_config = TFTTrainerConfig {
|
||||
use_int8_quantization: true, // INT8 enabled
|
||||
use_qat: true, // QAT enabled (the key difference!)
|
||||
qat_calibration_batches: opts.qat_calibration_batches,
|
||||
checkpoint_dir: format!("{}/qat", opts.output_dir),
|
||||
..fp32_config.clone()
|
||||
};
|
||||
|
||||
let storage = std::sync::Arc::new(FileSystemStorage::new(PathBuf::from(
|
||||
qat_config.checkpoint_dir.clone(),
|
||||
)));
|
||||
let mut qat_trainer = TFTTrainer::new(qat_config.clone(), storage)?;
|
||||
|
||||
let qat_start = std::time::Instant::now();
|
||||
let qat_metrics = qat_trainer.train_from_parquet(&opts.parquet_file).await?;
|
||||
let qat_duration = qat_start.elapsed();
|
||||
|
||||
info!("");
|
||||
info!("✅ QAT training complete!");
|
||||
info!(" Val loss: {:.6}", qat_metrics.val_loss);
|
||||
info!(" RMSE: {:.6}", qat_metrics.rmse);
|
||||
info!(" Time: {:.1}s", qat_duration.as_secs_f64());
|
||||
info!("");
|
||||
|
||||
// Print comparison table
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("📊 Accuracy Comparison Results");
|
||||
info!("━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━");
|
||||
info!("");
|
||||
info!("┌──────────┬─────────────┬───────────┬───────────┬─────────────┐");
|
||||
info!("│ Model │ Val Loss │ RMSE │ Time │ Memory │");
|
||||
info!("├──────────┼─────────────┼───────────┼───────────┼─────────────┤");
|
||||
info!(
|
||||
"│ FP32 │ {:.6} │ {:.6} │ {:>6.1}s │ ~1000MB │",
|
||||
fp32_metrics.val_loss,
|
||||
fp32_metrics.rmse,
|
||||
fp32_duration.as_secs_f64()
|
||||
);
|
||||
info!(
|
||||
"│ PTQ │ {:.6} │ {:.6} │ {:>6.1}s │ ~125MB │",
|
||||
ptq_metrics.val_loss,
|
||||
ptq_metrics.rmse,
|
||||
ptq_duration.as_secs_f64()
|
||||
);
|
||||
info!(
|
||||
"│ QAT │ {:.6} │ {:.6} │ {:>6.1}s │ ~125MB │",
|
||||
qat_metrics.val_loss,
|
||||
qat_metrics.rmse,
|
||||
qat_duration.as_secs_f64()
|
||||
);
|
||||
info!("└──────────┴─────────────┴───────────┴───────────┴─────────────┘");
|
||||
info!("");
|
||||
|
||||
// Calculate improvements
|
||||
let ptq_loss_delta =
|
||||
((ptq_metrics.val_loss - fp32_metrics.val_loss) / fp32_metrics.val_loss) * 100.0;
|
||||
let qat_loss_delta =
|
||||
((qat_metrics.val_loss - fp32_metrics.val_loss) / fp32_metrics.val_loss) * 100.0;
|
||||
let qat_vs_ptq_improvement =
|
||||
((ptq_metrics.val_loss - qat_metrics.val_loss) / ptq_metrics.val_loss) * 100.0;
|
||||
|
||||
info!("📈 Analysis:");
|
||||
info!("");
|
||||
info!(" PTQ vs FP32:");
|
||||
info!(" • Loss degradation: {:.2}%", ptq_loss_delta);
|
||||
info!(" • Memory reduction: 8x (1000MB → 125MB)");
|
||||
info!(" • Training time: Same as FP32");
|
||||
info!("");
|
||||
info!(" QAT vs FP32:");
|
||||
info!(" • Loss degradation: {:.2}%", qat_loss_delta);
|
||||
info!(" • Memory reduction: 8x (1000MB → 125MB)");
|
||||
info!(
|
||||
" • Training time: {:.1}x slower",
|
||||
qat_duration.as_secs_f64() / fp32_duration.as_secs_f64()
|
||||
);
|
||||
info!("");
|
||||
info!(" QAT vs PTQ:");
|
||||
info!(" • Accuracy improvement: {:.2}%", qat_vs_ptq_improvement);
|
||||
info!(" • Same memory footprint (~125MB)");
|
||||
info!(" • Training overhead: Worth it for production models!");
|
||||
info!("");
|
||||
info!("💡 Recommendation:");
|
||||
if qat_vs_ptq_improvement > 1.0 {
|
||||
info!(
|
||||
" ✅ Use QAT for production - {:.1}% better accuracy is worth the training time",
|
||||
qat_vs_ptq_improvement
|
||||
);
|
||||
} else {
|
||||
info!(
|
||||
" ⚠️ PTQ may be sufficient - QAT improvement is only {:.1}%",
|
||||
qat_vs_ptq_improvement
|
||||
);
|
||||
}
|
||||
info!("");
|
||||
info!("🎉 Comparison complete!");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_cli_parsing() {
|
||||
let args = vec!["train_tft_qat"];
|
||||
let opts = Opts::try_parse_from(args).expect("Failed to parse default args");
|
||||
|
||||
assert_eq!(opts.parquet_file, "test_data/ES_FUT_small.parquet");
|
||||
assert_eq!(opts.epochs, 20);
|
||||
assert_eq!(opts.qat_calibration_batches, 100);
|
||||
assert!(!opts.compare_accuracy);
|
||||
assert!(!opts.use_gpu);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cli_with_qat_options() {
|
||||
let args = vec![
|
||||
"train_tft_qat",
|
||||
"--qat-calibration-batches",
|
||||
"200",
|
||||
"--compare-accuracy",
|
||||
"--use-gpu",
|
||||
];
|
||||
|
||||
let opts = Opts::try_parse_from(args).expect("Failed to parse QAT args");
|
||||
|
||||
assert_eq!(opts.qat_calibration_batches, 200);
|
||||
assert!(opts.compare_accuracy);
|
||||
assert!(opts.use_gpu);
|
||||
}
|
||||
}
|
||||
@@ -1,536 +0,0 @@
|
||||
//! TGGN (Temporal Graph Gated Network) Training with Real DBN Market Data
|
||||
//!
|
||||
//! Trains a TGGN model using the UnifiedTrainable adapter on real market data
|
||||
//! from Databento DBN files. The adapter wraps a candle-based projection network
|
||||
//! (input -> hidden -> scalar prediction) trained with MSE loss via AdamW.
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Default: 20 epochs on ES data
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_tggn_dbn --release -- \
|
||||
//! --data-dir data/cache/futures-baseline --symbol ES
|
||||
//!
|
||||
//! # Custom configuration
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_tggn_dbn --release -- \
|
||||
//! --data-dir data/cache/futures-baseline \
|
||||
//! --symbol ES \
|
||||
//! --epochs 50 \
|
||||
//! --batch-size 64 \
|
||||
//! --learning-rate 0.0005 \
|
||||
//! --max-steps-per-epoch 2000
|
||||
//! ```
|
||||
//!
|
||||
//! # Output
|
||||
//!
|
||||
//! Checkpoints saved to `--output-dir` (default `ml/trained_models/tggn`):
|
||||
//! - `tggn_epoch_N.safetensors` every 5 epochs
|
||||
//! - `tggn_final.safetensors` at the end of training
|
||||
|
||||
#![allow(unused_crate_dependencies)]
|
||||
#![deny(
|
||||
clippy::unwrap_used,
|
||||
clippy::expect_used,
|
||||
clippy::panic,
|
||||
clippy::indexing_slicing
|
||||
)]
|
||||
|
||||
use std::path::PathBuf;
|
||||
use std::time::Instant;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::{Device, Tensor};
|
||||
use clap::Parser;
|
||||
use tracing::info;
|
||||
|
||||
use ml::tgnn::trainable_adapter::TGGNTrainableAdapter;
|
||||
use ml::tgnn::TGGNConfig;
|
||||
use ml::training::unified_trainer::UnifiedTrainable;
|
||||
|
||||
#[allow(unreachable_pub)]
|
||||
mod baseline_common;
|
||||
use baseline_common::{load_all_bars, spread_cost_bps};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CLI Arguments
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Train TGGN on real Databento OHLCV data via the UnifiedTrainable adapter.
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "train_tggn_dbn", about = "Train TGGN on real DBN market data")]
|
||||
struct Args {
|
||||
/// Directory containing per-symbol subdirectories of .dbn.zst files
|
||||
#[arg(long, default_value = "data/cache/futures-baseline")]
|
||||
data_dir: PathBuf,
|
||||
|
||||
/// Symbol subdirectory to load (e.g. ES, NQ, 6E, ZN)
|
||||
#[arg(long, default_value = "ES")]
|
||||
symbol: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value_t = 20)]
|
||||
epochs: usize,
|
||||
|
||||
/// Batch size for training
|
||||
#[arg(long, default_value_t = 64)]
|
||||
batch_size: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value_t = 0.001)]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Maximum training steps per epoch (0 = unlimited)
|
||||
#[arg(long, default_value_t = 2000)]
|
||||
max_steps_per_epoch: usize,
|
||||
|
||||
/// Output directory for checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models/tggn")]
|
||||
output_dir: PathBuf,
|
||||
|
||||
/// Round-trip commission cost in basis points
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
tx_cost_bps: f64,
|
||||
|
||||
/// Instrument tick size in price units (ES=0.25)
|
||||
#[arg(long, default_value_t = 0.25)]
|
||||
tick_size: f64,
|
||||
|
||||
/// Typical bid-ask spread in ticks
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
spread_ticks: f64,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Feature engineering helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Number of features we extract per bar (node_dim for TGGN).
|
||||
const NODE_DIM: usize = 8;
|
||||
|
||||
/// Extract a feature vector from a single OHLCV bar.
|
||||
///
|
||||
/// Features (8-dim):
|
||||
/// 0: log-return from previous close (0.0 for the first bar)
|
||||
/// 1: normalized range (high - low) / close
|
||||
/// 2: normalized body (close - open) / close
|
||||
/// 3: log(volume + 1) (clamped)
|
||||
/// 4: upper-wick ratio (high - max(open,close)) / (high - low + 1e-10)
|
||||
/// 5: lower-wick ratio (min(open,close) - low) / (high - low + 1e-10)
|
||||
/// 6: bar midpoint relative to close: (high + low) / 2 / close - 1
|
||||
/// 7: volume-price product (log scale)
|
||||
fn bar_features(bar: &ml::types::OHLCVBar, prev_close: Option<f64>) -> [f64; NODE_DIM] {
|
||||
let close = bar.close;
|
||||
let open = bar.open;
|
||||
let high = bar.high;
|
||||
let low = bar.low;
|
||||
let volume = bar.volume;
|
||||
|
||||
let log_return = match prev_close {
|
||||
Some(pc) if pc.abs() > 1e-12 => (close / pc).ln(),
|
||||
_ => 0.0,
|
||||
};
|
||||
|
||||
let range = high - low;
|
||||
let safe_range = if range.abs() < 1e-10 { 1e-10 } else { range };
|
||||
let safe_close = if close.abs() < 1e-10 { 1e-10 } else { close };
|
||||
|
||||
let norm_range = range / safe_close;
|
||||
let norm_body = (close - open) / safe_close;
|
||||
let log_vol = (volume + 1.0).ln();
|
||||
let upper_wick = (high - open.max(close)) / safe_range;
|
||||
let lower_wick = (open.min(close) - low) / safe_range;
|
||||
let mid_rel = (high + low) / 2.0 / safe_close - 1.0;
|
||||
let vol_price = (volume * close.abs() + 1.0).ln();
|
||||
|
||||
[
|
||||
log_return, norm_range, norm_body, log_vol, upper_wick, lower_wick, mid_rel, vol_price,
|
||||
]
|
||||
}
|
||||
|
||||
/// Build (input, target) pairs from OHLCV bars.
|
||||
///
|
||||
/// Each sample uses features from bar[i] as input and the log-return at bar[i+1]
|
||||
/// as the scalar target. Returns tensors on `device`.
|
||||
fn build_dataset(
|
||||
bars: &[ml::types::OHLCVBar],
|
||||
device: &Device,
|
||||
tx_cost_bps: f64,
|
||||
tick_size: f64,
|
||||
spread_ticks: f64,
|
||||
) -> Result<Vec<(Tensor, Tensor)>> {
|
||||
if bars.len() < 3 {
|
||||
anyhow::bail!("Need at least 3 bars to build training data, got {}", bars.len());
|
||||
}
|
||||
|
||||
let mut inputs: Vec<[f64; NODE_DIM]> = Vec::with_capacity(bars.len().saturating_sub(2));
|
||||
let mut targets: Vec<f64> = Vec::with_capacity(bars.len().saturating_sub(2));
|
||||
|
||||
// We need bar[i-1] for prev_close, bar[i] for features, bar[i+1] for target
|
||||
for i in 1..bars.len().saturating_sub(1) {
|
||||
let prev_bar = bars.get(i.wrapping_sub(1));
|
||||
let cur_bar = match bars.get(i) {
|
||||
Some(b) => b,
|
||||
None => continue,
|
||||
};
|
||||
let next_bar = match bars.get(i + 1) {
|
||||
Some(b) => b,
|
||||
None => continue,
|
||||
};
|
||||
|
||||
let prev_close = prev_bar.map(|b| b.close);
|
||||
let feats = bar_features(cur_bar, prev_close);
|
||||
let cur_close = cur_bar.close;
|
||||
|
||||
// Target: next-bar log return minus transaction and spread costs (in log-return units)
|
||||
let raw_return = if cur_close.abs() > 1e-12 {
|
||||
(next_bar.close / cur_close).ln()
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
let spread = spread_cost_bps(cur_close, tick_size, spread_ticks);
|
||||
let cost = (tx_cost_bps + spread) / 10_000.0; // convert bps to fractional
|
||||
let target = raw_return - cost;
|
||||
|
||||
inputs.push(feats);
|
||||
targets.push(target);
|
||||
}
|
||||
|
||||
if inputs.is_empty() {
|
||||
anyhow::bail!("No training samples generated from {} bars", bars.len());
|
||||
}
|
||||
|
||||
// Convert to flat f32 vecs for Tensor creation
|
||||
let n = inputs.len();
|
||||
let flat_inputs: Vec<f32> = inputs
|
||||
.iter()
|
||||
.flat_map(|f| f.iter().map(|&v| v as f32))
|
||||
.collect();
|
||||
let flat_targets: Vec<f32> = targets.iter().map(|&v| v as f32).collect();
|
||||
|
||||
let input_tensor =
|
||||
Tensor::from_vec(flat_inputs, &[n, NODE_DIM], device).map_err(|e| {
|
||||
anyhow::anyhow!("Failed to create input tensor: {}", e)
|
||||
})?;
|
||||
let target_tensor =
|
||||
Tensor::from_vec(flat_targets, &[n, 1], device).map_err(|e| {
|
||||
anyhow::anyhow!("Failed to create target tensor: {}", e)
|
||||
})?;
|
||||
|
||||
// Split into individual (batch=1) pairs for the validate() API, but keep
|
||||
// full tensors for batch training. We'll return a single pair per entry
|
||||
// only for validation; training uses sliced batches below.
|
||||
Ok(vec![(input_tensor, target_tensor)])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Main
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
fn main() -> Result<()> {
|
||||
// Initialize tracing
|
||||
tracing_subscriber::fmt()
|
||||
.with_max_level(tracing::Level::INFO)
|
||||
.with_target(false)
|
||||
.with_thread_ids(false)
|
||||
.init();
|
||||
|
||||
let args = Args::parse();
|
||||
|
||||
info!("================================================================");
|
||||
info!(" TGGN Training on Real DBN Market Data");
|
||||
info!("================================================================");
|
||||
info!(" Data dir: {}", args.data_dir.display());
|
||||
info!(" Symbol: {}", args.symbol);
|
||||
info!(" Epochs: {}", args.epochs);
|
||||
info!(" Batch size: {}", args.batch_size);
|
||||
info!(" Learning rate: {}", args.learning_rate);
|
||||
info!(" Max steps/epoch: {}", args.max_steps_per_epoch);
|
||||
info!(" Output dir: {}", args.output_dir.display());
|
||||
info!("================================================================");
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Device selection
|
||||
// -----------------------------------------------------------------------
|
||||
let device = match Device::new_cuda(0) {
|
||||
Ok(d) => {
|
||||
info!("Using CUDA GPU");
|
||||
d
|
||||
}
|
||||
Err(_) => {
|
||||
info!("CUDA unavailable, using CPU");
|
||||
Device::Cpu
|
||||
}
|
||||
};
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Load data
|
||||
// -----------------------------------------------------------------------
|
||||
info!("Loading OHLCV bars for symbol {} ...", args.symbol);
|
||||
let bars = load_all_bars(&args.data_dir, &args.symbol)?;
|
||||
info!("Loaded {} bars total", bars.len());
|
||||
|
||||
if bars.len() < 100 {
|
||||
anyhow::bail!(
|
||||
"Insufficient data: only {} bars for symbol {}. Need >= 100.",
|
||||
bars.len(),
|
||||
args.symbol
|
||||
);
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Split into train (80%) and validation (20%)
|
||||
// -----------------------------------------------------------------------
|
||||
let split_idx = bars.len() * 80 / 100;
|
||||
let train_bars = bars.get(..split_idx).ok_or_else(|| {
|
||||
anyhow::anyhow!("Failed to slice training bars")
|
||||
})?;
|
||||
let val_bars = bars.get(split_idx..).ok_or_else(|| {
|
||||
anyhow::anyhow!("Failed to slice validation bars")
|
||||
})?;
|
||||
|
||||
info!(
|
||||
"Split: {} training bars, {} validation bars",
|
||||
train_bars.len(),
|
||||
val_bars.len()
|
||||
);
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Build datasets
|
||||
// -----------------------------------------------------------------------
|
||||
info!("Building training dataset ...");
|
||||
let train_data = build_dataset(train_bars, &device, args.tx_cost_bps, args.tick_size, args.spread_ticks)?;
|
||||
let (train_inputs, train_targets) = match train_data.first() {
|
||||
Some(pair) => pair,
|
||||
None => anyhow::bail!("Empty training dataset"),
|
||||
};
|
||||
let n_train = train_inputs
|
||||
.dims()
|
||||
.first()
|
||||
.copied()
|
||||
.ok_or_else(|| anyhow::anyhow!("No dimensions on train tensor"))?;
|
||||
info!("Training samples: {}", n_train);
|
||||
|
||||
info!("Building validation dataset ...");
|
||||
let val_data = build_dataset(val_bars, &device, args.tx_cost_bps, args.tick_size, args.spread_ticks)?;
|
||||
let (val_inputs, val_targets) = match val_data.first() {
|
||||
Some(pair) => pair,
|
||||
None => anyhow::bail!("Empty validation dataset"),
|
||||
};
|
||||
let n_val = val_inputs
|
||||
.dims()
|
||||
.first()
|
||||
.copied()
|
||||
.ok_or_else(|| anyhow::anyhow!("No dimensions on val tensor"))?;
|
||||
info!("Validation samples: {}", n_val);
|
||||
|
||||
// Build validation pairs for the validate() trait method
|
||||
let val_pairs: Vec<(Tensor, Tensor)> = build_val_pairs(val_inputs, val_targets, &device)?;
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Create TGGN adapter
|
||||
// -----------------------------------------------------------------------
|
||||
let tggn_config = TGGNConfig {
|
||||
max_nodes: 64,
|
||||
max_edges: 128,
|
||||
node_dim: NODE_DIM,
|
||||
edge_dim: 4,
|
||||
hidden_dim: 32,
|
||||
num_layers: 2,
|
||||
temporal_decay: 0.99,
|
||||
update_frequency_ns: 1_000_000,
|
||||
use_simd: false,
|
||||
};
|
||||
|
||||
let mut adapter = TGGNTrainableAdapter::new(tggn_config, &device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create TGGN adapter: {}", e))?;
|
||||
|
||||
// Set learning rate
|
||||
adapter
|
||||
.set_learning_rate(args.learning_rate)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to set learning rate: {}", e))?;
|
||||
|
||||
info!(
|
||||
"TGGN adapter created (node_dim={}, hidden_dim=32, lr={})",
|
||||
NODE_DIM, args.learning_rate
|
||||
);
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Create output directory
|
||||
// -----------------------------------------------------------------------
|
||||
std::fs::create_dir_all(&args.output_dir)
|
||||
.with_context(|| format!("Failed to create output dir: {}", args.output_dir.display()))?;
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Training loop
|
||||
// -----------------------------------------------------------------------
|
||||
let training_start = Instant::now();
|
||||
let batch_size = args.batch_size;
|
||||
let max_steps = if args.max_steps_per_epoch == 0 {
|
||||
usize::MAX
|
||||
} else {
|
||||
args.max_steps_per_epoch
|
||||
};
|
||||
|
||||
let mut best_val_loss = f64::INFINITY;
|
||||
|
||||
for epoch in 0..args.epochs {
|
||||
let epoch_start = Instant::now();
|
||||
let mut epoch_loss_sum = 0.0_f64;
|
||||
let mut epoch_steps = 0_usize;
|
||||
let mut offset = 0_usize;
|
||||
|
||||
// Iterate over mini-batches
|
||||
while offset < n_train && epoch_steps < max_steps {
|
||||
let end = (offset + batch_size).min(n_train);
|
||||
let batch_len = end - offset;
|
||||
|
||||
// Slice input and target batches
|
||||
let batch_input = train_inputs
|
||||
.narrow(0, offset, batch_len)
|
||||
.map_err(|e| anyhow::anyhow!("Input narrow failed: {}", e))?;
|
||||
let batch_target = train_targets
|
||||
.narrow(0, offset, batch_len)
|
||||
.map_err(|e| anyhow::anyhow!("Target narrow failed: {}", e))?;
|
||||
|
||||
// Forward pass
|
||||
adapter
|
||||
.zero_grad()
|
||||
.map_err(|e| anyhow::anyhow!("zero_grad failed: {}", e))?;
|
||||
let predictions = adapter
|
||||
.forward(&batch_input)
|
||||
.map_err(|e| anyhow::anyhow!("forward failed: {}", e))?;
|
||||
|
||||
// Compute loss
|
||||
let loss = adapter
|
||||
.compute_loss(&predictions, &batch_target)
|
||||
.map_err(|e| anyhow::anyhow!("compute_loss failed: {}", e))?;
|
||||
|
||||
let loss_val: f32 = loss
|
||||
.to_scalar()
|
||||
.map_err(|e| anyhow::anyhow!("loss to_scalar failed: {}", e))?;
|
||||
|
||||
// Backward + optimizer step
|
||||
let _grad_norm = adapter
|
||||
.backward(&loss)
|
||||
.map_err(|e| anyhow::anyhow!("backward failed: {}", e))?;
|
||||
adapter
|
||||
.optimizer_step()
|
||||
.map_err(|e| anyhow::anyhow!("optimizer_step failed: {}", e))?;
|
||||
|
||||
epoch_loss_sum += loss_val as f64;
|
||||
epoch_steps += 1;
|
||||
offset = end;
|
||||
}
|
||||
|
||||
let avg_loss = if epoch_steps > 0 {
|
||||
epoch_loss_sum / epoch_steps as f64
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
// Validation
|
||||
let val_loss = adapter
|
||||
.validate(&val_pairs)
|
||||
.map_err(|e| anyhow::anyhow!("validation failed: {}", e))?;
|
||||
|
||||
let epoch_elapsed = epoch_start.elapsed();
|
||||
|
||||
info!(
|
||||
"Epoch {:3}/{}: train_loss={:.6}, val_loss={:.6}, steps={}, time={:.1}s",
|
||||
epoch + 1,
|
||||
args.epochs,
|
||||
avg_loss,
|
||||
val_loss,
|
||||
epoch_steps,
|
||||
epoch_elapsed.as_secs_f64()
|
||||
);
|
||||
|
||||
// Track best validation
|
||||
if val_loss < best_val_loss {
|
||||
best_val_loss = val_loss;
|
||||
info!(
|
||||
" New best validation loss: {:.6} (epoch {})",
|
||||
best_val_loss,
|
||||
epoch + 1
|
||||
);
|
||||
}
|
||||
|
||||
// Checkpoint every 5 epochs
|
||||
if (epoch + 1) % 5 == 0 {
|
||||
let ckpt_path = args
|
||||
.output_dir
|
||||
.join(format!("tggn_epoch_{}", epoch + 1));
|
||||
let ckpt_str = ckpt_path
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid checkpoint path"))?;
|
||||
adapter
|
||||
.save_checkpoint(ckpt_str)
|
||||
.map_err(|e| anyhow::anyhow!("save_checkpoint failed: {}", e))?;
|
||||
info!(" Checkpoint saved: {}", ckpt_str);
|
||||
}
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Save final checkpoint
|
||||
// -----------------------------------------------------------------------
|
||||
let final_path = args.output_dir.join("tggn_final");
|
||||
let final_str = final_path
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("Invalid final checkpoint path"))?;
|
||||
adapter
|
||||
.save_checkpoint(final_str)
|
||||
.map_err(|e| anyhow::anyhow!("save final checkpoint failed: {}", e))?;
|
||||
|
||||
let total_time = training_start.elapsed();
|
||||
let metrics = adapter.collect_metrics();
|
||||
|
||||
info!("================================================================");
|
||||
info!(" Training Complete");
|
||||
info!("================================================================");
|
||||
info!(" Total time: {:.1}s", total_time.as_secs_f64());
|
||||
info!(" Best val loss: {:.6}", best_val_loss);
|
||||
info!(" Total steps: {}", adapter.get_step());
|
||||
info!(" Final LR: {}", adapter.get_learning_rate());
|
||||
if let Some(gn) = metrics.grad_norm {
|
||||
info!(" Last grad norm: {:.6}", gn);
|
||||
}
|
||||
info!(" Final checkpoint: {}", final_str);
|
||||
info!("================================================================");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Split full validation tensors into individual (input, target) pairs for the
|
||||
/// `UnifiedTrainable::validate()` API which takes `&[(Tensor, Tensor)]`.
|
||||
///
|
||||
/// To avoid creating one pair per sample (which could be millions), we chunk
|
||||
/// the validation set into batches of 256 samples each.
|
||||
fn build_val_pairs(
|
||||
inputs: &Tensor,
|
||||
targets: &Tensor,
|
||||
_device: &Device,
|
||||
) -> Result<Vec<(Tensor, Tensor)>> {
|
||||
let n = inputs
|
||||
.dims()
|
||||
.first()
|
||||
.copied()
|
||||
.ok_or_else(|| anyhow::anyhow!("No dimensions on val input tensor"))?;
|
||||
|
||||
let chunk_size = 256_usize;
|
||||
let mut pairs = Vec::new();
|
||||
let mut offset = 0_usize;
|
||||
|
||||
while offset < n {
|
||||
let len = (offset + chunk_size).min(n) - offset;
|
||||
let inp = inputs
|
||||
.narrow(0, offset, len)
|
||||
.map_err(|e| anyhow::anyhow!("val input narrow failed: {}", e))?;
|
||||
let tgt = targets
|
||||
.narrow(0, offset, len)
|
||||
.map_err(|e| anyhow::anyhow!("val target narrow failed: {}", e))?;
|
||||
pairs.push((inp, tgt));
|
||||
offset += len;
|
||||
}
|
||||
|
||||
Ok(pairs)
|
||||
}
|
||||
@@ -1,302 +0,0 @@
|
||||
//! TLOB Training Example
|
||||
//!
|
||||
//! Trains a TLOB transformer model on Level-2 order book data and saves checkpoints to disk.
|
||||
//!
|
||||
//! # Prerequisites
|
||||
//!
|
||||
//! - Level-2 order book data (MBP-10) from Agent 71
|
||||
//! - Data directory: test_data/real/databento/ml_training_l2/
|
||||
//! - GPU: RTX 3050 Ti (optional, will fall back to CPU)
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Train with default parameters (500 epochs)
|
||||
//! cargo run -p ml --example train_tlob --release --features cuda
|
||||
//!
|
||||
//! # Custom epochs and output path
|
||||
//! cargo run -p ml --example train_tlob --release --features cuda -- \
|
||||
//! --epochs 1000 \
|
||||
//! --output ml/trained_models/tlob_model
|
||||
//!
|
||||
//! # Custom data directory and hyperparameters
|
||||
//! cargo run -p ml --example train_tlob --release --features cuda -- \
|
||||
//! --data-dir test_data/real/databento/ml_training_l2 \
|
||||
//! --epochs 500 \
|
||||
//! --batch-size 16 \
|
||||
//! --learning-rate 0.0001 \
|
||||
//! --seq-len 128
|
||||
//!
|
||||
//! # CPU-only training (slower but works without GPU)
|
||||
//! cargo run -p ml --example train_tlob --release -- \
|
||||
//! --no-gpu \
|
||||
//! --epochs 100
|
||||
//! ```
|
||||
//!
|
||||
//! # Expected Output
|
||||
//!
|
||||
//! - Training checkpoints: ml/trained_models/tlob_epoch_*.safetensors
|
||||
//! - Final model: ml/trained_models/tlob_final_epoch500.safetensors
|
||||
//! - Training time: 5-8 hours (GPU), 20-30 hours (CPU) for 500 epochs
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use clap::Parser;
|
||||
use std::path::PathBuf;
|
||||
use tracing::{info, warn};
|
||||
use tracing_subscriber::FmtSubscriber;
|
||||
|
||||
use ml::trainers::tlob::{TLOBHyperparameters, TLOBTrainer, TLOBTrainingMetrics};
|
||||
|
||||
#[derive(Debug, Parser)]
|
||||
#[command(
|
||||
name = "train_tlob",
|
||||
about = "Train TLOB transformer on Level-2 order book data"
|
||||
)]
|
||||
struct Opts {
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value = "500")]
|
||||
epochs: usize,
|
||||
|
||||
/// Learning rate
|
||||
#[arg(long, default_value = "0.0001")]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Batch size (max 32 for RTX 3050 Ti 4GB)
|
||||
#[arg(long, default_value = "16")]
|
||||
batch_size: usize,
|
||||
|
||||
/// Sequence length (number of order book snapshots)
|
||||
#[arg(long, default_value = "128")]
|
||||
seq_len: usize,
|
||||
|
||||
/// Transformer hidden dimension
|
||||
#[arg(long, default_value = "256")]
|
||||
d_model: usize,
|
||||
|
||||
/// Number of attention heads
|
||||
#[arg(long, default_value = "8")]
|
||||
num_heads: usize,
|
||||
|
||||
/// Number of transformer layers
|
||||
#[arg(long, default_value = "4")]
|
||||
num_layers: usize,
|
||||
|
||||
/// Dropout rate
|
||||
#[arg(long, default_value = "0.1")]
|
||||
dropout: f64,
|
||||
|
||||
/// Gradient clipping threshold
|
||||
#[arg(long, default_value = "1.0")]
|
||||
grad_clip: f64,
|
||||
|
||||
/// Weight decay for regularization
|
||||
#[arg(long, default_value = "0.0001")]
|
||||
weight_decay: f64,
|
||||
|
||||
/// Checkpoint save frequency (epochs)
|
||||
#[arg(long, default_value = "10")]
|
||||
checkpoint_frequency: usize,
|
||||
|
||||
/// Output directory for trained model
|
||||
#[arg(long, default_value = "ml/trained_models")]
|
||||
output_dir: String,
|
||||
|
||||
/// Data directory containing Level-2 order book files
|
||||
#[arg(long, default_value = "test_data/real/databento/ml_training_l2")]
|
||||
data_dir: String,
|
||||
|
||||
/// Disable GPU acceleration (use CPU only)
|
||||
#[arg(long)]
|
||||
no_gpu: bool,
|
||||
|
||||
/// Verbose logging
|
||||
#[arg(short, long)]
|
||||
verbose: bool,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<()> {
|
||||
// Parse CLI options
|
||||
let opts = Opts::parse();
|
||||
|
||||
// Setup logging
|
||||
let level = if opts.verbose {
|
||||
tracing::Level::DEBUG
|
||||
} else {
|
||||
tracing::Level::INFO
|
||||
};
|
||||
|
||||
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
|
||||
tracing::subscriber::set_global_default(subscriber)
|
||||
.context("Failed to set tracing subscriber")?;
|
||||
|
||||
info!("🚀 Starting TLOB Transformer Training");
|
||||
info!("Configuration:");
|
||||
info!(" • Epochs: {}", opts.epochs);
|
||||
info!(" • Learning rate: {}", opts.learning_rate);
|
||||
info!(" • Batch size: {}", opts.batch_size);
|
||||
info!(" • Sequence length: {}", opts.seq_len);
|
||||
info!(" • Hidden dimension: {}", opts.d_model);
|
||||
info!(" • Attention heads: {}", opts.num_heads);
|
||||
info!(" • Transformer layers: {}", opts.num_layers);
|
||||
info!(" • Dropout: {}", opts.dropout);
|
||||
info!(" • Gradient clipping: {}", opts.grad_clip);
|
||||
info!(" • Weight decay: {}", opts.weight_decay);
|
||||
info!(
|
||||
" • Checkpoint frequency: {} epochs",
|
||||
opts.checkpoint_frequency
|
||||
);
|
||||
info!(" • Output directory: {}", opts.output_dir);
|
||||
info!(" • Data directory: {}", opts.data_dir);
|
||||
info!(" • GPU enabled: {}", !opts.no_gpu);
|
||||
|
||||
// Check if data directory exists
|
||||
let data_path = PathBuf::from(&opts.data_dir);
|
||||
if !data_path.exists() {
|
||||
warn!("⚠️ Data directory not found: {}", opts.data_dir);
|
||||
warn!("⚠️ This is expected if Agent 71 hasn't completed yet.");
|
||||
warn!("⚠️ Training will use dummy data for testing purposes.");
|
||||
}
|
||||
|
||||
// Create output directory
|
||||
let output_path = PathBuf::from(&opts.output_dir);
|
||||
if !output_path.exists() {
|
||||
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
|
||||
info!("✅ Created output directory: {}", opts.output_dir);
|
||||
}
|
||||
|
||||
// Configure TLOB hyperparameters
|
||||
let hyperparams = TLOBHyperparameters {
|
||||
learning_rate: opts.learning_rate,
|
||||
batch_size: opts.batch_size,
|
||||
seq_len: opts.seq_len,
|
||||
num_price_levels: 10, // MBP-10
|
||||
d_model: opts.d_model,
|
||||
num_heads: opts.num_heads,
|
||||
num_layers: opts.num_layers,
|
||||
dropout: opts.dropout,
|
||||
epochs: opts.epochs,
|
||||
checkpoint_frequency: opts.checkpoint_frequency,
|
||||
grad_clip: opts.grad_clip,
|
||||
weight_decay: opts.weight_decay,
|
||||
};
|
||||
|
||||
// Create TLOB trainer
|
||||
let mut trainer = TLOBTrainer::new(hyperparams, &output_path, !opts.no_gpu)
|
||||
.context("Failed to create TLOB trainer")?;
|
||||
|
||||
info!("✅ TLOB trainer initialized");
|
||||
|
||||
// Track training progress
|
||||
let mut last_epoch = 0;
|
||||
let mut best_val_loss = f64::INFINITY;
|
||||
|
||||
// Create progress callback
|
||||
let progress_callback = |metrics: TLOBTrainingMetrics| {
|
||||
if metrics.epoch != last_epoch {
|
||||
last_epoch = metrics.epoch;
|
||||
|
||||
info!(
|
||||
"📊 Epoch {}/{}: train_loss={:.6}, val_loss={:.6}, mae={:.6}, grad_norm={:.6}",
|
||||
metrics.epoch,
|
||||
opts.epochs,
|
||||
metrics.train_loss,
|
||||
metrics.val_loss,
|
||||
metrics.avg_mae,
|
||||
metrics.gradient_norm
|
||||
);
|
||||
|
||||
if metrics.val_loss < best_val_loss {
|
||||
best_val_loss = metrics.val_loss;
|
||||
info!("🌟 New best validation loss: {:.6}", best_val_loss);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Train the model
|
||||
info!("\n🏋️ Starting training...\n");
|
||||
let start_time = std::time::Instant::now();
|
||||
|
||||
let metrics = trainer
|
||||
.train(&opts.data_dir, progress_callback)
|
||||
.await
|
||||
.context("Training failed")?;
|
||||
|
||||
let training_duration = start_time.elapsed();
|
||||
|
||||
// Print final metrics
|
||||
info!("\n✅ Training completed successfully!");
|
||||
info!("\n📊 Final Metrics:");
|
||||
info!(" • Final train loss: {:.6}", metrics.train_loss);
|
||||
info!(" • Final val loss: {:.6}", metrics.val_loss);
|
||||
info!(" • Best val loss: {:.6}", best_val_loss);
|
||||
info!(" • Final MAE: {:.6}", metrics.avg_mae);
|
||||
info!(" • Final gradient norm: {:.6}", metrics.gradient_norm);
|
||||
info!(" • Epochs trained: {}", metrics.epoch);
|
||||
info!(
|
||||
" • Training time: {:.1}s ({:.1} min, {:.1} hours)",
|
||||
training_duration.as_secs_f64(),
|
||||
training_duration.as_secs_f64() / 60.0,
|
||||
training_duration.as_secs_f64() / 3600.0
|
||||
);
|
||||
|
||||
// Calculate training speed
|
||||
let seconds_per_epoch = training_duration.as_secs_f64() / opts.epochs as f64;
|
||||
info!(" • Average time per epoch: {:.2}s", seconds_per_epoch);
|
||||
|
||||
// Save final model
|
||||
let final_model_path = output_path.join(format!("tlob_final_epoch{}.safetensors", opts.epochs));
|
||||
info!("\n💾 Saving final model to: {}", final_model_path.display());
|
||||
|
||||
// Get final model state
|
||||
let final_checkpoint_data = trainer
|
||||
.serialize_model()
|
||||
.await
|
||||
.context("Failed to serialize final model")?;
|
||||
|
||||
std::fs::write(&final_model_path, &final_checkpoint_data)
|
||||
.context("Failed to save final model")?;
|
||||
|
||||
info!(
|
||||
"✅ Final model saved: {} ({} bytes, {:.2} MB)",
|
||||
final_model_path.display(),
|
||||
final_checkpoint_data.len(),
|
||||
final_checkpoint_data.len() as f64 / 1_048_576.0
|
||||
);
|
||||
|
||||
// Print summary
|
||||
info!("\n🎉 TLOB training complete!");
|
||||
info!("📁 Model files saved to: {}", opts.output_dir);
|
||||
info!("\n📈 Training Summary:");
|
||||
info!(" • Best validation loss: {:.6}", best_val_loss);
|
||||
info!(
|
||||
" • Convergence: {}",
|
||||
if best_val_loss < 0.001 {
|
||||
"✅ Excellent"
|
||||
} else if best_val_loss < 0.01 {
|
||||
"✅ Good"
|
||||
} else {
|
||||
"⚠️ Needs more epochs"
|
||||
}
|
||||
);
|
||||
info!(" • Training efficiency: {:.2}s/epoch", seconds_per_epoch);
|
||||
|
||||
// Estimate production inference latency
|
||||
let estimated_inference_us = seconds_per_epoch * 1_000_000.0 / 1000.0; // Rough estimate
|
||||
info!("\n🚀 Production Inference Estimate:");
|
||||
info!(
|
||||
" • Expected latency: <{:.0}μs per prediction",
|
||||
estimated_inference_us.min(100.0)
|
||||
);
|
||||
info!(" • Target: <50μs (sub-50μs HFT requirement)");
|
||||
|
||||
// Next steps
|
||||
info!("\n📋 Next Steps:");
|
||||
info!(" 1. Validate model with test data");
|
||||
info!(" 2. Convert to ONNX for production inference");
|
||||
info!(" 3. Integrate with TLOB inference engine");
|
||||
info!(" 4. Benchmark inference latency (<50μs target)");
|
||||
info!(" 5. Deploy to ML Training Service");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,518 +0,0 @@
|
||||
//! xLSTM Training with Real DataBento DBN Market Data
|
||||
//!
|
||||
//! Standalone training script for the xLSTM (Extended Long Short-Term Memory)
|
||||
//! model on real OHLCV data loaded from Databento DBN files.
|
||||
//!
|
||||
//! The xLSTM architecture combines two cell types:
|
||||
//! - **sLSTM**: Exponential gating with scalar memory (sequential patterns)
|
||||
//! - **mLSTM**: Matrix memory with key-value association (higher capacity)
|
||||
//!
|
||||
//! The `slstm_ratio` parameter controls the mix of sLSTM vs mLSTM blocks.
|
||||
//!
|
||||
//! # Usage
|
||||
//!
|
||||
//! ```bash
|
||||
//! # Default: 30 epochs on ES.FUT
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_xlstm_dbn --release -- \
|
||||
//! --data-dir data/cache/futures-baseline --symbol ES.FUT
|
||||
//!
|
||||
//! # Custom configuration
|
||||
//! SQLX_OFFLINE=true cargo run -p ml --example train_xlstm_dbn --release -- \
|
||||
//! --data-dir data/cache/futures-baseline --symbol ES.FUT \
|
||||
//! --epochs 50 --batch-size 64 --learning-rate 0.0005 \
|
||||
//! --hidden-dim 64 --num-blocks 4 --num-heads 4 --slstm-ratio 0.5
|
||||
//! ```
|
||||
//!
|
||||
//! # Output
|
||||
//!
|
||||
//! Checkpoints saved to `--output-dir` (default `ml/trained_models/xlstm`):
|
||||
//! - `xlstm_epoch_N.safetensors` every 5 epochs
|
||||
//! - `xlstm_final.safetensors` at end of training
|
||||
|
||||
#![allow(unused_crate_dependencies)]
|
||||
#![deny(
|
||||
clippy::unwrap_used,
|
||||
clippy::expect_used,
|
||||
clippy::panic,
|
||||
clippy::indexing_slicing
|
||||
)]
|
||||
|
||||
use std::path::PathBuf;
|
||||
use std::time::Instant;
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use candle_core::{Device, Tensor};
|
||||
use clap::Parser;
|
||||
use tracing::info;
|
||||
|
||||
use ml::features::extraction::extract_ml_features;
|
||||
use ml::training::unified_trainer::UnifiedTrainable;
|
||||
use ml::types::OHLCVBar;
|
||||
use ml::xlstm::{XLSTMConfig, XLSTMTrainableAdapter};
|
||||
|
||||
#[allow(unreachable_pub)]
|
||||
mod baseline_common;
|
||||
use baseline_common::{load_all_bars, spread_cost_bps};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// CLI Arguments
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Train xLSTM on real OHLCV market data from DataBento DBN files.
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "train_xlstm_dbn", about = "Train xLSTM model on real DBN market data")]
|
||||
struct Args {
|
||||
/// Path to directory containing .dbn.zst files (with symbol subdirectories)
|
||||
#[arg(long, default_value = "data/cache/futures-baseline")]
|
||||
data_dir: PathBuf,
|
||||
|
||||
/// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT", "6E.FUT", "ZN.FUT")
|
||||
#[arg(long, default_value = "ES.FUT")]
|
||||
symbol: String,
|
||||
|
||||
/// Number of training epochs
|
||||
#[arg(long, default_value_t = 30)]
|
||||
epochs: usize,
|
||||
|
||||
/// Training batch size (number of samples per forward/backward pass)
|
||||
#[arg(long, default_value_t = 64)]
|
||||
batch_size: usize,
|
||||
|
||||
/// Learning rate for AdamW optimizer
|
||||
#[arg(long, default_value_t = 1e-3)]
|
||||
learning_rate: f64,
|
||||
|
||||
/// Max environment steps per epoch (0 = use all available bars)
|
||||
#[arg(long, default_value_t = 2000)]
|
||||
max_steps_per_epoch: usize,
|
||||
|
||||
/// Output directory for checkpoints
|
||||
#[arg(long, default_value = "ml/trained_models/xlstm")]
|
||||
output_dir: PathBuf,
|
||||
|
||||
/// xLSTM hidden dimension. OOM note: mLSTM matrix memory is
|
||||
/// (hidden_dim/num_heads)^2 per head per sample.
|
||||
#[arg(long, default_value_t = 64)]
|
||||
hidden_dim: usize,
|
||||
|
||||
/// Number of xLSTM blocks (interleaved sLSTM and mLSTM)
|
||||
#[arg(long, default_value_t = 4)]
|
||||
num_blocks: usize,
|
||||
|
||||
/// Number of attention heads for mLSTM blocks
|
||||
#[arg(long, default_value_t = 4)]
|
||||
num_heads: usize,
|
||||
|
||||
/// Ratio of sLSTM blocks (0.0 = all mLSTM, 1.0 = all sLSTM)
|
||||
#[arg(long, default_value_t = 0.5)]
|
||||
slstm_ratio: f64,
|
||||
|
||||
/// Dropout rate between blocks
|
||||
#[arg(long, default_value_t = 0.1)]
|
||||
dropout: f64,
|
||||
|
||||
/// Weight decay for AdamW regularization
|
||||
#[arg(long, default_value_t = 1e-4)]
|
||||
weight_decay: f64,
|
||||
|
||||
/// Gradient clipping max norm
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
grad_clip: f64,
|
||||
|
||||
/// Sequence length: number of time steps per input window
|
||||
#[arg(long, default_value_t = 32)]
|
||||
seq_len: usize,
|
||||
|
||||
/// Train/validation split ratio (fraction used for training)
|
||||
#[arg(long, default_value_t = 0.8)]
|
||||
train_split: f64,
|
||||
|
||||
/// Early stopping patience (epochs without improvement)
|
||||
#[arg(long, default_value_t = 10)]
|
||||
patience: usize,
|
||||
|
||||
/// Round-trip commission cost in basis points
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
tx_cost_bps: f64,
|
||||
|
||||
/// Instrument tick size in price units (ES=0.25)
|
||||
#[arg(long, default_value_t = 0.25)]
|
||||
tick_size: f64,
|
||||
|
||||
/// Typical bid-ask spread in ticks
|
||||
#[arg(long, default_value_t = 1.0)]
|
||||
spread_ticks: f64,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Feature preparation helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Feature dimension produced by `extract_ml_features` (51-dimensional).
|
||||
const FEATURE_DIM: usize = 51;
|
||||
|
||||
/// Build (input, target) tensor pairs from OHLCV bars for the xLSTM model.
|
||||
///
|
||||
/// Each sample is a sliding window of `seq_len` feature vectors as input,
|
||||
/// with the target being the normalised close-price return of the next bar.
|
||||
///
|
||||
/// Returns `Vec<(Tensor, Tensor)>` where:
|
||||
/// - input shape: `[1, seq_len, FEATURE_DIM]`
|
||||
/// - target shape: `[1, 1]`
|
||||
fn build_sequences(
|
||||
bars: &[OHLCVBar],
|
||||
seq_len: usize,
|
||||
max_steps: usize,
|
||||
device: &Device,
|
||||
tx_cost_bps: f64,
|
||||
tick_size: f64,
|
||||
spread_ticks: f64,
|
||||
) -> Result<Vec<(Tensor, Tensor)>> {
|
||||
// Extract 51-dim features (consumes a warmup period of ~50 bars)
|
||||
let features = extract_ml_features(bars).context("Feature extraction failed")?;
|
||||
let n_features = features.len();
|
||||
|
||||
if n_features < seq_len + 1 {
|
||||
anyhow::bail!(
|
||||
"Not enough feature vectors ({}) for seq_len={} + 1 target",
|
||||
n_features,
|
||||
seq_len
|
||||
);
|
||||
}
|
||||
|
||||
// Align bars to features (features skip warmup period)
|
||||
let warmup_offset = bars.len().saturating_sub(n_features);
|
||||
|
||||
let total_possible = n_features.saturating_sub(seq_len);
|
||||
let n_samples = if max_steps > 0 {
|
||||
total_possible.min(max_steps)
|
||||
} else {
|
||||
total_possible
|
||||
};
|
||||
|
||||
let mut sequences = Vec::with_capacity(n_samples);
|
||||
|
||||
for i in 0..n_samples {
|
||||
// Build input: [seq_len, FEATURE_DIM] flattened to f32
|
||||
let mut input_data = Vec::with_capacity(seq_len * FEATURE_DIM);
|
||||
for t in 0..seq_len {
|
||||
let feat = match features.get(i + t) {
|
||||
Some(f) => f,
|
||||
None => continue,
|
||||
};
|
||||
for &v in feat.iter() {
|
||||
input_data.push(v as f32);
|
||||
}
|
||||
}
|
||||
|
||||
// Target: normalised return of the bar after the sequence window
|
||||
let target_feat_idx = i + seq_len;
|
||||
let target_bar_idx = warmup_offset + target_feat_idx;
|
||||
|
||||
let close_prev = bars.get(target_bar_idx.saturating_sub(1)).map(|b| b.close).unwrap_or(1.0);
|
||||
let close_cur = bars.get(target_bar_idx).map(|b| b.close).unwrap_or(close_prev);
|
||||
|
||||
let raw_ret = if close_prev.abs() > 1e-10 {
|
||||
((close_cur - close_prev) / close_prev).clamp(-0.05, 0.05)
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
let spread = spread_cost_bps(close_prev, tick_size, spread_ticks);
|
||||
let cost = (tx_cost_bps + spread) / 10_000.0; // convert bps to fractional
|
||||
let ret = (raw_ret - cost) as f32;
|
||||
|
||||
// Only add if we built the full window
|
||||
if input_data.len() == seq_len * FEATURE_DIM {
|
||||
let input_tensor = Tensor::from_vec(input_data, &[1, seq_len, FEATURE_DIM], device)
|
||||
.map_err(|e| anyhow::anyhow!("input tensor: {e}"))?;
|
||||
let target_tensor = Tensor::from_vec(vec![ret], &[1, 1], device)
|
||||
.map_err(|e| anyhow::anyhow!("target tensor: {e}"))?;
|
||||
sequences.push((input_tensor, target_tensor));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(sequences)
|
||||
}
|
||||
|
||||
/// Concatenate a batch of (input, target) pairs along the batch dimension.
|
||||
///
|
||||
/// Returns `(batched_input, batched_target)` where:
|
||||
/// - batched_input shape: `[batch_size, seq_len, FEATURE_DIM]`
|
||||
/// - batched_target shape: `[batch_size, 1]`
|
||||
fn batch_samples(
|
||||
samples: &[(Tensor, Tensor)],
|
||||
) -> Result<(Tensor, Tensor)> {
|
||||
if samples.is_empty() {
|
||||
anyhow::bail!("Cannot batch empty sample slice");
|
||||
}
|
||||
|
||||
let inputs: Vec<&Tensor> = samples.iter().map(|(inp, _)| inp).collect();
|
||||
let targets: Vec<&Tensor> = samples.iter().map(|(_, tgt)| tgt).collect();
|
||||
|
||||
let batched_input = Tensor::cat(&inputs, 0)
|
||||
.map_err(|e| anyhow::anyhow!("batch inputs: {e}"))?;
|
||||
let batched_target = Tensor::cat(&targets, 0)
|
||||
.map_err(|e| anyhow::anyhow!("batch targets: {e}"))?;
|
||||
|
||||
Ok((batched_input, batched_target))
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Main
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
fn main() -> Result<()> {
|
||||
// Initialize tracing
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
|
||||
)
|
||||
.init();
|
||||
|
||||
let args = Args::parse();
|
||||
|
||||
info!("=== xLSTM Training on Real DBN Market Data ===");
|
||||
info!(" Symbol: {}", args.symbol);
|
||||
info!(" Data dir: {}", args.data_dir.display());
|
||||
info!(" Epochs: {}", args.epochs);
|
||||
info!(" Batch size: {}", args.batch_size);
|
||||
info!(" Learning rate: {:.1e}", args.learning_rate);
|
||||
info!(" Seq len: {}", args.seq_len);
|
||||
info!(" Hidden dim: {}", args.hidden_dim);
|
||||
info!(" Num blocks: {}", args.num_blocks);
|
||||
info!(" Num heads: {}", args.num_heads);
|
||||
info!(" sLSTM ratio: {:.2}", args.slstm_ratio);
|
||||
info!(" Dropout: {:.2}", args.dropout);
|
||||
info!(" Weight decay: {:.1e}", args.weight_decay);
|
||||
info!(" Grad clip: {:.1}", args.grad_clip);
|
||||
info!(" Max steps/epoch: {}", args.max_steps_per_epoch);
|
||||
info!(" Train split: {:.0}%", args.train_split * 100.0);
|
||||
info!(" Patience: {}", args.patience);
|
||||
info!(" Output dir: {}", args.output_dir.display());
|
||||
|
||||
// ---- Device selection ----
|
||||
let device = match Device::cuda_if_available(0) {
|
||||
Ok(dev) => {
|
||||
if dev.is_cuda() {
|
||||
info!(" Device: CUDA GPU");
|
||||
} else {
|
||||
info!(" Device: CPU (CUDA not available)");
|
||||
}
|
||||
dev
|
||||
}
|
||||
Err(_) => {
|
||||
info!(" Device: CPU (fallback)");
|
||||
Device::Cpu
|
||||
}
|
||||
};
|
||||
|
||||
// ---- Load OHLCV bars from DBN files ----
|
||||
info!("Step 1/5: Loading OHLCV bars from DBN files...");
|
||||
let all_bars = load_all_bars(&args.data_dir, &args.symbol)?;
|
||||
if all_bars.is_empty() {
|
||||
anyhow::bail!("No bars loaded from {}", args.data_dir.display());
|
||||
}
|
||||
info!(
|
||||
" Loaded {} bars ({} to {})",
|
||||
all_bars.len(),
|
||||
all_bars.first().map(|b| b.timestamp.to_string()).unwrap_or_default(),
|
||||
all_bars.last().map(|b| b.timestamp.to_string()).unwrap_or_default(),
|
||||
);
|
||||
|
||||
// ---- Build sequences ----
|
||||
info!("Step 2/5: Building input/target sequences (seq_len={})...", args.seq_len);
|
||||
let all_sequences = build_sequences(
|
||||
&all_bars, args.seq_len, args.max_steps_per_epoch, &device,
|
||||
args.tx_cost_bps, args.tick_size, args.spread_ticks,
|
||||
)?;
|
||||
if all_sequences.is_empty() {
|
||||
anyhow::bail!("No sequences produced from {} bars", all_bars.len());
|
||||
}
|
||||
info!(" Built {} sequences", all_sequences.len());
|
||||
|
||||
// ---- Train/val split ----
|
||||
info!("Step 3/5: Splitting data...");
|
||||
let split_idx = (all_sequences.len() as f64 * args.train_split) as usize;
|
||||
let split_idx = split_idx.max(1).min(all_sequences.len().saturating_sub(1));
|
||||
|
||||
let (train_sequences, val_sequences) = all_sequences.split_at(split_idx);
|
||||
info!(
|
||||
" Train: {} sequences, Val: {} sequences",
|
||||
train_sequences.len(),
|
||||
val_sequences.len()
|
||||
);
|
||||
|
||||
if train_sequences.is_empty() || val_sequences.is_empty() {
|
||||
anyhow::bail!(
|
||||
"Insufficient data for train/val split: {} total sequences with split at {}",
|
||||
all_sequences.len(),
|
||||
split_idx,
|
||||
);
|
||||
}
|
||||
|
||||
// ---- Create xLSTM model ----
|
||||
info!("Step 4/5: Initializing xLSTM model...");
|
||||
let xlstm_config = XLSTMConfig {
|
||||
input_dim: FEATURE_DIM,
|
||||
hidden_dim: args.hidden_dim,
|
||||
num_blocks: args.num_blocks,
|
||||
num_heads: args.num_heads,
|
||||
slstm_ratio: args.slstm_ratio,
|
||||
output_dim: 1, // regression: predict next-bar return
|
||||
dropout: args.dropout,
|
||||
learning_rate: args.learning_rate,
|
||||
weight_decay: args.weight_decay,
|
||||
grad_clip: args.grad_clip,
|
||||
};
|
||||
|
||||
let mut adapter = XLSTMTrainableAdapter::new(xlstm_config, &device)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create xLSTM adapter: {e}"))?;
|
||||
|
||||
info!(" Model type: {}", adapter.model_type());
|
||||
info!(
|
||||
" Config: input_dim={}, hidden_dim={}, blocks={}, heads={}, slstm_ratio={:.2}",
|
||||
FEATURE_DIM, args.hidden_dim, args.num_blocks, args.num_heads, args.slstm_ratio
|
||||
);
|
||||
|
||||
// Create output directory
|
||||
std::fs::create_dir_all(&args.output_dir)
|
||||
.with_context(|| format!("Failed to create output dir: {}", args.output_dir.display()))?;
|
||||
|
||||
// ---- Training loop ----
|
||||
info!("Step 5/5: Training...");
|
||||
let training_start = Instant::now();
|
||||
|
||||
let mut best_val_loss = f64::MAX;
|
||||
let mut epochs_without_improvement = 0_usize;
|
||||
|
||||
for epoch in 0..args.epochs {
|
||||
let epoch_start = Instant::now();
|
||||
|
||||
// --- Training pass ---
|
||||
let mut epoch_loss_sum = 0.0_f64;
|
||||
let mut epoch_steps = 0_usize;
|
||||
|
||||
// Process training data in mini-batches
|
||||
let mut batch_start = 0_usize;
|
||||
while batch_start < train_sequences.len() {
|
||||
let batch_end = (batch_start + args.batch_size).min(train_sequences.len());
|
||||
let batch_slice = match train_sequences.get(batch_start..batch_end) {
|
||||
Some(s) => s,
|
||||
None => break,
|
||||
};
|
||||
|
||||
if batch_slice.is_empty() {
|
||||
break;
|
||||
}
|
||||
|
||||
let (batched_input, batched_target) = batch_samples(batch_slice)?;
|
||||
|
||||
// Forward pass
|
||||
adapter.zero_grad()?;
|
||||
let predictions = adapter
|
||||
.forward(&batched_input)
|
||||
.map_err(|e| anyhow::anyhow!("forward: {e}"))?;
|
||||
|
||||
// Compute loss
|
||||
let loss = adapter
|
||||
.compute_loss(&predictions, &batched_target)
|
||||
.map_err(|e| anyhow::anyhow!("loss: {e}"))?;
|
||||
|
||||
let loss_val = loss
|
||||
.to_scalar::<f32>()
|
||||
.map_err(|e| anyhow::anyhow!("loss scalar: {e}"))?;
|
||||
|
||||
// Backward + optimizer step
|
||||
let _grad_norm = adapter
|
||||
.backward(&loss)
|
||||
.map_err(|e| anyhow::anyhow!("backward: {e}"))?;
|
||||
adapter
|
||||
.optimizer_step()
|
||||
.map_err(|e| anyhow::anyhow!("optimizer_step: {e}"))?;
|
||||
|
||||
epoch_loss_sum += loss_val as f64;
|
||||
epoch_steps += 1;
|
||||
batch_start = batch_end;
|
||||
}
|
||||
|
||||
let avg_train_loss = if epoch_steps > 0 {
|
||||
epoch_loss_sum / epoch_steps as f64
|
||||
} else {
|
||||
f64::MAX
|
||||
};
|
||||
|
||||
// --- Validation pass ---
|
||||
let val_loss = adapter
|
||||
.validate(val_sequences)
|
||||
.map_err(|e| anyhow::anyhow!("validate: {e}"))?;
|
||||
|
||||
let epoch_time = epoch_start.elapsed();
|
||||
let metrics = adapter.collect_metrics();
|
||||
|
||||
info!(
|
||||
" Epoch {}/{} -- train_loss={:.6} val_loss={:.6} lr={:.1e} grad_norm={:.4} step={} ({:.2}s)",
|
||||
epoch + 1,
|
||||
args.epochs,
|
||||
avg_train_loss,
|
||||
val_loss,
|
||||
metrics.learning_rate,
|
||||
metrics.grad_norm.unwrap_or(0.0),
|
||||
adapter.get_step(),
|
||||
epoch_time.as_secs_f64(),
|
||||
);
|
||||
|
||||
// --- Checkpoint every 5 epochs ---
|
||||
if (epoch + 1) % 5 == 0 {
|
||||
let ckpt_path = args.output_dir.join(format!("xlstm_epoch_{}", epoch + 1));
|
||||
match adapter.save_checkpoint(&ckpt_path.to_string_lossy()) {
|
||||
Ok(_) => info!(" Saved checkpoint: {}", ckpt_path.display()),
|
||||
Err(e) => info!(" Checkpoint save failed: {}", e),
|
||||
}
|
||||
}
|
||||
|
||||
// --- Early stopping ---
|
||||
if val_loss < best_val_loss {
|
||||
best_val_loss = val_loss;
|
||||
epochs_without_improvement = 0;
|
||||
|
||||
// Save best model
|
||||
let best_path = args.output_dir.join("xlstm_best");
|
||||
match adapter.save_checkpoint(&best_path.to_string_lossy()) {
|
||||
Ok(_) => info!(" New best model (val_loss={:.6})", val_loss),
|
||||
Err(e) => info!(" Best checkpoint save failed: {}", e),
|
||||
}
|
||||
} else {
|
||||
epochs_without_improvement += 1;
|
||||
if epochs_without_improvement >= args.patience {
|
||||
info!(
|
||||
" Early stopping at epoch {} (patience {} exhausted)",
|
||||
epoch + 1,
|
||||
args.patience,
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Save final checkpoint ----
|
||||
let final_path = args.output_dir.join("xlstm_final");
|
||||
match adapter.save_checkpoint(&final_path.to_string_lossy()) {
|
||||
Ok(_) => info!("Saved final checkpoint: {}", final_path.display()),
|
||||
Err(e) => info!("Final checkpoint save failed: {}", e),
|
||||
}
|
||||
|
||||
// ---- Summary ----
|
||||
let total_time = training_start.elapsed();
|
||||
let final_metrics = adapter.collect_metrics();
|
||||
|
||||
info!("=== Training Summary ===");
|
||||
info!(" Total time: {:.1}s ({:.1} min)", total_time.as_secs_f64(), total_time.as_secs_f64() / 60.0);
|
||||
info!(" Best val loss: {:.6}", best_val_loss);
|
||||
info!(" Final train loss: {:.6}", final_metrics.loss);
|
||||
info!(" Total steps: {}", adapter.get_step());
|
||||
info!(" Checkpoints: {}", args.output_dir.display());
|
||||
info!("=== xLSTM Training Complete ===");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
# GPU training image for ML model training Jobs
|
||||
# Usage: docker build -f infra/docker/Dockerfile.training .
|
||||
# Run: docker run --gpus all foxhunt-training train_dqn [args...]
|
||||
# Run: docker run --gpus all foxhunt-training train_baseline_supervised --model kan [args...]
|
||||
|
||||
# =============================================================================
|
||||
# Stage 1: Builder (CUDA dev image with Rust toolchain)
|
||||
@@ -59,26 +59,17 @@ COPY testing ./testing
|
||||
ENV SQLX_OFFLINE=true
|
||||
ENV CUDA_COMPUTE_CAP=90
|
||||
|
||||
# Build all training example binaries with CUDA support
|
||||
# Build training binaries with CUDA support (2 unified baselines + eval + hyperopt)
|
||||
RUN cargo build --release -p ml --features ml/cuda \
|
||||
--example train_dqn_es_fut \
|
||||
--example train_ppo_parquet \
|
||||
--example train_tft_dbn \
|
||||
--example train_mamba2_dbn \
|
||||
--example train_liquid_dbn \
|
||||
--example train_tggn_dbn \
|
||||
--example train_kan_dbn \
|
||||
--example train_xlstm_dbn \
|
||||
--example train_diffusion_dbn \
|
||||
--example train_tlob \
|
||||
--example train_baseline \
|
||||
--example train_baseline_rl \
|
||||
--example train_baseline_supervised \
|
||||
--example evaluate_baseline \
|
||||
--example hyperopt_dqn_demo \
|
||||
--example hyperopt_ppo_demo \
|
||||
--example hyperopt_tft_demo \
|
||||
--example hyperopt_mamba2_demo \
|
||||
&& mkdir -p /build/out \
|
||||
&& for bin in train_dqn_es_fut train_ppo_parquet train_tft_dbn train_mamba2_dbn train_liquid_dbn train_tggn_dbn train_kan_dbn train_xlstm_dbn train_diffusion_dbn train_tlob train_baseline evaluate_baseline hyperopt_dqn_demo hyperopt_ppo_demo hyperopt_tft_demo hyperopt_mamba2_demo; do \
|
||||
&& for bin in train_baseline_rl train_baseline_supervised evaluate_baseline hyperopt_dqn_demo hyperopt_ppo_demo hyperopt_tft_demo hyperopt_mamba2_demo; do \
|
||||
cp target/release/examples/${bin} /build/out/ && strip /build/out/${bin}; \
|
||||
done
|
||||
|
||||
|
||||
@@ -30,12 +30,10 @@ spec:
|
||||
- name: training
|
||||
image: rg.fr-par.scw.cloud/foxhunt/training:latest
|
||||
# Available binaries in /usr/local/bin/:
|
||||
# train_dqn_es_fut, train_ppo_parquet, train_tft_dbn,
|
||||
# train_mamba2_dbn, train_liquid_dbn, train_tggn_dbn,
|
||||
# train_kan_dbn, train_xlstm_dbn, train_diffusion_dbn,
|
||||
# train_tlob, train_baseline, evaluate_baseline,
|
||||
# hyperopt_dqn_demo, hyperopt_ppo_demo, hyperopt_tft_demo,
|
||||
# hyperopt_mamba2_demo
|
||||
# train_baseline_rl (for dqn, ppo)
|
||||
# train_baseline_supervised (for tft, mamba2, tggn, tlob, liquid, kan, xlstm, diffusion)
|
||||
# evaluate_baseline
|
||||
# hyperopt_dqn_demo, hyperopt_ppo_demo, hyperopt_tft_demo, hyperopt_mamba2_demo
|
||||
command: ["/usr/local/bin/$(TRAINING_BINARY)"]
|
||||
args:
|
||||
- "--symbol=ES.FUT"
|
||||
@@ -43,7 +41,7 @@ spec:
|
||||
- "--output-dir=/output"
|
||||
env:
|
||||
- name: TRAINING_BINARY
|
||||
value: train_dqn_es_fut
|
||||
value: train_baseline_supervised
|
||||
- name: RUST_LOG
|
||||
value: info
|
||||
- name: SQLX_OFFLINE
|
||||
|
||||
@@ -34,16 +34,16 @@ ALL_MODELS=(dqn ppo tft mamba2 tggn tlob liquid kan xlstm diffusion)
|
||||
|
||||
# --- Model-to-binary mapping ----------------------------------------------
|
||||
declare -A MODEL_BINARY=(
|
||||
[dqn]=train_dqn_es_fut
|
||||
[ppo]=train_ppo_parquet
|
||||
[tft]=train_tft_dbn
|
||||
[mamba2]=train_mamba2_dbn
|
||||
[tggn]=train_tggn_dbn
|
||||
[tlob]=train_tlob
|
||||
[liquid]=train_liquid_dbn
|
||||
[kan]=train_kan_dbn
|
||||
[xlstm]=train_xlstm_dbn
|
||||
[diffusion]=train_diffusion_dbn
|
||||
[dqn]=train_baseline_rl
|
||||
[ppo]=train_baseline_rl
|
||||
[tft]=train_baseline_supervised
|
||||
[mamba2]=train_baseline_supervised
|
||||
[tggn]=train_baseline_supervised
|
||||
[tlob]=train_baseline_supervised
|
||||
[liquid]=train_baseline_supervised
|
||||
[kan]=train_baseline_supervised
|
||||
[xlstm]=train_baseline_supervised
|
||||
[diffusion]=train_baseline_supervised
|
||||
)
|
||||
|
||||
EVAL_BINARY="evaluate_baseline"
|
||||
@@ -129,6 +129,8 @@ build_args() {
|
||||
local binary="${MODEL_BINARY[$model]}"
|
||||
local args=("$binary")
|
||||
|
||||
# Both unified binaries accept --model to select the specific model
|
||||
args+=("--model" "$model")
|
||||
args+=("--symbol" "$SYMBOL")
|
||||
args+=("--max-steps-per-epoch" "$MAX_STEPS")
|
||||
args+=("--data-dir" "$DATA_DIR")
|
||||
|
||||
Reference in New Issue
Block a user