// 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.. Unit {
let n = if a.length() < b.length() { a.length() } else { b.length() }
for i in 0.. Array[Float] {
let n = if a.length() < b.length() { a.length() } else { b.length() }
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0.. Array[Float] {
let n = if a.length() < b.length() { a.length() } else { b.length() }
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0..