Files
foxhunt/crates/ml/tests/validation/early_stopping_validation_tests.rs
jgrusewski ca4c38d921 fix(tests): CI GPU test stability, walltime reduction, BF16 tolerance
- Reduce CI GPU test datasets 16x for walltime reduction
- Reduce early-stop epochs 50→10, add --test-threads=1
- Serialize all GPU lib tests to prevent cuBLAS init race
- Align state_dim to 16 for BF16 tensor core HMMA dispatch
- BF16 precision tolerance in ml-dqn tests
- Enable branching DQN + tracing subscriber in smoke tests
- Prevent min_replay_size > buffer_size deadlock in early-stop tests
- Prevent AutoReplaySizer from breaking gradient collapse warmup
- Replace racy tokio::spawn checkpoint counter with AtomicUsize
- Set warmup_steps=0 and max_training_steps_per_epoch=300 in early-stop tests
- RealDataLoader respects TEST_DATA_DIR for CI PVC layout
- Add collapse_warmup_capacity to gpu_smoketest DQNConfig
- Drain CUDA context between test binaries
- Detached HEAD checkout prevents local branch corruption
- GPU pipeline tests: fix BF16 dtype and rank-1 squeeze assertions
- OOD input handling tests use use_gpu: true

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-15 12:00:13 +01:00

449 lines
17 KiB
Rust
Raw 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.
//! Validation Tests for Early Stopping with Real Data
//!
//! This module contains validation tests that use real market data to verify
//! early stopping works correctly in production-like scenarios.
//!
//! Test Scenarios:
//! 1. Convergence validation (quality preservation)
//! 2. Premature stopping prevention
//! 3. Late bloomer detection
//! 4. Plateau vs noise distinction
//! 5. Real-world resource savings measurement
use std::path::Path;
use tracing::info;
use tracing::warn;
// ============================================================================
// CONVERGENCE VALIDATION TESTS
// ============================================================================
/// Test that early stopping converges to near-optimal solution
#[test]
#[ignore = "Slow test: requires real data and ~10 minutes runtime"]
fn test_early_stopping_converges_to_optimal() {
// This test would run two identical hyperopt runs:
// 1. With early stopping enabled
// 2. Without early stopping (baseline)
//
// Success criteria:
// - Final loss within 5% (quality preserved)
// - Resource savings 30-50%
// - Best trial in both runs has similar performance
info!("Test: Early stopping convergence validation");
info!("Expected: Run hyperopt WITH and WITHOUT early stopping, compare final best loss (within 5%), measure resource savings (30-50%)");
// Test configuration
let test_config = ValidationTestConfig {
adapter_name: "DQN",
num_trials: 10,
max_epochs: 100,
data_path: "test_data/real/databento/ml_training/",
tolerance_pct: 5.0,
min_savings_pct: 30.0,
max_savings_pct: 70.0,
};
info!(
adapter = test_config.adapter_name,
trials = test_config.num_trials,
max_epochs = test_config.max_epochs,
tolerance_pct = test_config.tolerance_pct,
min_savings_pct = test_config.min_savings_pct,
max_savings_pct = test_config.max_savings_pct,
"Test configuration"
);
// Verify test data exists
if Path::new(test_config.data_path).exists() {
info!("Test data found");
} else {
info!("Test data not found - test would be skipped");
}
}
/// Test that early stopping preserves model quality
#[test]
#[ignore = "Slow test: requires real data"]
fn test_early_stopping_quality_preservation() {
// Verify that models trained with early stopping achieve similar
// validation metrics as models trained to completion
struct QualityMetrics {
train_loss: f64,
val_loss: f64,
accuracy: f64,
precision: f64,
recall: f64,
f1_score: f64,
}
// Expected metrics (baseline without early stopping)
let baseline = QualityMetrics {
train_loss: 0.75,
val_loss: 0.80,
accuracy: 0.65,
precision: 0.63,
recall: 0.67,
f1_score: 0.65,
};
// Simulated metrics with early stopping
let with_early_stop = QualityMetrics {
train_loss: 0.77, // Within 5%
val_loss: 0.82, // Within 5%
accuracy: 0.64, // Within 5%
precision: 0.62, // Within 5%
recall: 0.66, // Within 5%
f1_score: 0.64, // Within 5%
};
// Verify all metrics within tolerance
let tolerance = 0.05; // 5%
let train_loss_delta = (with_early_stop.train_loss - baseline.train_loss).abs() / baseline.train_loss;
let val_loss_delta = (with_early_stop.val_loss - baseline.val_loss).abs() / baseline.val_loss;
let accuracy_delta = (with_early_stop.accuracy - baseline.accuracy).abs() / baseline.accuracy;
info!(
train_loss_delta_pct = train_loss_delta * 100.0,
val_loss_delta_pct = val_loss_delta * 100.0,
accuracy_delta_pct = accuracy_delta * 100.0,
"Quality Metrics Comparison"
);
assert!(train_loss_delta <= tolerance, "Train loss delta exceeds tolerance");
assert!(val_loss_delta <= tolerance, "Val loss delta exceeds tolerance");
assert!(accuracy_delta <= tolerance, "Accuracy delta exceeds tolerance");
}
// ============================================================================
// PREMATURE STOPPING PREVENTION TESTS
// ============================================================================
/// Test that minimum epochs prevent premature stopping
#[test]
fn test_minimum_epochs_prevents_premature_stopping() {
// Simulate training with intentionally noisy early data
let noisy_early_losses = vec![
10.0, 8.0, 9.5, 7.5, 8.2, 6.8, 7.3, 6.2, 6.7, 5.9, // Epochs 0-9: noisy
5.5, 5.1, 4.8, 4.5, 4.2, 4.0, 3.8, 3.6, 3.5, 3.4, // Epochs 10-19: stabilizing
3.3, 3.2, 3.1, 3.0, 2.9, 2.8, 2.7, 2.6, 2.5, 2.4, // Epochs 20-29: converging
];
let min_epochs = 20;
let window = 5;
let min_improvement = 2.0;
// Check that early stopping doesn't trigger before min_epochs
for epoch in window*2..min_epochs {
let recent_avg = noisy_early_losses[epoch-window..epoch].iter().sum::<f64>() / window as f64;
let older_avg = noisy_early_losses[epoch-2*window..epoch-window].iter().sum::<f64>() / window as f64;
let improvement = ((older_avg - recent_avg) / older_avg * 100.0).abs();
if improvement < min_improvement {
// Would normally trigger, but min_epochs should prevent it
info!(epoch, improvement_pct = improvement, min_improvement, min_epochs, "Improvement below threshold but below min_epochs");
assert!(epoch < min_epochs, "Early stopping triggered before min_epochs!");
}
}
info!("Minimum epochs successfully prevented premature stopping");
}
/// Test that warmup period allows model to stabilize
#[test]
fn test_warmup_period_handling() {
// Simulate training with typical warmup behavior:
// - High initial loss
// - Rapid initial decrease
// - Oscillations during warmup
// - Stabilization after warmup
let losses_with_warmup = vec![
// Warmup: 0-19 (high variance, rapid changes)
50.0, 40.0, 35.0, 30.0, 28.0, 26.0, 25.0, 24.0, 23.5, 23.0,
22.0, 21.5, 21.0, 20.5, 20.0, 19.0, 18.0, 17.0, 16.0, 15.0,
// Post-warmup: 20-39 (stable convergence)
14.5, 14.0, 13.5, 13.0, 12.5, 12.0, 11.5, 11.0, 10.5, 10.0,
9.8, 9.6, 9.4, 9.2, 9.0, 8.9, 8.8, 8.7, 8.6, 8.5,
];
let warmup_epochs = 20;
// Calculate variance in warmup vs post-warmup
let warmup_variance = calculate_variance(&losses_with_warmup[0..warmup_epochs]);
let stable_variance = calculate_variance(&losses_with_warmup[warmup_epochs..]);
info!(
warmup_variance,
stable_variance,
ratio = warmup_variance / stable_variance,
"Warmup analysis"
);
// Warmup should have significantly higher variance
assert!(warmup_variance > stable_variance * 2.0,
"Warmup variance should be >2x stable variance");
info!("Warmup period correctly identified");
}
fn calculate_variance(values: &[f64]) -> f64 {
let mean = values.iter().sum::<f64>() / values.len() as f64;
values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / values.len() as f64
}
// ============================================================================
// LATE BLOOMER DETECTION TESTS
// ============================================================================
/// Test that late bloomers (slow starters) are not pruned prematurely
#[test]
fn test_late_bloomer_not_pruned() {
// Simulate a "late bloomer" trial:
// - Poor performance early
// - Breakthrough at epoch 40
// - Strong performance after breakthrough
let late_bloomer_losses = vec![
// Poor start: 0-39
5.0, 4.9, 4.8, 4.7, 4.6, 4.5, 4.4, 4.3, 4.2, 4.1,
4.0, 3.95, 3.9, 3.85, 3.8, 3.75, 3.7, 3.65, 3.6, 3.55,
3.5, 3.48, 3.46, 3.44, 3.42, 3.4, 3.38, 3.36, 3.34, 3.32,
3.3, 3.29, 3.28, 3.27, 3.26, 3.25, 3.24, 3.23, 3.22, 3.21,
// Breakthrough: 40-49
3.0, 2.7, 2.4, 2.1, 1.8, 1.5, 1.2, 0.9, 0.7, 0.6,
// Strong finish: 50-59
0.55, 0.52, 0.50, 0.48, 0.47, 0.46, 0.45, 0.44, 0.43, 0.42,
];
// Typical early bloomer for comparison
let early_bloomer_losses = vec![
// Fast start: 0-19
5.0, 3.5, 2.5, 1.8, 1.3, 1.0, 0.85, 0.75, 0.68, 0.63,
0.60, 0.58, 0.57, 0.56, 0.55, 0.54, 0.53, 0.52, 0.51, 0.50,
];
// At epoch 20, early bloomer looks much better
let late_bloomer_at_20 = late_bloomer_losses[20];
let early_bloomer_at_20 = early_bloomer_losses[19];
info!(
late_bloomer_at_20,
early_bloomer_at_20,
difference = late_bloomer_at_20 - early_bloomer_at_20,
"Epoch 20 comparison"
);
// At epoch 60, late bloomer catches up
let late_bloomer_final = *late_bloomer_losses.last().unwrap();
let early_bloomer_final = *early_bloomer_losses.last().unwrap();
info!(
late_bloomer_final,
early_bloomer_final,
late_bloomer_wins_by = early_bloomer_final - late_bloomer_final,
"Final comparison"
);
assert!(late_bloomer_final < early_bloomer_final,
"Late bloomer should achieve better final loss");
info!("Late bloomer achieved better final performance");
}
/// Test patience mechanism allows for recovery
#[test]
fn test_patience_allows_recovery() {
// Simulate training with temporary plateau followed by recovery
let losses_with_recovery = vec![
// Good progress: 0-19
5.0, 4.5, 4.0, 3.5, 3.0, 2.7, 2.5, 2.3, 2.1, 2.0,
1.9, 1.8, 1.7, 1.6, 1.5, 1.45, 1.4, 1.35, 1.3, 1.25,
// Plateau: 20-34 (15 epochs stuck)
1.24, 1.24, 1.23, 1.24, 1.23, 1.24, 1.23, 1.24, 1.23, 1.24,
1.23, 1.24, 1.23, 1.24, 1.23,
// Recovery: 35-44 (breakthrough!)
1.15, 1.0, 0.85, 0.7, 0.6, 0.55, 0.52, 0.50, 0.49, 0.48,
];
let patience = 20; // Must be > 15 to allow recovery
let mut no_improvement_count = 0;
let mut best_loss = f64::INFINITY;
for (epoch, &loss) in losses_with_recovery.iter().enumerate() {
if loss < best_loss {
best_loss = loss;
no_improvement_count = 0;
info!(epoch, best_loss, "New best loss");
} else {
no_improvement_count += 1;
if no_improvement_count % 5 == 0 {
info!(epoch, no_improvement_count, "Epochs without improvement");
}
}
if no_improvement_count >= patience {
panic!("Patience exhausted at epoch {} - missed recovery at epoch 35!", epoch);
}
}
info!(max_no_improvement = no_improvement_count, "Patience mechanism allowed recovery");
assert!(best_loss < 0.5, "Should achieve final loss < 0.5");
}
// ============================================================================
// PLATEAU VS NOISE DISTINCTION TESTS
// ============================================================================
/// Test that plateau detection is robust to noise
#[test]
fn test_plateau_detection_robust_to_noise() {
// Test Case 1: Real plateau with noise
let noisy_plateau = vec![
1.0, 1.01, 0.99, 1.02, 0.98, 1.01, 0.99, 1.00,
1.01, 0.99, 1.02, 0.98, 1.01, 0.99, 1.00, 1.01,
];
let window = 4;
let min_improvement = 2.0;
let recent_avg = noisy_plateau[12..16].iter().sum::<f64>() / 4.0;
let older_avg = noisy_plateau[8..12].iter().sum::<f64>() / 4.0;
let improvement = ((older_avg - recent_avg) / older_avg * 100.0).abs();
info!(recent_avg, older_avg, improvement_pct = improvement, "Noisy plateau analysis");
// Should detect as plateau despite noise
assert!(improvement < min_improvement, "Should detect plateau with noise");
// Test Case 2: Consistent improvement with noise
let noisy_improvement = vec![
1.0, 0.98, 1.01, 0.95, 0.99, 0.92, 0.96, 0.89,
0.93, 0.86, 0.90, 0.83, 0.87, 0.80, 0.84, 0.77,
];
let recent_avg2 = noisy_improvement[12..16].iter().sum::<f64>() / 4.0;
let older_avg2 = noisy_improvement[8..12].iter().sum::<f64>() / 4.0;
let improvement2 = ((older_avg2 - recent_avg2) / older_avg2 * 100.0).abs();
info!(recent_avg = recent_avg2, older_avg = older_avg2, improvement_pct = improvement2, "Noisy improvement analysis");
// Should NOT detect as plateau (real improvement)
assert!(improvement2 >= min_improvement, "Should detect continued improvement");
}
/// Test signal-to-noise ratio in loss trajectory
#[test]
fn test_signal_to_noise_ratio_analysis() {
// Calculate SNR for different training phases
struct PhaseAnalysis {
phase_name: &'static str,
losses: Vec<f64>,
mean: f64,
variance: f64,
snr_db: f64,
}
let early_phase = vec![5.0, 4.8, 4.9, 4.5, 4.7, 4.3, 4.5, 4.1, 4.3, 3.9];
let late_phase = vec![1.0, 1.01, 0.99, 1.02, 0.98, 1.01, 0.99, 1.00, 1.01, 0.99];
let early_mean = early_phase.iter().sum::<f64>() / early_phase.len() as f64;
let late_mean = late_phase.iter().sum::<f64>() / late_phase.len() as f64;
let early_var = calculate_variance(&early_phase);
let late_var = calculate_variance(&late_phase);
let early_snr = 10.0 * (early_mean / early_var.sqrt()).log10();
let late_snr = 10.0 * (late_mean / late_var.sqrt()).log10();
info!(early_mean, early_var, early_snr_db = early_snr, late_mean, late_var, late_snr_db = late_snr, "Signal-to-Noise Ratio Analysis");
// Late phase should have lower variance (more stable)
assert!(late_var < early_var, "Late phase should be more stable");
}
// ============================================================================
// REAL-WORLD RESOURCE SAVINGS TESTS
// ============================================================================
/// Test resource savings measurement methodology
#[test]
#[ignore = "Slow test: requires real hyperopt run"]
fn test_real_world_resource_savings_measurement() {
// This test would run a real hyperopt and measure:
// 1. Wall-clock time saved
// 2. Epochs saved
// 3. GPU memory-hours saved
// 4. Cost saved (RunPod pricing)
struct ResourceSavings {
trials_total: usize,
trials_stopped_early: usize,
epochs_without_es: usize,
epochs_with_es: usize,
time_without_es_mins: f64,
time_with_es_mins: f64,
cost_without_es_usd: f64,
cost_with_es_usd: f64,
}
// Simulated results (would come from real run)
let savings = ResourceSavings {
trials_total: 20,
trials_stopped_early: 14,
epochs_without_es: 2000, // 20 trials × 100 epochs
epochs_with_es: 1180, // Average 59 epochs/trial
time_without_es_mins: 60.0,
time_with_es_mins: 35.4,
cost_without_es_usd: 0.25, // RTX A4000 @ $0.25/hr
cost_with_es_usd: 0.15,
};
let epoch_savings_pct = (1.0 - savings.epochs_with_es as f64 / savings.epochs_without_es as f64) * 100.0;
let time_savings_pct = (1.0 - savings.time_with_es_mins / savings.time_without_es_mins) * 100.0;
let cost_savings_pct = (1.0 - savings.cost_with_es_usd / savings.cost_without_es_usd) * 100.0;
info!(
trials_total = savings.trials_total,
trials_stopped_early = savings.trials_stopped_early,
stopped_early_pct = savings.trials_stopped_early as f64 / savings.trials_total as f64 * 100.0,
epochs_without_es = savings.epochs_without_es,
epochs_with_es = savings.epochs_with_es,
epochs_saved = savings.epochs_without_es - savings.epochs_with_es,
epoch_savings_pct,
time_without_es_mins = savings.time_without_es_mins,
time_with_es_mins = savings.time_with_es_mins,
time_savings_pct,
cost_without_es_usd = savings.cost_without_es_usd,
cost_with_es_usd = savings.cost_with_es_usd,
cost_savings_pct,
"Resource Savings Analysis"
);
// Verify savings targets
assert!(epoch_savings_pct >= 30.0 && epoch_savings_pct <= 70.0,
"Epoch savings {}% outside 30-70% range", epoch_savings_pct);
assert!(time_savings_pct >= 30.0 && time_savings_pct <= 70.0,
"Time savings {}% outside 30-70% range", time_savings_pct);
assert!(cost_savings_pct >= 30.0 && cost_savings_pct <= 70.0,
"Cost savings {}% outside 30-70% range", cost_savings_pct);
}
// ============================================================================
// HELPER STRUCTURES
// ============================================================================
struct ValidationTestConfig {
adapter_name: &'static str,
num_trials: usize,
max_epochs: usize,
data_path: &'static str,
tolerance_pct: f64,
min_savings_pct: f64,
max_savings_pct: f64,
}