Files
foxhunt/ml/examples/train_diffusion_dbn.rs
jgrusewski 5fb84a02f0 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>
2026-02-24 20:47:50 +01:00

666 lines
21 KiB
Rust

//! **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(())
}