// gnn_train_demo.mbt -- end-to-end node-classification training
// (v0.155.0).
//
// Batches Y through AB shipped five backward passes and spent the last
// of them proving they were CORRECT. Correctness is not the same
// property as usefulness, though: a gradient can agree with a finite
// difference to seven digits and still leave a model that does not
// learn, because the composition can be right while the SIGNALS are
// wrong -- a head wired to the wrong half of a layer, a loss with the
// wrong sign, a learning rate that scales the update wrongly, a label
// convention inverted against the features. No gradient check sees any
// of that. This file is the check that does.
//
// WHAT IT FOUND ON ITS FIRST RUN. `graph_cross_entropy` had an
// inverted sign on its `logf(sum_exp)` term, so it returned a NEGATIVE
// loss for every confident and correct prediction. Its gradient was
// correct, so training worked and accuracy rose -- only the reported
// loss moved the wrong way. The gradient gate could not have caught it,
// because the gate deliberately checks a bounded quadratic objective
// instead of cross-entropy. A gate that exercises one objective cannot
// see a bug in another; this file is the second objective, and
// `gradcheck_ce_consistency` in gnn_gradcheck.mbt now pins the loss
// against its own gradient so the pair cannot drift apart again.
//
// The composite step is backbone -> node embeddings -> GraphLinear head
// -> node cross-entropy, not "backbone output is the logits". The
// reason is in `graph_net_train_step`'s own doc: GCN/GIN/MPNN apply
// their ReLU on EVERY layer including the last, so their raw output is
// non-negative by construction. Reading ReLU output as logits gives a
// classifier that can only ever push a class UP and never down, which
// still trains and still reports a falling loss -- so the demo would
// look fine while measuring a handicapped model. GAT is the opposite
// (its final layer is linear), which is exactly the kind of
// inconsistency a reader would otherwise have to discover by comparing
// five result rows. A head makes all five comparable, and it is also
// what you would actually put on a GNN.
//
// THE FIXTURE IS NOT THE GRADIENT-CHECK FIXTURE, and that is
// deliberate. `gradcheck_toy_graph` splits its labels by ring parity so
// that NEIGHBOURING NODES DISAGREE -- an anti-homophilous graph is
// exactly right for a gradient check, because a wrong mean aggregator
// then looks wrong instead of accidentally right. A training demo needs
// the opposite property: message passing can only work on a homophilous
// graph, so on the gradient-check fixture no amount of correct gradient
// would make any of these nets learn. One fixture cannot be both.
//
// WHY THERE ARE CONTROLS, AND WHY THEY USE HELD-OUT NODES. "Loss fell
// and accuracy rose" is a claim about a metric, and a metric that rises
// for free proves nothing. Two things make such a claim free:
//
// 1. MEMORISATION. Twenty-four training points are memorisable by a
// network with thousands of parameters, so a training-set accuracy
// of 100% distinguishes nothing. Every headline number here is a
// HELD-OUT accuracy, and the control below judges held-out
// accuracy too.
// 2. A TRIVIALLY SEPARABLE FIXTURE. The first version of this file
// used per-node noise of half-width 0.15 against a signal of 0.15,
// which puts them at exactly the threshold where the sign of the
// class-carrying feature is NEVER wrong. Every architecture, and a
// bare linear probe, hit 100% before a single training step. Half-
// width 0.60 against signal 0.30 puts the noise at twice the
// signal, which is the regime homophily actually helps in, and
// `gnn_train_fixture_separability` measures the resulting gap
// with NO TRAINING AT ALL so the claim rests on the data rather
// than on any model's good behaviour.
// ---------------------------------------------------------------------------
// Metrics
// ---------------------------------------------------------------------------
///|
/// Fraction of SELECTED nodes whose argmax logit equals their label.
/// `train` is a per-node mask; a node with `train[v] == false` is
/// skipped entirely. Ties break to the lowest class index, so this is
/// deterministic.
pub fn node_mask_accuracy(
logits : Array[Float],
labels : Array[Int],
train : Array[Bool],
n_nodes : Int,
num_classes : Int,
) -> Float {
if n_nodes == 0 || num_classes == 0 {
return 0.0F
}
let mut hit = 0
let mut seen = 0
for v in 0.. off { logits[off] } else { 0.0F }
for k in 1.. idx { logits[idx] } else { 0.0F }
if val > best_val {
best_val = val
best = k
}
}
let y = if v < labels.length() { labels[v] } else { 0 }
if best == y {
hit = hit + 1
}
seen = seen + 1
}
if seen == 0 {
return 0.0F
}
Float::from_int(hit) / Float::from_int(seen)
}
///|
/// Fraction of ALL nodes whose argmax logit equals their label.
pub fn graph_class_accuracy(
logits : Array[Float],
labels : Array[Int],
n_nodes : Int,
num_classes : Int,
) -> Float {
let all : Array[Bool] = Array::make(n_nodes, true)
node_mask_accuracy(logits, labels, all, n_nodes, num_classes)
}
///|
/// Mean cross-entropy over the SELECTED nodes, matching
/// `node_mask_accuracy`'s mask convention.
pub fn node_mask_ce_loss(
logits : Array[Float],
labels : Array[Int],
train : Array[Bool],
n_nodes : Int,
num_classes : Int,
) -> Float {
if n_nodes == 0 {
return 0.0F
}
let mut total = 0.0F
let mut seen = 0
for v in 0.. Array[Float] {
let out : Array[Float] = Array::make(n_nodes * num_classes, 0.0F)
if n_nodes == 0 {
return out
}
let mut seen = 0
for v in 0.. Int {
60
}
///|
/// Nodes per community.
pub fn gnn_train_per_community() -> Int {
30
}
///|
/// In-degree WITHIN a community: a circulant graph, each node joined to
/// the next four. This is the one fixture parameter that was not free.
///
/// The first version used two 20-cliques, i.e. in-degree 19. That is
/// not a realistic message-passing graph (Cora averages about 3) and it
/// broke GIN specifically: GIN's aggregator is an UNNORMALISED SUM, so
/// an embedding is ~19 times a single node's projection, the head
/// saturates, the cross-entropy underflows to ~5e-4 within five steps and
/// the gradient vanishes. GIN reached 75% held-out while the homophily
/// ceiling was 95%, and no learning rate fixes it -- the scale is set by
/// the degree. Degree 4 keeps every architecture in a regime where its
/// own aggregation is the thing being tested.
///
/// The circulant is generated by stride, not by an RNG, so the fixture's
/// STRUCTURE is byte-identical too and only the features carry noise.
pub fn gnn_train_in_degree() -> Int {
4
}
///|
/// Training nodes per community (the rest are held out).
pub fn gnn_train_per_community_train() -> Int {
18
}
///|
/// The training mask for `gnn_train_graph`. Transductive node
/// classification, the Cora setup: the held-out nodes stay IN the graph
/// and keep their edges and features, because a message-passing model
/// is supposed to use a node's neighbours whatever their labels are.
/// Only the LABELS are held out.
pub fn gnn_train_mask() -> Array[Bool] {
let n = gnn_train_n()
let per = gnn_train_per_community()
let n_train = gnn_train_per_community_train()
let mask : Array[Bool] = Array::make(n, false)
for v in 0.. Graph {
let n = gnn_train_n()
let per = gnn_train_per_community()
let deg = gnn_train_in_degree()
let feat_dim = 4
// Each community contributes per*deg edges; 2 cross edges on top.
let n_edges = 2 * per * deg + 2
let s : Array[Int] = Array::make(n_edges, 0)
let d : Array[Int] = Array::make(n_edges, 0)
let mut k = 0
for base in 0..<2 {
let off = base * per
for i in 0.. Array[Int] {
let n = gnn_train_n()
let per = gnn_train_per_community()
let out : Array[Int] = Array::make(n, 0)
for v in 0.. Array[Int] {
let n = gnn_train_n()
let per = gnn_train_per_community()
let truth = gnn_train_labels()
let out : Array[Int] = Array::make(n, 0)
for v in 0.. (Float, Float) {
let g = gnn_train_graph()
let labels = gnn_train_labels()
let n = g.n_nodes
let mut raw_hit = 0
let mut smooth_hit = 0
for v in 0..= 0.0F { 0 } else { 1 }
// Mean over in-neighbours plus the node itself: what a mean
// aggregator computes, on the one feature that carries the class.
let mut acc = x0
let mut cnt = 1.0F
for e in 0..= 0.0F { 0 } else { 1 }
if pred_raw == y {
raw_hit = raw_hit + 1
}
if pred_smooth == y {
smooth_hit = smooth_hit + 1
}
}
(
Float::from_int(raw_hit) / Float::from_int(n),
Float::from_int(smooth_hit) / Float::from_int(n),
)
}
///|
/// Build the classification head for a backbone: a `GraphLinear` from
/// the backbone's node-embedding width to `num_classes`.
///
/// The head's `in_dim` is read from `trainable_net_node_dim` rather
/// than assumed, because for GAT that value is `num_heads` times the
/// last per-head width. A head sized from anything else still builds,
/// still trains, and silently reads the wrong slice of every node row.
pub fn gnn_train_head(
net : TrainableGraphNet,
num_classes : Int,
seed : UInt64,
) -> GraphLinear {
GraphLinear::new(trainable_net_node_dim(net), num_classes, seed + 7777UL)
}
// ---------------------------------------------------------------------------
// The composite training step
// ---------------------------------------------------------------------------
///|
/// One full supervised step: backbone forward, head forward, mean
/// cross-entropy over the TRAINING nodes, backward through the head,
/// backward through the backbone, SGD on both. Returns
/// `(new_net, new_head, loss, train_acc, test_acc)`.
///
/// `loss` and the accuracies are the values at the CURRENT parameters,
/// i.e. BEFORE this step's update. Recording them before rather than
/// after is what makes "the loss fell" a measurement of the initial
/// model instead of a number that starts one step in.
pub fn gnn_train_step(
net : TrainableGraphNet,
head : GraphLinear,
g : Graph,
h : Array[Float],
labels : Array[Int],
train : Array[Bool],
num_classes : Int,
lr : Float,
) -> (TrainableGraphNet, GraphLinear, Float, Float, Float) {
let n = g.n_nodes
let emb = trainable_net_forward(net, g, h)
let logits = graph_linear_forward(head, emb, n)
let loss = node_mask_ce_loss(logits, labels, train, n, num_classes)
let train_acc = node_mask_accuracy(logits, labels, train, n, num_classes)
let held : Array[Bool] = Array::make(n, true)
let test_acc = node_mask_accuracy(
logits, labels, invert_mask(held, train), n, num_classes,
)
let d_logits = node_mask_ce_grad(logits, labels, train, n, num_classes)
// Backward through the head first: the backbone's gradient IS the
// head's input gradient, so this is the join between the two halves.
let (d_emb, head_grad) = graph_linear_backward(head, d_logits, n, emb)
let (net2, _) = trainable_net_sgd_step(net, g, h, d_emb, lr)
let head2 = graph_linear_sgd_step(head, head_grad, lr)
(net2, head2, loss, train_acc, test_acc)
}
///|
/// helper: the complement of a per-node mask, with the same length.
fn invert_mask(all : Array[Bool], train : Array[Bool]) -> Array[Bool] {
let out : Array[Bool] = Array::make(all.length(), true)
for v in 0.. (TrainableGraphNet, GraphLinear, Array[Float], Array[Float], Array[Float]) {
let n = g.n_nodes
let held = invert_mask(Array::make(n, true), train)
let losses : Array[Float] = Array::make(steps + 1, 0.0F)
let train_accs : Array[Float] = Array::make(steps + 1, 0.0F)
let test_accs : Array[Float] = Array::make(steps + 1, 0.0F)
let mut cur = net
let mut cur_head = head
let mut i = 0
while i <= steps {
let emb = trainable_net_forward(cur, g, h)
let logits = graph_linear_forward(cur_head, emb, n)
losses[i] = node_mask_ce_loss(logits, labels, train, n, num_classes)
train_accs[i] = node_mask_accuracy(logits, labels, train, n, num_classes)
test_accs[i] = node_mask_accuracy(logits, labels, held, n, num_classes)
if i < steps {
let (n2, h2, _, _, _) = gnn_train_step(
cur, cur_head, g, h, labels, train, num_classes, lr,
)
cur = n2
cur_head = h2
}
i = i + 1
}
(cur, cur_head, losses, train_accs, test_accs)
}
// ---------------------------------------------------------------------------
// Baseline: the same head with no message passing
// ---------------------------------------------------------------------------
///|
/// The no-message-passing baseline: a bare `GraphLinear` trained on the
/// RAW features with the same objective, learning rate and step count as
/// the GNNs, scored on the HELD-OUT nodes exactly as they are.
///
/// Returns `(held_out_accuracies, training_accuracies)`, both of length
/// `steps + 1`.
///
/// Scored held-out on purpose. Trained on 24 points this model
/// memorises them, so its TRAINING accuracy says nothing; the
/// held-out number is the one `gnn_train_fixture_separability` predicts
/// should be beaten, and on this fixture the prediction is that it
/// should NOT be.
pub fn gnn_linear_probe_curve(
g : Graph,
h : Array[Float],
labels : Array[Int],
train : Array[Bool],
num_classes : Int,
steps : Int,
lr : Float,
seed : UInt64,
) -> (Array[Float], Array[Float]) {
let n = g.n_nodes
let held = invert_mask(Array::make(n, true), train)
let head = GraphLinear::new(g.feat_dim, num_classes, seed + 9999UL)
let test_accs : Array[Float] = Array::make(steps + 1, 0.0F)
let train_accs : Array[Float] = Array::make(steps + 1, 0.0F)
let mut cur = head
let mut i = 0
while i <= steps {
let logits = graph_linear_forward(cur, h, n)
test_accs[i] = node_mask_accuracy(logits, labels, held, n, num_classes)
train_accs[i] = node_mask_accuracy(logits, labels, train, n, num_classes)
if i < steps {
let d = node_mask_ce_grad(logits, labels, train, n, num_classes)
let (_, grad) = graph_linear_backward(cur, d, n, h)
cur = graph_linear_sgd_step(cur, grad, lr)
}
i = i + 1
}
(test_accs, train_accs)
}
///|
/// Convenience: the five trainable backbones at a given embedding
/// width, with the graph each one needs. GAT uses TWO heads of width
/// `emb_dim / 2` so the multi-head stack -- concat, shared-input sum,
/// ELU mask -- is what actually trains here, not a degenerate one-head
/// configuration.
///
/// The GCN entry is built on the NORMALISED graph; the rest take the raw
/// graph, since `gcn_sgd_step` reads `edge_weight` and an unweighted
/// edge list is a legal but different model.
pub fn gnn_train_backbones(
emb_dim : Int,
seed : UInt64,
) -> (Array[TrainableGraphNet], Array[Graph]) {
let raw = gnn_train_graph()
let norm = normalise_adjacency(raw)
let gcn = GCN::new([4, emb_dim], 0.0F, seed)
let gin = GIN::new([4, emb_dim], 0.0F, seed)
let mpnn = MPnn::new([4, emb_dim], seed)
let pna = PNA::new(
[4, emb_dim], emb_dim, 0.001F, PNAScaler::Amplification, seed,
)
let heads = 2
let per_head = emb_dim / heads
let gat = GraphAttention::new([4, per_head], heads, seed)
let nets : Array[TrainableGraphNet] = Array::make(5, TrainableGraphNet::Gin(gin))
nets[0] = TrainableGraphNet::Gcn(gcn)
nets[1] = TrainableGraphNet::Gin(gin)
nets[2] = TrainableGraphNet::Mpnn(mpnn)
nets[3] = TrainableGraphNet::Pna(pna)
nets[4] = TrainableGraphNet::Gat(gat)
let graphs : Array[Graph] = Array::make(5, raw)
graphs[0] = norm
(nets, graphs)
}