Files
foxhunt/ml/examples/train_kan_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

616 lines
20 KiB
Rust

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