The 4-branch DQN (direction x magnitude) had 3 degenerate variants (Short25, Flat, Long25) that all mapped to 0.0 target exposure when direction=Flat, causing 82% Flat collapse. Collapse these into a single Flat variant, giving 7 levels (ShortSmall/Half/Full, Flat, LongSmall/Half/Full) and 63 total factored actions (7x3x3). - ExposureLevel enum: 9 variants -> 7 (add direction/magnitude/from_dir_mag) - FactoredAction: 81 -> 63 total actions, from_index/to_index updated - DQN epsilon-greedy: use from_dir_mag() instead of dir*3+mag indexing - DQN config: num_actions default 9 -> 7 - PPO action space: 45 -> 63 actions, action masking updated - Signal adapter CUDA kernel: 5-bin -> 7-bin exposure aggregation - All tests updated for new variant names and index ranges Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
390 lines
13 KiB
Rust
390 lines
13 KiB
Rust
#![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<usize> = 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::<Vec<_>>()
|
|
);
|
|
}
|
|
}
|
|
|
|
#[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);
|
|
}
|
|
}
|
|
}
|