// gin_backward.mbt -- GIN and MPNN backward passes (v0.146.0).
//
// The first two consumers of v0.145.0's scatter primitives. Both
// architectures have the same shape: a self term, a neighbour term,
// a Linear, and a ReLU. The backward therefore has three distinct
// stages, and the middle one is the one that was previously missing
// from the whole repository -- routing an upstream gradient back
// through the edge list into the *input* embedding.
//
// GIN's forward (v0.141.0):
//   a   = (1 + eps) * h + gin_aggregate(g, h)        <- unweighted sum
//   z1  = W1 . a + b1
//   h'  = ReLU(W2 . ReLU(z1) + b2)
//
// Backward, with G' the ReLU indicator and d_* the upstream grads:
//   d_z2  = G'_out . d_h'
//   (d_a, g2) = linear_backward(W2, d_z2, h1act)
//   d_z1  = G'_mid . d_a
//   (d_w1, g1) = linear_backward(W1, d_z1, comb)
//   d_h   = (1 + eps) * d_a + scatter_sum_backward(g, d_a)
//
// The last line is the whole point: `d_a` is scattered back along the
// edge list, so a node's input gradient is its own share PLUS the
// share of every node that reads it. On a hub that is the dominant
// term, which is exactly why a forward-only GNN cannot be trained
// without it.
//
// MPNN (v0.138.0) is the same with mean instead of sum and two
// Linears summed instead of a 2-layer MLP:
//   h' = ReLU( W_self . h + W_neigh . scatter_mean(g, h) )
// so its neighbour gradient carries the 1/deg factor from
// `scatter_mean_backward`.
//
// No forward cache: both backwards recompute what they need from `h`
// and `g`. The recompute is one extra aggregation pass, which is
// cheaper than the bookkeeping a cache would need across 5+ tensor
// shapes, and it keeps the forward signatures from v0.138.0 / v0.141.0
// untouched (no cache means the forward is still usable for inference
// with no extra allocations).

///|
/// Parameter gradients for one GINLayer.
pub struct GINLayerGrad {
  mlp1 : GraphLinearGrad
  mlp2 : GraphLinearGrad
}

///|
/// Parameter gradients for a whole GIN stack.
pub struct GINGrad {
  layers : Array[GINLayerGrad]
  num_layers : Int
}

///|
/// Zero-initialised GIN gradients for a stack.
pub fn GINGrad::zero(net : GIN) -> GINGrad {
  let layers : Array[GINLayerGrad] = Array::make(
    net.num_layers,
    GINLayerGrad::{
      mlp1: GraphLinearGrad::zero(net.layers[0].mlp1),
      mlp2: GraphLinearGrad::zero(net.layers[0].mlp2),
    },
  )
  for l in 0..