// optimizer_rmsprop.mbt — RMSprop optimiser (v0.16.1).
//
// Reference: G. Hinton, "Neural Networks for Machine Learning" lecture
// 6e (unpublished, 2012). Common modern form:
//
//   v <- alpha * v + (1 - alpha) * g^2
//   param <- param - lr * g / (sqrt(v) + eps)
//
// Unlike Adam, RMSprop has no first moment (`m`) and no bias correction
// — it scales the gradient by an exponential moving average of its
// squared magnitude. With `alpha = 0` it degenerates to plain SGD with
// per-parameter adaptive learning rate; with `alpha -> 1` it approaches
// pure `1/sqrt(g^2 + eps)` scaling (which is essentially AdaDelta's
// root form).
//
// Conventions match the project's bit-exact Float32 contract.

///|
/// RMSprop state. `v_w` / `v_b` are exponential moving averages of the
/// squared gradient for `weight` / `bias`. No step counter — RMSprop
/// has no bias correction term.
pub struct RMSpropState {
  v_w : Array[Float]
  v_b : Array[Float]
}

///|
/// Allocate zero-initialised RMSprop state.
pub fn rmsprop_init(weight_len : Int, bias_len : Int) -> RMSpropState {
  { v_w: Array::make(weight_len, 0.0F), v_b: Array::make(bias_len, 0.0F) }
}

///|
/// One RMSprop update at raw-array level.
pub fn rmsprop_update_arrays(
  weight : Array[Float],
  bias : Array[Float],
  d_weight : Array[Float],
  d_bias : Array[Float],
  state : RMSpropState,
  lr : Float,
  alpha : Float,
  eps : Float,
) -> (Array[Float], Array[Float], RMSpropState) {
  let w2 : 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, RMSpropState) {
  let (w, b, s2) = rmsprop_update_arrays(
    param.weight, param.bias, d_weight, d_bias, state, lr, alpha, eps,
  )
  let p2 = { weight: w, bias: b, in_features: param.in_features,
             out_features: param.out_features }
  (p2, s2)
}

///|
/// RMSprop wrapper for `Conv2dParam`.
pub fn rmsprop_update_conv(
  param : Conv2dParam,
  d_weight : Array[Float],
  d_bias : Array[Float],
  state : RMSpropState,
  lr : Float,
  alpha : Float,
  eps : Float,
) -> (Conv2dParam, RMSpropState) {
  let (w, b, s2) = rmsprop_update_arrays(
    param.weight, param.bias, d_weight, d_bias, state, lr, alpha, 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)
}