TFT create_model used TFTConfig::default() values for num_known_features(10) and num_unknown_features(210) totaling 220, but input_dim was 51 from the feature extractor. Set both explicitly: known=0, unknown=feature_dim. S3 uploader now uses path-style requests (required for Scaleway S3) and explicitly passes AWS credentials from env vars instead of relying on the instance metadata credential provider (unavailable on Kapsule). Also fix runner tags lost during session (kapsule, rust, docker restored). Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
799 lines
27 KiB
Rust
799 lines
27 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::completion::{write_failure_marker, write_success_marker, CompletionMetrics};
|
|
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,
|
|
num_known_features: 0,
|
|
num_unknown_features: feature_dim,
|
|
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 OHLCV market features),
|
|
/// 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)
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// Training orchestration
|
|
// ---------------------------------------------------------------------------
|
|
|
|
/// Container for per-model results collected during training.
|
|
struct TrainingResult {
|
|
model_name: String,
|
|
fold_results: Vec<(usize, f64)>,
|
|
total_epochs: usize,
|
|
}
|
|
|
|
/// Run the full walk-forward training pipeline.
|
|
///
|
|
/// Returns a vec of per-model results so `main()` can write completion markers.
|
|
fn run_training(args: &Args) -> Result<Vec<TrainingResult>> {
|
|
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());
|
|
|
|
let mut all_results = Vec::new();
|
|
|
|
// 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();
|
|
let mut total_epochs = 0_usize;
|
|
|
|
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) => {
|
|
total_epochs += args.epochs; // upper bound per fold
|
|
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());
|
|
|
|
all_results.push(TrainingResult {
|
|
model_name: (*model_name).to_owned(),
|
|
fold_results,
|
|
total_epochs,
|
|
});
|
|
}
|
|
|
|
Ok(all_results)
|
|
}
|
|
|
|
// ---------------------------------------------------------------------------
|
|
// 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();
|
|
|
|
// Ensure output directory exists before training so markers can always be written.
|
|
if let Err(e) = std::fs::create_dir_all(&args.output_dir) {
|
|
error!("Failed to create output dir {}: {}", args.output_dir.display(), e);
|
|
}
|
|
|
|
match run_training(&args) {
|
|
Ok(results) => {
|
|
for result in &results {
|
|
let best_val = result
|
|
.fold_results
|
|
.iter()
|
|
.map(|(_, loss)| *loss)
|
|
.fold(f64::MAX, f64::min);
|
|
|
|
let metrics = CompletionMetrics {
|
|
model: result.model_name.clone(),
|
|
symbol: args.symbol.clone(),
|
|
best_val_loss: if best_val < f64::MAX {
|
|
Some(best_val)
|
|
} else {
|
|
None
|
|
},
|
|
sharpe_ratio: None,
|
|
epochs_completed: result.total_epochs,
|
|
folds_completed: result.fold_results.len(),
|
|
};
|
|
write_success_marker(&args.output_dir, &metrics);
|
|
}
|
|
Ok(())
|
|
}
|
|
Err(e) => {
|
|
let msg = format!("{:#}", e);
|
|
error!("Training failed: {}", msg);
|
|
write_failure_marker(&args.output_dir, &msg);
|
|
Err(e)
|
|
}
|
|
}
|
|
}
|