///|
pub fn sgd_step(params : Array[@tensor.Tensor], lr : Double) -> Unit {
  for param in params {
    match param.grad() {
      Some(grad_tensor) => {
        let data = param.data()
        let grad = grad_tensor.data()
        for i in 0.. ()
    }
  }
}

///|
struct AdamWState {
  params : Array[@tensor.Tensor]
  m : Array[Array[Double]]
  v : Array[Array[Double]]
  mut lr : Double
  beta1 : Double
  beta2 : Double
  eps : Double
  weight_decays : Array[Double]
  mut step : Int
  mut beta1_power : Double
  mut beta2_power : Double
}

///|
pub struct AdamW {
  priv state : Ref[AdamWState]
}

///|
pub struct AdamWCheckpoint {
  priv m : Array[Array[Double]]
  priv v : Array[Array[Double]]
  priv lr : Double
  priv beta1 : Double
  priv beta2 : Double
  priv eps : Double
  priv weight_decays : Array[Double]
  priv step : Int
  priv beta1_power : Double
  priv beta2_power : Double
}

///|
pub struct AdamWConfig {
  priv lr : Double
  priv beta1 : Double
  priv beta2 : Double
  priv eps : Double
  priv weight_decay : Double
}

///|
fn copy_double_arrays(values : Array[Array[Double]]) -> Array[Array[Double]] {
  let copied : Array[Array[Double]] = []
  for value in values {
    copied.push(value.copy())
  }
  copied
}

///|
pub fn AdamWConfig::AdamWConfig(
  lr : Double,
  beta1? : Double = 0.9,
  beta2? : Double = 0.999,
  eps? : Double = 1.0e-8,
  weight_decay? : Double = 0.0,
) -> AdamWConfig {
  if lr <= 0.0 {
    abort("AdamW learning rate must be positive")
  }
  if beta1 < 0.0 || beta1 >= 1.0 || beta2 < 0.0 || beta2 >= 1.0 {
    abort("AdamW beta values must be in [0, 1)")
  }
  if eps <= 0.0 {
    abort("AdamW eps must be positive")
  }
  if weight_decay < 0.0 {
    abort("AdamW weight decay must not be negative")
  }
  { lr, beta1, beta2, eps, weight_decay }
}

///|
pub fn AdamWConfig::lr(self : AdamWConfig) -> Double {
  self.lr
}

///|
pub fn AdamWConfig::beta1(self : AdamWConfig) -> Double {
  self.beta1
}

///|
pub fn AdamWConfig::beta2(self : AdamWConfig) -> Double {
  self.beta2
}

///|
pub fn AdamWConfig::eps(self : AdamWConfig) -> Double {
  self.eps
}

///|
pub fn AdamWConfig::weight_decay(self : AdamWConfig) -> Double {
  self.weight_decay
}

///|
pub fn AdamW::AdamW(
  params : Array[@tensor.Tensor],
  config : AdamWConfig,
) -> AdamW {
  AdamW::with_parameter_weight_decays(
    params,
    config,
    Array::make(params.length(), config.weight_decay),
  )
}

///|
pub fn AdamW::with_parameter_weight_decays(
  params : Array[@tensor.Tensor],
  config : AdamWConfig,
  weight_decays : Array[Double],
) -> AdamW {
  if params.length() != weight_decays.length() {
    abort("AdamW weight decay count must match parameter count")
  }
  for weight_decay in weight_decays {
    if weight_decay < 0.0 {
      abort("AdamW weight decay must not be negative")
    }
  }
  let m : Array[Array[Double]] = []
  let v : Array[Array[Double]] = []
  for param in params {
    m.push(Array::make(param.numel(), 0.0))
    v.push(Array::make(param.numel(), 0.0))
  }
  {
    state: {
      val: {
        params: params.copy(),
        m,
        v,
        lr: config.lr,
        beta1: config.beta1,
        beta2: config.beta2,
        eps: config.eps,
        weight_decays: weight_decays.copy(),
        step: 0,
        beta1_power: 1.0,
        beta2_power: 1.0,
      },
    },
  }
}

///|
pub fn AdamW::step(self : AdamW) -> Unit {
  self.step_with_gradient_scale(1.0)
}

///|
pub fn AdamW::set_learning_rate(self : AdamW, lr : Double) -> Unit {
  if lr <= 0.0 {
    abort("AdamW learning rate must be positive")
  }
  self.state.val.lr = lr
}

///|
pub fn AdamW::step_with_grad_clip(self : AdamW, max_norm : Double) -> Unit {
  if max_norm < 0.0 {
    abort("AdamW max_norm must not be negative")
  }
  let scale = self.gradient_clip_scale(max_norm)
  self.step_with_gradient_scale(scale)
}

///|
fn AdamW::gradient_clip_scale(self : AdamW, max_norm : Double) -> Double {
  if max_norm == 0.0 {
    return 1.0
  }
  let state = self.state.val
  let mut squared_sum = 0.0
  for param in state.params {
    match param.grad() {
      Some(grad_tensor) => {
        let grad = grad_tensor.data()
        for value in grad {
          squared_sum += value * value
        }
      }
      None => ()
    }
  }
  let norm = squared_sum.sqrt()
  if norm > max_norm {
    max_norm / (norm + 1.0e-6)
  } else {
    1.0
  }
}

///|
fn AdamW::step_with_gradient_scale(self : AdamW, grad_scale : Double) -> Unit {
  let state = self.state.val
  state.step += 1
  state.beta1_power *= state.beta1
  state.beta2_power *= state.beta2
  let bias1 = 1.0 - state.beta1_power
  let bias2 = 1.0 - state.beta2_power
  for pi in 0.. {
        let data = param.data()
        let grad = grad_tensor.data()
        if data.length() != state.m[pi].length() {
          abort("AdamW state does not match parameter size")
        }
        for i in 0.. ()
    }
  }
}

///|
pub fn AdamW::step_count(self : AdamW) -> Int {
  self.state.val.step
}

///|
pub fn AdamW::checkpoint(self : AdamW) -> AdamWCheckpoint {
  let state = self.state.val
  {
    m: copy_double_arrays(state.m),
    v: copy_double_arrays(state.v),
    lr: state.lr,
    beta1: state.beta1,
    beta2: state.beta2,
    eps: state.eps,
    weight_decays: state.weight_decays.copy(),
    step: state.step,
    beta1_power: state.beta1_power,
    beta2_power: state.beta2_power,
  }
}

///|
pub fn AdamWCheckpoint::m(self : AdamWCheckpoint) -> Array[Array[Double]] {
  copy_double_arrays(self.m)
}

///|
pub fn AdamWCheckpoint::v(self : AdamWCheckpoint) -> Array[Array[Double]] {
  copy_double_arrays(self.v)
}

///|
pub fn AdamWCheckpoint::lr(self : AdamWCheckpoint) -> Double {
  self.lr
}

///|
pub fn AdamWCheckpoint::beta1(self : AdamWCheckpoint) -> Double {
  self.beta1
}

///|
pub fn AdamWCheckpoint::beta2(self : AdamWCheckpoint) -> Double {
  self.beta2
}

///|
pub fn AdamWCheckpoint::eps(self : AdamWCheckpoint) -> Double {
  self.eps
}

///|
pub fn AdamWCheckpoint::weight_decays(self : AdamWCheckpoint) -> Array[Double] {
  self.weight_decays.copy()
}

///|
pub fn AdamWCheckpoint::step(self : AdamWCheckpoint) -> Int {
  self.step
}

///|
pub fn AdamWCheckpoint::beta1_power(self : AdamWCheckpoint) -> Double {
  self.beta1_power
}

///|
pub fn AdamWCheckpoint::beta2_power(self : AdamWCheckpoint) -> Double {
  self.beta2_power
}