// optimizer_adamw.mbt — AdamW optimiser (v0.16.0).
//
// Reference: Loshchilov & Hutter, "Decoupled Weight Decay Regularization",
// ICLR 2019 (arXiv:1711.05101).
//
// AdamW differs from Adam by decoupling the weight-decay term from the
// gradient-based update. Adam conflates L2 regularisation with the
// gradient (adding `lambda * param` to `g` before the m/v accumulation),
// which interacts poorly with the adaptive moments. AdamW instead adds
// the decay directly to the parameter update:
//
//   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)
//                          + weight_decay * param )
//
// State shape matches `AdamState`; step is incremented inside the
// update (caller threads the returned state forward, as with Adam).
// `weight_decay` is a Float32 hyperparameter; setting it to 0 makes
// AdamW behaviourally equivalent to Adam (modulo the bias-correction
// convention — both implementations match the PyTorch default).

///|
/// AdamW state. Identical shape to `AdamState`; `step` is 1-based and
/// tracked implicitly via the returned state (not mutated in place).
pub struct AdamWState {
  m_w : Array[Float]
  v_w : Array[Float]
  m_b : Array[Float]
  v_b : Array[Float]
  step : Int
}

///|
/// Allocate zero-initialised AdamW state for a parameter of given
/// `weight_len` / `bias_len`. `step` starts at 0 (bias correction
/// uses `t = state.step + 1`).
pub fn adamw_init(weight_len : Int, bias_len : Int) -> AdamWState {
  {
    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,
  }
}

///|
/// One AdamW update at raw-array level. Returns the updated
/// `(weight, bias)` arrays plus the new state with `step` incremented.
/// `weight_decay` is applied additively to the parameter update
/// (`param <- param - lr*(...) + lr*weight_decay*param`), matching the
/// PyTorch `optim.AdamW` convention.
pub fn adamw_update_arrays(
  weight : Array[Float],
  bias : Array[Float],
  d_weight : Array[Float],
  d_bias : Array[Float],
  state : AdamWState,
  lr : Float,
  beta1 : Float,
  beta2 : Float,
  eps : Float,
  weight_decay : Float,
) -> (Array[Float], Array[Float], AdamWState) {
  let t = state.step + 1
  let bc1 = 1.0F - pow_beta(beta1, t)
  let bc2 = 1.0F - pow_beta(beta2, t)
  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, AdamWState) {
  let (w, b, s2) = adamw_update_arrays(
    param.weight, param.bias, d_weight, d_bias, state,
    lr, beta1, beta2, eps, weight_decay,
  )
  let p2 = { weight: w, bias: b, in_features: param.in_features,
             out_features: param.out_features }
  (p2, s2)
}

///|
/// AdamW wrapper for `Conv2dParam`.
pub fn adamw_update_conv(
  param : Conv2dParam,
  d_weight : Array[Float],
  d_bias : Array[Float],
  state : AdamWState,
  lr : Float,
  beta1 : Float,
  beta2 : Float,
  eps : Float,
  weight_decay : Float,
) -> (Conv2dParam, AdamWState) {
  let (w, b, s2) = adamw_update_arrays(
    param.weight, param.bias, d_weight, d_bias, state,
    lr, beta1, beta2, eps, weight_decay,
  )
  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)
}