feat: training pipeline for all 10 ML ensemble models

- Add standalone training binaries: TGGN, KAN, xLSTM, Diffusion (DBN data)
- Update Dockerfile.training: 6 → 16 binaries (all 10 models + hyperopt + baseline)
- Expand train.sh: 4 → 10 models, fix registry URL and GPU pool nodeSelector
- Add GPU overlay manifests for trading-service and ml-training-service
- Create training data PVC and upload pod manifests
- Expand web-gateway model validation: 4 → 10 types (training + tune routes)
- Extend dashboard: 10 model cards grouped by category (RL/Temporal/Graph/Generative)
- Add training image build job to Gitea CI workflow
- Update GPU taint controller to exclude inference pool from tainting
- Fix job-template nodeSelector: gpu → gpu-training

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-24 20:47:50 +01:00
parent e45bda9412
commit 5fb84a02f0
19 changed files with 3602 additions and 35 deletions

View File

@@ -0,0 +1,665 @@
//! **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;
// ---------------------------------------------------------------------------
// 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,
}
// ---------------------------------------------------------------------------
// Data preparation
// ---------------------------------------------------------------------------
/// Extract normalized close-price sequences from OHLCV bars.
///
/// Returns a vector of f32 sequences, each of length `seq_len`, created by
/// sliding a window over the close prices. Prices are z-score normalized
/// within each window to keep the diffusion model input in a stable range.
fn prepare_sequences(bars: &[OHLCVBar], seq_len: usize) -> Vec<Vec<f32>> {
if bars.len() < seq_len {
return Vec::new();
}
let closes: Vec<f64> = bars.iter().map(|b| b.close).collect();
let n_sequences = closes.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> = closes
.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);
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(())
}

View File

@@ -0,0 +1,615 @@
//! 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(())
}

View File

@@ -0,0 +1,518 @@
//! 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;
// ---------------------------------------------------------------------------
// 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,
}
// ---------------------------------------------------------------------------
// 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,
) -> 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
let target = if cur_close.abs() > 1e-12 {
(next_bar.close / cur_close).ln()
} else {
0.0
};
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)?;
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)?;
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)
}

View File

@@ -0,0 +1,497 @@
//! 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;
// ---------------------------------------------------------------------------
// 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,
}
// ---------------------------------------------------------------------------
// 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,
) -> 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 ret = if close_prev.abs() > 1e-10 {
((close_cur - close_prev) / close_prev).clamp(-0.05, 0.05) as f32
} else {
0.0_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)?;
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(())
}