//! MAMBA2-Specific Edge Case Tests for Hyperparameter Optimization //! //! This test suite covers MAMBA2-specific edge cases: //! 1. Async data loading edge cases //! 2. Sequence length and stride edge cases //! 3. Normalization parameter edge cases //! 4. SSM-specific numerical stability //! 5. Batch size clamping with GPU memory //! //! Purpose: Ensure MAMBA2 adapter handles all edge cases robustly use ml::hyperopt::adapters::mamba2::{Mamba2Params, Mamba2Trainer}; use ml::hyperopt::traits::{HyperparameterOptimizable, ParameterSpace}; use tempfile::TempDir; // ============================================================================ // TEST UTILITIES // ============================================================================ fn create_test_parquet(temp_dir: &TempDir, num_rows: usize, suffix: &str) -> String { use arrow::array::{Float64Array, PrimitiveArray, UInt64Array}; use arrow::datatypes::{DataType, Field, Schema, TimestampNanosecondType}; use arrow::record_batch::RecordBatch; use parquet::arrow::arrow_writer::ArrowWriter; use parquet::file::properties::WriterProperties; use std::fs::File; use std::sync::Arc; let schema = Arc::new(Schema::new(vec![ Field::new("ts_event", DataType::UInt64, false), Field::new("rtype", DataType::UInt8, false), Field::new("publisher_id", DataType::UInt16, false), Field::new("open", DataType::Float64, false), Field::new("high", DataType::Float64, false), Field::new("low", DataType::Float64, false), Field::new("close", DataType::Float64, false), Field::new("volume", DataType::UInt64, false), Field::new("symbol", DataType::Utf8, false), Field::new("timestamp", DataType::Timestamp(arrow::datatypes::TimeUnit::Nanosecond, None), false), ])); let file_path = temp_dir.path().join(format!("mamba2_test_{}.parquet", suffix)); let file = File::create(&file_path).unwrap(); let props = WriterProperties::builder().build(); let mut writer = ArrowWriter::try_new(file, schema.clone(), Some(props)).unwrap(); let base_price = 5000.0; let base_timestamp = 1700000000_000_000_000u64; let batch = RecordBatch::try_new( schema, vec![ Arc::new(UInt64Array::from( (0..num_rows).map(|i| base_timestamp + i as u64 * 60_000_000_000).collect::>() )), Arc::new(arrow::array::UInt8Array::from(vec![1u8; num_rows])), Arc::new(arrow::array::UInt16Array::from(vec![1u16; num_rows])), Arc::new(Float64Array::from( (0..num_rows).map(|i| base_price + (i as f64 * 0.1)).collect::>() )), Arc::new(Float64Array::from( (0..num_rows).map(|i| base_price + (i as f64 * 0.1) + 5.0).collect::>() )), Arc::new(Float64Array::from( (0..num_rows).map(|i| base_price + (i as f64 * 0.1) - 5.0).collect::>() )), Arc::new(Float64Array::from( (0..num_rows).map(|i| base_price + (i as f64 * 0.1) + 2.5).collect::>() )), Arc::new(UInt64Array::from(vec![1000u64; num_rows])), Arc::new(arrow::array::StringArray::from(vec!["ES.FUT"; num_rows])), Arc::new(PrimitiveArray::::from( (0..num_rows).map(|i| base_timestamp as i64 + i as i64 * 60_000_000_000).collect::>() )), ], ) .unwrap(); writer.write(&batch).unwrap(); writer.close().unwrap(); file_path.to_string_lossy().to_string() } // ============================================================================ // ASYNC DATA LOADING EDGE CASES // ============================================================================ #[test] fn test_async_loading_with_small_dataset() { // Async loading with dataset smaller than prefetch_count let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 100, "small_async"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer") .with_async_loading(true, 10); // Prefetch 10 batches, but dataset might be smaller let params = Mamba2Params::default(); let result = trainer.train_with_params(params); // Should complete successfully assert!( result.is_ok(), "Async loading should handle small datasets, got: {:?}", result ); } #[test] fn test_sync_vs_async_loading_consistency() { // Verify sync and async loading produce consistent results let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 200, "sync_async"); // Train with sync loading let mut sync_trainer = Mamba2Trainer::new(&parquet_file, 10) .expect("Failed to create sync trainer") .with_async_loading(false, 0); let params = Mamba2Params::default(); let sync_result = sync_trainer.train_with_params(params.clone()); // Train with async loading let mut async_trainer = Mamba2Trainer::new(&parquet_file, 10) .expect("Failed to create async trainer") .with_async_loading(true, 3); let async_result = async_trainer.train_with_params(params); // Both should succeed assert!(sync_result.is_ok() && async_result.is_ok()); let sync_metrics = sync_result.unwrap(); let async_metrics = async_result.unwrap(); // Metrics should be similar (within 10% tolerance due to different data ordering) let loss_diff = (sync_metrics.val_loss - async_metrics.val_loss).abs(); let max_loss = sync_metrics.val_loss.max(async_metrics.val_loss); assert!( loss_diff / max_loss < 0.1, "Sync and async losses should be similar: sync={}, async={}", sync_metrics.val_loss, async_metrics.val_loss ); } #[test] #[should_panic(expected = "Prefetch count must be >= 2")] fn test_async_loading_invalid_prefetch_count() { // Prefetch count < 2 should panic let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 100, "invalid_prefetch"); let _trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer") .with_async_loading(true, 1); // Should panic } // ============================================================================ // SEQUENCE LENGTH AND STRIDE EDGE CASES // ============================================================================ #[test] fn test_lookback_window_min_bound() { // Test minimum lookback_window (30) let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 100, "min_lookback"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.lookback_window = 30; // Minimum bound let result = trainer.train_with_params(params); // Should succeed assert!( result.is_ok(), "Minimum lookback_window should work, got: {:?}", result ); } #[test] fn test_lookback_window_max_bound() { // Test maximum lookback_window (120) let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 200, "max_lookback"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.lookback_window = 120; // Maximum bound let result = trainer.train_with_params(params); // Should succeed assert!( result.is_ok(), "Maximum lookback_window should work, got: {:?}", result ); } #[test] fn test_sequence_stride_min() { // Test minimum sequence_stride (1) let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "min_stride"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.sequence_stride = 1; // Minimum (non-overlapping) let result = trainer.train_with_params(params); assert!( result.is_ok(), "sequence_stride=1 should work, got: {:?}", result ); } #[test] fn test_sequence_stride_max() { // Test maximum sequence_stride (5) let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "max_stride"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.sequence_stride = 5; // Maximum (heavily overlapping) let result = trainer.train_with_params(params); assert!( result.is_ok(), "sequence_stride=5 should work, got: {:?}", result ); } #[test] fn test_lookback_exceeds_dataset_length() { // lookback_window > dataset length let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 50, "lookback_exceeds"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.lookback_window = 100; // Exceeds 50 rows let result = trainer.train_with_params(params); // Should error or return penalty match result { Err(e) => { let err_msg = format!("{:?}", e); assert!( err_msg.contains("insufficient") || err_msg.contains("empty"), "Expected insufficient data error, got: {}", err_msg ); } Ok(metrics) => { // Penalty loss assert!( metrics.val_loss >= 1000.0, "Expected penalty for excessive lookback, got: {}", metrics.val_loss ); } } } // ============================================================================ // NORMALIZATION PARAMETER EDGE CASES // ============================================================================ #[test] fn test_norm_eps_min_bound() { // Test minimum norm_eps (1e-6) let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "norm_eps_min"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.norm_eps = 1e-6; // Minimum bound let result = trainer.train_with_params(params); assert!( result.is_ok(), "Minimum norm_eps should work, got: {:?}", result ); } #[test] fn test_norm_eps_max_bound() { // Test maximum norm_eps (1e-4) let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "norm_eps_max"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.norm_eps = 1e-4; // Maximum bound let result = trainer.train_with_params(params); assert!( result.is_ok(), "Maximum norm_eps should work, got: {:?}", result ); } #[test] fn test_denormalize_before_training() { // Calling denormalize_prediction before training should panic let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "denorm_before"); let trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); // This should panic let result = std::panic::catch_unwind(|| { trainer.denormalize_prediction(0.5) }); assert!( result.is_err(), "denormalize_prediction before training should panic" ); } #[test] fn test_denormalize_after_training() { // Calling denormalize_prediction after training should work let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "denorm_after"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let params = Mamba2Params::default(); let result = trainer.train_with_params(params); assert!(result.is_ok(), "Training should succeed"); // Now denormalization should work let denormalized = trainer.denormalize_prediction(0.5); assert!(denormalized.is_finite(), "Denormalized value should be finite"); assert!(denormalized > 0.0, "Denormalized price should be positive"); } // ============================================================================ // SSM-SPECIFIC NUMERICAL STABILITY // ============================================================================ #[test] fn test_adam_epsilon_bounds() { // Test minimum and maximum adam_epsilon let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "adam_eps"); // Test minimum (1e-9) let mut trainer_min = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params_min = Mamba2Params::default(); params_min.adam_epsilon = 1e-9; let result_min = trainer_min.train_with_params(params_min); assert!(result_min.is_ok(), "Minimum adam_epsilon should work"); // Test maximum (1e-7) let mut trainer_max = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params_max = Mamba2Params::default(); params_max.adam_epsilon = 1e-7; let result_max = trainer_max.train_with_params(params_max); assert!(result_max.is_ok(), "Maximum adam_epsilon should work"); } #[test] fn test_grad_clip_bounds() { // Test gradient clipping bounds let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "grad_clip"); // Test minimum (0.5) let mut trainer_min = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params_min = Mamba2Params::default(); params_min.grad_clip = 0.5; let result_min = trainer_min.train_with_params(params_min); assert!(result_min.is_ok(), "Minimum grad_clip should work"); // Test maximum (5.0) let mut trainer_max = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params_max = Mamba2Params::default(); params_max.grad_clip = 5.0; let result_max = trainer_max.train_with_params(params_max); assert!(result_max.is_ok(), "Maximum grad_clip should work"); } #[test] fn test_adam_beta_bounds() { // Test Adam beta parameter bounds let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "adam_beta"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer"); let mut params = Mamba2Params::default(); params.adam_beta1 = 0.85; // Minimum params.adam_beta2 = 0.98; // Minimum let result = trainer.train_with_params(params); assert!(result.is_ok(), "Minimum Adam betas should work"); } // ============================================================================ // BATCH SIZE CLAMPING WITH GPU MEMORY // ============================================================================ #[test] fn test_batch_size_clamping_min() { // Test batch_size clamping to minimum bound let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "batch_clamp_min"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer") .with_batch_size_bounds(16.0, 128.0); let mut params = Mamba2Params::default(); params.batch_size = 4; // Below minimum (16) let result = trainer.train_with_params(params); // Should clamp to 16 and succeed assert!( result.is_ok(), "Batch size clamping to minimum should work, got: {:?}", result ); } #[test] fn test_batch_size_clamping_max() { // Test batch_size clamping to maximum bound let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "batch_clamp_max"); let mut trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer") .with_batch_size_bounds(4.0, 32.0); // RTX 3050 Ti constraints let mut params = Mamba2Params::default(); params.batch_size = 256; // Above maximum (32) let result = trainer.train_with_params(params); // Should clamp to 32 and succeed assert!( result.is_ok(), "Batch size clamping to maximum should work, got: {:?}", result ); } #[test] #[should_panic(expected = "Minimum batch size must be >= 1")] fn test_batch_size_bounds_invalid_min() { // Setting minimum batch size < 1 should panic let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "invalid_min"); let _trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer") .with_batch_size_bounds(0.0, 32.0); // Should panic } #[test] #[should_panic(expected = "Maximum batch size must be > minimum")] fn test_batch_size_bounds_invalid_max() { // Setting maximum <= minimum should panic let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 150, "invalid_max"); let _trainer = Mamba2Trainer::new(&parquet_file, 5) .expect("Failed to create trainer") .with_batch_size_bounds(32.0, 16.0); // Should panic } // ============================================================================ // INTEGRATION TESTS // ============================================================================ #[test] fn test_all_13_params_roundtrip() { // Verify all 13 MAMBA2 parameters survive roundtrip conversion let params = Mamba2Params { learning_rate: 5e-5, batch_size: 64, dropout: 0.15, weight_decay: 5e-5, grad_clip: 2.0, warmup_steps: 500, adam_beta1: 0.9, adam_beta2: 0.999, adam_epsilon: 1e-8, total_decay_steps: 10000, lookback_window: 90, sequence_stride: 2, norm_eps: 1e-5, }; let continuous = params.to_continuous(); assert_eq!(continuous.len(), 13, "Should have 13 continuous parameters"); let recovered = Mamba2Params::from_continuous(&continuous) .expect("Failed to recover params"); // Verify all parameters assert!((recovered.learning_rate - params.learning_rate).abs() < 1e-10); assert_eq!(recovered.batch_size, params.batch_size); assert!((recovered.dropout - params.dropout).abs() < 1e-10); assert!((recovered.weight_decay - params.weight_decay).abs() < 1e-10); assert!((recovered.grad_clip - params.grad_clip).abs() < 1e-6); assert_eq!(recovered.warmup_steps, params.warmup_steps); assert!((recovered.adam_beta1 - params.adam_beta1).abs() < 1e-10); assert!((recovered.adam_beta2 - params.adam_beta2).abs() < 1e-10); assert!((recovered.adam_epsilon - params.adam_epsilon).abs() < 1e-12); assert_eq!(recovered.total_decay_steps, params.total_decay_steps); assert_eq!(recovered.lookback_window, params.lookback_window); assert_eq!(recovered.sequence_stride, params.sequence_stride); assert!((recovered.norm_eps - params.norm_eps).abs() < 1e-12); } #[test] fn test_full_training_pipeline() { // End-to-end test: create data, train, denormalize predictions let temp_dir = TempDir::new().unwrap(); let parquet_file = create_test_parquet(&temp_dir, 200, "full_pipeline"); let mut trainer = Mamba2Trainer::new(&parquet_file, 10) .expect("Failed to create trainer") .with_batch_size_bounds(4.0, 32.0) .with_async_loading(true, 3) .with_train_split(0.8); let params = Mamba2Params::default(); let result = trainer.train_with_params(params); assert!(result.is_ok(), "Full training pipeline should succeed"); let metrics = result.unwrap(); // Verify metrics are reasonable assert!(metrics.val_loss.is_finite(), "Validation loss should be finite"); assert!(metrics.val_loss >= 0.0, "Validation loss should be non-negative"); assert!(metrics.directional_accuracy >= 0.0 && metrics.directional_accuracy <= 1.0); assert!(metrics.mae >= 0.0); assert!(metrics.rmse >= 0.0); assert!(metrics.r_squared >= -1.0 && metrics.r_squared <= 1.0); assert_eq!(metrics.epochs_completed, 10); // Test denormalization let pred = trainer.denormalize_prediction(0.5); assert!(pred.is_finite() && pred > 0.0, "Denormalized prediction should be valid"); }