#![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, )] //! Multi-episode sequence batching tests //! //! These tests verify that `to_sequence_ranges()` correctly handles episode boundaries //! and prevents hidden state contamination between unrelated episodes. //! //! The API was migrated from `to_sequences()` (returning materialized `Sequence` structs) //! to `to_sequence_ranges()` (returning `SequenceRange` index ranges into the parent batch). //! Tests now validate via index-based lookups into the `TrajectoryBatch` fields. use ml::ppo::trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep}; use ml_core::common::action::{ExposureLevel, FactoredAction, OrderType, Urgency}; /// Create a trajectory with specified number of steps fn create_trajectory(num_steps: usize, state_dim: usize) -> Trajectory { let mut trajectory = Trajectory::new(); for i in 0..num_steps { let state = vec![i as f32; state_dim]; let done = i == num_steps - 1; // Last step is done trajectory.add_step(TrajectoryStep { state, action: FactoredAction::new(ExposureLevel::Flat, OrderType::Market, Urgency::Normal), log_prob: -1.0, value: 0.5, reward: 1.0, done, }); } trajectory } #[test] fn test_sequence_length_zero_returns_empty() { // Test that sequence_length == 0 returns empty (API returns Vec, no panic) let episode = create_trajectory(10, 32); let advantages = vec![0.1; 10]; let returns = vec![0.5; 10]; let batch = TrajectoryBatch::from_trajectories(vec![episode], advantages, returns); let sequences = batch.to_sequence_ranges(0); assert!( sequences.is_empty(), "sequence_length=0 should return empty ranges" ); } #[test] fn test_multi_episode_no_contamination() { // Create 3 episodes with different lengths let episode1 = create_trajectory(5, 32); // 5 steps let episode2 = create_trajectory(17, 32); // 17 steps let episode3 = create_trajectory(3, 32); // 3 steps // Total: 25 steps let total_steps = 5 + 17 + 3; let advantages = vec![0.1; total_steps]; let returns = vec![0.5; total_steps]; let batch = TrajectoryBatch::from_trajectories( vec![episode1, episode2, episode3], advantages, returns ); let sequences = batch.to_sequence_ranges(8); // Verify: No sequence contains steps from multiple episodes. // Each sequence range maps into the batch's flattened arrays. // The done flags within a sequence should only appear at the last position // (if the sequence ends an episode). for seq in &sequences { let dones_in_range: Vec<(usize, bool)> = (seq.start..seq.end) .map(|i| (i - seq.start, batch.dones[i])) .collect(); let done_indices: Vec = dones_in_range .iter() .filter(|(_, done)| *done) .map(|(i, _)| *i) .collect(); // If a done flag exists, it should ONLY be at the last element if !done_indices.is_empty() { assert_eq!( done_indices.len(), 1, "Sequence should have at most 1 done flag (found {})", done_indices.len() ); assert_eq!( done_indices[0], seq.actual_length - 1, "Done flag should be at last position (index {}), found at index {}", seq.actual_length - 1, done_indices[0] ); } } } #[test] fn test_episode_boundary_state_values() { // Create 2 episodes with distinct state values to catch contamination let mut episode1 = Trajectory::new(); for i in 0..5 { episode1.add_step(TrajectoryStep { state: vec![100.0 + i as f32; 32], // States: 100, 101, 102, 103, 104 action: FactoredAction::new(ExposureLevel::LongFull, OrderType::Market, Urgency::Normal), log_prob: -1.0, value: 0.5, reward: 1.0, done: i == 4, // Last step done }); } let mut episode2 = Trajectory::new(); for i in 0..7 { episode2.add_step(TrajectoryStep { state: vec![200.0 + i as f32; 32], // States: 200, 201, 202, 203, 204, 205, 206 action: FactoredAction::new(ExposureLevel::ShortSmall, OrderType::Market, Urgency::Normal), log_prob: -1.0, value: 0.5, reward: 1.0, done: i == 6, // Last step done }); } let total_steps = 5 + 7; let advantages = vec![0.1; total_steps]; let returns = vec![0.5; total_steps]; let batch = TrajectoryBatch::from_trajectories( vec![episode1, episode2], advantages, returns ); let sequences = batch.to_sequence_ranges(8); // Verify: Episode 1 states (100-104) never mixed with Episode 2 states (200-206) for seq in &sequences { let first_state_value = batch.states[seq.start][0]; // Check that all states in this sequence are from the same episode for idx in seq.start..seq.end { let state_value = batch.states[idx][0]; if first_state_value >= 100.0 && first_state_value < 200.0 { // Sequence starts in episode 1 assert!( state_value >= 100.0 && state_value < 200.0, "Sequence starting in episode 1 (state {}) contains episode 2 state ({})", first_state_value, state_value ); } else { // Sequence starts in episode 2 assert!( state_value >= 200.0, "Sequence starting in episode 2 (state {}) contains episode 1 state ({})", first_state_value, state_value ); } } } } #[test] fn test_varying_episode_lengths_with_small_sequences() { // Test edge case: sequence_length > some episodes but < others let episode1 = create_trajectory(3, 32); // Shorter than seq_len let episode2 = create_trajectory(20, 32); // Longer than seq_len let episode3 = create_trajectory(5, 32); // Between let total_steps = 3 + 20 + 5; let advantages = vec![0.1; total_steps]; let returns = vec![0.5; total_steps]; let batch = TrajectoryBatch::from_trajectories( vec![episode1, episode2, episode3], advantages, returns ); let sequences = batch.to_sequence_ranges(8); // Verify: Each episode is properly segmented // Episode 1 (3 steps): 1 sequence of length 3 // Episode 2 (20 steps): 3 sequences (8 + 8 + 4) // Episode 3 (5 steps): 1 sequence of length 5 // Total: 5 sequences assert_eq!( sequences.len(), 5, "Expected 5 sequences (1+3+1), got {}", sequences.len() ); // Verify no sequence crosses episode boundaries for seq in &sequences { let done_count = (seq.start..seq.end) .filter(|&i| batch.dones[i]) .count(); assert!( done_count <= 1, "Sequence crosses episode boundary (has {} done flags)", done_count ); } } #[test] fn test_action_continuity_within_sequences() { // Create episodes with different action patterns let mut episode1 = Trajectory::new(); for i in 0..5 { episode1.add_step(TrajectoryStep { state: vec![0.0; 32], action: FactoredAction::new(ExposureLevel::LongFull, OrderType::Market, Urgency::Normal), // All Buy log_prob: -1.0, value: 0.5, reward: 1.0, done: i == 4, }); } let mut episode2 = Trajectory::new(); for i in 0..7 { episode2.add_step(TrajectoryStep { state: vec![0.0; 32], action: FactoredAction::new(ExposureLevel::ShortSmall, OrderType::Market, Urgency::Normal), // All Sell log_prob: -1.0, value: 0.5, reward: 1.0, done: i == 6, }); } let total_steps = 5 + 7; let advantages = vec![0.1; total_steps]; let returns = vec![0.5; total_steps]; let batch = TrajectoryBatch::from_trajectories( vec![episode1, episode2], advantages, returns ); let sequences = batch.to_sequence_ranges(8); // Verify: Actions within each sequence are consistent for seq in &sequences { if seq.start >= seq.end { continue; } let first_action = batch.actions[seq.start]; let all_same_action = (seq.start..seq.end) .all(|i| batch.actions[i] == first_action); // All actions should be the same (Buy or Sell, not mixed) assert!( all_same_action, "Sequence contains mixed actions from different episodes: {:?}", (seq.start..seq.end).map(|i| batch.actions[i]).collect::>() ); } } #[test] fn test_sequence_actual_length_correctness() { // Test that actual_length correctly reflects sequence boundaries let episode1 = create_trajectory(5, 32); let episode2 = create_trajectory(17, 32); let total_steps = 5 + 17; let advantages = vec![0.1; total_steps]; let returns = vec![0.5; total_steps]; let batch = TrajectoryBatch::from_trajectories( vec![episode1, episode2], advantages, returns ); let sequences = batch.to_sequence_ranges(8); // Verify: actual_length matches range span (end - start) for seq in &sequences { let range_len = seq.end - seq.start; assert_eq!( range_len, seq.actual_length, "Mismatch: range span={}, actual_length={}", range_len, seq.actual_length ); // Verify the range doesn't exceed batch bounds assert!( seq.end <= batch.total_steps(), "Sequence end {} exceeds batch total_steps {}", seq.end, batch.total_steps() ); // Verify all referenced fields are accessible for i in seq.start..seq.end { assert!(i < batch.states.len(), "State index {} out of bounds", i); assert!(i < batch.actions.len(), "Action index {} out of bounds", i); assert!(i < batch.log_probs.len(), "Log prob index {} out of bounds", i); assert!(i < batch.values.len(), "Value index {} out of bounds", i); assert!(i < batch.rewards.len(), "Reward index {} out of bounds", i); assert!(i < batch.dones.len(), "Done index {} out of bounds", i); assert!(i < batch.advantages.len(), "Advantage index {} out of bounds", i); assert!(i < batch.returns.len(), "Return index {} out of bounds", i); } } }