feat(validation): add validation harness orchestrator with verdict system

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-20 16:44:03 +01:00
parent 5c37db5d00
commit 2bd1208db5
3 changed files with 512 additions and 3 deletions

View File

@@ -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<f64>,
/// 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<RegimeType, RegimeMetrics>,
/// 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<S: ValidatableStrategy>(
&self,
strategy: &mut S,
data: &TimeSeriesData,
) -> Result<ValidationReport, MLError> {
// 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<f64> = Vec::with_capacity(folds.len());
let mut all_returns: Vec<f64> = Vec::new();
let mut all_features: Vec<Vec<f32>> = 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<Vec<f32>> = 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::<f64>() / n;
per_fold_sharpes
.iter()
.map(|s| (s - mean).powi(2))
.sum::<f64>()
/ (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<f64> = 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<Vec<f64>, 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<chrono::DateTime<chrono::Utc>> {
(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<Vec<f32>> {
(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");
}
}

View File

@@ -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,

View File

@@ -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,