// graph_backward.mbt -- Gradient primitives shared by every GNN
// backward (v0.145.0).
//
// Batches V (v0.133.0-v0.140.0) and X (v0.141.0-v0.144.0) shipped
// six message-passing architectures, all forward-only. They all defer
// to the SAME hard problem, and it is not the Linear algebra -- it is
// the scatter.
//
// A GNN forward is `out[d] = AGG over edges e ending at d`. Its
// backward is a *gather-scatter* in the opposite direction: every
// incoming edge of node d must READ d's upstream gradient and route a
// share of it back to its own source. For a node with in-degree d_v
// that is d_v writes into one row, which is why this is the part worth
// getting right once rather than seven times:
//
//   d_messages[s] += edge_weight[e] * d_out[d]        (sum)
//   d_messages[s] += edge_weight[e] * d_out[d] / deg[d] (mean)
//   d_messages[s] += d_out[d]        at the argmax only (max/min)
//
// The third is the only one that is not a pure scatter: the gradient
// of a max is a ROUTING, not a sum -- it must be routed to the single
// element that won, which means the forward pass has to remember where
// that was. Hence `scatter_max_forward_with_idx` below: the forward
// is re-run to recover the argmax, and the max backward is a *gather*
// (every winner reads its own gradient) rather than a scatter.
//
// Scope of v0.145.0:
//   - scatter_sum_backward / scatter_mean_backward (edge-list gradient
//     accumulation -- the shared primitive).
//   - scatter_max_forward_with_idx + scatter_max_backward /
//     scatter_min_forward_with_idx + scatter_min_backward (routing).
//   - GraphLinearGrad + graph_linear_backward.
//   - graph_linear_sgd_step: one vanilla SGD step on a GraphLinear.
//
// The argument-gradient correctness of these primitives is exactly
// the contract `linear_backward` in linear_backward.mbt already has
// for a dense layer, so the two are written to be read side by side.

///|
/// Backward of `scatter_sum`. Routes each destination's upstream
/// gradient back along its incoming edges:
///
///   d_messages[s] += edge_weight[e] * d_out[d]
///
/// Note this is the same edge-list walk as the forward, so a node
/// accumulates one contribution per incoming edge. Nodes with no
/// outgoing contribution keep a zero row.
///|
/// Backward of `scatter_sum` with an EXPLICIT message width.
///
/// The width parameter exists because a GNN layer's internal tensors
/// are not the node-feature width. `gcn_layer_support` produces
/// `[n_nodes x layer.in_dim]`, which for a hidden layer is wider than
/// `g.feat_dim`; reading it with `g.feat_dim` silently truncates the
/// gradient routed back to the previous layer, and the truncation
/// surfaces one layer later as an out-of-bounds read.
///
/// `scatter_sum_backward` below is the `g.feat_dim` special case.
pub fn scatter_sum_backward_w(
  g : Graph,
  d_out : Array[Float],
  width : Int,
) -> Array[Float] {
  let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
  for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
      continue
    }
    let w = g.edge_weight[e]
    let src_off = s * width
    let dst_off = d * width
    for k in 0.. Array[Float] {
  scatter_sum_backward_w(g, d_out, g.feat_dim)
}

///|
/// Backward of `scatter_mean` with an EXPLICIT message width. See
/// `scatter_sum_backward_w` for why the width cannot be assumed.
pub fn scatter_mean_backward_w(
  g : Graph,
  d_out : Array[Float],
  width : Int,
) -> Array[Float] {
  let deg = graph_in_degree(g)
  let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
  for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
      continue
    }
    if deg[d] == 0 {
      continue
    }
    let w = g.edge_weight[e]
    let scale = w / Float::from_int(deg[d])
    let src_off = s * width
    let dst_off = d * width
    for k in 0.. Array[Float] {
  scatter_mean_backward_w(g, d_out, g.feat_dim)
}

///|
/// Scatter-max forward that also records, per destination node and
/// per feature, WHICH incoming edge won. The backward needs the argmax
/// and cannot recover it from the forward's output alone (many
/// neighbourhoods share the same maximum value).
///
/// Ties go to the FIRST edge in list order, matching the forward's
/// `first || m > out[...]` tie-break, so the recorded index and the
/// value the forward actually produced always agree.
pub fn scatter_max_forward_with_idx(
  g : Graph,
  messages : Array[Float],
) -> (Array[Float], Array[Int]) {
  let dim = g.feat_dim
  let out : Array[Float] = Array::make(g.n_nodes * dim, 0.0F)
  let argmax : Array[Int] = Array::make(g.n_nodes * dim, -1)
  let seeded : Array[Int] = Array::make(g.n_nodes, 0)
  for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
      continue
    }
    let w = g.edge_weight[e]
    let src_off = s * dim
    let dst_off = d * dim
    let first = seeded[d] == 0
    seeded[d] = 1
    for k in 0.. out[dst_off + k] {
        out[dst_off + k] = m
        argmax[dst_off + k] = e
      }
    }
  }
  (out, argmax)
}

///|
/// Backward of `scatter_max`. This is a ROUTING, not a scatter: the
/// gradient goes only to the edge that produced the maximum.
///
///   for each (d, k):  d_messages[src(argmax[d,k])][k] += w * d_out[d][k]
///
/// `argmax` must come from `scatter_max_forward_with_idx` on the SAME
/// messages; passing an index array from a different forward (or from
/// the plain `scatter_max`) routes the gradient to the wrong edge.
pub fn scatter_max_backward(
  g : Graph,
  d_out : Array[Float],
  argmax : Array[Int],
) -> Array[Float] {
  let dim = g.feat_dim
  let out : Array[Float] = Array::make(g.n_nodes * dim, 0.0F)
  for d in 0..= g.n_edges {
        continue
      }
      let s = g.edge_src[e]
      if s < 0 || s >= g.n_nodes {
        continue
      }
      let scale = g.edge_weight[e]
      let src_off = s * dim
      out[src_off + k] = out[src_off + k] + scale * d_out[dst_off + k]
    }
  }
  out
}

///|
/// Scatter-min forward that records the argmin. Same tie-break (first
/// edge wins) as `scatter_max_forward_with_idx`.
pub fn scatter_min_forward_with_idx(
  g : Graph,
  messages : Array[Float],
) -> (Array[Float], Array[Int]) {
  let dim = g.feat_dim
  let out : Array[Float] = Array::make(g.n_nodes * dim, 0.0F)
  let argmin : Array[Int] = Array::make(g.n_nodes * dim, -1)
  let seeded : Array[Int] = Array::make(g.n_nodes, 0)
  for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
      continue
    }
    let w = g.edge_weight[e]
    let src_off = s * dim
    let dst_off = d * dim
    let first = seeded[d] == 0
    seeded[d] = 1
    for k in 0.. Array[Float] {
  scatter_max_backward(g, d_out, argmin)
}

///|
/// Parameter gradients for a `GraphLinear`. `d_w` mirrors the
/// (out_dim x in_dim) shape of the weight matrix; `d_b` matches the
/// bias length.
pub struct GraphLinearGrad {
  d_w : Array[Array[Float]]
  d_b : Array[Float]
}

///|
/// Zero-initialised gradients matching a GraphLinear's shapes.
///
/// The rows are built EXPLICITLY rather than with
/// `Array::make(out_dim, Array::make(in_dim, 0.0F))`. MoonBit arrays
/// are reference types, so the latter hands back `out_dim` aliases of
/// one shared inner array -- every row of `d_w` would be the same
/// memory, and accumulating into row o would corrupt all of them. The
/// symptom is a `d_w` that is right only in row 0.
pub fn GraphLinearGrad::zero(lin : GraphLinear) -> GraphLinearGrad {
  let dw : Array[Array[Float]] = Array::make(lin.out_dim, [])
  for o in 0.. (Array[Float], GraphLinearGrad) {
  let d_input : Array[Float] = Array::make(n_nodes * lin.in_dim, 0.0F)
  let grad = GraphLinearGrad::zero(lin)
  for v in 0.. Int {
  lin.in_dim * lin.out_dim + lin.out_dim
}

///|
/// One vanilla SGD step on a GraphLinear. Returns a FRESH GraphLinear
/// (the input is not mutated) so a parameter can be threaded through a
/// `for` loop the same way the forward outputs are.
pub fn graph_linear_sgd_step(
  lin : GraphLinear,
  grad : GraphLinearGrad,
  lr : Float,
) -> GraphLinear {
  let w : Array[Array[Float]] = Array::make(lin.out_dim, [])
  for o in 0..