// optimizer_adam.mbt — Adam optimiser (v0.15.1).
//
// Reference: Kingma & Ba, "Adam: A Method for Stochastic Optimization",
// ICLR 2015 (arXiv:1412.6980).
//
// Standard update rule (with bias correction):
//
//   m <- beta1 * m + (1 - beta1) * g
//   v <- beta2 * v + (1 - beta2) * g^2
//   m_hat <- m / (1 - beta1^t)
//   v_hat <- v / (1 - beta2^t)
//   param <- param - lr * m_hat / (sqrt(v_hat) + eps)
//
// State `step` is a 1-based step counter (incremented before the update,
// so the first step uses `t = 1`). Default hyper-parameters:
//   - beta1 = 0.9
//   - beta2 = 0.999
//   - eps   = 1e-8
//   - lr    = 1e-3
//
// Conventions match the project's bit-exact Float32 contract.

///|
/// Adam state. `m_w / v_w` track first / second moments of `d_weight`;
/// `m_b / v_b` do the same for `d_bias`. `step` is 1-based and is
/// incremented implicitly inside `adam_update_arrays` — the caller must
/// pass the state returned from the previous call into the next call.
/// (We do NOT mutate the input state in-place because MoonBit struct
/// fields share memory across function boundaries, which would corrupt
/// any previously-returned state object.)
pub struct AdamState {
  m_w : Array[Float]
  v_w : Array[Float]
  m_b : Array[Float]
  v_b : Array[Float]
  step : Int
}

///|
/// Allocate zero-initialised Adam state for a parameter of given
/// `weight_len` / `bias_len`. Initial `step = 0` (no bias correction yet).
pub fn adam_init(weight_len : Int, bias_len : Int) -> AdamState {
  {
    m_w: Array::make(weight_len, 0.0F),
    v_w: Array::make(weight_len, 0.0F),
    m_b: Array::make(bias_len, 0.0F),
    v_b: Array::make(bias_len, 0.0F),
    step: 0,
  }
}

///|
/// Pure accessor: returns `state.step + 1`. Used internally by
/// `adam_update_arrays` and exposed for callers that want to inspect
/// the next step counter without mutating the state.
pub fn adam_next_step(state : AdamState) -> Int {
  state.step + 1
}

///|
/// One Adam update step at raw-array level. Returns the updated
/// `(weight, bias)` arrays plus a new `AdamState` with the incremented
/// step counter. The input `state` is read-only; pass the returned
/// state into the next call.
pub fn adam_update_arrays(
  weight : Array[Float],
  bias : Array[Float],
  d_weight : Array[Float],
  d_bias : Array[Float],
  state : AdamState,
  lr : Float,
  beta1 : Float,
  beta2 : Float,
  eps : Float,
) -> (Array[Float], Array[Float], AdamState) {
  let t = state.step + 1
  // Bias-correction denominators.
  let bc1 = 1.0F - pow_beta(beta1, t)
  let bc2 = 1.0F - pow_beta(beta2, t)
  // Update weight moments + apply.
  let w2 : Array[Float] = Array::make(weight.length(), 0.0F)
  let new_m_w : Array[Float] = Array::make(weight.length(), 0.0F)
  let new_v_w : Array[Float] = Array::make(weight.length(), 0.0F)
  for i in 0.. (LinearParam, AdamState) {
  let (w, b, s2) = adam_update_arrays(
    param.weight, param.bias, d_weight, d_bias, state, lr, beta1, beta2, eps,
  )
  let p2 = { weight: w, bias: b, in_features: param.in_features,
             out_features: param.out_features }
  (p2, s2)
}

///|
/// Adam wrapper for `Conv2dParam`.
pub fn adam_update_conv(
  param : Conv2dParam,
  d_weight : Array[Float],
  d_bias : Array[Float],
  state : AdamState,
  lr : Float,
  beta1 : Float,
  beta2 : Float,
  eps : Float,
) -> (Conv2dParam, AdamState) {
  let (w, b, s2) = adam_update_arrays(
    param.weight, param.bias, d_weight, d_bias, state, lr, beta1, beta2, eps,
  )
  let p2 = { weight: w, bias: b, c_out: param.c_out, c_in: param.c_in,
             kh: param.kh, kw: param.kw, stride: param.stride, pad: param.pad }
  (p2, s2)
}

// ---------------------------------------------------------------------------
// Float32 helpers: pow(beta, t) and sqrt.
// Implemented as plain Float32 loops + libm FFI (no `Float::pow` builtin
// in moonbitlang/core/math).
// ---------------------------------------------------------------------------

///|
/// Compute `beta^t` for Float32 `beta` and Int `t >= 0` via
/// exponentiation by squaring. Returns 1.0 when `t == 0`.
fn pow_beta(beta : Float, t : Int) -> Float {
  let mut result = 1.0F
  let mut b = beta
  let mut n = t
  while n > 0 {
    if (n & 1) == 1 {
      result = result * b
    }
    b = b * b
    n = n >> 1
  }
  result
}

///|
/// Float32 sqrt via libm. We already have `math_sqrt_f32` exposed in
// `math_native.mbt`; reuse it. If unavailable, fall back to a Newton
// iteration (see below).
pub extern "C" fn sqrtf(x : Float) -> Float = "sqrtf"

///|
fn sqrt_f32(x : Float) -> Float {
  sqrtf(x)
}