Files
foxhunt/crates/ml/tests/ppo_sequence_batching_multi_episode_tests.rs
jgrusewski 9c3d741a08 refactor: restructure repo — crates/, bin/, testing/ layout
Move 17 library crates into crates/, CLI binary into bin/fxt,
consolidate 10 test crates into testing/, split config crate
from deployment config files.

Root directory reduced from 38+ to ~17 directories.
All Cargo.toml paths and build.rs proto refs updated.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-25 11:56:00 +01:00

290 lines
8.9 KiB
Rust

//! 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<usize> = 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);
}
}