//! Multi-episode sequence batching tests //! //! These tests verify that `to_sequences()` correctly handles episode boundaries //! and prevents hidden state contamination between unrelated episodes. //! //! This addresses the critical flaw identified by both GPT-5 Codex and Gemini 2.5 Pro. use ml::ppo::trajectories::{Trajectory, TrajectoryBatch, TrajectoryStep}; use ml::dqn::TradingAction; /// 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: TradingAction::Hold, log_prob: -1.0, value: 0.5, reward: 1.0, done, }); } trajectory } #[test] #[should_panic(expected = "sequence_length must be greater than 0")] fn test_sequence_length_zero_validation() { // Test that sequence_length == 0 is rejected (prevents infinite loop) 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); // This should panic let _sequences = batch.to_sequences(0); } #[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_sequences(8); // Verify: No sequence contains steps from multiple episodes for seq in &sequences { let done_indices: Vec = seq.dones .iter() .enumerate() .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: TradingAction::Buy, 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: TradingAction::Sell, 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_sequences(8); // Verify: Episode 1 states (100-104) never mixed with Episode 2 states (200-206) for seq in &sequences { let first_state_value = seq.states[0][0]; // Check that all states in this sequence are from the same episode for state in &seq.states { let state_value = state[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_sequences(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.dones.iter().filter(|&&d| d).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: TradingAction::Buy, // 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: TradingAction::Sell, // 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_sequences(8); // Verify: Actions within each sequence are consistent for seq in &sequences { if seq.actions.is_empty() { continue; } let first_action = seq.actions[0]; let all_same_action = seq.actions.iter().all(|&a| a == 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.actions ); } } #[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_sequences(8); // Verify: actual_length matches actual data length for seq in &sequences { assert_eq!( seq.states.len(), seq.actual_length, "Mismatch: states.len()={}, actual_length={}", seq.states.len(), seq.actual_length ); assert_eq!(seq.actions.len(), seq.actual_length); assert_eq!(seq.log_probs.len(), seq.actual_length); assert_eq!(seq.values.len(), seq.actual_length); assert_eq!(seq.rewards.len(), seq.actual_length); assert_eq!(seq.dones.len(), seq.actual_length); assert_eq!(seq.advantages.len(), seq.actual_length); assert_eq!(seq.returns.len(), seq.actual_length); } }