// autodiff_tape.mbt — Tape-based reverse-mode automatic differentiation (v0.31.0).
//
// Provides a `Tape[A]` recorder that captures a sequence of arithmetic
// operations during the forward pass, then computes gradients via a
// single reverse traversal. Supports both forward-mode (derivative of
// every intermediate w.r.t. one input) and backward-mode (gradient of
// one scalar output w.r.t. all inputs).
//
// API design follows moonbit-pilot/autodiff:
//   - `Loc[A] = Const(A) | Memory(Int)` — pointers to constant values
//     (no gradient) or memory slots produced by earlier instructions.
//   - `Inst[A] = Var(A) | Unary(...) | Binary(...)` — recorded ops.
//   - `Diffable` trait requires `Add + Sub` with `zero`/`one` literals.
//   - `BasicPrims[A]` wraps a Tape with algebraic ops
//     (`neg`, `add`, `sub`, `mul`, `div`).
//   - `MathPrims` wraps a `Tape[Float]` with transcendental ops
//     (`exp`, `ln`, `sin`, `cos`, `tanh`, `sigmoid`).
//
// Typical usage:
//   let tape : Tape[Float] = Tape::new()
//   let prims = BasicPrims::on(tape)
//   let x0 = tape.variable(2.0F)
//   let x1 = tape.variable(5.0F)
//   let res = prims.add(prims.neg(x0), prims.mul(x0, x1))
//   let mem = tape.eval()                              // forward pass
//   let grads = tape.diff_backward(mem)                // ∂res/∂x for all x
//
// Convention for backward functions:
//   - op1: `bwd(g, x)` returns `g · df/dx` evaluated at the input `x`.
//   - op2: `bwd_lhs(g, a, b)` returns `g · df/dlhs` at `(a, b)`;
//          `bwd_rhs(g, a, b)` returns `g · df/drhs` at `(a, b)`.

///|
/// Reference to a value in the tape — either a literal constant
/// (no gradient propagation) or a memory slot produced by an earlier
/// instruction.
pub enum Loc[A] {
  Const(A)
  Memory(Int)
}

///|
/// Recorded instruction in the tape. `Var` stores a literal value;
/// `Unary` and `Binary` capture operations with their forward and
/// backward callbacks.
pub enum Inst[A] {
  /// A variable (input) — the constant value is stored directly so
  /// `eval` can place it at the corresponding memory slot. Backward
  /// does not propagate to the input (variables are leaves).
  Var(A)
  /// A unary operation. `fwd` maps the input to the output;
  /// `bwd` maps `(g, original_input)` to the input's gradient
  /// (`g · df/dx` at the input).
  Unary((A) -> A, (A, A) -> A, Loc[A])
  /// A binary operation. `fwd` maps `(lhs, rhs)` to the output;
  /// `bwd_lhs(g, lhs, rhs)` returns the lhs gradient
  /// (`g · df/dlhs` at `(lhs, rhs)`); `bwd_rhs(g, lhs, rhs)` returns
  /// the rhs gradient (`g · df/drhs` at `(lhs, rhs)`).
  Binary((A, A) -> A, (A, A, A) -> A, (A, A, A) -> A, Loc[A], Loc[A])
}

///|
/// A linear sequence of instructions recorded during the forward
/// pass. Use `Tape::variable` to declare inputs, `Tape::op1`/`op2` to
/// record operations, `Tape::eval` to run forward, and
/// `Tape::diff_backward` (or `diff_forward`) to compute gradients.
pub struct Tape[A] {
  insts : Array[Inst[A]]
  names : Array[String]
}

///|
/// Trait required for any type that can be used with the tape.
/// Pre-implemented for `Float` (Float32) and `Double` (Float64).
pub(open) trait Diffable : Add + Sub {
  zero() -> Self
  one() -> Self
}

pub impl Diffable for Float with zero() {
  0.0F
}

pub impl Diffable for Float with one() {
  1.0F
}

pub impl Diffable for Double with zero() {
  0.0
}

pub impl Diffable for Double with one() {
  1.0
}

///|
/// Create an empty tape. All subsequent operations are appended to
/// this tape in order.
pub fn[A] Tape::new() -> Tape[A] {
  { insts: [], names: [] }
}

///|
/// Append a variable with the given value to the tape. Returns the
/// memory slot where the variable's value will be stored. Variables
/// are leaves of the computation graph — they receive gradient in
/// `diff_backward` but do not propagate further.
pub fn[A] Tape::variable(self : Tape[A], value : A) -> Loc[A] {
  let idx = self.insts.length()
  self.insts.push(Inst::Var(value))
  self.names.push("x\{idx}")
  Loc::Memory(idx)
}

///|
/// Wrap a literal value as a `Loc::Const`. Constants do not receive
/// gradient in backward mode — they are not in the parameter list.
pub fn[A] constant(value : A) -> Loc[A] {
  Loc::Const(value)
}

///|
/// Append a unary operation to the tape. Returns a function that
/// takes a `Loc[A]` input and produces a `Loc[A]` for the output.
///
/// `bwd : (A, A) -> A` takes the upstream gradient and the original
/// input value and returns the input's gradient
/// (`g · df/dx` evaluated at the input).
pub fn[A] Tape::op1(
  self : Tape[A],
  name : String,
  fwd : (A) -> A,
  bwd : (A, A) -> A,
) -> (Loc[A]) -> Loc[A] {
  fn(input : Loc[A]) -> Loc[A] {
    let idx = self.insts.length()
    self.insts.push(Inst::Unary(fwd, bwd, input))
    self.names.push(name)
    Loc::Memory(idx)
  }
}

///|
/// Append a binary operation to the tape. Returns a function that
/// takes two `Loc[A]` inputs and produces a `Loc[A]` for the output.
///
/// `bwd_lhs(g, lhs, rhs)` returns the lhs gradient
/// (`g · df/dlhs` at `(lhs, rhs)`); `bwd_rhs(g, lhs, rhs)` returns the
/// rhs gradient. Both functions receive BOTH original operands so
/// they can implement partial derivatives that depend on the partner
/// (e.g. `df/db = -a/b²` for `a / b`).
pub fn[A] Tape::op2(
  self : Tape[A],
  name : String,
  fwd : (A, A) -> A,
  bwd_lhs : (A, A, A) -> A,
  bwd_rhs : (A, A, A) -> A,
) -> (Loc[A], Loc[A]) -> Loc[A] {
  fn(lhs : Loc[A], rhs : Loc[A]) -> Loc[A] {
    let idx = self.insts.length()
    self.insts.push(Inst::Binary(fwd, bwd_lhs, bwd_rhs, lhs, rhs))
    self.names.push(name)
    Loc::Memory(idx)
  }
}

///|
/// Resolve a `Loc[A]` to its concrete value at runtime.
fn[A] resolve_loc(loc : Loc[A], mem : Array[A]) -> A {
  match loc {
    Const(c) => c
    Memory(j) => mem[j]
  }
}

///|
/// Run the forward pass, evaluating every instruction and writing its
/// result into the corresponding memory slot. Returns an array of
/// length `insts.length()` where `mem[i]` is the value of the i-th
/// instruction (variables first, then operations in order).
pub fn[A : Diffable] Tape::eval(self : Tape[A]) -> Array[A] {
  let n = self.insts.length()
  let mem : Array[A] = Array::make(n, A::zero())
  for i in 0.. mem[i] = v
      Unary(fwd, _bwd, input) => mem[i] = fwd(resolve_loc(input, mem))
      Binary(fwd, _blhs, _brhs, lhs, rhs) =>
        mem[i] = fwd(resolve_loc(lhs, mem), resolve_loc(rhs, mem))
    }
  }
  mem
}

///|
/// Compute backward-mode gradients. Returns an array `grads` where
/// `grads[i]` is `∂output / ∂inst[i]` — i.e. the partial derivative
/// of the last instruction (assumed to be the scalar output) with
/// respect to the i-th instruction. The gradient of the output w.r.t.
/// itself is seeded to `A::one()`.
pub fn[A : Diffable] Tape::diff_backward(
  self : Tape[A],
  mem : Array[A],
) -> Array[A] {
  let n = self.insts.length()
  let grads : Array[A] = Array::make(n, A::zero())
  if n == 0 {
    return grads
  }
  // dL/d(output) = 1
  grads[n - 1] = A::one()
  // Reverse traversal — propagate gradient from each instruction
  // back to its inputs.
  for k in 0.. () // variables are leaves; no further propagation
      Unary(_fwd, bwd, input) => {
        let x : A = resolve_loc(input, mem)
        match input {
          Memory(j) => {
            let g_in : A = bwd(g, x)
            grads[j] = grads[j] + g_in
          }
          Const(_) => ()
        }
      }
      Binary(_fwd, bwd_lhs, bwd_rhs, lhs, rhs) => {
        let a : A = resolve_loc(lhs, mem)
        let b : A = resolve_loc(rhs, mem)
        match lhs {
          Memory(j) => {
            let g_in : A = bwd_lhs(g, a, b)
            grads[j] = grads[j] + g_in
          }
          Const(_) => ()
        }
        match rhs {
          Memory(j) => {
            let g_in : A = bwd_rhs(g, a, b)
            grads[j] = grads[j] + g_in
          }
          Const(_) => ()
        }
      }
    }
  }
  grads
}

///|
/// Compute forward-mode derivative of every intermediate value with
/// respect to the variable at index `wrt`. Returns an array
/// `df[i] = ∂inst[i] / ∂inst[wrt]`. The seed `df[wrt] = 1`.
///
/// **Caveat**: the backward functions for our primitive ops expect
/// the *upstream gradient* as the first argument. To re-use them in
/// forward mode, we substitute `1` for the gradient — this gives the
/// **local derivative at the input** (since `bwd(1, x) = df/dx` for
/// our convention where `bwd(g, x) = g · df/dx`).
pub fn[A : Diffable + Mul + Add] Tape::diff_forward(
  self : Tape[A],
  mem : Array[A],
  wrt~ : Int = 0,
) -> Array[A] {
  let n = self.insts.length()
  let df : Array[A] = Array::make(n, A::zero())
  if n > wrt {
    df[wrt] = A::one()
  }
  for i in 0.. () // leaves — seed is already there
      Unary(_fwd, bwd, input) => {
        let x : A = resolve_loc(input, mem)
        let d_in : A = match input {
          Memory(j) => df[j]
          Const(_) => A::zero()
        }
        // Local derivative at the input: bwd(1, x) = df/dx.
        let local_d : A = bwd(A::one(), x)
        df[i] = local_d * d_in
      }
      Binary(_fwd, bwd_lhs, bwd_rhs, lhs, rhs) => {
        let a : A = resolve_loc(lhs, mem)
        let b : A = resolve_loc(rhs, mem)
        let da : A = match lhs {
          Memory(j) => df[j]
          Const(_) => A::zero()
        }
        let db : A = match rhs {
          Memory(j) => df[j]
          Const(_) => A::zero()
        }
        let local_lhs : A = bwd_lhs(A::one(), a, b)
        let local_rhs : A = bwd_rhs(A::one(), a, b)
        df[i] = local_lhs * da + local_rhs * db
      }
    }
  }
  df
}