- 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>
616 lines
20 KiB
Rust
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(())
|
|
}
|