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:
665
ml/examples/train_diffusion_dbn.rs
Normal file
665
ml/examples/train_diffusion_dbn.rs
Normal 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(())
|
||||
}
|
||||
615
ml/examples/train_kan_dbn.rs
Normal file
615
ml/examples/train_kan_dbn.rs
Normal 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(())
|
||||
}
|
||||
518
ml/examples/train_tggn_dbn.rs
Normal file
518
ml/examples/train_tggn_dbn.rs
Normal 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)
|
||||
}
|
||||
497
ml/examples/train_xlstm_dbn.rs
Normal file
497
ml/examples/train_xlstm_dbn.rs
Normal 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(())
|
||||
}
|
||||
Reference in New Issue
Block a user