// gat_backward.mbt -- GAT backward pass (v0.153.0).
//
// The last message-passing architecture without a gradient. GAT is
// harder than everything in v0.145.0-v0.148.0 for one structural
// reason, and that reason is the whole content of this file:
//
// **the attention score couples BOTH endpoints of one edge.**
//
// Every reducer already shipped -- sum, mean, max, min, std, GCN's
// weighted sum -- sends an edge's influence to exactly ONE end (the
// source, whose features are being aggregated). GAT's score
//
//   e_uv = a_src . (W h_u) + a_dst . (W h_v)
//
// reads both h_u and h_v, so dL/de_uv fans out to BOTH. That needs
// two scatter DIRECTIONS in one pass: the source-half of the score
// gradient accumulates on the source node, the destination-half on the
// destination node. The existing `scatter_*_backward` family only ever
// scatters destination -> source, because no other architecture here
// needed the other way.
//
// The backward, in order:
//   1. d_p[s]   += alpha_e * d_out[d]        weighted sum, to source
//      d_alpha_e  = 
//   2. softmax per destination d (over its IN-edges):
//      d_logits_e = alpha_e * (d_alpha_e - SUM_e' alpha_e' d_alpha_e')
//   3. d_tsrc[s] += d_logits_e                <- fan-out, source side
//      d_tdst[d] += d_logits_e                <- fan-out, dest side
//   4. t_src = leaky_relu(a_src . p):  d_a_src[o] += d_s[v] * p[v][o]
//                                    d_p[v][o]   += d_s[v] * a_src[o]
//      t_dst = a_dst . p (linear):   d_a_dst[o] += d_tdst[v] * p[v][o]
//                                    d_p[v][o]   += d_tdst[v] * a_dst[o]
//   5. d_W[o][k] += SUM_v d_p[v][o] * h[v][k]      (outer product, nodes)
//      d_h[v][k] += SUM_o d_p[v][o] * W[o][k]      (per node, dense)
//
// GATLayer carries NO bias (the paper's single-head form), so there is
// no d_b.

///|
/// Parameter gradients for one GAT head. `d_w` mirrors the layer's
/// (out_dim x in_dim) projection; there is no bias term.
pub struct GATLayerGrad {
  d_w : Array[Array[Float]]
  d_a_src : Array[Float]
  d_a_dst : Array[Float]
}

///|
/// Zero-initialised gradients matching a GATLayer. Rows are built
/// EXPLICITLY: `Array::make(out_dim, Array::make(in_dim, 0.0F))` hands
/// back out_dim aliases of one shared inner array in MoonBit, which
/// silently collapses every row of d_w into row 0.
pub fn GATLayerGrad::zero(layer : GATLayer) -> GATLayerGrad {
  let dw : Array[Array[Float]] = Array::make(layer.out_dim, [])
  for o in 0.. Array[Float] {
  let d = layer.out_dim
  let p : Array[Float] = Array::make(n_nodes * d, 0.0F)
  for v in 0.. (Array[Float], Array[Float]) {
  let d = layer.out_dim
  let t_src : Array[Float] = Array::make(n_nodes, 0.0F)
  let t_dst : Array[Float] = Array::make(n_nodes, 0.0F)
  for v in 0.. Array[Float] {
  let logits : Array[Float] = Array::make(g.n_edges, 0.0F)
  for e in 0..= t_src.length() || d < 0 || d >= t_dst.length() {
      logits[e] = 0.0F
    } else {
      logits[e] = t_src[s] + t_dst[d]
    }
  }
  let deg = graph_in_degree(g)
  let alpha : Array[Float] = Array::make(g.n_edges, 0.0F)
  for v in 0.. m {
        m = logits[e]
      }
    }
    let mut sum_exp = 0.0F
    for e in 0.. Array[Float] {
  let d_logits : Array[Float] = Array::make(g.n_edges, 0.0F)
  for v in 0.. source,
/// because no other architecture reads both endpoints of an edge.
pub fn scatter_add_by_src(
  g : Graph,
  per_edge : Array[Float],
) -> Array[Float] {
  let out : Array[Float] = Array::make(g.n_nodes, 0.0F)
  for e in 0..= g.n_nodes {
      continue
    }
    out[s] = out[s] + per_edge[e]
  }
  out
}

///|
/// Accumulate per-edge values onto their DESTINATION nodes.
pub fn scatter_add_by_dst(
  g : Graph,
  per_edge : Array[Float],
) -> Array[Float] {
  let out : Array[Float] = Array::make(g.n_nodes, 0.0F)
  for e in 0..= g.n_nodes {
      continue
    }
    out[d] = out[d] + per_edge[e]
  }
  out
}

///|
/// Backward of one GAT head. Returns `(d_input, grad)`.
///
/// `elu` selects the hidden-layer activation: pass true when this head
/// sits in a hidden layer of a multi-head stack (whose concatenated
/// output is passed through ELU) and false for the final layer, whose
/// output is linear and therefore needs no derivative mask.
///
/// The mask is applied HERE, not by the caller, because it needs the
/// head's own PRE-activation -- the same value `gat_layer_forward`
/// produces -- and a caller that had to recompute it would be one more
/// place for the forward and the backward to disagree. It also means
/// `elu = true` cannot silently degrade to a no-op, which is what this
/// parameter used to be: the single-head caller below reached the tail
/// of this function, found a placeholder, and returned an UNMASKED
/// gradient for every hidden head. An unmasked gradient is not a small
/// error -- it is the gradient of a different network.
pub fn gat_layer_backward(
  layer : GATLayer,
  g : Graph,
  h : Array[Float],
  d_out : Array[Float],
  elu : Bool,
) -> (Array[Float], GATLayerGrad) {
  let n = g.n_nodes
  let d = layer.out_dim
  let grad = GATLayerGrad::zero(layer)
  let p = gat_project(layer, h, n)
  let (t_src, t_dst) = gat_attn_terms(layer, p, n)
  let alpha = gat_alpha(g, t_src, t_dst)
  // ELU mask on the incoming gradient. `d_out` belongs to the caller,
  // so the masked copy is a fresh buffer; `d_out` is never mutated.
  let eff = if elu { gat_apply_elu_mask(layer, g, h, d_out) } else { d_out }
  // 1. Weighted sum: out[d] += alpha_e * p[s]. The gradient goes to
  //    the SOURCE (that is the only node the sum reads), and the
  //    coefficient's gradient is the inner product with p[s].
  let d_p : Array[Float] = Array::make(n * d, 0.0F)
  let d_alpha : Array[Float] = Array::make(g.n_edges, 0.0F)
  for e in 0..= n || dst < 0 || dst >= n {
      continue
    }
    let a = alpha[e]
    let src_off = s * d
    let dst_off = dst * d
    let mut dot = 0.0F
    for o in 0.. 0 and the negative slope 0.2 otherwise.
    let mut s = 0.0F
    for o in 0.. 0.0F { 1.0F } else { 0.2F }
    let d_s = d_tsrc[v] * slope
    let grad_asrc = grad.d_a_src
    let grad_adst = grad.d_a_dst
    for o in 0..