Files
foxhunt/ml/examples/train_dqn.rs
jgrusewski 86ed7af58f fix(ml): DQN early stopping checkpoint naming (Option B)
Added is_final parameter to checkpoint callback to distinguish final
checkpoints from regular epoch checkpoints. Early stopping now saves
as dqn_final_epoch{N}.safetensors instead of dqn_epoch_{N}.safetensors.

Changes:
- Updated callback signature: Fn(usize, Vec<u8>, bool)
- Early stopping passes is_final=true
- Regular checkpoints pass is_final=false
- Callback uses final naming when is_final=true

Fixes checkpoint overwrite bug where final model was indistinguishable
from regular epoch checkpoints.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-25 23:19:40 +02:00

320 lines
10 KiB
Rust
Raw Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! DQN Training Example
//!
//! Trains a DQN model on market data and saves checkpoints to disk.
//!
//! # Usage
//!
//! ```bash
//! # Train with default parameters (100 epochs)
//! cargo run -p ml --example train_dqn --release --features cuda
//!
//! # Custom epochs and output path
//! cargo run -p ml --example train_dqn --release --features cuda -- \
//! --epochs 500 \
//! --output ml/trained_models/dqn_model.safetensors
//!
//! # Custom data directory
//! cargo run -p ml --example train_dqn --release --features cuda -- \
//! --data-dir test_data/real/databento/ml_training \
//! --epochs 500
//! ```
// Use mimalloc allocator for 10-25% performance improvement
#[cfg(feature = "mimalloc-allocator")]
use mimalloc::MiMalloc;
#[cfg(feature = "mimalloc-allocator")]
#[global_allocator]
static GLOBAL: MiMalloc = MiMalloc;
use anyhow::{Context, Result};
use clap::Parser;
use std::path::PathBuf;
use tracing::info;
use tracing_subscriber::FmtSubscriber;
use ml::checkpoint::{CheckpointConfig, CheckpointManager};
use ml::data_loaders::BarSamplingMethod;
use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer};
/// Train DQN model on market data
#[derive(Debug, Parser)]
#[command(name = "train_dqn", about = "Train DQN model on market data")]
struct Opts {
/// Number of training epochs
#[arg(long, default_value = "100")]
epochs: usize,
/// Learning rate
#[arg(long, default_value = "0.0001")]
learning_rate: f64,
/// Batch size (max 230 for RTX 3050 Ti 4GB)
#[arg(long, default_value = "128")]
batch_size: usize,
/// Discount factor (gamma)
#[arg(long, default_value = "0.99")]
gamma: f64,
/// Checkpoint save frequency (epochs)
#[arg(long, default_value = "10")]
checkpoint_frequency: usize,
/// Output directory for trained model
#[arg(long, default_value = "ml/trained_models")]
output_dir: String,
/// Data directory containing DBN files
#[arg(long, default_value = "test_data/real/databento/ml_training")]
data_dir: String,
/// Parquet file path (overrides data_dir if specified)
#[arg(long)]
parquet_file: Option<String>,
/// Verbose logging
#[arg(short, long)]
verbose: bool,
/// Enable early stopping (recommended, use --no-early-stopping to disable)
#[arg(long)]
early_stopping: bool,
/// Disable early stopping
#[arg(long)]
no_early_stopping: bool,
/// Q-value floor threshold for early stopping
#[arg(long, default_value = "0.5")]
q_value_floor: f64,
/// Minimum loss improvement percentage for plateau detection
#[arg(long, default_value = "2.0")]
min_loss_improvement: f64,
/// Plateau detection window size (epochs)
#[arg(long, default_value = "30")]
plateau_window: usize,
/// Alternative bar sampling method (time, tick, volume, dollar, imbalance, run)
#[arg(long, default_value = "time")]
bar_method: String,
/// Bar sampling threshold (tick count, volume, dollar value, imbalance, or run length)
#[arg(long)]
bar_threshold: Option<f64>,
}
#[tokio::main]
async fn main() -> Result<()> {
// Parse CLI options
let opts = Opts::parse();
// Setup logging
let level = if opts.verbose {
tracing::Level::DEBUG
} else {
tracing::Level::INFO
};
let subscriber = FmtSubscriber::builder().with_max_level(level).finish();
tracing::subscriber::set_global_default(subscriber)
.context("Failed to set tracing subscriber")?;
#[cfg(feature = "mimalloc-allocator")]
info!("🚀 Using mimalloc allocator for improved performance");
#[cfg(not(feature = "mimalloc-allocator"))]
info!(" Using system allocator (consider --features mimalloc-allocator for 10-25% speedup)");
info!("🚀 Starting DQN Training");
info!("Configuration:");
info!(" • Epochs: {}", opts.epochs);
info!(" • Learning rate: {}", opts.learning_rate);
info!(" • Batch size: {}", opts.batch_size);
info!(" • Gamma: {}", opts.gamma);
info!(
" • Checkpoint frequency: {} epochs",
opts.checkpoint_frequency
);
info!(" • Output directory: {}", opts.output_dir);
info!(" • Data directory: {}", opts.data_dir);
info!(" • Bar sampling method: {}", opts.bar_method);
if let Some(threshold) = opts.bar_threshold {
info!(" • Bar threshold: {}", threshold);
}
// Determine early stopping (enabled by default, unless --no-early-stopping is specified)
let early_stopping_enabled = !opts.no_early_stopping;
info!(
" • Early stopping: {}",
if early_stopping_enabled {
"enabled"
} else {
"disabled"
}
);
if early_stopping_enabled {
info!(" - Q-value floor: {}", opts.q_value_floor);
info!(" - Min loss improvement: {}%", opts.min_loss_improvement);
info!(" - Plateau window: {} epochs", opts.plateau_window);
}
// Create output directory
let output_path = PathBuf::from(&opts.output_dir);
if !output_path.exists() {
std::fs::create_dir_all(&output_path).context("Failed to create output directory")?;
info!("✅ Created output directory: {}", opts.output_dir);
}
// Configure DQN hyperparameters
let hyperparams = DQNHyperparameters {
learning_rate: opts.learning_rate,
batch_size: opts.batch_size,
gamma: opts.gamma,
epsilon_start: 1.0,
epsilon_end: 0.01,
epsilon_decay: 0.995,
buffer_size: 100_000,
epochs: opts.epochs,
checkpoint_frequency: opts.checkpoint_frequency,
early_stopping_enabled,
q_value_floor: opts.q_value_floor,
min_loss_improvement_pct: opts.min_loss_improvement,
plateau_window: opts.plateau_window,
min_epochs_before_stopping: 50,
};
// Configure alternative bar sampling (Wave B)
let bar_sampling = match opts.bar_method.as_str() {
"tick" => BarSamplingMethod::TickBars(opts.bar_threshold.unwrap_or(100.0) as usize),
"volume" => BarSamplingMethod::VolumeBars(opts.bar_threshold.unwrap_or(10000.0)),
"dollar" => BarSamplingMethod::DollarBars(opts.bar_threshold.unwrap_or(2_000_000.0)),
"imbalance" => BarSamplingMethod::ImbalanceBars(opts.bar_threshold.unwrap_or(1000.0)),
"run" => BarSamplingMethod::RunBars(opts.bar_threshold.unwrap_or(50.0) as usize),
_ => BarSamplingMethod::TimeBars,
};
info!("✅ Bar sampling configured: {:?}", bar_sampling);
// Create DQN trainer
let mut trainer = DQNTrainer::new(hyperparams).context("Failed to create DQN trainer")?;
// Note: DQN trainer will need to accept bar_sampling parameter
// This requires updating DQNTrainer to use DbnSequenceLoader
info!("✅ DQN trainer initialized");
// Setup checkpoint manager
let checkpoint_config = CheckpointConfig {
base_dir: output_path.clone(),
max_checkpoints_per_model: 10,
auto_cleanup: true,
validate_checksums: true,
..Default::default()
};
let _checkpoint_manager =
CheckpointManager::new(checkpoint_config).context("Failed to create checkpoint manager")?;
info!("✅ Checkpoint manager initialized (max 10 checkpoints, auto-cleanup enabled)");
// Create checkpoint callback
let output_dir_for_callback = opts.output_dir.clone();
let checkpoint_callback = move |epoch: usize, model_data: Vec<u8>, is_final: bool| -> Result<String> {
let filename = if is_final {
format!("dqn_final_epoch{}.safetensors", epoch)
} else {
format!("dqn_epoch_{}.safetensors", epoch)
};
let checkpoint_path = PathBuf::from(&output_dir_for_callback).join(filename);
// Save checkpoint to disk
std::fs::write(&checkpoint_path, &model_data)
.context(format!("Failed to save checkpoint: {:?}", checkpoint_path))?;
info!(
"💾 Checkpoint saved: {} ({} bytes)",
checkpoint_path.display(),
model_data.len()
);
Ok(checkpoint_path.to_string_lossy().to_string())
};
// Train the model
info!("\n🏋️ Starting training...\n");
let start_time = std::time::Instant::now();
let metrics = if let Some(ref parquet_path) = opts.parquet_file {
info!("Using Parquet file: {}", parquet_path);
trainer
.train_from_parquet(parquet_path, checkpoint_callback)
.await
.context("Training from Parquet failed")?
} else {
info!("Using DBN directory: {}", opts.data_dir);
trainer
.train(&opts.data_dir, checkpoint_callback)
.await
.context("Training failed")?
};
let training_duration = start_time.elapsed();
// Print final metrics
info!("\n✅ Training completed successfully!");
info!("\n📊 Final Metrics:");
info!(" • Final loss: {:.6}", metrics.loss);
info!(" • Epochs trained: {}", metrics.epochs_trained);
info!(
" • Training time: {:.1}s ({:.1} min)",
metrics.training_time_seconds,
metrics.training_time_seconds / 60.0
);
info!(
" • Actual elapsed time: {:.1}s (includes data loading + overhead)",
training_duration.as_secs_f64()
);
info!(
" • Convergence: {}",
if metrics.convergence_achieved {
"✅ Yes"
} else {
"❌ No"
}
);
// Additional metrics from training
if let Some(avg_q_value) = metrics.additional_metrics.get("avg_q_value") {
info!(" • Average Q-value: {:.4}", avg_q_value);
}
if let Some(final_epsilon) = metrics.additional_metrics.get("final_epsilon") {
info!(" • Final epsilon: {:.4}", final_epsilon);
}
if let Some(grad_norm) = metrics.additional_metrics.get("avg_gradient_norm") {
info!(" • Average gradient norm: {:.6}", grad_norm);
}
// Save final model
let final_model_path = output_path.join(format!("dqn_final_epoch{}.safetensors", opts.epochs));
info!("\n💾 Saving final model to: {}", final_model_path.display());
// Get final model state
let final_checkpoint_data = trainer
.serialize_model()
.await
.context("Failed to serialize final model")?;
std::fs::write(&final_model_path, &final_checkpoint_data)
.context("Failed to save final model")?;
info!(
"✅ Final model saved: {} ({} bytes)",
final_model_path.display(),
final_checkpoint_data.len()
);
info!("\n🎉 DQN training complete!");
info!("📁 Model files saved to: {}", opts.output_dir);
Ok(())
}