refactor(validation): convert to directory module for validation stack

Move validation.rs to validation/financial.rs, create validation/mod.rs
with re-exports for backward compatibility, and remove orphan
numerical_tests.rs that referenced nonexistent types.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-20 16:12:19 +01:00
parent 7ce7e33115
commit d956add21c
3 changed files with 13 additions and 391 deletions

13
ml/src/validation/mod.rs Normal file
View File

@@ -0,0 +1,13 @@
//! Statistical and Financial Validation Module
//!
//! Provides two categories of validation:
//! - Financial validation (price/volume/quantity sanity checks)
//! - Statistical validation (walk-forward, DSR, PBO, permutation tests)
pub mod financial;
// Re-export financial validation for backward compatibility
pub use financial::{
validate_model_basic, validate_model_comprehensive, validate_type_conversions,
FinancialValidationResult, ValidationResult,
};

View File

@@ -1,391 +0,0 @@
//! Numerical equivalence tests for ML model accuracy validation
//!
//! Critical for HFT production deployment to ensure our Rust implementations
//! produce numerically equivalent results to reference Python implementations.
use std::collections::HashMap;
use anyhow::{Result, Context};
use serde::{Deserialize, Serialize};
use tokio::time::{Duration, timeout};
use crate::{
MLModel, ModelType, Features, ModelPrediction,
mamba::Mamba2SSM,
dqn::RainbowDQN,
tlob::TLOBTransformer,
tft::TemporalFusionTransformer,
};
/// Tolerance levels for numerical comparison in HFT context
#[derive(Debug, Clone)]
pub struct NumericalTolerance {
/// Absolute tolerance for exact comparisons
pub absolute: f64,
/// Relative tolerance for proportional comparisons
pub relative: f64,
/// Maximum acceptable difference for HFT decisions
pub decision_threshold: f64,
}
impl Default for NumericalTolerance {
fn default() -> Self {
Self {
absolute: 1e-10, // 10 decimal places precision
relative: 1e-8, // 8 significant figures
decision_threshold: 1e-6, // 1 millionth for trading decisions
}
}
}
/// Test case for numerical equivalence validation
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NumericalTestCase {
pub name: String,
pub model_type: ModelType,
pub input_features: Features,
pub expected_output: ModelPrediction,
pub tolerance: NumericalTolerance,
pub description: String,
}
/// Results of numerical equivalence testing
#[derive(Debug, Clone)]
pub struct NumericalTestResult {
pub test_name: String,
pub passed: bool,
pub rust_output: ModelPrediction,
pub python_reference: ModelPrediction,
pub absolute_error: f64,
pub relative_error: f64,
pub latency_ns: u64,
pub error_message: Option<String>,
}
/// Comprehensive numerical validation framework
pub struct NumericalValidator {
test_cases: Vec<NumericalTestCase>,
models: HashMap<ModelType, Box<dyn MLModel>>,
python_bridge: Option<PythonModelBridge>,
}
impl NumericalValidator {
/// Create new validator with comprehensive test suite
pub fn new() -> Self {
Self {
test_cases: Self::create_standard_test_cases(),
models: HashMap::new(),
python_bridge: None,
}
}
/// Register Rust model for testing
pub fn register_model(&mut self, model_type: ModelType, model: Box<dyn MLModel>) {
self.models.insert(model_type, model);
}
/// Initialize Python bridge for reference comparisons
pub fn with_python_bridge(mut self, bridge: PythonModelBridge) -> Self {
self.python_bridge = Some(bridge);
self
}
/// Run complete numerical validation suite
pub async fn validate_all_models(&mut self) -> Result<Vec<NumericalTestResult>> {
let mut results = Vec::new();
for test_case in &self.test_cases {
match self.validate_single_test(test_case).await {
Ok(result) => results.push(result),
Err(e) => {
results.push(NumericalTestResult {
test_name: test_case.name.clone(),
passed: false,
rust_output: ModelPrediction::default(),
python_reference: ModelPrediction::default(),
absolute_error: f64::INFINITY,
relative_error: f64::INFINITY,
latency_ns: 0,
error_message: Some(e.to_string()),
});
}
}
}
Ok(results)
}
/// Validate single test case with timing
async fn validate_single_test(&mut self, test_case: &NumericalTestCase) -> Result<NumericalTestResult> {
let model = self.models.get(&test_case.model_type)
.context("Model not registered for testing")?;
// Time the Rust inference
let start = std::time::Instant::now();
let rust_output = timeout(
Duration::from_millis(100), // 100ms timeout for HFT
model.predict(&test_case.input_features)
)
.await
.context("Inference timeout")?
.context("Inference failed")?;
let latency_ns = start.elapsed().as_nanos() as u64;
// Get Python reference if available
let python_reference = if let Some(ref bridge) = self.python_bridge {
bridge.predict(&test_case.model_type, &test_case.input_features).await?
} else {
test_case.expected_output.clone()
};
// Calculate numerical differences
let absolute_error = self.calculate_absolute_error(&rust_output, &python_reference)?;
let relative_error = self.calculate_relative_error(&rust_output, &python_reference)?;
// Determine if test passed
let passed = absolute_error <= test_case.tolerance.absolute &&
relative_error <= test_case.tolerance.relative &&
absolute_error <= test_case.tolerance.decision_threshold;
Ok(NumericalTestResult {
test_name: test_case.name.clone(),
passed,
rust_output,
python_reference,
absolute_error,
relative_error,
latency_ns,
error_message: None,
})
}
/// Calculate absolute error between predictions
fn calculate_absolute_error(&self, rust: &ModelPrediction, python: &ModelPrediction) -> Result<f64> {
match (rust, python) {
(ModelPrediction::Price(r), ModelPrediction::Price(p)) => {
Ok((r.value - p.value).abs())
},
(ModelPrediction::Direction(r), ModelPrediction::Direction(p)) => {
Ok(if r.direction == p.direction { 0.0 } else { 1.0 })
},
(ModelPrediction::Portfolio(r), ModelPrediction::Portfolio(p)) => {
let mut total_error = 0.0;
for (symbol, rust_pos) in &r.positions {
if let Some(python_pos) = p.positions.get(symbol) {
total_error += (rust_pos.quantity - python_pos.quantity).abs();
}
}
Ok(total_error)
},
_ => Ok(f64::INFINITY), // Type mismatch
}
}
/// Calculate relative error between predictions
fn calculate_relative_error(&self, rust: &ModelPrediction, python: &ModelPrediction) -> Result<f64> {
let absolute_error = self.calculate_absolute_error(rust, python)?;
match python {
ModelPrediction::Price(p) => {
if p.value.abs() < 1e-15 {
Ok(absolute_error)
} else {
Ok(absolute_error / p.value.abs())
}
},
ModelPrediction::Direction(_) => Ok(absolute_error),
ModelPrediction::Portfolio(p) => {
let total_value: f64 = p.positions.values()
.map(|pos| pos.quantity.abs())
.sum();
if total_value < 1e-15 {
Ok(absolute_error)
} else {
Ok(absolute_error / total_value)
}
},
}
}
/// Create comprehensive test cases for all models
fn create_standard_test_cases() -> Vec<NumericalTestCase> {
vec![
// MAMBA-2 SSM Tests
NumericalTestCase {
name: "mamba_sequence_modeling".to_string(),
model_type: ModelType::MAMBA,
input_features: Features::create_time_series_features(vec![1.0, 2.0, 3.0, 4.0, 5.0]),
expected_output: ModelPrediction::new(
"mamba_test".to_string(),
6.0,
0.95,
),
tolerance: NumericalTolerance::default(),
description: "MAMBA-2 sequence modeling accuracy".to_string(),
},
// Rainbow DQN Tests
NumericalTestCase {
name: "dqn_action_values".to_string(),
model_type: ModelType::DQN,
input_features: Features::create_market_state_features(
vec![100.0, 101.0, 99.5, 100.5], // OHLC
vec![1000.0, 1500.0], // Volume, Spread
),
expected_output: ModelPrediction::Direction(crate::types::DirectionPrediction {
direction: crate::types::TradeDirection::Buy,
confidence: 0.87,
expected_return: 0.02,
}),
tolerance: NumericalTolerance::default(),
description: "Rainbow DQN action value estimation".to_string(),
},
// TLOB Transformer Tests
NumericalTestCase {
name: "tlob_order_book_prediction".to_string(),
model_type: ModelType::TLOB,
input_features: Features::create_order_book_features(
vec![(100.0, 1000.0), (100.1, 2000.0)], // Bids
vec![(100.2, 1500.0), (100.3, 1000.0)], // Asks
),
expected_output: ModelPrediction::new(
"tlob_test".to_string(),
100.15,
0.92,
),
tolerance: NumericalTolerance {
absolute: 1e-8,
relative: 1e-6,
decision_threshold: 1e-5, // Tighter for price predictions
},
description: "TLOB Transformer order book prediction".to_string(),
},
// TFT Tests
NumericalTestCase {
name: "tft_temporal_fusion".to_string(),
model_type: ModelType::TFT,
input_features: Features::create_multivariate_features(
vec![
vec![1.0, 2.0, 3.0],
vec![4.0, 5.0, 6.0],
vec![7.0, 8.0, 9.0],
]
),
expected_output: ModelPrediction::new(
"tft_test".to_string(),
10.5,
0.89,
),
tolerance: NumericalTolerance::default(),
description: "TFT temporal fusion accuracy".to_string(),
},
]
}
}
/// Bridge to Python models for reference testing
pub struct PythonModelBridge {
// Would connect to Python process running reference implementations
// For now, this is a placeholder for the interface
}
impl PythonModelBridge {
pub fn new() -> Result<Self> {
// Known limitation: Python bridge not implemented
// Production should use PyO3 for in-process Python interop or
// spawn subprocess running reference implementation for validation
Ok(Self {})
}
pub async fn predict(&self, model_type: &ModelType, features: &Features) -> Result<ModelPrediction> {
// Known limitation: Python reference implementation not connected
// Production should:
// 1. Serialize features to JSON/numpy format
// 2. Call Python reference model via PyO3 or subprocess
// 3. Compare predictions for numerical accuracy validation
// 4. Return reference predictions for cross-validation testing
let _ = (model_type, features); // Suppress unused warnings
Ok(ModelPrediction::default())
}
}
/// Generate comprehensive test report
pub fn generate_test_report(results: &[NumericalTestResult]) -> String {
let mut report = String::new();
report.push_str("# ML Numerical Equivalence Test Report\n\n");
let passed = results.iter().filter(|r| r.passed).count();
let total = results.len();
let pass_rate = (passed as f64 / total as f64) * 100.0;
report.push_str(&format!("## Summary\n"));
report.push_str(&format!("- **Tests Passed**: {}/{} ({:.1}%)\n", passed, total, pass_rate));
report.push_str(&format!("- **Production Ready**: {}\n\n", if pass_rate >= 95.0 { "✅ YES" } else { "❌ NO" }));
// Latency analysis
let avg_latency: f64 = results.iter()
.map(|r| r.latency_ns as f64)
.sum::<f64>() / results.len() as f64;
report.push_str(&format!("## Performance\n"));
report.push_str(&format!("- **Average Latency**: {:.1}μs\n", avg_latency / 1000.0));
report.push_str(&format!("- **HFT Ready**: {}\n\n", if avg_latency < 50_000.0 { "✅ Sub-50μs" } else { "⚠️ Above 50μs" }));
// Detailed results
report.push_str("## Detailed Results\n\n");
for result in results {
let status = if result.passed { "✅ PASS" } else { "❌ FAIL" };
report.push_str(&format!("### {} - {}\n", result.test_name, status));
report.push_str(&format!("- **Absolute Error**: {:.2e}\n", result.absolute_error));
report.push_str(&format!("- **Relative Error**: {:.2e}\n", result.relative_error));
report.push_str(&format!("- **Latency**: {:.1}μs\n", result.latency_ns as f64 / 1000.0));
if let Some(ref error) = result.error_message {
report.push_str(&format!("- **Error**: {}\n", error));
}
report.push_str("\n");
}
report
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_numerical_validator_creation() {
let validator = NumericalValidator::new();
assert!(!validator.test_cases.is_empty());
}
#[test]
fn test_tolerance_defaults() {
let tolerance = NumericalTolerance::default();
assert!(tolerance.absolute > 0.0);
assert!(tolerance.relative > 0.0);
assert!(tolerance.decision_threshold > 0.0);
}
#[test]
fn test_report_generation() {
let results = vec![
NumericalTestResult {
test_name: "test1".to_string(),
passed: true,
rust_output: ModelPrediction::default(),
python_reference: ModelPrediction::default(),
absolute_error: 1e-12,
relative_error: 1e-10,
latency_ns: 25_000,
error_message: None,
}
];
let report = generate_test_report(&results);
assert!(report.contains("✅ YES"));
assert!(report.contains("✅ Sub-50μs"));
}
}