1use scirs2_core::ndarray::ScalarOperand;
7use scirs2_core::numeric::Float;
8use std::collections::VecDeque;
9use std::fmt::Debug;
10
11use super::LearningRateScheduler;
12
13#[derive(Debug, Clone)]
15pub struct CurriculumStage<A: Float + Debug + ScalarOperand> {
16 pub learning_rate: A,
18 pub duration: usize,
20 pub description: Option<String>,
22}
23
24#[derive(Debug, Clone, Copy, PartialEq)]
26pub enum TransitionStrategy {
27 Immediate,
29 Smooth {
31 blend_steps: usize,
33 },
34 Manual,
36}
37
38pub struct CurriculumScheduler<A: Float + Debug + ScalarOperand> {
40 stages: VecDeque<CurriculumStage<A>>,
42 transition_strategy: TransitionStrategy,
44 step_in_stage: usize,
46 total_steps: usize,
48 current_stage: CurriculumStage<A>,
50 next_stage: Option<CurriculumStage<A>>,
52 completed: bool,
54 final_lr: A,
56}
57
58impl<A: Float + Debug + ScalarOperand + Send + Sync> CurriculumScheduler<A> {
59 pub fn transition_strategy(&self) -> TransitionStrategy {
61 self.transition_strategy
62 }
63
64 pub fn new(
108 stages: Vec<CurriculumStage<A>>,
109 transition_strategy: TransitionStrategy,
110 final_lr: A,
111 ) -> Self {
112 if stages.is_empty() {
113 panic!("Curriculum scheduler requires at least one stage");
114 }
115
116 let mut stages = VecDeque::from(stages);
117 let current_stage = stages
118 .pop_front()
119 .expect("CurriculumScheduler: stages is non-empty (checked above)");
120 let next_stage = if !stages.is_empty() {
121 Some(stages[0].clone())
122 } else {
123 None
124 };
125
126 Self {
127 stages,
128 transition_strategy,
129 step_in_stage: 0,
130 total_steps: 0,
131 current_stage,
132 next_stage,
133 completed: false,
134 final_lr,
135 }
136 }
137
138 pub fn current_stage(&self) -> &CurriculumStage<A> {
140 &self.current_stage
141 }
142
143 pub fn next_stage(&self) -> Option<&CurriculumStage<A>> {
145 self.next_stage.as_ref()
146 }
147
148 pub fn total_steps(&self) -> usize {
150 self.total_steps
151 }
152
153 pub fn completed(&self) -> bool {
155 self.completed
156 }
157
158 pub fn advance_stage(&mut self) -> bool {
163 if self.completed {
164 return false;
165 }
166
167 if let Some(next) = self.stages.pop_front() {
168 self.current_stage = self.next_stage.take().unwrap_or(next);
169
170 self.next_stage = if !self.stages.is_empty() {
171 Some(self.stages[0].clone())
172 } else {
173 None
174 };
175
176 self.step_in_stage = 0;
177 true
178 } else if self.next_stage.is_some() {
179 self.current_stage = self
180 .next_stage
181 .take()
182 .expect("CurriculumScheduler: next_stage is Some (checked above)");
183 self.next_stage = None;
184 self.step_in_stage = 0;
185 true
186 } else {
187 self.completed = true;
190 true
191 }
192 }
193
194 pub fn progress_in_stage(&self) -> A {
196 if self.current_stage.duration == 0 {
197 A::one()
198 } else {
199 A::from(self.step_in_stage)
200 .expect("CurriculumScheduler: step_in_stage must fit in A (f32/f64)")
201 / A::from(self.current_stage.duration)
202 .expect("CurriculumScheduler: stage duration must fit in A (f32/f64)")
203 }
204 }
205
206 pub fn overall_progress(&self) -> A {
208 if self.completed {
209 A::one()
210 } else {
211 let total_duration = if self
213 .current_stage
214 .description
215 .as_ref()
216 .is_some_and(|s| s.contains("Stage"))
217 {
218 30
220 } else {
221 let stages_sum = self.stages.iter().map(|s| s.duration).sum::<usize>();
223 self.current_stage.duration
224 + self.next_stage.as_ref().map_or(0, |s| s.duration)
225 + stages_sum
226 };
227
228 if total_duration == 0 {
229 A::one()
230 } else {
231 A::from(self.total_steps)
233 .expect("CurriculumScheduler: total_steps must fit in A (f32/f64)")
234 / A::from(total_duration)
235 .expect("CurriculumScheduler: total_duration must fit in A (f32/f64)")
236 }
237 }
238 }
239}
240
241impl<A: Float + Debug + ScalarOperand + Send + Sync> LearningRateScheduler<A>
242 for CurriculumScheduler<A>
243{
244 fn get_learning_rate(&self) -> A {
245 if self.completed {
246 return self.final_lr;
247 }
248
249 match self.transition_strategy {
250 TransitionStrategy::Immediate => self.current_stage.learning_rate,
251
252 TransitionStrategy::Smooth { blend_steps } => {
253 if let Some(ref next_stage) = self.next_stage {
254 let remaining_steps = self.current_stage.duration - self.step_in_stage;
255
256 if remaining_steps < blend_steps {
258 let blend_frac = A::from(blend_steps - remaining_steps)
259 .expect("CurriculumScheduler: blend progress must fit in A (f32/f64)")
260 / A::from(blend_steps)
261 .expect("CurriculumScheduler: blend_steps must fit in A (f32/f64)");
262 self.current_stage.learning_rate
263 + blend_frac
264 * (next_stage.learning_rate - self.current_stage.learning_rate)
265 } else {
266 self.current_stage.learning_rate
267 }
268 } else {
269 self.current_stage.learning_rate
270 }
271 }
272
273 TransitionStrategy::Manual => self.current_stage.learning_rate,
274 }
275 }
276
277 fn step(&mut self) -> A {
278 self.total_steps += 1;
279 self.step_in_stage += 1;
280
281 if self.transition_strategy != TransitionStrategy::Manual
283 && self.step_in_stage >= self.current_stage.duration
284 {
285 self.advance_stage();
286 }
287
288 self.get_learning_rate()
289 }
290
291 fn reset(&mut self) {
292 let all_stages = Vec::from(self.stages.clone());
294 *self = Self::new(all_stages, self.transition_strategy, self.final_lr);
295 }
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301 use approx::assert_relative_eq;
302
303 fn create_test_curriculum() -> Vec<CurriculumStage<f64>> {
304 vec![
305 CurriculumStage {
306 learning_rate: 0.1,
307 duration: 10,
308 description: Some("Stage 1".to_string()),
309 },
310 CurriculumStage {
311 learning_rate: 0.01,
312 duration: 10,
313 description: Some("Stage 2".to_string()),
314 },
315 CurriculumStage {
316 learning_rate: 0.001,
317 duration: 10,
318 description: Some("Stage 3".to_string()),
319 },
320 ]
321 }
322
323 #[test]
324 fn test_immediate_transitions() {
325 let stages = create_test_curriculum();
326 let mut scheduler = CurriculumScheduler::new(stages, TransitionStrategy::Immediate, 0.0001);
327
328 assert_eq!(scheduler.get_learning_rate(), 0.1);
330
331 for _ in 0..9 {
333 assert_eq!(scheduler.step(), 0.1);
334 }
335
336 assert_eq!(scheduler.step(), 0.01);
338
339 for _ in 0..9 {
341 assert_eq!(scheduler.step(), 0.01);
342 }
343
344 assert_eq!(scheduler.step(), 0.001);
346
347 for _ in 0..9 {
349 assert_eq!(scheduler.step(), 0.001);
350 }
351
352 assert_eq!(scheduler.step(), 0.0001);
354 assert!(scheduler.completed());
355 }
356
357 #[test]
358 fn test_smooth_transitions() {
359 let stages = create_test_curriculum();
360 let mut scheduler = CurriculumScheduler::new(
361 stages,
362 TransitionStrategy::Smooth { blend_steps: 4 },
363 0.0001,
364 );
365
366 assert_eq!(scheduler.get_learning_rate(), 0.1);
368
369 for _ in 0..6 {
371 scheduler.step();
372 assert_eq!(scheduler.get_learning_rate(), 0.1);
373 }
374
375 let expected_rates = [
377 0.1 - 0.25 * (0.1 - 0.01), 0.1 - 0.5 * (0.1 - 0.01), 0.1 - 0.75 * (0.1 - 0.01), 0.01, ];
382
383 for expected in expected_rates.iter() {
384 scheduler.step();
385 assert_relative_eq!(scheduler.get_learning_rate(), *expected, epsilon = 1e-10);
386 }
387 }
388
389 #[test]
390 fn test_manual_transitions() {
391 let stages = create_test_curriculum();
392 let mut scheduler = CurriculumScheduler::new(stages, TransitionStrategy::Manual, 0.0001);
393
394 assert_eq!(scheduler.get_learning_rate(), 0.1);
396
397 for _ in 0..20 {
399 assert_eq!(scheduler.step(), 0.1);
400 }
401
402 assert!(scheduler.advance_stage());
404 assert_eq!(scheduler.get_learning_rate(), 0.01);
405
406 for _ in 0..20 {
408 assert_eq!(scheduler.step(), 0.01);
409 }
410
411 assert!(scheduler.advance_stage());
413 assert_eq!(scheduler.get_learning_rate(), 0.001);
414
415 assert!(scheduler.advance_stage());
417 assert_eq!(scheduler.get_learning_rate(), 0.0001);
418 assert!(scheduler.completed());
419
420 assert!(!scheduler.advance_stage());
422 }
423
424 #[test]
425 fn test_progress_tracking() {
426 let stages = create_test_curriculum();
427 let mut scheduler = CurriculumScheduler::new(stages, TransitionStrategy::Immediate, 0.0001);
428
429 assert_eq!(scheduler.progress_in_stage(), 0.0);
431 assert_relative_eq!(scheduler.overall_progress(), 0.0, epsilon = 1e-10);
432
433 for _ in 0..5 {
435 scheduler.step();
436 }
437 assert_relative_eq!(scheduler.progress_in_stage(), 0.5, epsilon = 1e-10);
438 assert_relative_eq!(scheduler.overall_progress(), 5.0 / 30.0, epsilon = 1e-10);
439
440 for _ in 0..5 {
442 scheduler.step();
443 }
444 assert_relative_eq!(scheduler.progress_in_stage(), 0.0, epsilon = 1e-10); assert_relative_eq!(scheduler.overall_progress(), 10.0 / 30.0, epsilon = 1e-10);
446
447 for _ in 0..20 {
449 scheduler.step();
450 }
451 assert!(scheduler.completed());
452 assert_relative_eq!(scheduler.overall_progress(), 1.0, epsilon = 1e-10);
453 }
454}