200 lines
6.2 KiB
Rust
200 lines
6.2 KiB
Rust
use mem_core::{Trajectory, CorpusStats};
|
|||
|
|
|
||
|
|
/// M5.3 Integration Tests — Training Corpus Export
|
||
|
|
///
|
||
|
|
/// Verifies export to verl format:
|
||
|
|
/// - Trajectories group turns by run
|
||
|
|
/// - Rewards computed correctly
|
||
|
|
/// - Prompts are exact byte recordings
|
||
|
|
/// - Format strict (any unparsed turn = r_format 0)
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a1_trajectory_grouping() {
|
||
|
|
let mut traj = Trajectory::new("run_001".to_string());
|
||
|
|
traj.add_turn(1, "Q1".to_string(), "A1".to_string(), 1, true);
|
||
|
|
traj.add_turn(2, "Q2".to_string(), "A2".to_string(), -1, true);
|
||
|
|
traj.add_turn(3, "Q3".to_string(), "A3".to_string(), 1, true);
|
||
|
|
|
||
|
|
assert_eq!(traj.turns.len(), 3);
|
||
|
|
assert_eq!(traj.turns[0].t, 1);
|
||
|
|
assert_eq!(traj.turns[1].t, 2);
|
||
|
|
assert_eq!(traj.turns[2].t, 3);
|
||
|
|
|
||
|
|
// All turns in ascending order
|
||
|
|
for i in 1..traj.turns.len() {
|
||
|
|
assert!(traj.turns[i].t > traj.turns[i - 1].t);
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a2_r_update_signs() {
|
||
|
|
let mut traj = Trajectory::new("run_001".to_string());
|
||
|
|
|
||
|
|
// Correct prediction
|
||
|
|
traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||
|
|
assert_eq!(traj.turns[0].r_update, 1);
|
||
|
|
|
||
|
|
// Incorrect prediction
|
||
|
|
traj.add_turn(2, "Q".to_string(), "B".to_string(), -1, true);
|
||
|
|
assert_eq!(traj.turns[1].r_update, -1);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a3_r_format_strict() {
|
||
|
|
let mut traj = Trajectory::new("run_001".to_string());
|
||
|
|
|
||
|
|
// First turn OK
|
||
|
|
traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||
|
|
assert_eq!(traj.r_format, 1.0);
|
||
|
|
|
||
|
|
// Second turn unparsed
|
||
|
|
traj.add_turn(2, "Q".to_string(), "malformed_response".to_string(), -1, false);
|
||
|
|
|
||
|
|
// ENTIRE trajectory marked as unparsed
|
||
|
|
assert_eq!(traj.r_format, 0.0, "Any unparsed turn makes whole trajectory unparsed");
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a4_r_exit_distribution() {
|
||
|
|
let mut perfect = Trajectory::new("run_001".to_string());
|
||
|
|
perfect.set_exit_reward(5, 5);
|
||
|
|
assert_eq!(perfect.r_exit, 0.0);
|
||
|
|
|
||
|
|
let mut early = Trajectory::new("run_002".to_string());
|
||
|
|
early.set_exit_reward(3, 5);
|
||
|
|
assert_eq!(early.r_exit, -0.75);
|
||
|
|
|
||
|
|
let mut late = Trajectory::new("run_003".to_string());
|
||
|
|
late.set_exit_reward(7, 5);
|
||
|
|
assert_eq!(late.r_exit, -0.5);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a5_prompt_exact_bytes() {
|
||
|
|
let prompt_bytes = "Question: what is X?\n\nContext: Y is Z".to_string();
|
||
|
|
let response_bytes = "Answer: X is Y".to_string();
|
||
|
|
|
||
|
|
let mut traj = Trajectory::new("run_001".to_string());
|
||
|
|
traj.add_turn(1, prompt_bytes.clone(), response_bytes.clone(), 1, true);
|
||
|
|
|
||
|
|
// Prompts should be exact recordings, not re-assembled
|
||
|
|
assert_eq!(traj.turns[0].prompt, prompt_bytes);
|
||
|
|
assert_eq!(traj.turns[0].response, response_bytes);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a6_corpus_stats_aggregation() {
|
||
|
|
let mut traj1 = Trajectory::new("run_001".to_string());
|
||
|
|
traj1.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||
|
|
traj1.add_turn(2, "Q".to_string(), "A".to_string(), -1, true);
|
||
|
|
|
||
|
|
let mut traj2 = Trajectory::new("run_002".to_string());
|
||
|
|
traj2.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||
|
|
|
||
|
|
let mut traj3 = Trajectory::new("run_003".to_string());
|
||
|
|
traj3.add_turn(1, "Q".to_string(), "A".to_string(), 1, false); // Unparsed
|
||
|
|
|
||
|
|
let stats = CorpusStats::from_trajectories(&[traj1, traj2, traj3]);
|
||
|
|
|
||
|
|
assert_eq!(stats.total_trajectories, 3);
|
||
|
|
assert_eq!(stats.total_turns, 4);
|
||
|
|
assert_eq!(stats.positive_r_update, 3);
|
||
|
|
assert_eq!(stats.negative_r_update, 1);
|
||
|
|
|
||
|
|
// 2 out of 3 have all parsed turns
|
||
|
|
assert!((stats.r_format_pass_rate - 2.0/3.0).abs() < 0.01);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a7_r_outcome_null() {
|
||
|
|
let traj = Trajectory::new("run_001".to_string());
|
||
|
|
|
||
|
|
// Should have no answer correctness signal
|
||
|
|
assert!(traj.r_outcome.is_none(), "r_outcome should be null (no answer-correctness signal)");
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a8_trajectory_ordering() {
|
||
|
|
let mut traj = Trajectory::new("run_001".to_string());
|
||
|
|
|
||
|
|
// Add in order
|
||
|
|
for t in 1..=10 {
|
||
|
|
traj.add_turn(t, format!("Q{}", t), format!("A{}", t), if t % 2 == 0 { 1 } else { -1 }, true);
|
||
|
|
}
|
||
|
|
|
||
|
|
// Check order preserved
|
||
|
|
for (i, turn) in traj.turns.iter().enumerate() {
|
||
|
|
assert_eq!(turn.t, i + 1);
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a9_multiple_trajectories() {
|
||
|
|
let mut trajs = Vec::new();
|
||
|
|
|
||
|
|
for run_id in 0..5 {
|
||
|
|
let mut traj = Trajectory::new(format!("run_{:03}", run_id));
|
||
|
|
for t in 1..=3 {
|
||
|
|
traj.add_turn(t, format!("Q{}", t), format!("A{}", t), 1, true);
|
||
|
|
}
|
||
|
|
trajs.push(traj);
|
||
|
|
}
|
||
|
|
|
||
|
|
assert_eq!(trajs.len(), 5);
|
||
|
|
assert_eq!(trajs[0].turns.len(), 3);
|
||
|
|
assert_eq!(trajs[0].trajectory_id, "run_000");
|
||
|
|
assert_eq!(trajs[4].trajectory_id, "run_004");
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a10_trajectory_serde_roundtrip() {
|
||
|
|
let mut traj = Trajectory::new("run_001".to_string());
|
||
|
|
traj.add_turn(1, "prompt1".to_string(), "response1".to_string(), 1, true);
|
||
|
|
traj.add_turn(2, "prompt2".to_string(), "response2".to_string(), -1, true);
|
||
|
|
traj.set_exit_reward(2, 1);
|
||
|
|
|
||
|
|
let json = serde_json::to_string(&traj).expect("Should serialize");
|
||
|
|
let restored: Trajectory = serde_json::from_str(&json)
|
||
|
|
.expect("Should deserialize");
|
||
|
|
|
||
|
|
assert_eq!(restored.trajectory_id, "run_001");
|
||
|
|
assert_eq!(restored.turns.len(), 2);
|
||
|
|
assert_eq!(restored.r_exit, -0.5);
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a11_corpus_stats_structure() {
|
||
|
|
let traj = Trajectory::new("run_001".to_string());
|
||
|
|
let stats = CorpusStats::from_trajectories(&[traj]);
|
||
|
|
|
||
|
|
// All fields present
|
||
|
|
assert!(stats.total_trajectories >= 0);
|
||
|
|
assert!(stats.total_turns >= 0);
|
||
|
|
assert!(stats.positive_r_update >= 0);
|
||
|
|
assert!(stats.negative_r_update >= 0);
|
||
|
|
assert!(stats.r_format_pass_rate >= 0.0 && stats.r_format_pass_rate <= 1.0);
|
||
|
|
assert!(!stats.r_exit_distribution.is_empty());
|
||
|
|
}
|
||
|
|
|
||
|
|
#[test]
|
||
|
|
fn a12_mixed_exit_rewards() {
|
||
|
|
let mut trajs = Vec::new();
|
||
|
|
|
||
|
|
let mut perfect = Trajectory::new("perfect".to_string());
|
||
|
|
perfect.set_exit_reward(3, 3);
|
||
|
|
trajs.push(perfect);
|
||
|
|
|
||
|
|
let mut early = Trajectory::new("early".to_string());
|
||
|
|
early.set_exit_reward(1, 3);
|
||
|
|
trajs.push(early);
|
||
|
|
|
||
|
|
let mut late = Trajectory::new("late".to_string());
|
||
|
|
late.set_exit_reward(5, 3);
|
||
|
|
trajs.push(late);
|
||
|
|
|
||
|
|
let stats = CorpusStats::from_trajectories(&trajs);
|
||
|
|
|
||
|
|
assert_eq!(stats.total_trajectories, 3);
|
||
|
|
assert_eq!(stats.r_exit_distribution.len(), 3); // Should have all 3 types
|
||
|
|
}
|