//! Curriculum schedule implementations. use std::collections::HashMap; /// Trait for curriculum scheduling pub trait Schedule: Send + Sync { /// Get difficulty threshold at given training step fn get_difficulty_at_step(&self, step: usize) -> f32; } /// Linear curriculum schedule #[derive(Debug, Clone)] pub struct LinearSchedule { initial_difficulty: f32, final_difficulty: f32, total_steps: usize, } impl LinearSchedule { pub fn new(initial_difficulty: f32, final_difficulty: f32, total_steps: usize) -> Self { Self { initial_difficulty, final_difficulty, total_steps, } } } impl Schedule for LinearSchedule { fn get_difficulty_at_step(&self, step: usize) -> f32 { if step >= self.total_steps { return self.final_difficulty; } let progress = step as f32 / self.total_steps as f32; self.initial_difficulty + progress * (self.final_difficulty - self.initial_difficulty) } } /// Exponential curriculum schedule #[derive(Debug, Clone)] pub struct ExponentialSchedule { initial_difficulty: f32, final_difficulty: f32, growth_rate: f32, } impl ExponentialSchedule { pub fn new(initial_difficulty: f32, final_difficulty: f32, growth_rate: f32) -> Self { Self { initial_difficulty, final_difficulty, growth_rate, } } } impl Schedule for ExponentialSchedule { fn get_difficulty_at_step(&self, step: usize) -> f32 { let exp_factor = 1.0 - (-self.growth_rate * step as f32).exp(); let difficulty = self.initial_difficulty + exp_factor * (self.final_difficulty - self.initial_difficulty); difficulty.min(self.final_difficulty) } } /// Adaptive curriculum schedule (adjusts based on performance) #[derive(Debug, Clone)] pub struct AdaptiveSchedule { initial_difficulty: f32, final_difficulty: f32, adaptation_rate: f32, target_performance: f32, performance_window: Vec, current_difficulty: f32, } impl AdaptiveSchedule { pub fn new( initial_difficulty: f32, final_difficulty: f32, adaptation_rate: f32, target_performance: f32, ) -> Self { Self { initial_difficulty, final_difficulty, adaptation_rate, target_performance, performance_window: Vec::new(), current_difficulty: initial_difficulty, } } pub fn update_performance(&mut self, performance: f32) { self.performance_window.push(performance); if self.performance_window.len() > 10 { self.performance_window.remove(0); } } fn current_performance(&self) -> f32 { if self.performance_window.is_empty() { self.target_performance } else { self.performance_window.iter().sum::() / self.performance_window.len() as f32 } } } impl Schedule for AdaptiveSchedule { fn get_difficulty_at_step(&self, _step: usize) -> f32 { let current_perf = self.current_performance(); let adjustment = if current_perf > self.target_performance { self.adaptation_rate // Increase difficulty } else { -self.adaptation_rate // Decrease difficulty }; (self.current_difficulty + adjustment).clamp(self.initial_difficulty, self.final_difficulty) } } /// Cyclic curriculum schedule #[derive(Debug, Clone)] pub struct CyclicSchedule { min_difficulty: f32, max_difficulty: f32, cycle_length: usize, } impl CyclicSchedule { pub fn new(min_difficulty: f32, max_difficulty: f32, cycle_length: usize) -> Self { Self { min_difficulty, max_difficulty, cycle_length, } } } impl Schedule for CyclicSchedule { fn get_difficulty_at_step(&self, step: usize) -> f32 { let cycle_position = (step % self.cycle_length) as f32; let normalized_position = cycle_position / self.cycle_length as f32; // Triangle wave: goes from min to max and back let triangle_wave = if normalized_position <= 0.5 { normalized_position * 2.0 // 0 -> 1 } else { 2.0 - normalized_position * 2.0 // 1 -> 0 }; self.min_difficulty + triangle_wave * (self.max_difficulty - self.min_difficulty) } } /// Multi-task curriculum schedule coordinator #[derive(Debug)] pub struct MultiTaskSchedule { task_schedules: HashMap>, coordination_weight: f32, } impl MultiTaskSchedule { pub fn new() -> Self { Self { task_schedules: HashMap::new(), coordination_weight: 0.0, } } pub fn add_task(&mut self, task_name: &str, schedule: Box) { self.task_schedules.insert(task_name.to_string(), schedule); } pub fn set_coordination_weight(&mut self, weight: f32) { self.coordination_weight = weight; } pub fn get_task_difficulty(&self, task_name: &str, step: usize) -> Option { self.task_schedules.get(task_name).map(|schedule| schedule.get_difficulty_at_step(step)) } pub fn get_coordinated_difficulty(&self, task_name: &str, step: usize) -> Option { let task_difficulty = self.get_task_difficulty(task_name, step)?; if self.coordination_weight <= 0.0 { return Some(task_difficulty); } // Calculate average difficulty across all tasks let total_difficulty: f32 = self.task_schedules .values() .map(|schedule| schedule.get_difficulty_at_step(step)) .sum(); let avg_difficulty = total_difficulty / self.task_schedules.len() as f32; // Blend task-specific and average difficulty Some(task_difficulty * (1.0 - self.coordination_weight) + avg_difficulty * self.coordination_weight) } } impl Default for MultiTaskSchedule { fn default() -> Self { Self::new() } }