// 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.. (Array[Float], Array[Float]) {
let dim = layer.in_dim
let agg = gin_aggregate_w(g, h, dim)
let comb_size = g.n_nodes * dim
let comb : Array[Float] = Array::make(comb_size, 0.0F)
for i in 0.. 0` test makes
/// the subgradient 0 there, which is what this reproduces.
fn relu_mask(pre : Array[Float]) -> Array[Float] {
let m : Array[Float] = Array::make(pre.length(), 0.0F)
for i in 0.. 0.0F { 1.0F } else { 0.0F }
}
m
}
///|
/// Backward of one GIN layer. Returns `(d_input, grad)`.
pub fn gin_layer_backward(
layer : GINLayer,
g : Graph,
h : Array[Float],
d_out : Array[Float],
) -> (Array[Float], GINLayerGrad) {
let (comb, mid) = gin_layer_intermediates(layer, g, h)
// Recompute the pre-activations for the derivative masks.
let z2_pre = graph_linear_forward(layer.mlp2, mid, g.n_nodes)
let z1_pre = graph_linear_forward(layer.mlp1, comb, g.n_nodes)
// Stage 3: second Linear.
let d_z2 = elem_mul(d_out, relu_mask(z2_pre))
let (d_mid, g2) = graph_linear_backward(layer.mlp2, d_z2, g.n_nodes, mid)
// Stage 2: first Linear, through its ReLU.
let d_z1 = elem_mul(d_mid, relu_mask(z1_pre))
let (d_comb, g1) = graph_linear_backward(layer.mlp1, d_z1, g.n_nodes, comb)
// Stage 1: the self/neighbour split, and the edge-list scatter.
// The message width is `layer.in_dim`, not `g.feat_dim`: a hidden
// layer's aggregate is wider than the node features, and using
// feat_dim truncates the gradient handed to the previous layer.
let width = layer.in_dim
let d_from_agg = scatter_sum_backward_w(g, d_comb, width)
let d_input : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
for i in 0.. (Array[Float], Array[Array[Float]]) {
let inputs : Array[Array[Float]] = Array::make(net.num_layers, [])
let mut acc = h
for l in 0.. (Array[Float], GINGrad) {
let grads = GINGrad::zero(net)
let mut d = d_out
for i in 0.. (GIN, Array[Float]) {
let (out, inputs) = gin_forward_with_inputs(net, g, h)
let (d_input, grads) = gin_backward(net, g, inputs, d_out)
ignore(out)
let layers : Array[GINLayer] = Array::make(net.num_layers, net.layers[0])
for l in 0.. MPnnGrad {
let layers : Array[MPnnLayerGrad] = Array::make(
net.num_layers,
MPnnLayerGrad::{
w_self: GraphLinearGrad::zero(net.layers[0].w_self),
w_neigh: GraphLinearGrad::zero(net.layers[0].w_neigh),
},
)
for l in 0.. (Array[Float], MPnnLayerGrad) {
// Recompute the two pre-ReLU branches.
let self_pre = graph_linear_forward(layer.w_self, h, g.n_nodes)
let neigh_mean = scatter_mean_w(g, h, layer.in_dim)
let neigh_pre = graph_linear_forward(layer.w_neigh, neigh_mean, g.n_nodes)
let d_input : Array[Float] = Array::make(g.n_nodes * layer.in_dim, 0.0F)
// The two branches sum before the activation, so both receive the
// same masked upstream gradient.
let mask = relu_mask(elem_add(self_pre, neigh_pre))
let d_self_branch = elem_mul(d_out, mask)
let d_neigh_branch = elem_mul(d_out, mask)
// d_h from the self branch. The branch width is layer.in_dim.
let width = layer.in_dim
let (d_from_self, g_self) = graph_linear_backward(
layer.w_self, d_self_branch, g.n_nodes, h,
)
add_into(d_input, d_from_self)
// d_h from the neighbour branch: two hops -- through the Linear,
// then back through the mean aggregation. The mean backward needs
// the same explicit width, for the same reason.
let (d_from_neigh_pre, g_neigh) = graph_linear_backward(
layer.w_neigh, d_neigh_branch, g.n_nodes, neigh_mean,
)
let d_from_agg = scatter_mean_backward_w(g, d_from_neigh_pre, width)
add_into(d_input, d_from_agg)
(d_input, MPnnLayerGrad::{ w_self: g_self, w_neigh: g_neigh, })
}
///|
/// Forward the MPNN stack, caching each layer's input.
pub fn mpnn_forward_with_inputs(
net : MPnn,
g : Graph,
h : Array[Float],
) -> (Array[Float], Array[Array[Float]]) {
let inputs : Array[Array[Float]] = Array::make(net.num_layers, [])
let mut acc = h
for l in 0.. (Array[Float], MPnnGrad) {
let grads = MPnnGrad::zero(net)
let mut d = d_out
for i in 0.. (MPnn, Array[Float]) {
let (out, inputs) = mpnn_forward_with_inputs(net, g, h)
let (d_input, grads) = mpnn_backward(net, g, inputs, d_out)
ignore(out)
let layers : Array[MPnnLayer] = Array::make(net.num_layers, net.layers[0])
for l in 0..