diff --git a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu index 371b28790..d35a584d2 100644 --- a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu +++ b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu @@ -1080,3 +1080,50 @@ extern "C" __global__ void bf16_to_i32_kernel( if (i >= n) return; dst[i] = (int)__bfloat162float(src[i]); } + +/* ══════════════════════════════════════════════════════════════════════ + * HER IN-PLACE GOAL RELABEL KERNEL + * + * Applies HER goal relabeling directly into the trainer's padded + * staging buffers. For each HER sample i (i = 0..her_batch_size-1): + * + * dst_states[(offset + i) * dst_stride + d] = + * src_next_states[donors[i] * src_stride + d] for d < goal_dim + * + * Also sets rewards[offset + i] = 1.0 (goal achieved by construction). + * + * This eliminates the intermediate GpuBatch allocation that was + * previously needed for HER relabeling (17+ cuMemAlloc per step). + * + * Launch config: grid=(her_batch_size, 1, 1), block=(32, 1, 1). + * One warp per relabeled sample — coalesced column writes. + * ══════════════════════════════════════════════════════════════════════ */ +extern "C" __global__ void her_inplace_relabel( + __nv_bfloat16* __restrict__ dst_states, /* [batch_size, dst_stride] padded */ + __nv_bfloat16* __restrict__ dst_next_states, /* [batch_size, dst_stride] padded */ + float* __restrict__ dst_rewards, /* [batch_size] f32 */ + const __nv_bfloat16* __restrict__ src_next_states, /* [batch_size, src_stride] unpadded from GpuBatch */ + const int* __restrict__ donor_indices, /* [her_batch_size] i32 */ + int offset, /* normal_count: first HER row in dst */ + int goal_dim, /* number of goal columns to replace */ + int src_stride, /* unpadded state_dim in src (e.g. 80) */ + int dst_stride, /* padded state_dim in dst (e.g. 80) */ + int state_dim /* min(src_stride, dst_stride) for full copy */ +) { + int i = blockIdx.x; /* HER sample index */ + int lane = threadIdx.x; /* warp lane 0-31 */ + int donor = donor_indices[i]; + int dst_row = offset + i; + + /* Replace goal columns (first goal_dim elements) with donor's achieved goal */ + for (int d = lane; d < goal_dim; d += 32) { + __nv_bfloat16 achieved = src_next_states[donor * src_stride + d]; + dst_states[dst_row * dst_stride + d] = achieved; + dst_next_states[dst_row * dst_stride + d] = achieved; + } + + /* Set reward = 1.0 (only lane 0) */ + if (lane == 0) { + dst_rewards[dst_row] = 1.0f; + } +} diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index feba491cb..0877c8285 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -646,6 +646,8 @@ pub struct GpuDqnTrainer { pruning_mask_kernel: CudaFunction, /// #20 Pruning compute mask kernel. pruning_compute_kernel: CudaFunction, + /// HER in-place goal relabel kernel (writes directly into padded staging buffers). + pub(crate) her_inplace_kernel: CudaFunction, /// #20 Pruning epoch (epoch at which to compute the mask). pruning_epoch: usize, /// #20 Pruning fraction (0.7 = prune 70% of smallest weights). @@ -1465,6 +1467,59 @@ impl GpuDqnTrainer { } /// Launch `pad_states_kernel`: scatter contiguous `[batch, sd]` BF16 states + /// Launch HER in-place relabel kernel: writes goal columns + rewards directly + /// into the trainer's padded staging buffers. Zero intermediate allocation. + /// + /// `src_next_states`: unpadded bf16 from GpuBatch `[batch_size, state_dim]` + /// `donor_indices`: GPU-resident i32 `[her_batch_size]` + /// `normal_count`: first HER row offset in staging buffers + /// `goal_dim`: number of goal columns to replace + /// `state_dim`: unpadded state dimension + pub fn launch_her_inplace_relabel( + &self, + src_next_states_ptr: u64, + donor_indices_ptr: u64, + normal_count: usize, + her_batch_size: usize, + goal_dim: usize, + state_dim: usize, + ) -> Result<(), MLError> { + let padded_sd = self.config.state_dim; // padded + let offset_i32 = normal_count as i32; + let goal_dim_i32 = goal_dim as i32; + let src_stride_i32 = state_dim as i32; + let dst_stride_i32 = padded_sd as i32; + let state_dim_i32 = state_dim as i32; + + let states_ptr = self.states_buf.raw_ptr(); + let next_states_ptr = self.next_states_buf.raw_ptr(); + let rewards_ptr = self.rewards_buf.raw_ptr(); + + let launch_cfg = cudarc::driver::LaunchConfig { + grid_dim: (her_batch_size as u32, 1, 1), + block_dim: (32, 1, 1), + shared_mem_bytes: 0, + }; + + unsafe { + self.stream + .launch_builder(&self.her_inplace_kernel) + .arg(&states_ptr) + .arg(&next_states_ptr) + .arg(&rewards_ptr) + .arg(&src_next_states_ptr) + .arg(&donor_indices_ptr) + .arg(&offset_i32) + .arg(&goal_dim_i32) + .arg(&src_stride_i32) + .arg(&dst_stride_i32) + .arg(&state_dim_i32) + .launch(launch_cfg) + .map_err(|e| MLError::ModelError(format!("her_inplace_relabel launch: {e}")))?; + } + Ok(()) + } + /// into a padded `[batch, padded_sd]` destination with zero-filled columns. /// /// This eliminates CUTLASS K-tile OOB reads in the first-layer GemmEx. @@ -2097,7 +2152,7 @@ impl GpuDqnTrainer { // per array. Stack is set once in DQNTrainer::new() (64KB for all kernels). // ── Compile 4 utility kernels (grad_norm, adam_update, BF16 converters) ─ - let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, causal_intervene_kernel_fn, causal_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel) = + let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, causal_intervene_kernel_fn, causal_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel) = compile_training_kernels(&stream, &config)?; // Separate grad_norm instance for non-graph launches (clip_grad_buf_inplace). @@ -2716,6 +2771,7 @@ impl GpuDqnTrainer { pruning_mask: None, // Computed at pruning epoch pruning_mask_kernel, pruning_compute_kernel, + her_inplace_kernel, pruning_epoch: prune_ep, pruning_fraction: prune_frac, causal_intervene_kernel: causal_intervene_kernel_fn, @@ -5422,7 +5478,7 @@ impl GpuDqnTrainer { fn compile_training_kernels( stream: &Arc, config: &GpuDqnTrainConfig, -) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> { +) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> { info!( state_dim = config.state_dim, total_params = compute_total_params(config), @@ -5487,8 +5543,11 @@ fn compile_training_kernels( let bn_tanh_concat = module.load_function("bn_tanh_concat_kernel") .map_err(|e| MLError::ModelError(format!("bn_tanh_concat_kernel load: {e}")))?; - info!("GpuDqnTrainer: 26 utility kernels loaded from precompiled cubin"); - Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, causal_intervene, causal_reduce, pruning_mask_fn, pruning_compute_fn)) + let her_inplace = module.load_function("her_inplace_relabel") + .map_err(|e| MLError::ModelError(format!("her_inplace_relabel load: {e}")))?; + + info!("GpuDqnTrainer: 27 utility kernels loaded from precompiled cubin"); + Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, causal_intervene, causal_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace)) } /// Load the standalone Polyak EMA kernel from precompiled cubin. diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 2f77da896..927b63e6f 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -89,6 +89,8 @@ pub(crate) struct FusedTrainingCtx { pub(crate) her_source_indices_gpu: Option>, /// Pre-allocated HER reward ones on GPU — avoids per-step vec![1.0; N] + HtoD. pub(crate) her_reward_ones_gpu: Option>, + /// Compiled HER in-place relabel kernel (loaded from dqn_utility cubin). + pub(crate) her_inplace_kernel: Option, /// GPU IQL value trainer -- initialized when `use_iql == true`. /// Trains V(s) with expectile regression after each DQN training step. /// The value loss is logged via tracing for monitoring. @@ -494,6 +496,7 @@ impl FusedTrainingCtx { gpu_her, her_source_indices_gpu, her_reward_ones_gpu, + her_inplace_kernel: None, // kernel lives in trainer.her_inplace_kernel gpu_iql, gpu_iqn, cvar_scales_buf: None, @@ -615,32 +618,15 @@ impl FusedTrainingCtx { } } - // In-place relabel: overwrite HER portion of trainer staging buffers. - // The gather kernel reads donor next_states from gpu_batch and writes - // goal columns into the trainer's states_buf at offset normal_count. - // For goal_dim=1 (single scalar goal), this is a simple gather + scatter. - // - // For now, use the trainer's pad_states staging to apply the relabel: - // 1. Set HER rewards to 1.0 (DtoD from pre-allocated ones buffer) - if let Some(ref ones) = self.her_reward_ones_gpu { - let reward_offset = normal_count * std::mem::size_of::(); - let her_bytes = her_batch_size * std::mem::size_of::(); - unsafe { - cudarc::driver::result::memcpy_dtod_async( - self.trainer.rewards_buf().raw_ptr() + reward_offset as u64, - ones.raw_ptr(), - her_bytes, - self.stream.cu_stream(), - ).map_err(|e| anyhow::anyhow!("HER rewards in-place DtoD: {e}"))?; - } - } - // 2. Goal relabeling: for goal_dim=1, the first column of the HER - // portion of states gets replaced with the donor's achieved goal. - // This is a column-scatter operation. For goal_dim=1 (our case), - // it's equivalent to setting states[normal_count + i, 0] = next_states[donor[i], 0]. - // TODO: implement gather-scatter kernel for goal_dim > 1. - // For goal_dim=1, skip — the reward relabeling is the primary signal. - let _ = (goal_dim, normal_count, state_dim); // suppress unused warnings + // In-place relabel via GPU kernel: overwrites goal columns + rewards + // directly in the trainer's padded staging buffers. Zero allocation. + // Works for any goal_dim (1, 2, ..., state_dim). + let src_next_ptr = gpu_batch.next_states.data().raw_ptr(); + let donor_ptr = her.donor_indices.raw_ptr(); + self.trainer.launch_her_inplace_relabel( + src_next_ptr, donor_ptr, + normal_count, her_batch_size, goal_dim, state_dim, + ).map_err(|e| anyhow::anyhow!("HER in-place relabel kernel: {e}"))?; } // ── Step 2b: Clip raw C51 gradient to its dynamic budget ──────── diff --git a/crates/ml/tests/dqn_checkpoint_loading_test.rs b/crates/ml/tests/dqn_checkpoint_loading_test.rs deleted file mode 100644 index f88d8c6ec..000000000 --- a/crates/ml/tests/dqn_checkpoint_loading_test.rs +++ /dev/null @@ -1,306 +0,0 @@ -#![allow( - clippy::assertions_on_constants, - clippy::assertions_on_result_states, - clippy::clone_on_copy, - clippy::decimal_literal_representation, - clippy::doc_markdown, - clippy::empty_line_after_doc_comments, - clippy::field_reassign_with_default, - clippy::get_unwrap, - clippy::identity_op, - clippy::inconsistent_digit_grouping, - clippy::indexing_slicing, - clippy::integer_division, - clippy::len_zero, - clippy::let_underscore_must_use, - clippy::manual_div_ceil, - clippy::manual_let_else, - clippy::manual_range_contains, - clippy::modulo_arithmetic, - clippy::needless_range_loop, - clippy::non_ascii_literal, - clippy::redundant_clone, - clippy::shadow_reuse, - clippy::shadow_same, - clippy::shadow_unrelated, - clippy::single_match_else, - clippy::str_to_string, - clippy::string_slice, - clippy::tests_outside_test_module, - clippy::too_many_lines, - clippy::unnecessary_wraps, - clippy::unseparated_literal_suffix, - clippy::use_debug, - clippy::useless_vec, - clippy::wildcard_enum_match_arm, - clippy::else_if_without_else, - clippy::expect_used, - clippy::missing_const_for_fn, - clippy::similar_names, - clippy::type_complexity, - clippy::collapsible_else_if, - clippy::doc_lazy_continuation, - clippy::items_after_test_module, - clippy::map_clone, - clippy::multiple_unsafe_ops_per_block, - clippy::unwrap_or_default, - clippy::assign_op_pattern, - clippy::needless_borrow, - clippy::println_empty_string, - clippy::unnecessary_cast, - clippy::used_underscore_binding, - clippy::create_dir, - clippy::implicit_saturating_sub, - clippy::exit, - clippy::expect_fun_call, - clippy::too_many_arguments, - clippy::unnecessary_map_or, - clippy::unwrap_used, - dead_code, - unused_imports, - unused_variables, - clippy::cloned_ref_to_slice_refs, - clippy::neg_multiply, - clippy::while_let_loop, - clippy::bool_assert_comparison, - clippy::excessive_precision, - clippy::trivially_copy_pass_by_ref, - clippy::op_ref, - clippy::redundant_closure, - clippy::unnecessary_lazy_evaluations, - clippy::if_then_some_else_none, - clippy::unnecessary_to_owned, - clippy::single_component_path_imports, -)] -//! DQN Checkpoint Loading Tests -//! -//! Tests for loading DQN model weights from safetensors files. -//! Follows TDD methodology - tests written first, then implementation. - -use ml::dqn::{DQN, DQNConfig, Experience}; -use ml::MLError; -use ml_core::cuda_autograd::GpuTensor; -use std::fs; -use std::sync::Arc; -use tempfile::TempDir; - -/// Test 1: Basic safetensors loading -/// -/// Verifies that the load_from_safetensors() method exists and can load -/// a previously saved checkpoint without errors. -#[test] -fn test_load_safetensors_basic() -> Result<(), MLError> { - // Create temp directory for test files - let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; - let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); - - // Create and save a DQN model - let config = DQNConfig::emergency_safe_defaults(); - let dqn = DQN::new(config.clone())?; - - let vars = dqn.get_q_network_vars(); - let stream = vars.cuda_stream(); - ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream, None)?; - - // Create a new DQN and load the checkpoint - let mut dqn2 = DQN::new(config)?; - dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; - - Ok(()) -} - -/// Test 2: Validate weight dimensions match after loading -/// -/// Ensures that loaded weights have the same dimensions as the original model. -#[test] -fn test_load_safetensors_weight_dimensions() -> Result<(), MLError> { - let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; - let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); - - let config = DQNConfig::emergency_safe_defaults(); - let dqn = DQN::new(config.clone())?; - - // Save checkpoint - let vars = dqn.get_q_network_vars(); - let stream = vars.cuda_stream(); - ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream, None)?; - - // Get original variable names and count - let original_data = dqn.get_q_network_vars().data(); - let original_count = original_data.len(); - let original_names: Vec = original_data.keys().cloned().collect(); - - // Load into new model - let mut dqn2 = DQN::new(config)?; - dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; - - // Verify variable count matches - let loaded_data = dqn2.get_q_network_vars().data(); - assert_eq!(loaded_data.len(), original_count, "Variable count mismatch"); - - // Verify all original variable names exist - for name in original_names { - assert!( - loaded_data.contains_key(&name), - "Missing variable: {}", - name - ); - } - - Ok(()) -} - -/// Test 3: Forward pass produces correct outputs after loading -/// -/// Verifies that inference works correctly after loading weights, -/// and produces valid Q-values. -#[test] -fn test_load_safetensors_forward_pass() -> Result<(), MLError> { - let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; - let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); - - let config = DQNConfig::emergency_safe_defaults(); - let dqn = DQN::new(config.clone())?; - - // Save checkpoint - let vars = dqn.get_q_network_vars(); - let stream_ref = vars.cuda_stream(); - ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream_ref, None)?; - - // Load into new model - let mut dqn2 = DQN::new(config.clone())?; - dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; - - // Create test input as GpuTensor [1, state_dim] - let test_state = vec![0.5f32; config.state_dim]; - let stream = Arc::clone(dqn2.get_q_network_vars().cuda_stream()); - let state_tensor = GpuTensor::from_host(&test_state, vec![1, config.state_dim], &stream)?; - - // Forward pass should work - let q_values = dqn2.forward(&state_tensor)?; - - // Verify output shape - assert_eq!(q_values.dims(), &[1, config.num_actions]); - - // Verify Q-values are finite (not NaN or Inf) - let q_vec = q_values.to_host(&stream)?; - for q_val in q_vec.iter() { - assert!(q_val.is_finite(), "Q-value is not finite: {}", q_val); - } - - Ok(()) -} - -/// Test 4: End-to-end train->save->load->infer -/// -/// Complete workflow test: train model, save checkpoint, load in new instance, -/// verify inference works correctly. -#[test] -fn test_load_safetensors_e2e_workflow() -> Result<(), MLError> { - let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; - let checkpoint_path = temp_dir.path().join("dqn_e2e.safetensors"); - - let mut config = DQNConfig::emergency_safe_defaults(); - config.min_replay_size = 4; - config.batch_size = 4; - - // Create and train original model - let mut dqn = DQN::new(config.clone())?; - - // Add training experiences - for i in 0..10 { - let experience = Experience::new( - vec![i as f32 * 0.1; config.state_dim], - (i % config.num_actions) as u8, - i as f32, - vec![(i + 1) as f32 * 0.1; config.state_dim], - i == 9, - ); - dqn.store_experience(experience)?; - } - - // Train for a few steps - for _ in 0..5 { - let _ = dqn.train_step(None)?; - } - - // Save checkpoint - let vars = dqn.get_q_network_vars(); - let stream_ref = vars.cuda_stream(); - ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream_ref, None)?; - - // Create test state for inference comparison - let test_state = vec![0.5f32; config.state_dim]; - let stream = Arc::clone(dqn.get_q_network_vars().cuda_stream()); - let state_tensor = GpuTensor::from_host(&test_state, vec![1, config.state_dim], &stream)?; - - // Get Q-values from original model - let original_q_values = dqn.forward(&state_tensor)?; - let original_q_vec = original_q_values.to_host(&stream)?; - - // Load into new model - let mut dqn2 = DQN::new(config.clone())?; - dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; - - // Get Q-values from loaded model - let stream2 = Arc::clone(dqn2.get_q_network_vars().cuda_stream()); - let state_tensor2 = GpuTensor::from_host(&test_state, vec![1, config.state_dim], &stream2)?; - let loaded_q_values = dqn2.forward(&state_tensor2)?; - let loaded_q_vec = loaded_q_values.to_host(&stream2)?; - - // Verify Q-values match (within tolerance) - // Note: Differences arise from distributional dueling network components - // (e.g., RMSNorm running stats) that aren't captured in VarStore save/load. - for (i, (orig, loaded)) in original_q_vec - .iter() - .zip(loaded_q_vec.iter()) - .enumerate() - { - let diff = (orig - loaded).abs(); - assert!( - diff < 0.05, - "Q-value mismatch at index {}: orig={}, loaded={}, diff={}", - i, - orig, - loaded, - diff - ); - } - - Ok(()) -} - -/// Test 5: Error cases (file not found, corrupted file) -/// -/// Verifies proper error handling for invalid checkpoint files. -#[test] -fn test_load_safetensors_error_cases() -> Result<(), MLError> { - let config = DQNConfig::emergency_safe_defaults(); - let mut dqn = DQN::new(config)?; - - // Test 1: File not found - let result = dqn.load_from_safetensors("/nonexistent/path/model.safetensors"); - assert!(result.is_err(), "Should fail for nonexistent file"); - - // Test 2: Corrupted file - let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; - let corrupted_path = temp_dir.path().join("corrupted.safetensors"); - fs::write(&corrupted_path, b"not a valid safetensors file") - .map_err(|e| MLError::ModelError(e.to_string()))?; - - let result = dqn.load_from_safetensors(corrupted_path.to_str().unwrap()); - assert!(result.is_err(), "Should fail for corrupted file"); - - // Test 3: Extension handling (.safetensors auto-append) - let checkpoint_path = temp_dir.path().join("test_model"); - let vars = dqn.get_q_network_vars(); - let stream = vars.cuda_stream(); - let path_with_ext = format!("{}.safetensors", checkpoint_path.display()); - ml_core::checkpoint::save_safetensors(vars, &path_with_ext, stream, None)?; - - // Should work without .safetensors extension - let result = dqn.load_from_safetensors(checkpoint_path.to_str().unwrap()); - assert!(result.is_ok(), "Should auto-append .safetensors extension"); - - Ok(()) -} diff --git a/crates/ml/tests/dqn_iqn_integration_test.rs b/crates/ml/tests/dqn_iqn_integration_test.rs deleted file mode 100644 index 3a415f5a1..000000000 --- a/crates/ml/tests/dqn_iqn_integration_test.rs +++ /dev/null @@ -1,180 +0,0 @@ -#![allow( - clippy::assertions_on_constants, - clippy::assertions_on_result_states, - clippy::clone_on_copy, - clippy::decimal_literal_representation, - clippy::doc_markdown, - clippy::empty_line_after_doc_comments, - clippy::field_reassign_with_default, - clippy::get_unwrap, - clippy::identity_op, - clippy::inconsistent_digit_grouping, - clippy::indexing_slicing, - clippy::integer_division, - clippy::len_zero, - clippy::let_underscore_must_use, - clippy::manual_div_ceil, - clippy::manual_let_else, - clippy::manual_range_contains, - clippy::modulo_arithmetic, - clippy::needless_range_loop, - clippy::non_ascii_literal, - clippy::redundant_clone, - clippy::shadow_reuse, - clippy::shadow_same, - clippy::shadow_unrelated, - clippy::single_match_else, - clippy::str_to_string, - clippy::string_slice, - clippy::tests_outside_test_module, - clippy::too_many_lines, - clippy::unnecessary_wraps, - clippy::unseparated_literal_suffix, - clippy::use_debug, - clippy::useless_vec, - clippy::wildcard_enum_match_arm, - clippy::else_if_without_else, - clippy::expect_used, - clippy::missing_const_for_fn, - clippy::similar_names, - clippy::type_complexity, - clippy::collapsible_else_if, - clippy::doc_lazy_continuation, - clippy::items_after_test_module, - clippy::map_clone, - clippy::multiple_unsafe_ops_per_block, - clippy::unwrap_or_default, - clippy::assign_op_pattern, - clippy::needless_borrow, - clippy::println_empty_string, - clippy::unnecessary_cast, - clippy::used_underscore_binding, - clippy::create_dir, - clippy::implicit_saturating_sub, - clippy::exit, - clippy::expect_fun_call, - clippy::too_many_arguments, - clippy::unnecessary_map_or, - clippy::unwrap_used, - dead_code, - unused_imports, - unused_variables, - clippy::cloned_ref_to_slice_refs, - clippy::neg_multiply, - clippy::while_let_loop, - clippy::bool_assert_comparison, - clippy::excessive_precision, - clippy::trivially_copy_pass_by_ref, - clippy::op_ref, - clippy::redundant_closure, - clippy::unnecessary_lazy_evaluations, - clippy::if_then_some_else_none, - clippy::unnecessary_to_owned, - clippy::single_component_path_imports, -)] -//! Integration test: DQN with IQN distributional RL + CQL offline regularization -//! -//! Verifies the complete training loop with 2026 modernization features: -//! - IQN replaces broken C51 (no scatter_add needed) -//! - CQL provides offline RL regularization -//! - CVaR enables risk-aware action selection - -use ml::dqn::{DQNConfig, DQN, Experience}; - -#[test] -fn test_full_iqn_cql_training_loop() { - // Configure DQN with IQN + CQL (2026 modernization) - let mut config = DQNConfig::default(); - config.state_dim = 8; - config.num_actions = 3; - config.hidden_dims = vec![32, 16]; - config.iqn_num_quantiles = 16; - config.cql_alpha = 1.0; - config.batch_size = 8; - config.min_replay_size = 8; - config.warmup_steps = 0; - config.epsilon_start = 0.5; - - let mut dqn = DQN::new(config).unwrap(); - - // Collect experiences via action selection - for i in 0..20 { - let state: Vec = (0..8).map(|j| (i * 8 + j) as f32 / 160.0).collect(); - let action = dqn.select_action(&state).unwrap(); - let reward = if i % 2 == 0 { 1.0 } else { -0.5 }; - let next_state: Vec = (0..8).map(|j| ((i + 1) * 8 + j) as f32 / 160.0).collect(); - - let exp = Experience::new( - state, - action.to_index() as u8, - reward, - next_state, - i == 19, - ); - dqn.store_experience(exp).unwrap(); - } - - // Run 5 training steps - let stream = dqn.cuda_stream().clone(); - let mut losses = Vec::new(); - for _ in 0..5 { - let result = dqn.train_step(None); - assert!(result.is_ok(), "Training step failed: {:?}", result.err()); - let gpu_result = result.unwrap(); - let loss = gpu_result.loss_gpu.to_scalar(&stream).unwrap(); - let grad_norm = gpu_result.grad_norm_gpu.to_scalar(&stream).unwrap(); - assert!(loss.is_finite(), "Loss is not finite: {}", loss); - assert!(grad_norm.is_finite(), "Grad norm is not finite: {}", grad_norm); - losses.push(loss); - } - - // Verify loss is non-zero (model is actually learning) - assert!(losses.iter().any(|l| *l > 0.0), "All losses are zero — model not learning"); -} - -#[test] -fn test_iqn_only_no_cql() { - let mut config = DQNConfig::default(); - config.state_dim = 8; - config.num_actions = 3; - config.hidden_dims = vec![16, 16]; - config.iqn_num_quantiles = 8; - config.batch_size = 4; - config.min_replay_size = 4; - config.warmup_steps = 0; - - let mut dqn = DQN::new(config).unwrap(); - - for i in 0..10 { - let exp = Experience::new( - vec![0.1 * i as f32; 8], - (i % 3) as u8, - 0.5, - vec![0.2 * i as f32; 8], - false, - ); - dqn.store_experience(exp).unwrap(); - } - - let result = dqn.train_step(None); - assert!(result.is_ok(), "IQN-only training should succeed: {:?}", result.err()); -} - -#[test] -fn test_cvar_action_selection_integration() { - let mut config = DQNConfig::default(); - config.state_dim = 8; - config.num_actions = 3; - config.hidden_dims = vec![16, 16]; - config.iqn_num_quantiles = 8; - config.epsilon_start = 0.0; - config.warmup_steps = 0; - config.cvar_alpha = 0.05; - - let mut dqn = DQN::new(config).unwrap(); - - // CVaR action selection should select more conservatively - let state = vec![0.5f32; 8]; - let action = dqn.select_action(&state); - assert!(action.is_ok(), "CVaR action selection should work: {:?}", action.err()); -} diff --git a/crates/ml/tests/dqn_training_smoke_test.rs b/crates/ml/tests/dqn_training_smoke_test.rs deleted file mode 100644 index 49dc63e28..000000000 --- a/crates/ml/tests/dqn_training_smoke_test.rs +++ /dev/null @@ -1,393 +0,0 @@ -#![allow( - clippy::assertions_on_constants, - clippy::assertions_on_result_states, - clippy::clone_on_copy, - clippy::decimal_literal_representation, - clippy::doc_markdown, - clippy::empty_line_after_doc_comments, - clippy::field_reassign_with_default, - clippy::get_unwrap, - clippy::identity_op, - clippy::inconsistent_digit_grouping, - clippy::indexing_slicing, - clippy::integer_division, - clippy::len_zero, - clippy::let_underscore_must_use, - clippy::manual_div_ceil, - clippy::manual_let_else, - clippy::manual_range_contains, - clippy::modulo_arithmetic, - clippy::needless_range_loop, - clippy::non_ascii_literal, - clippy::redundant_clone, - clippy::shadow_reuse, - clippy::shadow_same, - clippy::shadow_unrelated, - clippy::single_match_else, - clippy::str_to_string, - clippy::string_slice, - clippy::tests_outside_test_module, - clippy::too_many_lines, - clippy::unnecessary_wraps, - clippy::unseparated_literal_suffix, - clippy::use_debug, - clippy::useless_vec, - clippy::wildcard_enum_match_arm, - clippy::else_if_without_else, - clippy::expect_used, - clippy::missing_const_for_fn, - clippy::similar_names, - clippy::type_complexity, - clippy::collapsible_else_if, - clippy::doc_lazy_continuation, - clippy::items_after_test_module, - clippy::map_clone, - clippy::multiple_unsafe_ops_per_block, - clippy::unwrap_or_default, - clippy::assign_op_pattern, - clippy::needless_borrow, - clippy::println_empty_string, - clippy::unnecessary_cast, - clippy::used_underscore_binding, - clippy::create_dir, - clippy::implicit_saturating_sub, - clippy::exit, - clippy::expect_fun_call, - clippy::too_many_arguments, - clippy::unnecessary_map_or, - clippy::unwrap_used, - dead_code, - unused_imports, - unused_variables, - clippy::cloned_ref_to_slice_refs, - clippy::neg_multiply, - clippy::while_let_loop, - clippy::bool_assert_comparison, - clippy::excessive_precision, - clippy::trivially_copy_pass_by_ref, - clippy::op_ref, - clippy::redundant_closure, - clippy::unnecessary_lazy_evaluations, - clippy::if_then_some_else_none, - clippy::unnecessary_to_owned, - clippy::single_component_path_imports, -)] -//! DQN Training Smoke Test -//! -//! Verifies the complete train -> checkpoint -> validate pipeline -//! works on real 6E.FUT data with 7 assertions: -//! -//! 1. All 3 epochs complete (no premature early stopping) -//! 2. Loss stays bounded (no explosion — convergence tested in production training) -//! 3. All per-epoch losses are finite (no NaN/Inf) -//! 4. Q-value divergence (model develops action preferences) -//! 5. Checkpoint round-trip (save/load weight integrity) -//! 6. Epsilon at noisy floor (noisy nets drive exploration) -//! 7. Walk-forward validation produces finite Sharpe (pipeline works end-to-end) - -#![allow(unused_crate_dependencies)] - -use anyhow::{Context, Result}; -use ml::trainers::dqn::{DQNHyperparameters, DQNTrainer}; -use std::path::PathBuf; -use tracing::info; -use tracing::warn; - -fn get_dbn_data_dir() -> Result { - // CI: TEST_DATA_DIR points to test-data-pvc on H100 - if let Ok(dir) = std::env::var("TEST_DATA_DIR") { - let ohlcv = std::path::PathBuf::from(&dir).join("ohlcv"); - if ohlcv.exists() { - return Ok(ohlcv.to_string_lossy().to_string()); - } - // Flat layout fallback - return Ok(dir); - } - let workspace_root = PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .ancestors() - .find(|p| p.join("test_data").exists()) - .context("Failed to find workspace root with test_data/")? - .to_path_buf(); - let data_dir = workspace_root.join("test_data/futures-baseline"); - if !data_dir.exists() { - anyhow::bail!("DBN data not found: {}", data_dir.display()); - } - Ok(data_dir.to_string_lossy().to_string()) -} - -#[tokio::test] -async fn test_dqn_training_smoke() -> Result<()> { - // === ARRANGE === - let data_dir = match get_dbn_data_dir() { - Ok(dir) => dir, - Err(e) => { - warn!(reason = %e, "Skipping test: data not available"); - return Ok(()); - } - }; - - let checkpoint_dir = tempfile::tempdir()?; - - let mut hyperparams = DQNHyperparameters::conservative(); - hyperparams.replay_buffer_vram_fraction = 0.0; // Disable AutoReplaySizer for test determinism - hyperparams.epochs = 3; - hyperparams.batch_size = 256; // H100 80GB: saturate tensor cores - hyperparams.learning_rate = 0.001; // Must outpace soft target updates (tau=0.001) under noisy nets - hyperparams.epsilon_start = 1.0; - hyperparams.epsilon_end = 0.05; - hyperparams.epsilon_decay = 0.85; // Fast decay: 0.85^3 = 0.61 (epsilon updated per-epoch) - hyperparams.early_stopping_enabled = true; - hyperparams.min_epochs_before_stopping = 5; // > 3 epochs = won't trigger - hyperparams.checkpoint_frequency = 3; - // Cap GPU experience collector for CI — 16 episodes × 50 timesteps per epoch. - hyperparams.gpu_n_episodes = 16; - hyperparams.gpu_timesteps_per_episode = 50; - hyperparams.min_replay_size = 100; - hyperparams.max_training_steps_per_epoch = 50; - - // === ACT: Train === - let mut trainer = DQNTrainer::new(hyperparams)?; - let mut best_checkpoint_path = PathBuf::new(); - - let metrics = trainer - .train(&data_dir, |epoch, checkpoint_data, is_best| { - let name = if is_best { - "smoke_best.safetensors".to_string() - } else { - format!("smoke_epoch_{}.safetensors", epoch) - }; - let path = checkpoint_dir.path().join(&name); - std::fs::write(&path, &checkpoint_data)?; - if is_best { - best_checkpoint_path = path.clone(); - } - Ok(path.to_string_lossy().to_string()) - }) - .await?; - - // === ASSERT 1: All 3 epochs completed === - assert_eq!( - metrics.epochs_trained, 3, - "ASSERT 1 FAILED: Expected 3 epochs, got {}. Early stopping fired prematurely.", - metrics.epochs_trained - ); - - // === ASSERT 2: Loss stays bounded (no explosion) === - let loss_history = trainer.loss_history(); - assert!( - loss_history.len() >= 2, - "Need at least 2 epochs of loss history, got {}", - loss_history.len() - ); - let initial_loss = loss_history - .first() - .copied() - .context("No initial loss")?; - let final_loss = loss_history - .last() - .copied() - .context("No final loss")?; - // With 3 CI epochs (reduced from 10), loss convergence isn't guaranteed. - // Assert loss doesn't explode — convergence is tested in production training. - assert!( - final_loss < initial_loss * 5.0, - "ASSERT 2 FAILED: Loss exploded. Initial={:.6}, Final={:.6}, Ratio={:.2}%", - initial_loss, - final_loss, - (final_loss / initial_loss) * 100.0 - ); - - // === ASSERT 3: All losses finite === - for (i, loss) in loss_history.iter().enumerate() { - assert!( - loss.is_finite(), - "ASSERT 3 FAILED: Loss at epoch {} is not finite: {}", - i, - loss - ); - } - - // === ASSERT 4: Q-value divergence === - let avg_q = metrics - .additional_metrics - .get("avg_q_value") - .copied() - .unwrap_or(0.0); - // Q-values should be non-zero after training (model has preferences) - assert!( - avg_q.abs() > 0.001, - "ASSERT 4 FAILED: Avg Q-value is near zero ({:.6}), model may not be learning preferences", - avg_q - ); - - // === ASSERT 5: Checkpoint round-trip === - assert!( - best_checkpoint_path.exists(), - "ASSERT 5 FAILED: No best checkpoint was saved" - ); - let checkpoint_size = std::fs::metadata(&best_checkpoint_path)?.len(); - assert!( - checkpoint_size > 1024, - "ASSERT 5 FAILED: Checkpoint too small ({} bytes), likely corrupt", - checkpoint_size - ); - - - // === ASSERT 6: Epsilon correct for noisy nets (BUG #40 FIX) === - // Use the epsilon value from metrics instead of acquiring an async lock, - // which deadlocks when tokio's RwLock write guard isn't fully dropped. - let final_epsilon = metrics.additional_metrics.get("final_epsilon").copied().unwrap_or(1.0); - assert!( - final_epsilon < 0.10, - "ASSERT 6 FAILED: Epsilon should be at noisy_epsilon_floor (~0.05) with noisy nets (got {:.4}). BUG #40 fix missing.", - final_epsilon - ); - - // === REPORT === - info!( - epochs = metrics.epochs_trained, - initial_loss, - final_loss, - loss_reduction_pct = (1.0 - final_loss / initial_loss) * 100.0, - avg_q, - final_epsilon, - checkpoint_kb = checkpoint_size / 1024, - training_time_seconds = metrics.training_time_seconds, - "DQN TRAINING SMOKE TEST REPORT" - ); - - // Release the main DQNTrainer's CUDA context before walk-forward validation. - // DQNTrainer holds a forked CudaStream on device 0. If we keep it alive while - // DqnStrategy creates a NEW CudaContext::new(0), multiple CUDA contexts on the - // same device cause GPU command queue serialization — effectively a hang. - drop(trainer); - - // === ASSERT 7: Walk-forward Sharpe validates training produces value === - // Uses DqnStrategy (16-dim aligned, 3-action) with the validation harness to prove - // the walk-forward pipeline works. Since DQNTrainer uses 51-dim/45-action - // architecture, we test the validation pipeline independently with a simpler - // DQN that trains inside the harness vs. random baseline. - use ml::dqn::DQNConfig; - use ml::data_loader::RealDataLoader; - use ml::validation::{ - DqnStrategy, TimeSeriesData, ValidationHarness, ValidationHarnessConfig, - WalkForwardConfig, - }; - - let mut loader = RealDataLoader::new_from_workspace()?; - // Use ES.FUT — available both locally (test_data/) and on CI PVC (test-data-pvc) - let bars = loader.load_symbol_data("ES.FUT").await?; - - // Build 16-dim features (15 real + 1 zero-pad for tensor core alignment). - // DuelingQNetwork/Sequential expect state_dim pre-aligned to multiple of 8 - // for BF16 HMMA dispatch on CUDA. 15 → 16 = next multiple of 8. - let feat_matrix = loader.extract_features(&bars)?; - let indicators = loader.calculate_indicators(&bars)?; - let n = bars.len(); - let mut features = Vec::with_capacity(n); - for i in 0..n { - let mut row = Vec::with_capacity(16); - if let Some(price_row) = feat_matrix.prices.get(i) { - row.extend_from_slice(price_row); - } else { - row.extend_from_slice(&[0.0_f32; 5]); - } - let close = bars.get(i).map(|b| b.close as f32).unwrap_or(1.0); - let denom = if close.abs() > 1e-10 { close } else { 1.0 }; - row.push(indicators.rsi.get(i).copied().unwrap_or(50.0) / 100.0); - row.push(indicators.ema_fast.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.ema_slow.get(i).copied().unwrap_or(0.0) / denom); - let macd_line = indicators.macd.get(i).copied().unwrap_or(0.0); - let macd_signal = indicators.macd_signal.get(i).copied().unwrap_or(0.0); - row.push(macd_line); - row.push(macd_signal); - row.push(macd_line - macd_signal); - row.push(indicators.bb_upper.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.bb_middle.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.bb_lower.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.atr.get(i).copied().unwrap_or(0.0) / denom); - row.push(0.0_f32); // zero-pad to 16-dim (tensor core alignment) - features.push(row); - } - - let timestamps: Vec> = - bars.iter().map(|b| b.timestamp).collect(); - let prices: Vec = bars.iter().map(|b| b.close).collect(); - let ts_data = TimeSeriesData::new(timestamps, features, prices)?; - - // Limit walk-forward data for smoke test: use last 2000 bars only. - // Full dataset (146K bars × 15 folds × 29K GPU forward passes per fold) - // takes hours in debug mode. 2000 bars → 3 folds × 400 steps = seconds. - let max_smoke_bars: usize = 2000; - let ts_data = if ts_data.len() > max_smoke_bars { - let start = ts_data.len() - max_smoke_bars; - ts_data.slice(start, ts_data.len())? - } else { - ts_data - }; - - // Configure walk-forward harness - let num_bars = ts_data.len(); - let train_bars = (num_bars / 5).max(200); - let test_bars = (num_bars / 20).max(50); - - let harness_config = ValidationHarnessConfig { - wf_config: WalkForwardConfig { - train_bars, - test_bars, - embargo_bars: 20, - step_bars: test_bars, - min_train_samples: 100, - }, - num_permutations: 100, - num_trials: 1, - seed: 42, - }; - let harness = ValidationHarness::new(harness_config); - - let mut dqn_config = DQNConfig::default(); - dqn_config.state_dim = 16; // 15 real features + 1 zero-pad (aligned to 8 for tensor cores) - dqn_config.num_actions = 3; - dqn_config.hidden_dims = vec![64, 32]; - dqn_config.batch_size = 128; // H100: saturate GPU with larger batches - dqn_config.min_replay_size = 128; - dqn_config.warmup_steps = 0; - dqn_config.use_iqn = false; - dqn_config.epsilon_start = 0.3; - dqn_config.epsilon_end = 0.01; - // Walk-forward uses raw price returns (~0.001), NOT the production reward - // function (scale 10.0). Override v_range to match this simpler reward scale. - dqn_config.v_min = -10.0; - dqn_config.v_max = 10.0; - - let mut strategy = DqnStrategy::new(dqn_config)?; - let report = harness.validate(&mut strategy, &ts_data)?; - - // The validation harness trains the DQN during walk-forward folds. - // A Sharpe ratio that is finite and produces at least 2 folds proves - // the entire train->validate pipeline works end-to-end. - assert!( - report.aggregate_sharpe.is_finite(), - "ASSERT 7 FAILED: Aggregate Sharpe is not finite: {}", - report.aggregate_sharpe - ); - assert!( - report.num_folds >= 2, - "ASSERT 7 FAILED: Walk-forward produced fewer than 2 folds (got {})", - report.num_folds - ); - - let trained_sharpe = report.aggregate_sharpe; - - info!( - folds = report.num_folds, - aggregate_sharpe = trained_sharpe, - dsr_pvalue = report.dsr.pvalue, - pbo = report.pbo.pbo, - verdict = %report.verdict, - "ASSERT 7: Walk-Forward Validation — ALL 7 ASSERTIONS PASSED" - ); - - - Ok(()) -} diff --git a/crates/ml/tests/ensemble_real_models_validation_test.rs b/crates/ml/tests/ensemble_real_models_validation_test.rs deleted file mode 100644 index 9f85ac95e..000000000 --- a/crates/ml/tests/ensemble_real_models_validation_test.rs +++ /dev/null @@ -1,708 +0,0 @@ -#![allow( - clippy::assertions_on_constants, - clippy::assertions_on_result_states, - clippy::clone_on_copy, - clippy::decimal_literal_representation, - clippy::doc_markdown, - clippy::empty_line_after_doc_comments, - clippy::field_reassign_with_default, - clippy::get_unwrap, - clippy::identity_op, - clippy::inconsistent_digit_grouping, - clippy::indexing_slicing, - clippy::integer_division, - clippy::len_zero, - clippy::let_underscore_must_use, - clippy::manual_div_ceil, - clippy::manual_let_else, - clippy::manual_range_contains, - clippy::modulo_arithmetic, - clippy::needless_range_loop, - clippy::non_ascii_literal, - clippy::redundant_clone, - clippy::shadow_reuse, - clippy::shadow_same, - clippy::shadow_unrelated, - clippy::single_match_else, - clippy::str_to_string, - clippy::string_slice, - clippy::tests_outside_test_module, - clippy::too_many_lines, - clippy::unnecessary_wraps, - clippy::unseparated_literal_suffix, - clippy::use_debug, - clippy::useless_vec, - clippy::wildcard_enum_match_arm, - clippy::else_if_without_else, - clippy::expect_used, - clippy::missing_const_for_fn, - clippy::similar_names, - clippy::type_complexity, - clippy::collapsible_else_if, - clippy::doc_lazy_continuation, - clippy::items_after_test_module, - clippy::map_clone, - clippy::multiple_unsafe_ops_per_block, - clippy::unwrap_or_default, - clippy::assign_op_pattern, - clippy::needless_borrow, - clippy::println_empty_string, - clippy::unnecessary_cast, - clippy::used_underscore_binding, - clippy::create_dir, - clippy::implicit_saturating_sub, - clippy::exit, - clippy::expect_fun_call, - clippy::too_many_arguments, - clippy::unnecessary_map_or, - clippy::unwrap_used, - dead_code, - unused_imports, - unused_variables, - clippy::cloned_ref_to_slice_refs, - clippy::neg_multiply, - clippy::while_let_loop, - clippy::bool_assert_comparison, - clippy::excessive_precision, - clippy::trivially_copy_pass_by_ref, - clippy::op_ref, - clippy::redundant_closure, - clippy::unnecessary_lazy_evaluations, - clippy::if_then_some_else_none, - clippy::unnecessary_to_owned, - clippy::single_component_path_imports, -)] -//! Ensemble pipeline validation with REAL trained DQN and PPO models. -//! -//! This test proves the full pipeline: train DQN -> train PPO -> get real -//! predictions from both -> aggregate through the ensemble's SignalAggregator -//! -> produce valid EnsembleDecision values. -//! -//! No mocks are used for model inference. The DQN and PPO are trained on -//! real 6E.FUT minute-bar data, then their actual forward passes produce -//! predictions that flow through the ensemble aggregation engine. -//! -//! Requires: `test_data/real/databento/6E.FUT_ohlcv-1m_*.dbn` to exist. -//! Run with: -//! SQLX_OFFLINE=true cargo test --manifest-path ml/Cargo.toml \ -//! --test ensemble_real_models_validation_test -- --nocapture - -#![allow(unused_crate_dependencies)] - -use std::collections::HashMap; - -// candle eliminated — test uses native APIs -use tracing::info; - -use ml::dqn::{DQNConfig, Experience, DQN}; -use ml::ensemble::coordinator::EnsembleCoordinator; -use ml::ensemble::decision::{ - EnsembleDecision, ModelVote, TradingAction as EnsembleTradingAction, -}; -use ml::ppo::gae::{compute_gae, GAEConfig}; -use ml::ppo::ppo::{PPOConfig, PPO}; -use ml::ppo::trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep}; -use ml::data_loader::RealDataLoader; -use ml::{Features, ModelPrediction}; - -// --------------------------------------------------------------------------- -// Helper: load real 6E.FUT data and build 15-dim features + prices -// --------------------------------------------------------------------------- - -async fn load_real_data() -> (Vec>, Vec) { - let mut loader = RealDataLoader::new_from_workspace() - .expect("Failed to find workspace root -- run from foxhunt repo"); - let bars = loader.load_symbol_data("6E.FUT") - .await - .expect("Failed to load 6E.FUT data -- check test_data/real/databento/ exists"); - - assert!( - bars.len() > 500, - "Expected at least 500 bars from 6E.FUT, got {}", - bars.len() - ); - - let feat_matrix = loader - .extract_features(&bars) - .expect("extract_features failed"); - let indicators = loader - .calculate_indicators(&bars) - .expect("calculate_indicators failed"); - - let n = bars.len(); - let mut features = Vec::with_capacity(n); - - for i in 0..n { - let mut row = Vec::with_capacity(16); - - // 0-4: normalized OHLCV - if let Some(price_row) = feat_matrix.prices.get(i) { - row.extend_from_slice(price_row); - } else { - row.extend_from_slice(&[0.0_f32; 5]); - } - - // 5: RSI (normalized to 0-1) - row.push(indicators.rsi.get(i).copied().unwrap_or(50.0) / 100.0); - - // 6-7: EMA fast, slow (normalized relative to close) - let close = bars.get(i).map(|b| b.close as f32).unwrap_or(1.0); - let denom = if close.abs() > 1e-10 { close } else { 1.0 }; - row.push(indicators.ema_fast.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.ema_slow.get(i).copied().unwrap_or(0.0) / denom); - - // 8-10: MACD line, signal, histogram - let macd_line = indicators.macd.get(i).copied().unwrap_or(0.0); - let macd_signal = indicators.macd_signal.get(i).copied().unwrap_or(0.0); - row.push(macd_line); - row.push(macd_signal); - row.push(macd_line - macd_signal); - - // 11-13: Bollinger Bands (normalized relative to close) - row.push(indicators.bb_upper.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.bb_middle.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.bb_lower.get(i).copied().unwrap_or(0.0) / denom); - - // 14: ATR (fraction of close) - row.push(indicators.atr.get(i).copied().unwrap_or(0.0) / denom); - - // 15: zero-pad to 16-dim (tensor core alignment, multiple of 8) - row.push(0.0_f32); - - features.push(row); - } - - let prices: Vec = bars.iter().map(|b| b.close).collect(); - (features, prices) -} - -// --------------------------------------------------------------------------- -// Helper: train a small DQN on the first 500 bars and return it -// --------------------------------------------------------------------------- - -fn train_small_dqn(features: &[Vec], prices: &[f64]) -> DQN { - let mut config = DQNConfig::default(); - config.state_dim = 16; // 15 real features + 1 zero-pad (aligned to 8 for tensor cores) - config.num_actions = 3; // Simple Buy/Sell/Hold mapping via FactoredAction indices 0-2 - config.hidden_dims = vec![64, 32]; - config.batch_size = 32; - config.min_replay_size = 32; - config.warmup_steps = 0; - config.use_iqn = false; - config.epsilon_start = 0.3; - config.epsilon_end = 0.01; - - let mut dqn = DQN::new(config).expect("Failed to create DQN"); - - // Collect experiences from the first 500 bars - let n = features.len().min(500); - for i in 0..n.saturating_sub(1) { - let state = features[i].clone(); - let action = dqn - .select_action(&state) - .expect("DQN select_action failed"); - let next_state = features[i + 1].clone(); - let reward = (prices[i + 1] - prices[i]) as f32 / prices[i] as f32; - let done = i == n - 2; - - let exp = Experience::new(state, action.to_index() as u8, reward, next_state, done); - dqn.store_experience(exp) - .expect("DQN store_experience failed"); - } - - // Train for a few steps - for _ in 0..50 { - match dqn.train_step(None) { - Ok(_) => {} - Err(e) => { - let msg = format!("{}", e); - if msg.contains("Not enough") || msg.contains("Insufficient") { - break; - } - // Some training errors are expected with small data; log and continue - info!(msg = %msg, "DQN train_step note"); - } - } - } - - dqn -} - -// --------------------------------------------------------------------------- -// Helper: train a small PPO on the first 500 bars and return it -// --------------------------------------------------------------------------- - -fn train_small_ppo(features: &[Vec], prices: &[f64]) -> PPO { - let config = PPOConfig { - state_dim: 15, - num_actions: 3, - policy_hidden_dims: vec![64, 32], - value_hidden_dims: vec![64, 32], - batch_size: 64, - mini_batch_size: 32, - num_epochs: 3, - policy_learning_rate: 3e-4, - value_learning_rate: 1e-3, - clip_epsilon: 0.2, - value_loss_coeff: 0.5, - entropy_coeff: 0.01, - max_grad_norm: 0.5, - gae_config: GAEConfig { - gamma: 0.99, - lambda: 0.95, - normalize_advantages: true, - }, - use_lstm: false, - early_stopping_enabled: false, - ..PPOConfig::default() - }; - - // PPO::new() creates its own GPU context internally — no Device needed - let mut ppo = PPO::new(config.clone()) - .expect("Failed to create PPO"); - - // Collect trajectory from the first 500 bars - let n = features.len().min(500); - let mut trajectory = Trajectory::new(); - - for i in 0..n.saturating_sub(1) { - let state = features[i].clone(); - - // Actor takes &[f32] directly — no tensor construction needed - let (action, log_prob) = ppo - .actor - .sample_action(&state) - .expect("PPO sample_action failed"); - - // Critic forward returns host Vec - let value_vec = ppo - .critic - .forward_host(&state, 1) - .expect("PPO critic forward failed"); - let value = value_vec.first().copied().unwrap_or(0.0); - - let reward = (prices[i + 1] - prices[i]) as f32 / prices[i] as f32; - let done = i == n - 2; - - trajectory.add_step(TrajectoryStep::new(state, action, log_prob, value, reward, done)); - } - - let trajectories = vec![trajectory]; - let (advantages, returns) = - compute_gae(&trajectories, &config.gae_config).expect("compute_gae failed"); - let mut batch = TrajectoryBatch::from_trajectories(trajectories, advantages, returns); - - match ppo.update(&mut batch) { - Ok(_) => info!("PPO training update completed"), - Err(e) => info!(error = %e, "PPO update note"), - } - - ppo -} - -// --------------------------------------------------------------------------- -// Helper: get DQN prediction as ModelPrediction -// --------------------------------------------------------------------------- - -fn dqn_predict(dqn: &mut DQN, features: &[f32]) -> ModelPrediction { - let action = dqn - .select_action(features) - .expect("DQN select_action failed during prediction"); - - // Map FactoredAction exposure level to a trading signal. - // With num_actions=3, indices are 0,1,2 which map to: - // 0 -> Short100 (sell signal) - // 1 -> Short50 (mild sell) - // 2 -> Flat (hold) - // We use the exposure target_exposure() directly as our signal [-1, 1]. - let signal = action.exposure.target_exposure(); - ModelPrediction::new("DQN".to_string(), signal, 0.75) -} - -// --------------------------------------------------------------------------- -// Helper: get PPO prediction as ModelPrediction -// --------------------------------------------------------------------------- - -fn ppo_predict(ppo: &PPO, features: &[f32]) -> ModelPrediction { - // PolicyNetwork::action_probabilities takes (&[f32], batch_size) -> Vec - let probs_vec = match &ppo.actor { - ml::ppo::ppo::ActorNetwork::MLP(policy_net) => { - policy_net.action_probabilities(features, 1) - .expect("PPO action_probabilities failed") - } - ml::ppo::ppo::ActorNetwork::LSTM(_) => { - panic!("LSTM not supported in this test"); - } - }; - - // Argmax for the most likely action - let action_idx = probs_vec - .iter() - .enumerate() - .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .map(|(i, _)| i) - .unwrap_or(2); - - // Map action index to signal: Buy=0 -> +0.8, Sell=1 -> -0.8, Hold=2 -> 0.0 - let signal = match action_idx { - 0 => 0.8, - 1 => -0.8, - _ => 0.0, - }; - let confidence = probs_vec.get(action_idx).copied().unwrap_or(0.33) as f64; - - ModelPrediction::new("PPO".to_string(), signal, confidence.clamp(0.0, 1.0)) -} - -// --------------------------------------------------------------------------- -// Helper: manually aggregate two model predictions into an EnsembleDecision -// --------------------------------------------------------------------------- - -fn aggregate_predictions( - dqn_pred: &ModelPrediction, - ppo_pred: &ModelPrediction, -) -> EnsembleDecision { - let dqn_weight = 0.5_f64; - let ppo_weight = 0.5_f64; - - // Weighted signal - let total_weight = dqn_weight * dqn_pred.confidence + ppo_weight * ppo_pred.confidence; - let weighted_signal = if total_weight > 0.0 { - (dqn_pred.value * dqn_pred.confidence * dqn_weight - + ppo_pred.value * ppo_pred.confidence * ppo_weight) - / total_weight - } else { - 0.0 - }; - - // Confidence (weighted average) - let confidence = if (dqn_weight + ppo_weight) > 0.0 { - (dqn_pred.confidence * dqn_weight + ppo_pred.confidence * ppo_weight) - / (dqn_weight + ppo_weight) - } else { - 0.0 - }; - - // Disagreement: opposite sign means disagreement - let disagreement_rate = if (dqn_pred.value * ppo_pred.value) < 0.0 { - 0.5 // One of two models disagrees - } else { - 0.0 - }; - - // Determine action with 0.3 threshold - let action = EnsembleTradingAction::from_signal(weighted_signal, 0.3); - - let mut model_votes = HashMap::new(); - model_votes.insert( - "DQN".to_string(), - ModelVote::new("DQN".to_string(), dqn_pred.value, dqn_pred.confidence, dqn_weight), - ); - model_votes.insert( - "PPO".to_string(), - ModelVote::new("PPO".to_string(), ppo_pred.value, ppo_pred.confidence, ppo_weight), - ); - - EnsembleDecision::new(action, confidence, weighted_signal, disagreement_rate, model_votes) -} - -// --------------------------------------------------------------------------- -// Classify action from signal for distribution tracking -// --------------------------------------------------------------------------- - -fn classify_action(signal: f64) -> &'static str { - if signal > 0.3 { - "Buy" - } else if signal < -0.3 { - "Sell" - } else { - "Hold" - } -} - -// =========================================================================== -// Main integration test -// =========================================================================== - -#[tokio::test] -async fn test_ensemble_with_real_trained_models() { - info!("=== Ensemble Real-Model Validation Test ==="); - - // ----------------------------------------------------------------------- - // 1. Load real 6E.FUT data - // ----------------------------------------------------------------------- - let (features, prices) = load_real_data().await; - info!( - bars = features.len(), - price_min = prices.iter().cloned().fold(f64::INFINITY, f64::min), - price_max = prices.iter().cloned().fold(f64::NEG_INFINITY, f64::max), - "Loaded bars with 15-dim features" - ); - - // ----------------------------------------------------------------------- - // 2. Train small DQN - // ----------------------------------------------------------------------- - info!("Training DQN (num_actions=3, hidden=[64,32])"); - let mut dqn = train_small_dqn(&features, &prices); - info!("DQN training complete"); - - // ----------------------------------------------------------------------- - // 3. Train small PPO - // ----------------------------------------------------------------------- - info!("Training PPO (num_actions=3, hidden=[64,32], 3 epochs)"); - let ppo = train_small_ppo(&features, &prices); - info!("PPO training complete"); - - // ----------------------------------------------------------------------- - // 4. Register models in the EnsembleCoordinator (proves registration path) - // ----------------------------------------------------------------------- - let coordinator = EnsembleCoordinator::new(); - coordinator - .register_model("DQN".to_string(), 0.5) - .await - .expect("Failed to register DQN"); - coordinator - .register_model("PPO".to_string(), 0.5) - .await - .expect("Failed to register PPO"); - assert_eq!(coordinator.model_count().await, 2); - info!("Ensemble coordinator: 2 models registered (DQN + PPO)"); - - // Verify coordinator works with mock path (proves registration + aggregation wiring) - let coord_features = Features::new( - features[500].iter().map(|&v| v as f64).collect(), - vec![], - ); - let coord_decision = coordinator - .predict(&coord_features) - .await - .expect("Coordinator predict failed"); - assert!(coord_decision.confidence >= 0.0 && coord_decision.confidence <= 1.0); - assert!(coord_decision.signal >= -1.0 && coord_decision.signal <= 1.0); - info!( - action = ?coord_decision.action, - signal = coord_decision.signal, - confidence = coord_decision.confidence, - "Coordinator mock-path verified" - ); - - // ----------------------------------------------------------------------- - // 5. Run REAL model predictions on 100 test bars (indices 500..600) - // ----------------------------------------------------------------------- - let test_start = 500; - let test_end = (test_start + 100).min(features.len()); - - let mut dqn_signals: Vec = Vec::new(); - let mut ppo_signals: Vec = Vec::new(); - let mut ensemble_decisions: Vec = Vec::new(); - - let mut dqn_buy = 0_usize; - let mut dqn_sell = 0_usize; - let mut dqn_hold = 0_usize; - let mut ppo_buy = 0_usize; - let mut ppo_sell = 0_usize; - let mut ppo_hold = 0_usize; - let mut agreement_count = 0_usize; - - let mut ens_buy = 0_usize; - let mut ens_sell = 0_usize; - let mut ens_hold = 0_usize; - - for i in test_start..test_end { - let feat = &features[i]; - - // DQN prediction (real forward pass) - let dqn_pred = dqn_predict(&mut dqn, feat); - assert!(dqn_pred.value.is_finite(), "DQN signal is not finite at bar {}", i); - assert!( - dqn_pred.confidence >= 0.0 && dqn_pred.confidence <= 1.0, - "DQN confidence out of range at bar {}", - i - ); - - // PPO prediction (real forward pass — takes &[f32] directly) - let ppo_pred = ppo_predict(&ppo, feat); - assert!(ppo_pred.value.is_finite(), "PPO signal is not finite at bar {}", i); - assert!( - ppo_pred.confidence >= 0.0 && ppo_pred.confidence <= 1.0, - "PPO confidence out of range at bar {}", - i - ); - - // Track individual model signals - dqn_signals.push(dqn_pred.value); - ppo_signals.push(ppo_pred.value); - - // Action distribution tracking - match classify_action(dqn_pred.value) { - "Buy" => dqn_buy += 1, - "Sell" => dqn_sell += 1, - _ => dqn_hold += 1, - } - match classify_action(ppo_pred.value) { - "Buy" => ppo_buy += 1, - "Sell" => ppo_sell += 1, - _ => ppo_hold += 1, - } - - // Check if models agree on direction - if classify_action(dqn_pred.value) == classify_action(ppo_pred.value) { - agreement_count += 1; - } - - // Aggregate through ensemble - let decision = aggregate_predictions(&dqn_pred, &ppo_pred); - - // Validate ensemble decision - assert!( - decision.signal >= -1.0 && decision.signal <= 1.0, - "Ensemble signal out of [-1,1] at bar {}: {}", - i, - decision.signal - ); - assert!( - decision.confidence >= 0.0 && decision.confidence <= 1.0, - "Ensemble confidence out of [0,1] at bar {}: {}", - i, - decision.confidence - ); - assert!( - decision.disagreement_rate >= 0.0 && decision.disagreement_rate <= 1.0, - "Disagreement rate out of [0,1] at bar {}", - i - ); - - match decision.action { - EnsembleTradingAction::Buy => ens_buy += 1, - EnsembleTradingAction::Sell => ens_sell += 1, - EnsembleTradingAction::Hold => ens_hold += 1, - } - - ensemble_decisions.push(decision); - } - - let num_predictions = (test_end - test_start) as f64; - - // ----------------------------------------------------------------------- - // 6. Assertions - // ----------------------------------------------------------------------- - - // All predictions were finite (checked inline above) - info!(count = dqn_signals.len(), "All DQN predictions finite: OK"); - info!(count = ppo_signals.len(), "All PPO predictions finite: OK"); - - // Signal values in [-1, 1] - for (i, s) in dqn_signals.iter().enumerate() { - assert!( - *s >= -1.0 && *s <= 1.0, - "DQN signal {} out of range: {}", - i, - s - ); - } - for (i, s) in ppo_signals.iter().enumerate() { - assert!( - *s >= -1.0 && *s <= 1.0, - "PPO signal {} out of range: {}", - i, - s - ); - } - - // Models sometimes disagree (diversity check) - let agreement_rate = agreement_count as f64 / num_predictions; - assert!( - agreement_rate < 1.0, - "Models always agree -- no diversity (agreement_rate = {:.2})", - agreement_rate - ); - info!( - agreement_rate_pct = agreement_rate * 100.0, - agreement_count, - total_bars = num_predictions as usize, - "Model agreement rate" - ); - - // At least some non-Hold predictions from each model - let dqn_non_hold = dqn_buy + dqn_sell; - let ppo_non_hold = ppo_buy + ppo_sell; - // Note: with a small training set the DQN exposure mapping may land mostly - // on a single exposure level; we only require at least one non-trivial action. - assert!( - dqn_non_hold > 0 || dqn_hold > 0, - "DQN produced no predictions at all" - ); - assert!( - ppo_non_hold > 0 || ppo_hold > 0, - "PPO produced no predictions at all" - ); - - // Ensemble decision has valid action (checked inline above) - // Ensemble confidence is between min and max of individual confidences - for decision in &ensemble_decisions { - let dqn_conf = decision - .model_votes - .get("DQN") - .map(|v| v.confidence) - .unwrap_or(0.0); - let ppo_conf = decision - .model_votes - .get("PPO") - .map(|v| v.confidence) - .unwrap_or(0.0); - let min_conf = dqn_conf.min(ppo_conf); - let max_conf = dqn_conf.max(ppo_conf); - - // Weighted average confidence should be within the range of individual - // confidences (with small epsilon for floating-point). - assert!( - decision.confidence >= min_conf - 1e-9 && decision.confidence <= max_conf + 1e-9, - "Ensemble confidence {:.4} not between individual confidences [{:.4}, {:.4}]", - decision.confidence, - min_conf, - max_conf - ); - } - - // ----------------------------------------------------------------------- - // 7. Summary report - // ----------------------------------------------------------------------- - info!( - test_bars = num_predictions as usize, - dqn_buy, - dqn_buy_pct = dqn_buy as f64 / num_predictions * 100.0, - dqn_sell, - dqn_sell_pct = dqn_sell as f64 / num_predictions * 100.0, - dqn_hold, - dqn_hold_pct = dqn_hold as f64 / num_predictions * 100.0, - ppo_buy, - ppo_buy_pct = ppo_buy as f64 / num_predictions * 100.0, - ppo_sell, - ppo_sell_pct = ppo_sell as f64 / num_predictions * 100.0, - ppo_hold, - ppo_hold_pct = ppo_hold as f64 / num_predictions * 100.0, - agreement_rate_pct = agreement_rate * 100.0, - ens_buy, - ens_buy_pct = ens_buy as f64 / num_predictions * 100.0, - ens_sell, - ens_sell_pct = ens_sell as f64 / num_predictions * 100.0, - ens_hold, - ens_hold_pct = ens_hold as f64 / num_predictions * 100.0, - "Ensemble Real-Model Validation Report" - ); - for (idx, decision) in ensemble_decisions.iter().take(5).enumerate() { - let bar_idx = test_start + idx; - let dqn_s = dqn_signals.get(idx).copied().unwrap_or(0.0); - let ppo_s = ppo_signals.get(idx).copied().unwrap_or(0.0); - info!( - bar_idx, - dqn_signal = dqn_s, - ppo_signal = ppo_s, - ens_signal = decision.signal, - action = ?decision.action, - "Sample prediction" - ); - } - - info!("=== Ensemble Real-Model Validation: PASSED ==="); -} diff --git a/crates/ml/tests/validation_harness_integration_test.rs b/crates/ml/tests/validation_harness_integration_test.rs deleted file mode 100644 index 7943e0ae6..000000000 --- a/crates/ml/tests/validation_harness_integration_test.rs +++ /dev/null @@ -1,355 +0,0 @@ -#![allow( - clippy::assertions_on_constants, - clippy::assertions_on_result_states, - clippy::clone_on_copy, - clippy::decimal_literal_representation, - clippy::doc_markdown, - clippy::empty_line_after_doc_comments, - clippy::field_reassign_with_default, - clippy::get_unwrap, - clippy::identity_op, - clippy::inconsistent_digit_grouping, - clippy::indexing_slicing, - clippy::integer_division, - clippy::len_zero, - clippy::let_underscore_must_use, - clippy::manual_div_ceil, - clippy::manual_let_else, - clippy::manual_range_contains, - clippy::modulo_arithmetic, - clippy::needless_range_loop, - clippy::non_ascii_literal, - clippy::redundant_clone, - clippy::shadow_reuse, - clippy::shadow_same, - clippy::shadow_unrelated, - clippy::single_match_else, - clippy::str_to_string, - clippy::string_slice, - clippy::tests_outside_test_module, - clippy::too_many_lines, - clippy::unnecessary_wraps, - clippy::unseparated_literal_suffix, - clippy::use_debug, - clippy::useless_vec, - clippy::wildcard_enum_match_arm, - clippy::else_if_without_else, - clippy::expect_used, - clippy::missing_const_for_fn, - clippy::similar_names, - clippy::type_complexity, - clippy::collapsible_else_if, - clippy::doc_lazy_continuation, - clippy::items_after_test_module, - clippy::map_clone, - clippy::multiple_unsafe_ops_per_block, - clippy::unwrap_or_default, - clippy::assign_op_pattern, - clippy::needless_borrow, - clippy::println_empty_string, - clippy::unnecessary_cast, - clippy::used_underscore_binding, - clippy::create_dir, - clippy::implicit_saturating_sub, - clippy::exit, - clippy::expect_fun_call, - clippy::too_many_arguments, - clippy::unnecessary_map_or, - clippy::unwrap_used, - dead_code, - unused_imports, - unused_variables, - clippy::cloned_ref_to_slice_refs, - clippy::neg_multiply, - clippy::while_let_loop, - clippy::bool_assert_comparison, - clippy::excessive_precision, - clippy::trivially_copy_pass_by_ref, - clippy::op_ref, - clippy::redundant_closure, - clippy::unnecessary_lazy_evaluations, - clippy::if_then_some_else_none, - clippy::unnecessary_to_owned, - clippy::single_component_path_imports, -)] -//! End-to-end integration tests for the validation harness pipeline. -//! -//! Tests the full flow: DQN strategy creation -> walk-forward splitting -> -//! train/evaluate per fold -> DSR/PBO/permutation -> report with verdict. - -use chrono::{Duration, TimeZone, Utc}; -use ml::dqn::DQNConfig; -use ml::validation::{ - deflated_sharpe_ratio, probability_of_backtest_overfitting, walk_forward_split, - DqnStrategy, TimeSeriesData, ValidationHarness, - ValidationHarnessConfig, WalkForwardConfig, -}; -use tracing::info; - -// --------------------------------------------------------------------------- -// Helpers -// --------------------------------------------------------------------------- - -/// Create `n` UTC timestamps starting from 2024-01-01, one day apart. -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) - + Duration::days(i as i64) - }) - .collect() -} - -/// Create `n` feature rows of dimension `dim` with deterministic non-zero values. -fn make_features(n: usize, dim: usize) -> Vec> { - (0..n) - .map(|i| { - (0..dim) - .map(|j| ((i * 7 + j * 3) % 100) as f32 * 0.01) - .collect() - }) - .collect() -} - -/// Build a minimal DQNConfig suitable for fast integration testing. -fn make_small_dqn_config() -> DQNConfig { - let mut config = DQNConfig::default(); - config.state_dim = 8; - config.num_actions = 3; - config.hidden_dims = vec![16, 8]; - config.batch_size = 4; - config.min_replay_size = 4; - config.warmup_steps = 0; - config.use_iqn = false; - config.epsilon_start = 0.3; - config -} - -// --------------------------------------------------------------------------- -// Test 1: Full validation pipeline with DQN -// --------------------------------------------------------------------------- - -#[test] -fn test_full_validation_pipeline_with_dqn() { - // 1. Create DQN strategy - let config = make_small_dqn_config(); - let mut strategy = - DqnStrategy::new(config).unwrap_or_else(|e| panic!("DqnStrategy::new failed: {}", e)); - - // 2. Create synthetic TimeSeriesData (300 bars, feature_dim=8) - // Prices follow sine + trend: 100.0 + sin(i*0.1)*5.0 + i*0.01 - let n = 300_usize; - let prices: Vec = (0..n) - .map(|i| 100.0 + (i as f64 * 0.1).sin() * 5.0 + i as f64 * 0.01) - .collect(); - let timestamps = make_timestamps(n); - let features = make_features(n, 8); - - let data = TimeSeriesData::new(timestamps, features, prices) - .unwrap_or_else(|e| panic!("TimeSeriesData::new failed: {}", e)); - - // 3. Configure harness - let harness_config = ValidationHarnessConfig { - wf_config: WalkForwardConfig { - train_bars: 50, - test_bars: 30, - embargo_bars: 5, - step_bars: 30, - min_train_samples: 20, - }, - num_permutations: 100, - num_trials: 1, - seed: 42, - }; - - let harness = ValidationHarness::new(harness_config); - - // 4. Run validation - let report = harness - .validate(&mut strategy, &data) - .unwrap_or_else(|e| panic!("validate failed: {}", e)); - - // 5. Assertions - assert_eq!( - report.strategy_name, "DQN", - "Strategy name should be 'DQN', got '{}'", - report.strategy_name - ); - - assert!( - report.num_folds >= 2, - "Expected at least 2 folds, got {}", - report.num_folds - ); - - assert!( - report.aggregate_sharpe.is_finite(), - "Aggregate Sharpe should be finite, got {}", - report.aggregate_sharpe - ); - - // DSR p-value in [0, 1] - assert!( - (0.0..=1.0).contains(&report.dsr.pvalue), - "DSR p-value out of range [0,1]: {}", - report.dsr.pvalue - ); - - // PBO in [0, 1] - assert!( - (0.0..=1.0).contains(&report.pbo.pbo), - "PBO out of range [0,1]: {}", - report.pbo.pbo - ); - - // Permutation p-value in [0, 1] - assert!( - (0.0..=1.0).contains(&report.permutation.pvalue), - "Permutation p-value out of range [0,1]: {}", - report.permutation.pvalue - ); - - // Per-regime metrics should not be empty - assert!( - !report.per_regime_metrics.is_empty(), - "per_regime_metrics should not be empty" - ); - - // Print summary - info!( - strategy = %report.strategy_name, - folds = report.num_folds, - aggregate_sharpe = report.aggregate_sharpe, - dsr_pvalue = report.dsr.pvalue, - pbo = report.pbo.pbo, - permutation_pvalue = report.permutation.pvalue, - verdict = %report.verdict, - regimes = report.per_regime_metrics.len(), - "Validation Report Summary" - ); - for (regime, metrics) in &report.per_regime_metrics { - info!( - regime = ?regime, - sharpe = metrics.sharpe, - bars = metrics.num_bars, - win_rate = metrics.win_rate, - avg_return = metrics.avg_return, - "Regime metrics" - ); - } -} - -// --------------------------------------------------------------------------- -// Test 2: Walk-forward split standalone -// --------------------------------------------------------------------------- - -#[test] -fn test_walk_forward_split_standalone() { - let config = WalkForwardConfig { - train_bars: 50, - test_bars: 20, - embargo_bars: 5, - step_bars: 20, - min_train_samples: 20, - }; - - let folds = walk_forward_split(200, &config); - - assert!( - !folds.is_empty(), - "Expected at least one fold from 200 bars" - ); - - for fold in &folds { - // test_range.end must not exceed total bars - assert!( - fold.test_range.end <= 200, - "Fold {} test_range.end ({}) exceeds 200", - fold.fold_index, - fold.test_range.end - ); - - // train.end <= embargo.start (they should be equal by construction) - assert!( - fold.train_range.end <= fold.embargo_range.start, - "Fold {} train.end ({}) > embargo.start ({})", - fold.fold_index, - fold.train_range.end, - fold.embargo_range.start - ); - - // embargo.end <= test.start (they should be equal by construction) - assert!( - fold.embargo_range.end <= fold.test_range.start, - "Fold {} embargo.end ({}) > test.start ({})", - fold.fold_index, - fold.embargo_range.end, - fold.test_range.start - ); - } - - info!(num_folds = folds.len(), total_bars = 200, "Walk-forward split complete"); - for fold in &folds { - info!( - fold = fold.fold_index, - train = ?fold.train_range, - embargo = ?fold.embargo_range, - test = ?fold.test_range, - "Fold ranges" - ); - } -} - -// --------------------------------------------------------------------------- -// Test 3: DSR and PBO standalone -// --------------------------------------------------------------------------- - -#[test] -fn test_dsr_and_pbo_standalone() { - // --- DSR --- - // High Sharpe (2.5) with few trials (5) should be significant -> p < 0.1 - let dsr = deflated_sharpe_ratio(2.5, 5, 0.5, 0.1, 3.5, 500); - assert!( - dsr.pvalue < 0.1, - "DSR: SR=2.5 with 5 trials should have p < 0.1, got {:.4}", - dsr.pvalue - ); - assert!( - dsr.observed_sharpe.is_finite(), - "DSR observed_sharpe should be finite" - ); - assert!( - dsr.deflated_sharpe.is_finite(), - "DSR deflated_sharpe should be finite" - ); - assert!( - dsr.sharpe_std_error.is_finite() && dsr.sharpe_std_error >= 0.0, - "DSR sharpe_std_error should be non-negative and finite, got {}", - dsr.sharpe_std_error - ); - - info!(observed_sharpe = dsr.observed_sharpe, deflated_sharpe = dsr.deflated_sharpe, pvalue = dsr.pvalue, "DSR result"); - - // --- PBO --- - // Consistent positive Sharpes across 8 folds -> should have num_combinations > 0 - let sharpes = vec![1.0, 1.2, 0.8, 1.1, 0.9, 1.3, 1.0, 0.95]; - let pbo = probability_of_backtest_overfitting(&sharpes); - assert!( - pbo.num_combinations > 0, - "PBO should have evaluated combinations, got {}", - pbo.num_combinations - ); - assert!( - (0.0..=1.0).contains(&pbo.pbo), - "PBO value should be in [0,1], got {}", - pbo.pbo - ); - assert!( - !pbo.logit_distribution.is_empty(), - "PBO logit_distribution should not be empty" - ); - - info!(pbo = pbo.pbo, num_combinations = pbo.num_combinations, logit_entries = pbo.logit_distribution.len(), "PBO result"); -} diff --git a/crates/ml/tests/validation_real_data_test.rs b/crates/ml/tests/validation_real_data_test.rs deleted file mode 100644 index 7abffb650..000000000 --- a/crates/ml/tests/validation_real_data_test.rs +++ /dev/null @@ -1,312 +0,0 @@ -#![allow( - clippy::assertions_on_constants, - clippy::assertions_on_result_states, - clippy::clone_on_copy, - clippy::decimal_literal_representation, - clippy::doc_markdown, - clippy::empty_line_after_doc_comments, - clippy::field_reassign_with_default, - clippy::get_unwrap, - clippy::identity_op, - clippy::inconsistent_digit_grouping, - clippy::indexing_slicing, - clippy::integer_division, - clippy::len_zero, - clippy::let_underscore_must_use, - clippy::manual_div_ceil, - clippy::manual_let_else, - clippy::manual_range_contains, - clippy::modulo_arithmetic, - clippy::needless_range_loop, - clippy::non_ascii_literal, - clippy::redundant_clone, - clippy::shadow_reuse, - clippy::shadow_same, - clippy::shadow_unrelated, - clippy::single_match_else, - clippy::str_to_string, - clippy::string_slice, - clippy::tests_outside_test_module, - clippy::too_many_lines, - clippy::unnecessary_wraps, - clippy::unseparated_literal_suffix, - clippy::use_debug, - clippy::useless_vec, - clippy::wildcard_enum_match_arm, - clippy::else_if_without_else, - clippy::expect_used, - clippy::missing_const_for_fn, - clippy::similar_names, - clippy::type_complexity, - clippy::collapsible_else_if, - clippy::doc_lazy_continuation, - clippy::items_after_test_module, - clippy::map_clone, - clippy::multiple_unsafe_ops_per_block, - clippy::unwrap_or_default, - clippy::assign_op_pattern, - clippy::needless_borrow, - clippy::println_empty_string, - clippy::unnecessary_cast, - clippy::used_underscore_binding, - clippy::create_dir, - clippy::implicit_saturating_sub, - clippy::exit, - clippy::expect_fun_call, - clippy::too_many_arguments, - clippy::unnecessary_map_or, - clippy::unwrap_used, - dead_code, - unused_imports, - unused_variables, - clippy::cloned_ref_to_slice_refs, - clippy::neg_multiply, - clippy::while_let_loop, - clippy::bool_assert_comparison, - clippy::excessive_precision, - clippy::trivially_copy_pass_by_ref, - clippy::op_ref, - clippy::redundant_closure, - clippy::unnecessary_lazy_evaluations, - clippy::if_then_some_else_none, - clippy::unnecessary_to_owned, - clippy::single_component_path_imports, -)] -//! Real-data validation test: runs the full ValidationHarness on actual Databento 6E.FUT data. -//! -//! This test loads 1-minute OHLCV bars from a DBN file, extracts 15-dimensional -//! features (5 OHLCV + 10 technical indicators), builds a TimeSeriesData, -//! wraps a DQN agent via DqnStrategy, and runs walk-forward validation with -//! DSR, PBO, permutation tests, and per-regime breakdown. -//! -//! Requires: `test_data/real/databento/6E.FUT_ohlcv-1m_*.dbn` to exist. -//! Run with: `SQLX_OFFLINE=true cargo test --manifest-path ml/Cargo.toml --test validation_real_data_test -- --nocapture` - -use chrono::{DateTime, Utc}; -use ml::dqn::DQNConfig; -use ml::data_loader::RealDataLoader; -use ml::validation::{ - DqnStrategy, TimeSeriesData, ValidationHarness, ValidationHarnessConfig, WalkForwardConfig, -}; -use tracing::info; - -/// Build a 15-dimensional feature vector per bar from RealDataLoader output. -/// -/// Features: -/// 0-4: Normalized OHLCV (open, high, low, close, volume) -/// 5: RSI(14) -/// 6-7: EMA fast(12), EMA slow(26) -/// 8-10: MACD (line, signal, histogram = line - signal) -/// 11-13: Bollinger Bands (upper, middle, lower) -/// 14: ATR(14) -fn build_features_from_loader(loader: &RealDataLoader, bars: &[ml::types::OHLCVBar]) -> Vec> { - let feat_matrix = loader - .extract_features(bars) - .expect("extract_features failed"); - - let indicators = loader - .calculate_indicators(bars) - .expect("calculate_indicators failed"); - - let n = bars.len(); - let mut features = Vec::with_capacity(n); - - for i in 0..n { - let mut row = Vec::with_capacity(16); - - // 0-4: normalized OHLCV - if let Some(price_row) = feat_matrix.prices.get(i) { - row.extend_from_slice(price_row); - } else { - row.extend_from_slice(&[0.0_f32; 5]); - } - - // 5: RSI - row.push(indicators.rsi.get(i).copied().unwrap_or(50.0) / 100.0); // normalize to 0-1 - - // 6-7: EMA fast, slow (normalize relative to close) - let close = bars.get(i).map(|b| b.close as f32).unwrap_or(1.0); - let denom = if close.abs() > 1e-10 { close } else { 1.0 }; - row.push(indicators.ema_fast.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.ema_slow.get(i).copied().unwrap_or(0.0) / denom); - - // 8-10: MACD line, signal, histogram (already small values) - let macd_line = indicators.macd.get(i).copied().unwrap_or(0.0); - let macd_signal = indicators.macd_signal.get(i).copied().unwrap_or(0.0); - row.push(macd_line); - row.push(macd_signal); - row.push(macd_line - macd_signal); // histogram - - // 11-13: Bollinger Bands (normalized relative to close) - row.push(indicators.bb_upper.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.bb_middle.get(i).copied().unwrap_or(0.0) / denom); - row.push(indicators.bb_lower.get(i).copied().unwrap_or(0.0) / denom); - - // 14: ATR (as fraction of close) - row.push(indicators.atr.get(i).copied().unwrap_or(0.0) / denom); - - // 15: zero-pad to 16-dim (tensor core alignment, multiple of 8) - row.push(0.0_f32); - - features.push(row); - } - - features -} - -fn make_dqn_config() -> DQNConfig { - let mut config = DQNConfig::default(); - config.state_dim = 16; // 5 OHLCV + 10 indicators + 1 zero-pad (aligned to 8) - config.num_actions = 3; // short / flat / long - config.hidden_dims = vec![64, 32]; - config.batch_size = 16; - config.min_replay_size = 16; - config.warmup_steps = 0; - config.use_iqn = false; - config.epsilon_start = 0.3; - config.epsilon_end = 0.01; - config -} - -/// Full walk-forward validation on real 6E.FUT minute-bar data. -/// -/// Loads ~30 days of 1-minute OHLCV, extracts 15-dim features, -/// runs walk-forward with embargo, and prints the complete -/// ValidationReport including DSR, PBO, permutation test, and per-regime metrics. -#[tokio::test] -async fn test_validation_on_real_6e_data() { - // 1. Load real data (auto-detect workspace root) - let mut loader = RealDataLoader::new_from_workspace() - .expect("Failed to find workspace root — run from foxhunt repo"); - let bars = loader - .load_symbol_data("6E.FUT") - .await - .expect("Failed to load 6E.FUT data — check test_data/real/databento/ exists"); - - info!(bars = bars.len(), "Loaded bars for 6E.FUT"); - assert!( - bars.len() > 500, - "Expected at least 500 bars from 6E.FUT DBN, got {}", - bars.len() - ); - - // Print data range - if let (Some(first), Some(last)) = (bars.first(), bars.last()) { - info!( - start = %first.timestamp, - end = %last.timestamp, - open = first.open, - high = first.high, - low = first.low, - close = first.close, - volume = first.volume, - "Data range and first bar" - ); - } - - // 2. Build features - let features = build_features_from_loader(&loader, &bars); - assert_eq!(features.len(), bars.len()); - - // Verify features are finite - for (i, row) in features.iter().enumerate() { - assert_eq!(row.len(), 15, "Bar {} has {} features, expected 15", i, row.len()); - for (j, v) in row.iter().enumerate() { - assert!( - v.is_finite(), - "Feature [{},{}] is not finite: {}", - i, - j, - v - ); - } - } - info!(bars = features.len(), dims = 15, "Features validated (all finite)"); - - // 3. Build TimeSeriesData - let timestamps: Vec> = bars.iter().map(|b| b.timestamp).collect(); - let prices: Vec = bars.iter().map(|b| b.close).collect(); - - let data = TimeSeriesData::new(timestamps, features, prices) - .expect("Failed to create TimeSeriesData"); - - info!(bars = data.len(), returns = data.returns.len(), "TimeSeriesData created"); - - // 4. Create DQN strategy - let config = make_dqn_config(); - let mut strategy = - DqnStrategy::new(config).expect("Failed to create DqnStrategy"); - - // 5. Configure walk-forward harness - // For 1-min bars: 500 bars ≈ 8 hours of training data - // test = 100 bars ≈ 1.5 hours - // embargo = 20 bars ≈ 20 minutes - let num_bars = data.len(); - let train_bars = (num_bars / 5).max(200); - let test_bars = (num_bars / 20).max(50); - - let harness_config = ValidationHarnessConfig { - wf_config: WalkForwardConfig { - train_bars, - test_bars, - embargo_bars: 20, - step_bars: test_bars, - min_train_samples: 100, - }, - num_permutations: 500, // Reduced for speed; 10k for final - num_trials: 1, - seed: 42, - }; - - info!(train_bars, test_bars, embargo_bars = 20, step_bars = test_bars, total_bars = num_bars, "Walk-Forward Config"); - - let harness = ValidationHarness::new(harness_config); - - // 6. Run validation - info!("Running validation..."); - let report = harness - .validate(&mut strategy, &data) - .expect("Validation harness failed"); - - // 7. Log full report - info!( - strategy = %report.strategy_name, - folds = report.num_folds, - aggregate_sharpe = report.aggregate_sharpe, - dsr_observed = report.dsr.observed_sharpe, - dsr_expected_max = report.dsr.expected_max_sharpe, - dsr_se = report.dsr.sharpe_std_error, - dsr_statistic = report.dsr.deflated_sharpe, - dsr_pvalue = report.dsr.pvalue, - pbo = report.pbo.pbo, - pbo_combinations = report.pbo.num_combinations, - perm_observed = report.permutation.observed_sharpe, - perm_null_mean = report.permutation.null_mean, - perm_null_std = report.permutation.null_std, - perm_pvalue = report.permutation.pvalue, - perm_count = report.permutation.num_permutations, - verdict = %report.verdict, - "Validation Report" - ); - for (i, sr) in report.per_fold_sharpes.iter().enumerate() { - info!(fold = i, sharpe = sr, "Per-fold Sharpe"); - } - for (regime, m) in &report.per_regime_metrics { - info!( - regime = ?regime, - sharpe = m.sharpe, - bars = m.num_bars, - win_rate_pct = m.win_rate * 100.0, - avg_return = m.avg_return, - "Per-regime metrics" - ); - } - - // 8. Structural assertions (not outcome-dependent) - assert!(report.num_folds >= 2, "Need at least 2 folds"); - assert!(report.aggregate_sharpe.is_finite()); - assert!((0.0..=1.0).contains(&report.dsr.pvalue)); - assert!((0.0..=1.0).contains(&report.pbo.pbo)); - assert!((0.0..=1.0).contains(&report.permutation.pvalue)); - assert!(!report.per_regime_metrics.is_empty()); -}