// graph_classifier.mbt -- Graph-level readout and classification (v0.140.0).
//
// Node-level GNNs produce one embedding per node, but graph-level
// tasks (molecule property prediction, graph classification) need a
// single vector per graph. A READOUT function pools the node
// embeddings into a graph embedding.
//
// Three readouts ship here (all from the GIN / GraphCL line of work):
//
// mean_pool h_g = (1/n) sum_v h_v
// sum_pool h_g = sum_v h_v
// max_pool h_g = max_v h_v
//
// plus an MLP classification head on top.
//
// Scope of v0.140.0:
// - mean_pool / sum_pool / max_pool graph readouts.
// - GraphClassifier: a GCN or MPNN backbone + readout + MLP head.
// - graph_classifier_forward: graph -> class logits.
// - graph_classifier_predict: argmax.
// - cross_entropy_on_graphs + accuracy helper.
//
// Reference: Xu et al. 2019 (GIN); You et al. 2020 (GraphCL).
///|
/// Mean readout: average of all node embeddings.
pub fn mean_pool(h : Array[Float], n_nodes : Int, dim : Int) -> Array[Float] {
let out : Array[Float] = Array::make(dim, 0.0F)
if n_nodes == 0 {
return out
}
for v in 0.. Array[Float] {
let out : Array[Float] = Array::make(dim, 0.0F)
for v in 0.. Array[Float] {
let out : Array[Float] = Array::make(dim, 0.0F)
if n_nodes == 0 {
return out
}
for k in 0.. best {
best = x
}
}
out[k] = best
}
out
}
///|
/// Which readout a GraphClassifier uses.
pub(all) enum GraphReadout {
Mean
Sum
Max
} derive(Eq, Debug)
///|
/// GraphClassifier: a message-passing backbone, a graph readout, and
/// an MLP classification head.
pub struct GraphClassifier {
gcn : GCN
readout : GraphReadout
// Head: Linear(node_dim -> hidden) + ReLU + Linear(hidden -> classes).
head1 : GraphLinear
head2 : GraphLinear
node_dim : Int
hidden_dim : Int
num_classes : Int
}
///|
/// Build a GraphClassifier. `gcn` must output `node_dim`-wide node
/// embeddings (i.e. its last GCN layer has out_dim == node_dim).
pub fn GraphClassifier::new(
gcn : GCN,
readout : GraphReadout,
num_classes : Int,
hidden_dim : Int,
seed : UInt64,
) -> GraphClassifier {
let node_dim = gcn.out_dim
{
gcn,
readout,
head1: GraphLinear::new(node_dim, hidden_dim, seed + 70UL),
head2: GraphLinear::new(hidden_dim, num_classes, seed + 80UL),
node_dim,
hidden_dim,
num_classes,
}
}
///|
/// Forward: a single graph -> class logits [num_classes].
/// `g` must already carry the normalised adjacency weights.
pub fn graph_classifier_forward(
clf : GraphClassifier,
g : Graph,
h : Array[Float],
) -> Array[Float] {
// 1. Message passing over the graph.
let node_emb = gcn_forward(clf.gcn, g, h)
// 2. Readout: pool node embeddings into a graph embedding.
let graph_emb = match clf.readout {
Mean => mean_pool(node_emb, g.n_nodes, clf.node_dim)
Sum => sum_pool(node_emb, g.n_nodes, clf.node_dim)
Max => max_pool(node_emb, g.n_nodes, clf.node_dim)
}
// 3. MLP head.
let hid : Array[Float] = Array::make(1 * clf.hidden_dim, 0.0F)
let pre = graph_linear_forward(clf.head1, graph_emb, 1)
for i in 0.. 0.0F { pre[i] } else { 0.0F }
}
graph_linear_forward(clf.head2, hid, 1)
}
///|
/// Argmax over the class logits.
pub fn graph_classifier_predict(
clf : GraphClassifier,
g : Graph,
h : Array[Float],
) -> Int {
let logits = graph_classifier_forward(clf, g, h)
let mut best = 0
let mut best_val = logits[0]
for k in 1.. best_val {
best_val = logits[k]
best = k
}
}
best
}
///|
/// Cross-entropy loss on one graph.
///
/// -log softmax(logits)[target]
/// = -(logits[p] - m) + log(sum_exp) (max-subtracted)
/// = m - logits[p] + logf(sum_exp)
///
/// The `logf` term is ADDED. It was SUBTRACTED until v0.155.0, which
/// made this return a negative number for every confident and correct
/// prediction: with `p` the argmax, `sum_exp` is just above 1, so
/// `logf(sum_exp)` is a small positive number and the whole expression
/// came out as minus it. The magnitude was right and only the sign was
/// wrong, which is the worst possible shape for this bug:
///
/// * `graph_cross_entropy_grad` was (and still is) the correct
/// `softmax - onehot`, so TRAINING WORKED. Accuracy went up, the
/// optimiser's direction was right.
/// * The REPORTED loss therefore moved the wrong way, and a loop
/// that early-stops, schedules, or simply asserts "the loss fell"
/// on this function was reading a number with the wrong sign.
///
/// It survived v0.149.0 - v0.154.0 because the gradient gate
/// deliberately checks a bounded QUADRATIC objective instead of
/// cross-entropy (CE's loss grows with the logit scale, which made an
/// absolute tolerance meaningless). A gate that exercises one
/// objective cannot see a bug in another. The end-to-end training demo
/// found this in a single run because it is the first thing in the
/// package to actually TRAIN on this function, and
/// `gradcheck_ce_consistency` in gnn_gradcheck.mbt now pins the loss
/// against its own gradient by finite difference so it cannot come back.
pub fn graph_cross_entropy(logits : Array[Float], target : Int) -> Float {
let n = logits.length()
if n == 0 {
return 0.0F
}
let mut m = logits[0]
for i in 1.. m {
m = logits[i]
}
}
let mut sum_exp = 0.0F
for i in 0..= 0 && target < n { target } else { 0 }
m - logits[p] + logf(sum_exp)
}
///|
/// Mean cross-entropy over a list of (graph, features) pairs.
pub fn graph_batch_loss(
clf : GraphClassifier,
graphs : Array[Graph],
feats : Array[Array[Float]],
targets : Array[Int],
) -> Float {
let n = graphs.length()
if n == 0 {
return 0.0F
}
let mut total = 0.0F
for i in 0.. Float {
let n = graphs.length()
if n == 0 {
return 0.0F
}
let mut correct = 0.0F
for i in 0.. Float {
if n_nodes == 0 {
return 0.0F
}
let mut correct = 0.0F
for i in 0..