Files
foxhunt/crates/ml/examples/evaluate_supervised.rs
jgrusewski bfd2253a9d fix: wire real GPU backprop + SPSA gradients, fix checkpoint loading, eliminate candle from examples
- GpuAdamW: add grad_scale param to CUDA kernel — gradient clipping was computed but never applied
- PPO load_checkpoint: load .actor.bin/.critic.bin weights (was Xavier re-init with TODO)
- CudaLinear::set_weights(): new method for checkpoint weight import
- TLOB/KAN/TGGN/Liquid backward: real GPU backprop via GpuLinear::backward() + GpuAdamW
- Mamba2 backward: SPSA gradient estimation replacing random pseudo-gradients (Spall 1992)
- Mamba2 adapter: wire SPSA backward with GPU-cached input/target/loss tensors
- TFT/xLSTM/Diffusion backward: explicit errors routing to native train() methods
- TLOB load_checkpoint: load .weights.json via GpuVarStore::import_from_host()
- train_baseline_supervised: 30 candle→native API fixes (Tensor/Device eliminated)
- evaluate_baseline: 38 candle→native API fixes (DQN/PPO/supervised GPU eval paths)
- evaluate_supervised: candle→native fixes (forward_loss instead of forward+compute_loss)
- cuda_test: rewrite to cudarc 0.19 (MlDevice, CudaSlice, memcpy)
- train_baseline_rl: Device→CudaContext for GPU double-buffer
- hyperopt_baseline_rl: CudaContext→MlDevice::cuda() for device pool
- xLSTM deterministic test: fix for stateful LSTM (hidden state changes between predictions)
- Liquid early stopping test: deterministic data for reliable convergence
- Mamba2Config: add spsa_epsilon field (default 0.01, serde backward-compatible)
- Clean stale candle comments from trainer, inference_validator, mamba optimizer

1853 tests pass (302+359+168+169+855), 0 failures, 0 clippy warnings, 8/8 examples compile.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-20 15:41:42 +01:00

902 lines
31 KiB
Rust
Raw Permalink Blame History

This file contains ambiguous Unicode characters
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.
#![allow(
clippy::assertions_on_constants,
clippy::assertions_on_result_states,
clippy::clone_on_copy,
clippy::decimal_literal_representation,
clippy::doc_markdown,
clippy::empty_line_after_doc_comments,
clippy::field_reassign_with_default,
clippy::get_unwrap,
clippy::identity_op,
clippy::inconsistent_digit_grouping,
clippy::indexing_slicing,
clippy::integer_division,
clippy::len_zero,
clippy::let_underscore_must_use,
clippy::manual_div_ceil,
clippy::manual_let_else,
clippy::manual_range_contains,
clippy::modulo_arithmetic,
clippy::needless_range_loop,
clippy::non_ascii_literal,
clippy::redundant_clone,
clippy::shadow_reuse,
clippy::shadow_same,
clippy::shadow_unrelated,
clippy::single_match_else,
clippy::str_to_string,
clippy::string_slice,
clippy::tests_outside_test_module,
clippy::too_many_lines,
clippy::unnecessary_wraps,
clippy::unseparated_literal_suffix,
clippy::use_debug,
clippy::useless_vec,
clippy::wildcard_enum_match_arm,
clippy::else_if_without_else,
clippy::expect_used,
clippy::missing_const_for_fn,
clippy::similar_names,
clippy::type_complexity,
clippy::collapsible_else_if,
clippy::doc_lazy_continuation,
clippy::items_after_test_module,
clippy::map_clone,
clippy::multiple_unsafe_ops_per_block,
clippy::unwrap_or_default,
clippy::assign_op_pattern,
clippy::needless_borrow,
clippy::println_empty_string,
clippy::unnecessary_cast,
clippy::used_underscore_binding,
clippy::create_dir,
clippy::implicit_saturating_sub,
clippy::exit,
clippy::expect_fun_call,
clippy::too_many_arguments,
clippy::unnecessary_map_or,
clippy::unwrap_used,
dead_code,
unused_imports,
unused_variables,
clippy::cloned_ref_to_slice_refs,
clippy::neg_multiply,
clippy::while_let_loop,
clippy::bool_assert_comparison,
clippy::excessive_precision,
clippy::trivially_copy_pass_by_ref,
clippy::op_ref,
clippy::redundant_closure,
clippy::unnecessary_lazy_evaluations,
clippy::if_then_some_else_none,
clippy::unnecessary_to_owned,
clippy::single_component_path_imports,
)]
//! Walk-forward evaluation binary for supervised baseline models.
//!
//! Loads trained model checkpoints, runs inference on walk-forward test data,
//! converts directional predictions to trading signals, computes financial
//! metrics (Sharpe, drawdown, win rate, profit factor), and generates a JSON
//! report.
//!
//! # Usage
//!
//! ```bash
//! SQLX_OFFLINE=true cargo run -p ml --example evaluate_supervised -- \
//! --model tft --models-dir ml/trained_models \
//! --data-dir test_data/futures-baseline \
//! --output ml/trained_models/supervised_eval_report.json
//! ```
#![allow(unused_crate_dependencies)]
#![deny(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
use std::path::PathBuf;
use anyhow::{Context, Result};
// candle eliminated — test uses native APIs
use clap::Parser;
use serde::Serialize;
use tracing::{error, info, warn};
use common::metrics::{server as metrics_server, training_metrics as tm};
use ml::features::extraction::extract_ml_features;
use ml::features::extraction::FeatureVector;
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};
// Model adapter imports (same as train_baseline_supervised)
use ml::diffusion::{DiffusionConfig, DiffusionTrainableAdapter};
use ml::kan::{KANConfig, KANTrainableAdapter};
use ml::liquid::{CfCTrainConfig, DeviceConfig, LiquidTrainableAdapter};
use ml::mamba::{Mamba2Config, trainable_adapter::Mamba2TrainableAdapter};
use ml::tft::{TFTConfig, TrainableTFT};
use ml::tgnn::trainable_adapter::TGGNTrainableAdapter;
use ml::tgnn::TGGNConfig;
use ml::tlob::{TLOBAdapterConfig, TLOBTrainableAdapter};
use ml::xlstm::{XLSTMConfig, XLSTMTrainableAdapter};
// ---------------------------------------------------------------------------
// CLI Arguments
// ---------------------------------------------------------------------------
/// Walk-forward evaluation binary for supervised baseline models.
#[derive(Parser, Debug)]
#[command(
name = "evaluate_supervised",
about = "Evaluate trained supervised model checkpoints with walk-forward test data"
)]
struct Args {
/// Directory containing trained model checkpoints
#[arg(long, default_value = "ml/trained_models")]
models_dir: PathBuf,
/// Path to directory containing .dbn.zst files (env: FOXHUNT_DATA_DIR)
#[arg(long, env = "FOXHUNT_DATA_DIR")]
data_dir: PathBuf,
/// Output path for evaluation report JSON
#[arg(long, default_value = "ml/trained_models/supervised_eval_report.json")]
output: PathBuf,
/// Which model to evaluate: tft, mamba2, liquid, tggn, tlob, kan, xlstm, diffusion
#[arg(long)]
model: String,
/// Feature dimension (must match `extract_ml_features` 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,
/// Symbol subdirectory to load (e.g. "ES.FUT", "NQ.FUT")
#[arg(long, default_value = "ES.FUT")]
symbol: String,
/// 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 (1 bps = 0.01%)
#[arg(long, default_value_t = 1.0)]
tx_cost_bps: f64,
/// Instrument tick size in price units (ES=0.25, NQ=0.25, ZN=1/64)
#[arg(long, default_value_t = 0.25)]
tick_size: f64,
/// Typical bid-ask spread in ticks (ES=1.0, ZN=1.0, 6E=2.0)
#[arg(long, default_value_t = 1.0)]
spread_ticks: f64,
/// Minimum prediction magnitude (in bps) to trigger a trade; below this → HOLD
#[arg(long, default_value_t = 0.5)]
signal_threshold_bps: f64,
}
// ---------------------------------------------------------------------------
// Report Data Types
// ---------------------------------------------------------------------------
#[derive(Debug, Serialize)]
struct FoldMetrics {
fold: usize,
model: String,
sharpe_ratio: f64,
max_drawdown_pct: f64,
win_rate_pct: f64,
profit_factor: f64,
total_return_pct: f64,
num_trades: usize,
directional_accuracy_pct: f64,
test_start: String,
test_end: String,
}
#[derive(Debug, Serialize)]
struct AggregateMetrics {
avg_sharpe: f64,
avg_drawdown: f64,
avg_win_rate: f64,
avg_directional_accuracy: f64,
avg_profit_factor: f64,
}
#[derive(Debug, Serialize)]
struct SanityChecks {
beats_random: bool,
action_diversity: bool,
fold_consistency: bool,
}
#[derive(Debug, Serialize)]
struct EvaluationReport {
model: String,
folds: Vec<FoldMetrics>,
aggregate: AggregateMetrics,
sanity_checks: SanityChecks,
}
// ---------------------------------------------------------------------------
// Financial Metrics (same as evaluate_baseline)
// ---------------------------------------------------------------------------
struct ComputedMetrics {
sharpe_ratio: f64,
max_drawdown_pct: f64,
win_rate_pct: f64,
profit_factor: f64,
total_return_pct: f64,
num_trades: usize,
}
fn compute_metrics(returns: &[f64]) -> ComputedMetrics {
let n = returns.len();
if n == 0 {
return ComputedMetrics {
sharpe_ratio: 0.0,
max_drawdown_pct: 0.0,
win_rate_pct: 0.0,
profit_factor: 0.0,
total_return_pct: 0.0,
num_trades: 0,
};
}
let num_trades = returns.iter().filter(|&&r| r.abs() > 1e-12).count();
let sum: f64 = returns.iter().sum();
let mean = sum / n as f64;
let variance: f64 = returns.iter().map(|&r| (r - mean).powi(2)).sum::<f64>() / n as f64;
let std = variance.sqrt();
// Annualized Sharpe: 1-min bars, ~1380 bars/day × 252 days
let bars_per_year: f64 = 252.0 * 1380.0;
let sharpe_ratio = if std > 1e-12 {
(mean / std) * bars_per_year.sqrt()
} else {
0.0
};
let mut equity = 1.0_f64;
let mut peak = 1.0_f64;
let mut max_drawdown = 0.0_f64;
for &ret in returns {
equity += ret;
if equity > peak {
peak = equity;
}
let drawdown = if peak > 1e-12 {
(peak - equity) / peak
} else {
0.0
};
if drawdown > max_drawdown {
max_drawdown = drawdown;
}
}
let wins = returns.iter().filter(|&&r| r > 0.0).count();
let win_rate_pct = if num_trades > 0 {
(wins as f64 / num_trades as f64) * 100.0
} else {
0.0
};
let gross_profit: f64 = returns.iter().filter(|&&r| r > 0.0).sum();
let gross_loss: f64 = returns.iter().filter(|&&r| r < 0.0).map(|&r| r.abs()).sum();
let profit_factor = if gross_loss > 1e-12 {
gross_profit / gross_loss
} else if gross_profit > 0.0 {
f64::INFINITY
} else {
0.0
};
ComputedMetrics {
sharpe_ratio,
max_drawdown_pct: max_drawdown * 100.0,
win_rate_pct,
profit_factor,
total_return_pct: sum * 100.0,
num_trades,
}
}
// ---------------------------------------------------------------------------
// Model Factory (same defaults as train_baseline_supervised)
// ---------------------------------------------------------------------------
fn create_model(
name: &str,
feature_dim: usize,
) -> Result<Box<dyn UnifiedTrainable>> {
let lr = 1e-3; // Default; doesn't matter for eval (no training)
match name {
"tft" => {
let config = TFTConfig {
input_dim: feature_dim,
hidden_dim: 128,
num_heads: 4,
num_layers: 2,
num_quantiles: 3,
num_static_features: 0,
num_known_features: 0,
num_unknown_features: feature_dim,
sequence_length: 1,
prediction_horizon: 1,
dropout_rate: 0.1,
..TFTConfig::default()
};
let mut adapter = TrainableTFT::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create TFT: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set TFT learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"mamba2" => {
let config = Mamba2Config {
d_model: 128,
num_layers: 4,
d_state: 16,
max_seq_len: 60,
..Mamba2Config::default()
};
let native_device = ml_core::native_types::NativeDevice::Cuda(0);
let mut adapter = Mamba2TrainableAdapter::from_native_device(config, &native_device)
.map_err(|e| anyhow::anyhow!("Failed to create Mamba2: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set Mamba2 learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"liquid" => {
let config = CfCTrainConfig {
input_size: feature_dim,
hidden_size: 128,
output_size: 1,
backbone_hidden_sizes: vec![128, 64],
learning_rate: lr,
device: DeviceConfig::Auto,
..CfCTrainConfig::default()
};
let adapter = LiquidTrainableAdapter::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create Liquid: {}", e))?;
Ok(Box::new(adapter))
}
"tggn" => {
let config = TGGNConfig {
node_dim: feature_dim,
hidden_dim: 32,
num_layers: 2,
max_nodes: 64,
max_edges: 128,
edge_dim: 4,
temporal_decay: 0.99,
update_frequency_ns: 1_000_000,
use_simd: false,
};
let mut adapter = TGGNTrainableAdapter::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create TGGN: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set TGGN learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"tlob" => {
let config = TLOBAdapterConfig {
d_model: 128,
num_heads: 4,
num_layers: 2,
seq_len: 1,
feature_dim,
};
let mut adapter = TLOBTrainableAdapter::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create TLOB: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set TLOB learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"kan" => {
let config = KANConfig {
grid_size: 5,
spline_order: 4,
layer_widths: vec![feature_dim, 32, 16, 1],
learning_rate: lr,
weight_decay: 1e-4,
grad_clip: 1.0,
};
let adapter = KANTrainableAdapter::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create KAN: {}", e))?;
Ok(Box::new(adapter))
}
"xlstm" => {
let config = XLSTMConfig {
input_dim: feature_dim,
hidden_dim: 128,
..XLSTMConfig::default()
};
let mut adapter = XLSTMTrainableAdapter::new(config)
.map_err(|e| anyhow::anyhow!("Failed to create xLSTM: {}", e))?;
adapter
.set_learning_rate(lr)
.map_err(|e| anyhow::anyhow!("Failed to set xLSTM learning rate: {}", e))?;
Ok(Box::new(adapter))
}
"diffusion" => {
let config = DiffusionConfig {
feature_dim,
hidden_dim: 128,
..DiffusionConfig::default()
};
let ctx = cudarc::driver::CudaContext::new(0)
.map_err(|e| anyhow::anyhow!("CUDA context for Diffusion: {}", e))?;
let stream = ctx
.new_stream()
.map_err(|e| anyhow::anyhow!("CUDA stream for Diffusion: {}", e))?;
let adapter = DiffusionTrainableAdapter::new(config, &stream)
.map_err(|e| anyhow::anyhow!("Failed to create Diffusion: {}", e))?;
Ok(Box::new(adapter))
}
_ => anyhow::bail!("Unknown model: {}", name),
}
}
// ---------------------------------------------------------------------------
// Supervised Evaluation
// ---------------------------------------------------------------------------
/// Run supervised inference on test features and return per-bar trade returns,
/// action counts, and directional accuracy.
#[allow(clippy::cognitive_complexity)]
fn evaluate_fold(
fold: usize,
model_name: &str,
test_features: &[FeatureVector],
test_bars: &[OHLCVBar],
args: &Args,
) -> Result<(Vec<f64>, [usize; 3], f64)> {
let ckpt_path = args
.models_dir
.join(model_name)
.join(format!("{}_fold{}_best", model_name, fold));
// Check if checkpoint metadata exists
let meta_path = format!("{}.json", ckpt_path.display());
if !std::path::Path::new(&meta_path).exists() {
anyhow::bail!(
"Checkpoint not found: {} (looked for {})",
ckpt_path.display(),
meta_path
);
}
// Create model with same config as training, then load checkpoint
let mut model = create_model(model_name, args.feature_dim)?;
let ckpt_str = ckpt_path.to_str().unwrap_or("checkpoint");
model
.load_checkpoint(ckpt_str)
.map_err(|e| anyhow::anyhow!("Failed to load {} checkpoint fold {}: {}", model_name, fold, e))?;
info!(
" [{}] Loaded checkpoint: {}",
model_name.to_uppercase(),
ckpt_path.display()
);
let n = test_features.len();
let mut returns = Vec::with_capacity(n);
let mut action_counts = [0_usize; 3]; // [buy, sell, hold]
let mut correct_direction = 0_usize;
let mut total_predictions = 0_usize;
let threshold = args.signal_threshold_bps;
for i in 0..n.saturating_sub(1) {
let Some(feat) = test_features.get(i) else {
continue;
};
// Convert feature vector to f32 slice for forward_loss
let input_f32: Vec<f32> = feat.iter().map(|&v| v as f32).collect();
// Use a dummy zero target — we only care about the returned loss as a
// proxy for the directional prediction magnitude.
let dummy_target: Vec<f32> = vec![0.0_f32];
// Run forward_loss; the returned loss value acts as the signed
// prediction (models return MSE-style loss against the zero target,
// which encodes the prediction magnitude).
let pred_value = match model.forward_loss(&input_f32, &dummy_target) {
Ok(v) => v,
Err(e) => {
warn!(" [{}] forward error at step {}: {}", model_name, i, e);
continue;
}
};
// Convert prediction to action: positive → buy, negative → sell, small → hold
let action: u8 = if pred_value > threshold {
0 // BUY
} else if pred_value < -threshold {
1 // SELL
} else {
2 // HOLD
};
if let Some(count) = action_counts.get_mut(action as usize) {
*count += 1;
}
// Compute actual return
let close_cur = test_bars.get(i).map(|b| b.close).unwrap_or(0.0);
let close_next = test_bars.get(i + 1).map(|b| b.close).unwrap_or(close_cur);
let pct_change = if close_cur.abs() > 1e-12 {
((close_next - close_cur) / close_cur).clamp(-args.max_bar_return, args.max_bar_return)
} else {
0.0
};
// Directional accuracy: did the sign of prediction match the sign of actual return?
let actual_bps = pct_change * 10_000.0;
if actual_bps.abs() > 1e-6 {
total_predictions += 1;
if (pred_value > 0.0 && actual_bps > 0.0) || (pred_value < 0.0 && actual_bps < 0.0) {
correct_direction += 1;
}
}
// Apply transaction costs
let total_cost_bps =
args.tx_cost_bps + spread_cost_bps(close_cur, args.tick_size, args.spread_ticks);
let total_cost = total_cost_bps * 0.0001;
let ret = match action {
0 => pct_change - total_cost, // BUY
1 => -pct_change - total_cost, // SELL
_ => 0.0, // HOLD
};
returns.push(ret);
}
let directional_accuracy = if total_predictions > 0 {
(correct_direction as f64 / total_predictions as f64) * 100.0
} else {
0.0
};
Ok((returns, action_counts, directional_accuracy))
}
// ---------------------------------------------------------------------------
// Main
// ---------------------------------------------------------------------------
#[allow(clippy::cognitive_complexity, clippy::too_many_lines)]
fn main() -> Result<()> {
// Initialize tracing with optional OTLP export to Tempo
let otlp_endpoint = std::env::var("OTEL_EXPORTER_OTLP_ENDPOINT").ok();
if let Err(e) = common::observability::init_observability(
"evaluate_supervised",
otlp_endpoint.as_deref(),
) {
eprintln!("Observability init failed (non-fatal): {e}");
}
tm::init();
metrics_server::start_metrics_server(9094);
tm::set_active_workers(1.0);
let args = Args::parse();
info!("=== Walk-Forward Supervised Evaluation ===");
info!(" Model: {}", args.model);
info!(" Symbol: {}", args.symbol);
info!(" Models dir: {}", args.models_dir.display());
info!(" Data dir: {}", args.data_dir.display());
info!(" Output: {}", args.output.display());
info!(" Feature dim: {}", args.feature_dim);
info!(" Signal threshold: {:.1} bps", args.signal_threshold_bps);
info!(
" Tx cost: {:.1} bps commission + {:.1} tick spread (tick_size={:.4})",
args.tx_cost_bps, args.spread_ticks, args.tick_size
);
// 1. Load all OHLCV bars
info!("Step 1/4: Loading OHLCV bars from DBN files...");
let data_load_start = std::time::Instant::now();
let bars = load_all_bars(&args.data_dir, &args.symbol)?;
tm::record_data_load(&args.model, data_load_start.elapsed().as_secs_f64());
if bars.is_empty() {
anyhow::bail!("No bars loaded from {}", args.data_dir.display());
}
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(),
);
// 2. Generate walk-forward windows
info!("Step 2/4: Generating 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(&bars, &wf_config);
if windows.is_empty() {
anyhow::bail!(
"No walk-forward windows generated. Need at least {} months of data.",
wf_config.initial_train_months + wf_config.val_months + wf_config.test_months
);
}
info!(" Generated {} walk-forward folds", windows.len());
// 3. Evaluate each fold
info!("Step 3/4: Evaluating {} on test data...", args.model);
let mut all_fold_metrics: Vec<FoldMetrics> = Vec::new();
let mut all_action_counts: Vec<[usize; 3]> = Vec::new();
for window in &windows {
info!(
"--- Fold {} --- Test: {} bars ({} to {})",
window.fold,
window.test.len(),
window.test
.first()
.map(|b| b.timestamp.to_string())
.unwrap_or_default(),
window.test
.last()
.map(|b| b.timestamp.to_string())
.unwrap_or_default(),
);
// Load NormStats from training
let norm_path = args
.models_dir
.join(&args.model)
.join(format!("norm_stats_fold{}.json", window.fold));
let norm_stats: NormStats = if norm_path.exists() {
let norm_json = std::fs::read_to_string(&norm_path)
.with_context(|| format!("Failed to read {}", norm_path.display()))?;
serde_json::from_str(&norm_json)
.with_context(|| format!("Failed to parse {}", norm_path.display()))?
} else {
anyhow::bail!(
"NormStats not found at {} - cannot evaluate without training-set statistics. \
Re-run training to generate norm_stats files.",
norm_path.display()
);
};
// Extract and normalize test features
let test_features = match extract_ml_features(&window.test) {
Ok(f) => f,
Err(e) => {
warn!(
" Fold {} -- test feature extraction failed: {}",
window.fold, e
);
continue;
}
};
if test_features.is_empty() {
warn!(" Fold {} -- empty test features, skipping", window.fold);
continue;
}
let test_norm = norm_stats.normalize_batch(&test_features);
// Align bars to features
let warmup_offset = window.test.len().saturating_sub(test_norm.len());
let test_bars_aligned = if warmup_offset < window.test.len() {
window.test.get(warmup_offset..).unwrap_or(&window.test)
} else {
&window.test
};
let test_start = test_bars_aligned
.first()
.map(|b| b.timestamp.format("%Y-%m-%d").to_string())
.unwrap_or_default();
let test_end = test_bars_aligned
.last()
.map(|b| b.timestamp.format("%Y-%m-%d").to_string())
.unwrap_or_default();
let model_name = &args.model;
match evaluate_fold(
window.fold,
model_name,
&test_norm,
test_bars_aligned,
&args,
) {
Ok((returns, action_counts, directional_accuracy)) => {
let metrics = compute_metrics(&returns);
let fold_str = window.fold.to_string();
tm::set_epoch(model_name, &fold_str, window.fold as f64);
tm::set_eval_metrics(
model_name,
&fold_str,
directional_accuracy / 100.0,
metrics.sharpe_ratio,
metrics.profit_factor,
metrics.total_return_pct / 100.0,
);
info!(
" [{}] Fold {} -- Sharpe={:.4} MaxDD={:.2}% WinRate={:.1}% PF={:.2} Return={:.4}% Trades={} DirAcc={:.1}%",
args.model.to_uppercase(),
window.fold,
metrics.sharpe_ratio,
metrics.max_drawdown_pct,
metrics.win_rate_pct,
metrics.profit_factor,
metrics.total_return_pct,
metrics.num_trades,
directional_accuracy,
);
info!(
" [{}] Actions -- BUY={} SELL={} HOLD={}",
args.model.to_uppercase(),
action_counts.first().copied().unwrap_or(0),
action_counts.get(1).copied().unwrap_or(0),
action_counts.get(2).copied().unwrap_or(0),
);
all_fold_metrics.push(FoldMetrics {
fold: window.fold,
model: args.model.clone(),
sharpe_ratio: metrics.sharpe_ratio,
max_drawdown_pct: metrics.max_drawdown_pct,
win_rate_pct: metrics.win_rate_pct,
profit_factor: metrics.profit_factor,
total_return_pct: metrics.total_return_pct,
num_trades: metrics.num_trades,
directional_accuracy_pct: directional_accuracy,
test_start: test_start.clone(),
test_end: test_end.clone(),
});
all_action_counts.push(action_counts);
}
Err(e) => {
error!(
" [{}] Fold {} evaluation failed: {}",
args.model.to_uppercase(),
window.fold,
e
);
}
}
}
// 4. Compute aggregate and write report
info!("Step 4/4: Computing aggregate metrics...");
let n_folds = all_fold_metrics.len() as f64;
let aggregate = if n_folds > 0.0 {
AggregateMetrics {
avg_sharpe: all_fold_metrics.iter().map(|f| f.sharpe_ratio).sum::<f64>() / n_folds,
avg_drawdown: all_fold_metrics
.iter()
.map(|f| f.max_drawdown_pct)
.sum::<f64>()
/ n_folds,
avg_win_rate: all_fold_metrics.iter().map(|f| f.win_rate_pct).sum::<f64>() / n_folds,
avg_directional_accuracy: all_fold_metrics
.iter()
.map(|f| f.directional_accuracy_pct)
.sum::<f64>()
/ n_folds,
avg_profit_factor: all_fold_metrics
.iter()
.map(|f| if f.profit_factor.is_finite() { f.profit_factor } else { 0.0 })
.sum::<f64>()
/ n_folds,
}
} else {
AggregateMetrics {
avg_sharpe: 0.0,
avg_drawdown: 0.0,
avg_win_rate: 0.0,
avg_directional_accuracy: 0.0,
avg_profit_factor: 0.0,
}
};
info!(
" {} -- avg Sharpe={:.4} avg MaxDD={:.2}% avg WinRate={:.1}% avg DirAcc={:.1}%",
args.model.to_uppercase(),
aggregate.avg_sharpe,
aggregate.avg_drawdown,
aggregate.avg_win_rate,
aggregate.avg_directional_accuracy,
);
// Sanity checks
let beats_random = all_fold_metrics.iter().any(|f| f.sharpe_ratio > 0.0);
let mut total_actions = [0_usize; 3];
for counts in &all_action_counts {
for (total, &count) in total_actions.iter_mut().zip(counts.iter()) {
*total += count;
}
}
let action_diversity = total_actions.iter().all(|&c| c > 0);
let sharpe_values: Vec<f64> = all_fold_metrics.iter().map(|f| f.sharpe_ratio).collect();
let fold_consistency = if sharpe_values.is_empty() {
false
} else {
let n = sharpe_values.len() as f64;
let mean_sharpe = sharpe_values.iter().sum::<f64>() / n;
let var = sharpe_values
.iter()
.map(|&s| (s - mean_sharpe).powi(2))
.sum::<f64>()
/ n;
let std_sharpe = var.sqrt();
std_sharpe < 2.0 * mean_sharpe.abs()
};
let report = EvaluationReport {
model: args.model.clone(),
folds: all_fold_metrics,
aggregate,
sanity_checks: SanityChecks {
beats_random,
action_diversity,
fold_consistency,
},
};
// Save report
if let Some(parent) = args.output.parent() {
std::fs::create_dir_all(parent)
.with_context(|| format!("Failed to create output dir: {}", parent.display()))?;
}
let report_json = serde_json::to_string_pretty(&report)
.context("Failed to serialize evaluation report")?;
std::fs::write(&args.output, &report_json)
.with_context(|| format!("Failed to write report to {}", args.output.display()))?;
info!("=== Evaluation Complete ===");
info!(" Report saved to: {}", args.output.display());
info!(" Total fold evaluations: {}", report.folds.len());
tm::set_active_workers(0.0);
// Push final metrics to pushgateway so they persist after pod termination
if let Err(e) = metrics_server::push_to_gateway(None, "evaluate_supervised") {
tracing::warn!("Failed to push metrics to gateway (non-fatal): {e}");
}
Ok(())
}