// optim.mbt

///|
/// Stochastic Gradient Descent (SGD) optimizer.
pub struct SGD {
  params : Array[Tensor]
  lr : Double
}

///|
/// Create a new SGD optimizer.
pub fn SGD::new(params : Array[Tensor], lr : Double) -> SGD {
  { params, lr }
}

///|
/// Perform a single optimization step (updating parameter data).
pub fn SGD::step(self : SGD) -> Unit {
  for param in self.params {
    match param.grad {
      None => ()
      Some(g) => {
        let size = param.data.length()
        for i in 0.. Unit {
  for param in self.params {
    match param.grad {
      None => ()
      Some(g) => {
        let size = g.length()
        for i in 0.. Unit {
  for param in params {
    match param.grad {
      None => ()
      Some(g) =>
        for i in 0.. Double {
  let mut total = 0.0
  for param in params {
    match param.grad {
      None => ()
      Some(g) =>
        for value in g {
          total = total + value * value
        }
    }
  }
  total.sqrt()
}

///|
/// Scale gradients in place when their combined norm exceeds max_norm.
pub fn clip_grad_norm(params : Array[Tensor], max_norm : Double) -> Double {
  if max_norm <= 0.0 {
    panic()
  }
  let norm = grad_norm(params)
  if norm > max_norm && norm > 0.0 {
    let scale = max_norm / norm
    for param in params {
      match param.grad {
        None => ()
        Some(g) =>
          for i in 0.. MomentumSGD {
  if lr <= 0.0 || momentum < 0.0 || momentum >= 1.0 {
    panic()
  }
  let velocities : Array[Array[Double]] = Array::make(params.length(), [])
  for i in 0.. Unit {
  for p in 0.. ()
      Some(g) => {
        let velocity = self.velocities[p]
        for i in 0.. Unit {
  zero_grad(self.params)
}

///|
/// Return the number of parameters managed by an optimizer.
pub fn MomentumSGD::parameter_count(self : MomentumSGD) -> Int {
  let mut count = 0
  for param in self.params {
    count = count + param.data.length()
  }
  count
}