diff --git a/ml/src/validation/harness.rs b/ml/src/validation/harness.rs new file mode 100644 index 000000000..883772b1f --- /dev/null +++ b/ml/src/validation/harness.rs @@ -0,0 +1,504 @@ +//! Validation harness orchestrator. +//! +//! Coordinates walk-forward splitting, statistical tests (DSR, PBO, permutation), +//! and per-regime analysis into a single [`ValidationReport`] with a +//! [`ValidationVerdict`] summarising whether the strategy passes muster. + +use std::collections::HashMap; +use std::fmt; + +use serde::{Deserialize, Serialize}; + +use crate::dqn::RegimeType; +use crate::MLError; + +use super::{ + deflated_sharpe_ratio, excess_kurtosis, per_regime_breakdown, permutation_test, + probability_of_backtest_overfitting, sharpe_ratio, skewness, walk_forward_split, + DsrResult, PboResult, PermutationResult, RegimeMetrics, TimeSeriesData, + ValidatableStrategy, WalkForwardConfig, +}; + +// --------------------------------------------------------------------------- +// ValidationVerdict +// --------------------------------------------------------------------------- + +/// Overall verdict for a validated strategy. +/// +/// Criteria: +/// - **Pass**: DSR p < 0.05 AND PBO < 0.25 AND permutation p < 0.05 +/// - **Marginal**: at least one test passes at a relaxed threshold +/// - **Fail**: no statistical evidence of real skill +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum ValidationVerdict { + /// All three tests pass at strict thresholds. + Pass, + /// At least one test shows marginal evidence. + Marginal, + /// No statistical evidence of strategy skill. + Fail, +} + +impl fmt::Display for ValidationVerdict { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Pass => write!(f, "Pass"), + Self::Marginal => write!(f, "Marginal"), + Self::Fail => write!(f, "Fail"), + } + } +} + +// --------------------------------------------------------------------------- +// ValidationReport +// --------------------------------------------------------------------------- + +/// Full report produced by [`ValidationHarness::validate`]. +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ValidationReport { + /// Human-readable strategy name. + pub strategy_name: String, + /// Sharpe ratio computed on each fold's test returns. + pub per_fold_sharpes: Vec, + /// Sharpe ratio over all concatenated test returns. + pub aggregate_sharpe: f64, + /// Number of walk-forward folds evaluated. + pub num_folds: usize, + /// Deflated Sharpe Ratio result. + pub dsr: DsrResult, + /// Probability of Backtest Overfitting result. + pub pbo: PboResult, + /// Monte Carlo permutation test result. + pub permutation: PermutationResult, + /// Per-regime performance breakdown. + pub per_regime_metrics: HashMap, + /// Overall verdict. + pub verdict: ValidationVerdict, +} + +// --------------------------------------------------------------------------- +// ValidationHarnessConfig +// --------------------------------------------------------------------------- + +/// Configuration for the validation harness. +#[derive(Debug, Clone)] +pub struct ValidationHarnessConfig { + /// Walk-forward cross-validation configuration. + pub wf_config: WalkForwardConfig, + /// Number of Monte Carlo permutations for the permutation test. + pub num_permutations: usize, + /// Number of independent strategy trials (for DSR multiple-testing correction). + pub num_trials: usize, + /// RNG seed for permutation test reproducibility. + pub seed: u64, +} + +impl Default for ValidationHarnessConfig { + fn default() -> Self { + Self { + wf_config: WalkForwardConfig::default(), + num_permutations: 10_000, + num_trials: 1, + seed: 42, + } + } +} + +// --------------------------------------------------------------------------- +// ValidationHarness +// --------------------------------------------------------------------------- + +/// Orchestrates the full validation pipeline for a trading strategy. +/// +/// Given a [`ValidatableStrategy`] and [`TimeSeriesData`], the harness: +/// 1. Splits data via walk-forward cross-validation +/// 2. Trains/evaluates the strategy on each fold +/// 3. Computes DSR, PBO, and permutation tests +/// 4. Breaks down performance by market regime +/// 5. Issues a [`ValidationVerdict`] +#[derive(Debug)] +pub struct ValidationHarness { + config: ValidationHarnessConfig, +} + +impl ValidationHarness { + /// Create a new harness with the given configuration. + pub fn new(config: ValidationHarnessConfig) -> Self { + Self { config } + } + + /// Run the full validation pipeline. + /// + /// # Errors + /// + /// Returns [`MLError::InsufficientData`] when the data cannot produce at + /// least one walk-forward fold, or propagates errors from + /// `strategy.train()` / `strategy.evaluate()` / `data.slice()`. + pub fn validate( + &self, + strategy: &mut S, + data: &TimeSeriesData, + ) -> Result { + // 1. Walk-forward split + let folds = walk_forward_split(data.len(), &self.config.wf_config); + if folds.is_empty() { + return Err(MLError::InsufficientData(format!( + "No walk-forward folds could be generated from {} bars with config \ + (train={}, test={}, embargo={}, step={}, min_train={})", + data.len(), + self.config.wf_config.train_bars, + self.config.wf_config.test_bars, + self.config.wf_config.embargo_bars, + self.config.wf_config.step_bars, + self.config.wf_config.min_train_samples, + ))); + } + + // 2. Train/evaluate each fold, collecting per-fold returns + let mut per_fold_sharpes: Vec = Vec::with_capacity(folds.len()); + let mut all_returns: Vec = Vec::new(); + let mut all_features: Vec> = Vec::new(); + + for fold in &folds { + strategy.reset()?; + + let train_slice = data.slice(fold.train_range.start, fold.train_range.end)?; + strategy.train(&train_slice)?; + + let test_slice = data.slice(fold.test_range.start, fold.test_range.end)?; + let fold_returns = strategy.evaluate(&test_slice)?; + + // 3. Per-fold Sharpe + let fold_sharpe = sharpe_ratio(&fold_returns); + per_fold_sharpes.push(fold_sharpe); + + // Collect features matching the fold returns length (safe iterator, no indexing) + let fold_features: Vec> = test_slice + .features + .iter() + .take(fold_returns.len()) + .cloned() + .collect(); + all_features.extend(fold_features); + all_returns.extend(fold_returns); + } + + // 4. Aggregate statistics + let aggregate_sharpe = sharpe_ratio(&all_returns); + let sharpe_var = if per_fold_sharpes.len() >= 2 { + let n = per_fold_sharpes.len() as f64; + let mean = per_fold_sharpes.iter().sum::() / n; + per_fold_sharpes + .iter() + .map(|s| (s - mean).powi(2)) + .sum::() + / (n - 1.0) + } else { + 1.0 // fallback variance when only 1 fold + }; + let skew = skewness(&all_returns); + let kurt = excess_kurtosis(&all_returns); + let n_obs = all_returns.len(); + + // 5. Deflated Sharpe Ratio + let dsr = deflated_sharpe_ratio( + aggregate_sharpe, + self.config.num_trials, + sharpe_var, + skew, + kurt, + n_obs, + ); + + // 6. PBO (requires even number of folds; drop last if odd) + let pbo_sharpes: Vec = if per_fold_sharpes.len() % 2 != 0 { + per_fold_sharpes + .iter() + .take(per_fold_sharpes.len().saturating_sub(1)) + .copied() + .collect() + } else { + per_fold_sharpes.clone() + }; + let pbo = probability_of_backtest_overfitting(&pbo_sharpes); + + // 7. Permutation test + let permutation = permutation_test( + &all_returns, + self.config.num_permutations, + self.config.seed, + ); + + // 8. Per-regime breakdown + let per_regime_metrics = per_regime_breakdown(&all_returns, &all_features); + + // 9. Verdict + let verdict = determine_verdict(&dsr, &pbo, &permutation); + + Ok(ValidationReport { + strategy_name: strategy.name().to_string(), + per_fold_sharpes, + aggregate_sharpe, + num_folds: folds.len(), + dsr, + pbo, + permutation, + per_regime_metrics, + verdict, + }) + } +} + +// --------------------------------------------------------------------------- +// Verdict logic +// --------------------------------------------------------------------------- + +/// Determine the overall validation verdict from the three statistical tests. +/// +/// - **Pass**: `dsr.pvalue < 0.05` AND `pbo.pbo < 0.25` AND `permutation.pvalue < 0.05` +/// - **Marginal**: `dsr.pvalue < 0.10` OR `pbo.pbo < 0.50` OR `permutation.pvalue < 0.10` +/// - **Fail**: otherwise +fn determine_verdict( + dsr: &DsrResult, + pbo: &PboResult, + permutation: &PermutationResult, +) -> ValidationVerdict { + let dsr_strict = dsr.pvalue < 0.05; + let pbo_strict = pbo.pbo < 0.25; + let perm_strict = permutation.pvalue < 0.05; + + if dsr_strict && pbo_strict && perm_strict { + return ValidationVerdict::Pass; + } + + let dsr_marginal = dsr.pvalue < 0.10; + let pbo_marginal = pbo.pbo < 0.50; + let perm_marginal = permutation.pvalue < 0.10; + + if dsr_marginal || pbo_marginal || perm_marginal { + return ValidationVerdict::Marginal; + } + + ValidationVerdict::Fail +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use chrono::{TimeZone, Utc}; + + // ── Mock strategy ─────────────────────────────────────────────────── + + struct MockStrategy { + name: String, + returns_per_bar: f64, + } + + impl MockStrategy { + fn new(name: &str, returns_per_bar: f64) -> Self { + Self { + name: name.to_string(), + returns_per_bar, + } + } + } + + impl ValidatableStrategy for MockStrategy { + fn train(&mut self, _data: &TimeSeriesData) -> Result<(), MLError> { + Ok(()) + } + + fn evaluate(&self, data: &TimeSeriesData) -> Result, MLError> { + Ok(vec![self.returns_per_bar; data.len()]) + } + + fn name(&self) -> &str { + &self.name + } + + fn reset(&mut self) -> Result<(), MLError> { + Ok(()) + } + } + + // ── Helpers ───────────────────────────────────────────────────────── + + fn make_timestamps(n: usize) -> Vec> { + (0..n) + .map(|i| { + Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0) + .single() + .unwrap_or_else(Utc::now) + + chrono::Duration::days(i as i64) + }) + .collect() + } + + fn make_features(n: usize, dim: usize) -> Vec> { + (0..n).map(|_| vec![0.0_f32; dim]).collect() + } + + /// Build a `TimeSeriesData` with a gentle upward random walk. + fn make_data(n: usize) -> TimeSeriesData { + let mut prices = Vec::with_capacity(n); + let mut price = 100.0_f64; + for i in 0..n { + prices.push(price); + // Small deterministic drift to avoid non-positive prices + price += 0.01 + 0.005 * ((i as f64).sin()); + } + TimeSeriesData::new(make_timestamps(n), make_features(n, 50), prices) + .unwrap_or_else(|e| panic!("Failed to create test data: {}", e)) + } + + // ── Tests ─────────────────────────────────────────────────────────── + + #[test] + fn test_harness_runs_end_to_end() { + let data = make_data(500); + let config = ValidationHarnessConfig { + wf_config: WalkForwardConfig { + train_bars: 100, + test_bars: 50, + embargo_bars: 5, + step_bars: 50, + min_train_samples: 50, + }, + num_permutations: 100, + num_trials: 1, + seed: 42, + }; + + let harness = ValidationHarness::new(config); + let mut strategy = MockStrategy::new("test-mock", 0.001); + + let report = harness.validate(&mut strategy, &data); + assert!(report.is_ok(), "Expected Ok, got {:?}", report.err()); + let report = report.unwrap_or_else(|e| panic!("validate failed: {}", e)); + + // Should have at least one fold + assert!( + report.num_folds > 0, + "Expected at least 1 fold, got {}", + report.num_folds + ); + assert_eq!(report.per_fold_sharpes.len(), report.num_folds); + + // Aggregate Sharpe should be finite + assert!( + report.aggregate_sharpe.is_finite(), + "Aggregate Sharpe should be finite, got {}", + report.aggregate_sharpe + ); + + // p-values in [0, 1] + assert!( + (0.0..=1.0).contains(&report.dsr.pvalue), + "DSR p-value out of range: {}", + report.dsr.pvalue + ); + assert!( + (0.0..=1.0).contains(&report.pbo.pbo), + "PBO out of range: {}", + report.pbo.pbo + ); + assert!( + (0.0..=1.0).contains(&report.permutation.pvalue), + "Permutation p-value out of range: {}", + report.permutation.pvalue + ); + + // Strategy name preserved + assert_eq!(report.strategy_name, "test-mock"); + } + + #[test] + fn test_harness_insufficient_data_errors() { + let data = make_data(20); + let config = ValidationHarnessConfig { + wf_config: WalkForwardConfig { + train_bars: 100, + test_bars: 50, + embargo_bars: 5, + step_bars: 50, + min_train_samples: 50, + }, + num_permutations: 100, + num_trials: 1, + seed: 42, + }; + + let harness = ValidationHarness::new(config); + let mut strategy = MockStrategy::new("too-short", 0.001); + + let result = harness.validate(&mut strategy, &data); + assert!( + result.is_err(), + "Expected error for 20-bar data, got Ok" + ); + let err_msg = format!("{}", result.unwrap_err()); + assert!( + err_msg.contains("Insufficient data") || err_msg.contains("No walk-forward folds"), + "Expected insufficient data error, got: {}", + err_msg + ); + } + + #[test] + fn test_verdict_all_pass() { + let dsr = DsrResult { + observed_sharpe: 2.5, + expected_max_sharpe: 1.0, + sharpe_std_error: 0.3, + deflated_sharpe: 5.0, + pvalue: 0.001, // < 0.05 + }; + let pbo = PboResult { + pbo: 0.10, // < 0.25 + num_combinations: 70, + logit_distribution: vec![-1.0, -0.5, -0.3], + }; + let permutation = PermutationResult { + observed_sharpe: 2.5, + pvalue: 0.01, // < 0.05 + num_permutations: 10_000, + null_mean: 0.0, + null_std: 0.5, + }; + + let verdict = determine_verdict(&dsr, &pbo, &permutation); + assert_eq!(verdict, ValidationVerdict::Pass); + assert_eq!(format!("{}", verdict), "Pass"); + } + + #[test] + fn test_verdict_fail() { + let dsr = DsrResult { + observed_sharpe: 0.1, + expected_max_sharpe: 2.0, + sharpe_std_error: 0.5, + deflated_sharpe: -3.8, + pvalue: 0.99, // > 0.10 + }; + let pbo = PboResult { + pbo: 0.90, // > 0.50 + num_combinations: 70, + logit_distribution: vec![2.0, 1.5, 3.0], + }; + let permutation = PermutationResult { + observed_sharpe: 0.1, + pvalue: 0.85, // > 0.10 + num_permutations: 10_000, + null_mean: 0.0, + null_std: 0.5, + }; + + let verdict = determine_verdict(&dsr, &pbo, &permutation); + assert_eq!(verdict, ValidationVerdict::Fail); + assert_eq!(format!("{}", verdict), "Fail"); + } +} diff --git a/ml/src/validation/mod.rs b/ml/src/validation/mod.rs index 62b913264..9711b2a16 100644 --- a/ml/src/validation/mod.rs +++ b/ml/src/validation/mod.rs @@ -5,6 +5,7 @@ //! - Statistical validation (walk-forward, DSR, PBO, permutation tests) pub mod financial; +pub mod harness; pub mod regime_analysis; pub mod statistical; pub mod types; @@ -25,6 +26,9 @@ pub use walk_forward::{walk_forward_split, Fold, WalkForwardConfig}; // Re-export per-regime performance analysis pub use regime_analysis::{per_regime_breakdown, RegimeMetrics}; +// Re-export validation harness +pub use harness::{ValidationHarness, ValidationHarnessConfig, ValidationReport, ValidationVerdict}; + // Re-export statistical validation tests pub use statistical::{ deflated_sharpe_ratio, excess_kurtosis, normal_cdf, normal_ppf, permutation_test, diff --git a/ml/src/validation/statistical.rs b/ml/src/validation/statistical.rs index 2bd01b5ad..1fd7f6412 100644 --- a/ml/src/validation/statistical.rs +++ b/ml/src/validation/statistical.rs @@ -9,6 +9,7 @@ use rand::seq::SliceRandom; use rand::SeedableRng; use rand_chacha::ChaCha8Rng; +use serde::{Deserialize, Serialize}; // ─── Section 1: Math helpers ──────────────────────────────────────────────── @@ -199,7 +200,7 @@ pub fn excess_kurtosis(returns: &[f64]) -> f64 { // ─── Section 2: Deflated Sharpe Ratio ─────────────────────────────────────── /// Result of the Deflated Sharpe Ratio test. -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct DsrResult { /// The observed Sharpe ratio being tested. pub observed_sharpe: f64, @@ -287,7 +288,7 @@ pub fn deflated_sharpe_ratio( // ─── Section 3: Monte Carlo Permutation Test ──────────────────────────────── /// Result of a Monte Carlo permutation test. -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct PermutationResult { /// The observed Sharpe ratio of the original return series. pub observed_sharpe: f64, @@ -372,7 +373,7 @@ pub fn permutation_test( // ─── Section 4: Probability of Backtest Overfitting ───────────────────────── /// Result of the Probability of Backtest Overfitting (PBO) test. -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Serialize, Deserialize)] pub struct PboResult { /// Estimated probability of backtest overfitting [0, 1]. pub pbo: f64,