Initial commit
This commit is contained in:
@@ -0,0 +1,208 @@
|
||||
//! 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<f32>,
|
||||
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::<f32>() / 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<String, Box<dyn Schedule>>,
|
||||
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<dyn Schedule>) {
|
||||
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<f32> {
|
||||
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<f32> {
|
||||
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()
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user