// autograd.mbt

///|
/// Enum representing operators in the computation graph.
pub enum Op {
  Add(Tensor, Tensor)
  Sub(Tensor, Tensor)
  Mul(Tensor, Tensor)
  Div(Tensor, Tensor)
  MatMul(Tensor, Tensor)
  Reshape(Tensor, Array[Int])
  Transpose(Tensor, Int, Int)
  Slice(Tensor, Int, Int, Int)
  SliceMulti(Tensor, Array[(Int, Int)])
  ReLU(Tensor)
  Sigmoid(Tensor)
  Tanh(Tensor)
  MSELoss(Tensor, Tensor)
  ReduceSum(Tensor, Int?, Bool, Bool)
  Softmax(Tensor, Int)
  LogSoftmax(Tensor, Int)
  CrossEntropy(Tensor, Tensor)
  Abs(Tensor)
  Sqrt(Tensor)
  Exp(Tensor)
  Log(Tensor)
  Pow(Tensor, Double)
  Clamp(Tensor, Double, Double)
}

///|
/// Helper to get the inputs of an operation.
fn get_op_inputs(op : Op) -> Array[Tensor] {
  match op {
    Add(a, b) => [a, b]
    Sub(a, b) => [a, b]
    Mul(a, b) => [a, b]
    Div(a, b) => [a, b]
    MatMul(a, b) => [a, b]
    Reshape(a, _) => [a]
    Transpose(a, _, _) => [a]
    Slice(a, _, _, _) => [a]
    SliceMulti(a, _) => [a]
    ReLU(a) => [a]
    Sigmoid(a) => [a]
    Tanh(a) => [a]
    MSELoss(a, b) => [a, b]
    ReduceSum(a, _, _, _) => [a]
    Softmax(a, _) => [a]
    LogSoftmax(a, _) => [a]
    CrossEntropy(a, b) => [a, b]
    Abs(a) => [a]
    Sqrt(a) => [a]
    Exp(a) => [a]
    Log(a) => [a]
    Pow(a, _) => [a]
    Clamp(a, _, _) => [a]
  }
}

///|
/// Performs topological sort starting from the root tensor.
fn topological_sort(root : Tensor) -> Array[Tensor] {
  let visited : Array[Tensor] = Array::new()
  let result : Array[Tensor] = Array::new()

  fn visit(node : Tensor) {
    let mut found = false
    for v in visited {
      if v.id == node.id {
        found = true
        break
      }
    }
    if !found {
      visited.push(node)
      match node.creator {
        Some(op) => {
          let inputs = get_op_inputs(op)
          for input in inputs {
            visit(input)
          }
        }
        None => ()
      }
      result.push(node)
    }
  }

  visit(root)
  result
}

///|
/// Sum-reduce gradient delta along broadcasted dimensions.
fn reduce_grad(
  delta : Array[Double],
  out_shape : Array[Int],
  target_shape : Array[Int],
) -> Array[Double] {
  let target_len = target_shape.length()
  let mut target_size = 1
  for dim in target_shape {
    target_size = target_size * dim
  }
  if target_len == 0 {
    target_size = 1
  }

  let target_grad = Array::make(target_size, 0.0)
  let out_strides = shape_to_strides(out_shape)
  let target_strides = shape_to_strides(target_shape)

  let size = delta.length()
  for i in 0.. Unit {
  if !t.requires_grad {
    return
  }
  match t.grad {
    None => {
      let g = Array::make(t.data.length(), 0.0)
      for i in 0..
      for i in 0.. Unit {
  if !t.requires_grad {
    return
  }
  let reduced = if t.shape != out_shape {
    reduce_grad(delta, out_shape, t.shape)
  } else {
    delta
  }

  match t.grad {
    None => {
      let g = Array::make(t.data.length(), 0.0)
      for i in 0..
      for i in 0.. Unit {
  match op {
    Add(a, b) => {
      accumulate_grad(a, grad, out_shape)
      accumulate_grad(b, grad, out_shape)
    }
    Sub(a, b) => {
      accumulate_grad(a, grad, out_shape)
      let neg_grad = Array::make(grad.length(), 0.0)
      for i in 0.. {
      let out_strides = shape_to_strides(out_shape)
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0.. {
      let out_strides = shape_to_strides(out_shape)
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0.. {
      let m = a.shape[0]
      let k = a.shape[1]
      let n = b.shape[1]
      if a.requires_grad {
        let grad_a = Array::make(m * k, 0.0)
        for r in 0.. accumulate_grad_direct(a, grad)
    Transpose(a, dim0, dim1) =>
      if a.requires_grad {
        let size = grad.length()
        let grad_a = Array::make(a.data.length(), 0.0)
        let len = a.shape.length()
        let out_strides = shape_to_strides(out_shape)

        for i in 0..
      if a.requires_grad {
        let size = grad.length()
        let grad_a = Array::make(a.data.length(), 0.0)
        let len = a.shape.length()
        let out_strides = shape_to_strides(out_shape)

        for i in 0..
      if a.requires_grad {
        let size = grad.length()
        let grad_a = Array::make(a.data.length(), 0.0)
        let len = a.shape.length()
        let out_strides = shape_to_strides(out_shape)

        for i in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0.. 0.0 { grad[i] } else { 0.0 }
        }
        accumulate_grad(a, grad_a, out_shape)
      }
    Sigmoid(a) =>
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0.. 0.0 {
            1.0
          } else if a.data[i] < 0.0 {
            -1.0
          } else {
            0.0
          }
          grad_a[i] = grad[i] * sign
        }
        accumulate_grad_direct(a, grad_a)
      }
    Sqrt(a) =>
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0..= 0 {
                for _ in 0..
      if a.requires_grad {
        let grad_a = Array::make(grad.length(), 0.0)
        for i in 0.. lower && value < upper { grad[i] } else { 0.0 }
        }
        accumulate_grad_direct(a, grad_a)
      }
    ReduceSum(a, dim, keepdim, mean) =>
      if a.requires_grad {
        let grad_a = Array::make(a.data.length(), 0.0)
        match dim {
          None => {
            let factor = if mean {
              1.0 / a.data.length().to_double()
            } else {
              1.0
            }
            for i in 0.. {
            let count = a.shape[d].to_double()
            let factor = if mean { 1.0 / count } else { 1.0 }
            let output_strides = shape_to_strides(out_shape)
            for i in 0..
      if a.requires_grad {
        let values = softmax_data(a, dim)
        let width = a.shape[dim]
        let stride = a.strides[dim]
        let grad_a = Array::make(a.data.length(), 0.0)
        for i in 0..
      if a.requires_grad {
        let values = log_softmax_data(a, dim)
        let width = a.shape[dim]
        let stride = a.strides[dim]
        let grad_a = Array::make(a.data.length(), 0.0)
        for i in 0..
      if logits.requires_grad {
        let batch = logits.shape[0]
        let classes = logits.shape[1]
        let grad_logits = Array::make(logits.data.length(), 0.0)
        for row in 0.. maximum {
              maximum = logits.data[offset + col]
            }
          }
          let mut denominator = 0.0
          for col in 0.. {
      let size = pred.data.length()
      let factor = 2.0 / size.to_double() * grad[0]
      if pred.requires_grad {
        let grad_pred = Array::make(size, 0.0)
        for i in 0.. Unit {
  if self.grad is None {
    self.grad = Some(Array::make(self.data.length(), 1.0))
  }

  let sorted = topological_sort(self)
  let len = sorted.length()
  for i in 0.. ()
      Some(op) =>
        match t.grad {
          None => ()
          Some(g) => propagate_op_grad(op, g, t.shape)
        }
    }
  }
}