Files
foxhunt/crates/ml/examples/train_baseline_supervised.rs
jgrusewski 6e339316cf feat(ml): add manually-triggered GitLab CI training pipeline
Adds a parent/child GitLab CI pipeline for ML model training:

- Generator script produces per-model hyperopt/train/evaluate jobs
- Parent pipeline (.gitlab-ci-training.yml) with manual trigger
- NFS-backed ReadWriteMany PVC for shared training outputs
- Hyperopt params wired into training binaries (DQN, PPO, TFT, Mamba2)
- Shared DBN loader eliminates duplicate code across hyperopt adapters
- Supervised hyperopt unified to DBN data (was parquet-only)

Pipeline: hyperopt (4 models) → train (10 models) → evaluate ensemble

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-26 09:04:58 +01:00

729 lines
25 KiB
Rust

//! 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 test_data/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 serde_json::Value;
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 = "test_data/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,
/// Optional path to hyperopt results JSON -- overrides matching config fields
#[arg(long)]
hyperopt_params: Option<PathBuf>,
}
// ---------------------------------------------------------------------------
// 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(", ")
);
}
}
// ---------------------------------------------------------------------------
// Hyperopt parameter loading
// ---------------------------------------------------------------------------
/// Load best_params from a hyperopt results JSON file.
///
/// Expected format: `{ "model_key": { "best_params": { ... }, ... } }`
/// Returns `None` if the file doesn't exist or can't be parsed.
fn load_hyperopt_params(hp_path: &Option<PathBuf>, model_key: &str) -> Option<Value> {
let file_path = hp_path.as_ref()?;
if !file_path.exists() {
info!("Hyperopt params file not found: {}, using defaults", file_path.display());
return None;
}
let contents = match std::fs::read_to_string(file_path) {
Ok(c) => c,
Err(e) => {
warn!("Failed to read hyperopt params {}: {}", file_path.display(), e);
return None;
}
};
let json: Value = match serde_json::from_str(&contents) {
Ok(v) => v,
Err(e) => {
warn!("Failed to parse hyperopt params JSON: {}", e);
return None;
}
};
let params = json.get(model_key)
.and_then(|m| m.get("best_params"))
.cloned();
if params.is_some() {
info!("Loaded hyperopt params for '{}' from {}", model_key, file_path.display());
} else {
warn!("No best_params found for '{}' in {}", model_key, file_path.display());
}
params
}
fn hp_f64(params: &Option<Value>, key: &str) -> Option<f64> {
params.as_ref()?.get(key)?.as_f64()
}
fn hp_usize(params: &Option<Value>, key: &str) -> Option<usize> {
params.as_ref()?.get(key)?.as_u64().map(|v| v as usize)
}
// ---------------------------------------------------------------------------
// 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,
hp: &Option<Value>,
) -> Result<Box<dyn UnifiedTrainable>> {
let lr = hp_f64(hp, "learning_rate").unwrap_or(learning_rate);
match name {
"tft" => {
let config = TFTConfig {
input_dim: feature_dim,
hidden_dim: hp_usize(hp, "hidden_size").unwrap_or(128),
num_heads: hp_usize(hp, "num_heads").unwrap_or(4),
num_layers: 2,
num_quantiles: 3,
num_static_features: 0,
dropout_rate: hp_f64(hp, "dropout").unwrap_or(0.1),
..TFTConfig::default()
};
let mut adapter = TrainableTFT::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create TFT: {}", e))?;
adapter
.set_learning_rate(lr)
.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(lr)
.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: hp_usize(hp, "hidden_size").unwrap_or(128),
output_size: 1,
backbone_hidden_sizes: vec![128, 64],
learning_rate: lr,
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: hp_usize(hp, "hidden_dim").unwrap_or(32),
num_layers: hp_usize(hp, "num_layers").unwrap_or(2),
max_nodes: 64,
max_edges: 128,
edge_dim: 4,
temporal_decay: hp_f64(hp, "temporal_decay").unwrap_or(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(lr)
.map_err(|e| anyhow::anyhow!("Failed to set TGGN learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"tlob" => {
let config = TLOBAdapterConfig {
d_model: hp_usize(hp, "d_model").unwrap_or(128),
num_heads: hp_usize(hp, "num_heads").unwrap_or(4),
num_layers: hp_usize(hp, "num_layers").unwrap_or(2),
seq_len: hp_usize(hp, "seq_len").unwrap_or(1),
feature_dim,
};
let mut adapter = TLOBTrainableAdapter::new(config, device)
.map_err(|e| anyhow::anyhow!("Failed to create TLOB: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set TLOB learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"kan" => {
let config = KANConfig {
grid_size: hp_usize(hp, "grid_size").unwrap_or(5),
spline_order: hp_usize(hp, "spline_order").unwrap_or(4),
layer_widths: vec![feature_dim, 32, 16, 1],
learning_rate: lr,
weight_decay: hp_f64(hp, "weight_decay").unwrap_or(1e-4),
grad_clip: hp_f64(hp, "grad_clip").unwrap_or(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: hp_usize(hp, "hidden_dim").unwrap_or(128),
..XLSTMConfig::default()
};
let mut adapter = XLSTMTrainableAdapter::new(config, device)
.map_err(|e| anyhow::anyhow!("Failed to create xLSTM: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set xLSTM learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"diffusion" => {
let config = DiffusionConfig {
feature_dim,
hidden_dim: hp_usize(hp, "hidden_dim").unwrap_or(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,
hp: &Option<Value>,
) -> 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, hp)?;
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 hp = load_hyperopt_params(&args.hyperopt_params, 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,
&hp,
) {
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(())
}