// scheduler.mbt — Learning-rate schedulers (v0.18.0 + v0.18.1).
//
// All schedulers expose a small uniform API:
//   - Constructor takes the hyper-parameters (no `step` counter —
//     it's tracked via returned state to avoid MoonBit's "struct fields
//     share memory across functions" pitfall).
//   - `step(s)` returns the current learning rate (pure read).
//   - `next(s)` returns a fresh state with the step counter incremented.
//
// The caller pattern is:
//   for step_idx in 0.. StepLR {
  { base_lr, step_size, gamma, current_step: 0 }
}

///|
/// Compute the current learning rate.
pub fn StepLR::step(self : StepLR) -> Float {
  // Number of completed step intervals.
  let n_drops = self.current_step / self.step_size
  let mut factor = 1.0F
  for _ in 0.. StepLR {
  { base_lr: self.base_lr, step_size: self.step_size, gamma: self.gamma,
    current_step: self.current_step + 1 }
}

// ---------------------------------------------------------------------------
// ExponentialLR
// ---------------------------------------------------------------------------

///|
/// ExponentialLR state. lr_t = base_lr * gamma^t.
pub struct ExponentialLR {
  base_lr : Float
  gamma : Float
  current_step : Int
}

///|
/// Build an ExponentialLR scheduler.
pub fn ExponentialLR::new(base_lr : Float, gamma : Float) -> ExponentialLR {
  { base_lr, gamma, current_step: 0 }
}

///|
/// Compute the current learning rate.
pub fn ExponentialLR::step(self : ExponentialLR) -> Float {
  self.base_lr * pow_beta(self.gamma, self.current_step)
}

///|
/// Increment the step counter; returns a fresh state.
pub fn ExponentialLR::next(self : ExponentialLR) -> ExponentialLR {
  { base_lr: self.base_lr, gamma: self.gamma,
    current_step: self.current_step + 1 }
}

// ---------------------------------------------------------------------------
// CosineAnnealingLR
// ---------------------------------------------------------------------------

///|
/// CosineAnnealingLR state. lr follows a half-cosine from eta_max to
/// eta_min over `t_max` steps, then jumps back to eta_max and repeats.
pub struct CosineAnnealingLR {
  eta_max : Float
  eta_min : Float
  t_max : Int
  current_step : Int
}

///|
/// Build a CosineAnnealingLR scheduler.
pub fn CosineAnnealingLR::new(
  eta_max : Float,
  eta_min : Float,
  t_max : Int,
) -> CosineAnnealingLR {
  { eta_max, eta_min, t_max, current_step: 0 }
}

///|
/// Compute the current learning rate.
pub fn CosineAnnealingLR::step(self : CosineAnnealingLR) -> Float {
  // Modulo into the cycle: at t = t_max we wrap back to t = 0.
  let t = self.current_step % self.t_max
  let t_f = Float::from_int(t)
  let t_max_f = Float::from_int(self.t_max)
  let angle = 3.14159265F * t_f / t_max_f
  let c = cos_f32(angle)
  self.eta_min + 0.5F * (self.eta_max - self.eta_min) * (1.0F + c)
}

///|
/// Increment the step counter; returns a fresh state.
pub fn CosineAnnealingLR::next(self : CosineAnnealingLR) -> CosineAnnealingLR {
  { eta_max: self.eta_max, eta_min: self.eta_min, t_max: self.t_max,
    current_step: self.current_step + 1 }
}

///|
/// Float32 cos via libm. Wraps the `cosf` symbol exposed by moonbit's
/// built-in math.
extern "C" fn cosf(x : Float) -> Float = "cosf"

///|
fn cos_f32(x : Float) -> Float {
  cosf(x)
}

// ---------------------------------------------------------------------------
// ReduceLROnPlateau
// ---------------------------------------------------------------------------

///|
/// ReduceLROnPlateau state. Tracks the best metric seen so far and the
/// number of consecutive non-improving steps. lr is multiplied by
/// `factor` when patience runs out.
pub struct ReduceLROnPlateau {
  base_lr : Float
  factor : Float
  patience : Int
  threshold : Float
  best_metric : Float
  num_bad_epochs : Int
  current_lr : Float
}

///|
/// Build a ReduceLROnPlateau scheduler. `base_lr` is the initial lr;
/// `factor` is the decay multiplier applied on plateau (typically 0.1);
/// `patience` is the number of consecutive non-improving steps before
/// decay; `threshold` is the minimum relative improvement (in the
/// metric's own units) to count as "improved".
pub fn ReduceLROnPlateau::new(
  base_lr : Float,
  factor : Float,
  patience : Int,
  threshold : Float,
) -> ReduceLROnPlateau {
  { base_lr, factor, patience, threshold, best_metric: 1000000000.0F,
    num_bad_epochs: 0, current_lr: base_lr }
}

///|
/// Feed the current step's metric to the scheduler. Returns the new lr
/// and the updated state. The metric convention is "lower is better"
/// (loss-like); the metric counts as "improved" if it is at least
/// `threshold` below the previous best.
pub fn ReduceLROnPlateau::step(
  self : ReduceLROnPlateau,
  metric : Float,
) -> (Float, ReduceLROnPlateau) {
  // First call: no comparison yet.
  if self.best_metric == 1000000000.0F {
    return (self.current_lr, {
      base_lr: self.base_lr,
      factor: self.factor,
      patience: self.patience,
      threshold: self.threshold,
      best_metric: metric,
      num_bad_epochs: 0,
      current_lr: self.current_lr,
    })
  }
  // Improved by at least `threshold` -> reset bad counter.
  if metric < self.best_metric - self.threshold {
    return (self.current_lr, {
      base_lr: self.base_lr,
      factor: self.factor,
      patience: self.patience,
      threshold: self.threshold,
      best_metric: metric,
      num_bad_epochs: 0,
      current_lr: self.current_lr,
    })
  }
  // Did not improve enough -> increment bad counter.
  let next_bad = self.num_bad_epochs + 1
  if next_bad > self.patience {
    // Plateau: decay lr, reset counter.
    let new_lr = self.current_lr * self.factor
    return (new_lr, {
      base_lr: self.base_lr,
      factor: self.factor,
      patience: self.patience,
      threshold: self.threshold,
      best_metric: self.best_metric,
      num_bad_epochs: 0,
      current_lr: new_lr,
    })
  }
  // Plateau not yet triggered.
  (self.current_lr, {
    base_lr: self.base_lr,
    factor: self.factor,
    patience: self.patience,
    threshold: self.threshold,
    best_metric: self.best_metric,
    num_bad_epochs: next_bad,
    current_lr: self.current_lr,
  })
}