use super::*;
#[repr(u8)]
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub enum SubStepStrategy {
#[default]
DampedEuler,
KStageRK {
beta: f32,
},
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[repr(u8)]
pub enum IterationMode {
#[default]
Block,
Layer,
}
#[repr(u8)]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum CacheStrategy {
#[default]
Last,
First,
}
#[derive(Clone, Debug)]
pub struct TrainingFreeLoopConfig {
pub window_start: usize,
pub window_end: usize,
pub loop_count: usize,
pub strategy: SubStepStrategy,
pub iteration_mode: IterationMode,
pub cache_strategy: CacheStrategy,
}
impl Default for TrainingFreeLoopConfig {
fn default() -> Self {
Self {
window_start: 0,
window_end: 0,
loop_count: 2,
strategy: SubStepStrategy::KStageRK { beta: 0.5 },
iteration_mode: IterationMode::Block,
cache_strategy: CacheStrategy::First,
}
}
}
impl TrainingFreeLoopConfig {
pub fn from_config(config: &Config) -> Self {
let n = config.n_layer;
let (window_start, window_end) = if n <= 4 {
(0, n.saturating_sub(1))
} else {
let center = (n as f32 * 0.48) as usize;
(center.saturating_sub(1), (center + 2).min(n - 1))
};
Self {
window_start,
window_end,
loop_count: 2,
strategy: SubStepStrategy::KStageRK { beta: 0.5 },
iteration_mode: IterationMode::Block,
cache_strategy: CacheStrategy::First,
}
}
}