feat(M5.3): Add training corpus export infrastructure for verl
M5.3 — Training Corpus Export (verl format):
- Trajectory struct: trajectory_id, turns[], r_exit, r_format, r_outcome
- TrajectoryTurn: t, prompt, response, r_update, parsed
- CorpusStats: total_trajectories, total_turns, positive/negative split,
r_format pass rate, r_exit distribution
Reward computation:
- r_update_t: +1 if label matches U_t, -1 if mismatch (per turn)
- r_exit: 0 if exit == last_evidence_t, -0.75 if earlier, -0.5 if later
- r_format: 1.0 if all turns parsed, 0.0 if any unparsed (strict)
- r_outcome: null (no answer correctness signal available)
Files created:
crates/mem-core/src/trajectory.rs (280 LOC)
- Trajectory construction and reward calculation
- CorpusStats aggregation from trajectories
- Serialization for JSONL output
tests/it_export.rs (280 LOC, 12 tests)
- a1: Trajectory grouping by run
- a2: r_update signs correct
- a3: r_format strict (any unparsed = 0)
- a4: r_exit distribution (perfect/early/late)
- a5: Prompts are exact byte recordings
- a6: CorpusStats aggregation
- a7: r_outcome null
- a8: Turn ordering preserved
- a9: Multiple trajectories
- a10: Serde roundtrip
- a11: CorpusStats structure complete
- a12: Mixed exit rewards
Unit tests:
- crates/mem-core/src/trajectory.rs: 8/8 passing
Integration tests:
- tests/it_export.rs: 12/12 passing
Architecture:
Log + Labels → Trajectories → JSONL for verl
Each trajectory = one run with multiple turns
Per-turn rewards enable trajectory-level loss + turn-level loss
Blocks: M5.4 (vLLM setup), M5.5 (verl training)
Depends: M5.1 ✓, M5.2 ✓
This commit is contained in:
@@ -6,6 +6,7 @@ pub mod gate_parser;
|
||||
pub mod gated_loop;
|
||||
pub mod query_executor;
|
||||
pub mod shingle;
|
||||
pub mod trajectory;
|
||||
|
||||
pub use gate_parser::{GateResponse, ParseError, parse_gate_response};
|
||||
|
||||
@@ -19,3 +20,4 @@ pub use lesson::{
|
||||
pub use query::{Query, QuerySet, SynthesisQuery};
|
||||
pub use prompt::PromptBuilder;
|
||||
pub use shingle::{jaccard_similarity, matches_artifact, Shingle, ShingleConfig};
|
||||
pub use trajectory::{Trajectory, TrajectoryTurn, CorpusStats};
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
// M5.3 — Training Corpus Export
|
||||
//
|
||||
// Converts log + labels into trajectories for verl training.
|
||||
// A trajectory = one run, multiple turns, with per-turn rewards.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// A single turn in a trajectory
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TrajectoryTurn {
|
||||
/// Turn number
|
||||
pub t: usize,
|
||||
|
||||
/// Exact prompt bytes sent to model
|
||||
pub prompt: String,
|
||||
|
||||
/// Exact response bytes from model
|
||||
pub response: String,
|
||||
|
||||
/// Per-turn reward: +1 if U_t matches label, -1 if mismatch
|
||||
pub r_update: i32,
|
||||
|
||||
/// True if turn parsed successfully
|
||||
pub parsed: bool,
|
||||
}
|
||||
|
||||
/// A trajectory = one episode/run with multiple turns
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct Trajectory {
|
||||
/// Unique identifier for this run
|
||||
pub trajectory_id: String,
|
||||
|
||||
/// All turns in order
|
||||
pub turns: Vec<TrajectoryTurn>,
|
||||
|
||||
/// Exit reward: 0 if exit == last_evidence, -0.75 if earlier, -0.5 if later
|
||||
pub r_exit: f32,
|
||||
|
||||
/// Format reward: 1.0 if all turns parsed, 0.0 if any unparsed
|
||||
pub r_format: f32,
|
||||
|
||||
/// Outcome reward: null (no correctness signal available)
|
||||
pub r_outcome: Option<f32>,
|
||||
}
|
||||
|
||||
impl Trajectory {
|
||||
pub fn new(trajectory_id: String) -> Self {
|
||||
Self {
|
||||
trajectory_id,
|
||||
turns: Vec::new(),
|
||||
r_exit: 0.0,
|
||||
r_format: 1.0,
|
||||
r_outcome: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a turn to the trajectory
|
||||
pub fn add_turn(&mut self, t: usize, prompt: String, response: String, r_update: i32, parsed: bool) {
|
||||
self.turns.push(TrajectoryTurn {
|
||||
t,
|
||||
prompt,
|
||||
response,
|
||||
r_update,
|
||||
parsed,
|
||||
});
|
||||
|
||||
// Update r_format: 0 if any turn is unparsed
|
||||
if !parsed {
|
||||
self.r_format = 0.0;
|
||||
}
|
||||
}
|
||||
|
||||
/// Set exit reward based on when evidence appeared
|
||||
pub fn set_exit_reward(&mut self, exit_at: usize, last_evidence_at: usize) {
|
||||
self.r_exit = if exit_at == last_evidence_at {
|
||||
0.0 // Perfect: exited right after finding evidence
|
||||
} else if exit_at < last_evidence_at {
|
||||
-0.75 // Bad: exited before finding evidence
|
||||
} else {
|
||||
-0.5 // Moderate: exited after finding evidence (continued searching)
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
/// Summary of corpus statistics
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CorpusStats {
|
||||
pub total_trajectories: usize,
|
||||
pub total_turns: usize,
|
||||
pub positive_r_update: usize,
|
||||
pub negative_r_update: usize,
|
||||
pub r_format_pass_rate: f32,
|
||||
pub r_exit_distribution: HashMap<String, usize>,
|
||||
}
|
||||
|
||||
impl CorpusStats {
|
||||
pub fn from_trajectories(trajectories: &[Trajectory]) -> Self {
|
||||
let total_trajectories = trajectories.len();
|
||||
let mut total_turns = 0;
|
||||
let mut positive_r_update = 0;
|
||||
let mut negative_r_update = 0;
|
||||
let mut r_format_passes = 0;
|
||||
let mut r_exit_dist = HashMap::new();
|
||||
|
||||
for traj in trajectories {
|
||||
total_turns += traj.turns.len();
|
||||
|
||||
for turn in &traj.turns {
|
||||
match turn.r_update {
|
||||
1 => positive_r_update += 1,
|
||||
-1 => negative_r_update += 1,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if traj.r_format > 0.99 {
|
||||
r_format_passes += 1;
|
||||
}
|
||||
|
||||
let exit_key = if traj.r_exit > -0.1 {
|
||||
"perfect".to_string()
|
||||
} else if traj.r_exit < -0.6 {
|
||||
"too_early".to_string()
|
||||
} else {
|
||||
"too_late".to_string()
|
||||
};
|
||||
|
||||
*r_exit_dist.entry(exit_key).or_insert(0) += 1;
|
||||
}
|
||||
|
||||
let r_format_pass_rate = if total_trajectories > 0 {
|
||||
r_format_passes as f32 / total_trajectories as f32
|
||||
} else {
|
||||
0.0
|
||||
};
|
||||
|
||||
Self {
|
||||
total_trajectories,
|
||||
total_turns,
|
||||
positive_r_update,
|
||||
negative_r_update,
|
||||
r_format_pass_rate,
|
||||
r_exit_distribution: r_exit_dist,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_trajectory_creation() {
|
||||
let traj = Trajectory::new("run_001".to_string());
|
||||
assert_eq!(traj.trajectory_id, "run_001");
|
||||
assert_eq!(traj.turns.len(), 0);
|
||||
assert_eq!(traj.r_format, 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_add_turn() {
|
||||
let mut traj = Trajectory::new("run_001".to_string());
|
||||
traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||||
|
||||
assert_eq!(traj.turns.len(), 1);
|
||||
assert_eq!(traj.turns[0].t, 1);
|
||||
assert_eq!(traj.turns[0].r_update, 1);
|
||||
assert!(traj.turns[0].parsed);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_r_format_unparsed() {
|
||||
let mut traj = Trajectory::new("run_001".to_string());
|
||||
traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||||
assert_eq!(traj.r_format, 1.0);
|
||||
|
||||
traj.add_turn(2, "Q2".to_string(), "malformed".to_string(), -1, false);
|
||||
assert_eq!(traj.r_format, 0.0, "Any unparsed turn sets r_format to 0");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exit_reward_perfect() {
|
||||
let mut traj = Trajectory::new("run_001".to_string());
|
||||
traj.set_exit_reward(5, 5);
|
||||
assert_eq!(traj.r_exit, 0.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exit_reward_early() {
|
||||
let mut traj = Trajectory::new("run_001".to_string());
|
||||
traj.set_exit_reward(3, 5);
|
||||
assert_eq!(traj.r_exit, -0.75);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exit_reward_late() {
|
||||
let mut traj = Trajectory::new("run_001".to_string());
|
||||
traj.set_exit_reward(7, 5);
|
||||
assert_eq!(traj.r_exit, -0.5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_corpus_stats() {
|
||||
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(), "B".to_string(), 1, true);
|
||||
|
||||
let stats = CorpusStats::from_trajectories(&[traj1, traj2]);
|
||||
|
||||
assert_eq!(stats.total_trajectories, 2);
|
||||
assert_eq!(stats.total_turns, 3);
|
||||
assert_eq!(stats.positive_r_update, 2);
|
||||
assert_eq!(stats.negative_r_update, 1);
|
||||
assert_eq!(stats.r_format_pass_rate, 1.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trajectory_serde() {
|
||||
let mut traj = Trajectory::new("run_001".to_string());
|
||||
traj.add_turn(1, "Q".to_string(), "A".to_string(), 1, true);
|
||||
|
||||
let json = serde_json::to_string(&traj).expect("Should serialize");
|
||||
let deserialized: Trajectory = serde_json::from_str(&json)
|
||||
.expect("Should deserialize");
|
||||
|
||||
assert_eq!(deserialized.trajectory_id, "run_001");
|
||||
assert_eq!(deserialized.turns.len(), 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user