// gnn_gradcheck.mbt -- Central-difference gradient checks for the
// message-passing backwards (v0.149.0).
//
// Batches V, W, X and Y shipped nine message-passing architectures
// and four backward passes, and EVERY gradient in them was verified
// by reading the code. That is a real limitation, not a formality: a
// wrong edge-list gradient is invisible to inspection because the
// forward is correct and the shape of the backward is plausible. The
// only thing that catches it is a finite difference.
//
// The reason this file exists rather than a `_test.mbt` file is the
// Windows CreateProcessW 32K command-line limit: this package has
// 511 .mbt files, so `moon test` cannot link. A sub-package executable
// under `examples/` CAN be built and run (`moon run examples/`),
// because its command line only contains its own translation unit
// plus the already-built library. So the checks ship as library
// functions plus a runner, and the runner is what produces evidence.
//
// The check itself is the standard central difference:
//
// dL/dx_i ~= ( L(x + eps e_i) - L(x - eps e_i) ) / (2 eps)
//
// with eps = 1e-3. In Float32 the noise floor of a central difference
// is about 1e-3 relative for a loss of order 1, so the default
// tolerance is deliberately loose (1e-2 absolute). A tight tolerance
// here would report noise as failure and train the reader to ignore
// the output, which is worse than no check at all.
//
// NEGATIVE CONTROL. A checker that reports PASS for every input is
// worthless, so `gradcheck_negative_control` deliberately hands the
// checker a gradient scaled by 1.5 and asserts that the resulting
// max-abs-difference EXCEEDS the tolerance. If that control ever
// passes "cleanly", the checker is broken, not the code.
///|
/// The outcome of one gradient check.
pub struct GradCheckReport {
name : String
/// Largest absolute difference between analytic and numerical.
max_diff : Float
/// Tolerance the difference was judged against.
tol : Float
/// Index of the worst component in the analytic array.
worst_index : Int
/// True when max_diff <= tol.
passed : Bool
}
///|
/// The default central-difference step: 1e-4.
///
/// 1e-3 is the textbook choice for Float64, but these networks are
/// ReLU-activated, and ReLU is non-differentiable at exactly 0. A
/// central difference of half-width eps straddles the kink whenever a
/// pre-activation lands within eps of zero, and then it reports
/// roughly HALF the one-sided slope -- an O(gradient) error at a
/// handful of indices that has nothing to do with the code being
/// checked. Shrinking eps shrinks that band proportionally, and the
/// round-off cost grows only linearly: at eps = 1e-4 the Float32 noise
/// floor is ~3e-4 on a loss of order 1, still two orders of magnitude
/// below the 1e-2 tolerance. So 1e-4 trades a much larger band for a
/// much smaller noise floor, which is the right side of the trade.
///
/// Measured on the toy graph: at eps = 1e-3 the GCN input check
/// reported max_diff = 0.107 concentrated at 2 of 12 indices (the
/// kink), against a noise floor of ~3e-5.
pub fn gradcheck_eps() -> Float {
0.0001F
}
///|
/// The default tolerance. 1e-2 RELATIVE when judged against
/// `gradcheck_input_relative`, i.e. `max|analytic| * 1e-2`.
///
/// Absolute tolerances are the wrong instrument here: the six
/// architectures produce gradients spanning two orders of magnitude on
/// the same graph (GIN's logits are unnormalised and reach ~5e3 while
/// GCN's reach ~7), so one absolute number cannot be tight for the
/// small cases and loose enough for the large ones. 1e-2 relative sits
/// ~30x above the Float32 central-difference noise floor and ~30x below
/// the error a genuinely wrong routing produces.
pub fn gradcheck_tol() -> Float {
0.01F
}
///|
/// Scalar node-level cross-entropy of a tagged network. This is the
/// objective every check differentiates: a function of the input
/// features `h`, which is what makes dL/dh meaningful.
pub fn trainable_net_loss(
net : TrainableGraphNet,
g : Graph,
h : Array[Float],
labels : Array[Int],
num_classes : Int,
) -> Float {
let logits = trainable_net_forward(net, g, h)
node_ce_loss(logits, labels, g.n_nodes, num_classes)
}
///|
/// Analytic gradient of `trainable_net_loss` with respect to the input
/// node features. Returns `(d_input, loss)`.
///
/// This calls `trainable_net_sgd_step` with `lr = 0.0F` and uses only
/// its `d_input` return. At lr = 0 the SGD step computes
/// `w - 0 * grad`, which is bit-identical to `w` for any finite
/// gradient, so the returned network is unchanged -- this is a way to
/// reach the input gradient through one dispatch rather than writing
/// four near-identical `*_backward` wrappers for GCN / GIN / PNA /
/// MPNN. (If a gradient were already NaN or Inf, `0 * NaN` would
/// propagate it, which is a bug worth surfacing anyway.)
pub fn trainable_net_input_grad(
net : TrainableGraphNet,
g : Graph,
h : Array[Float],
labels : Array[Int],
num_classes : Int,
) -> (Array[Float], Float) {
let logits = trainable_net_forward(net, g, h)
let loss = node_ce_loss(logits, labels, g.n_nodes, num_classes)
let d_logits = node_ce_grad(logits, labels, g.n_nodes, num_classes)
let (_, d_input) = trainable_net_sgd_step(net, g, h, d_logits, 0.0F)
(d_input, loss)
}
///|
/// Build a report from an analytic and a numerical gradient.
pub fn gradcheck_make_report(
name : String,
analytic : Array[Float],
numeric : Array[Float],
tol : Float,
) -> GradCheckReport {
let (max_diff, worst) = max_abs_diff(analytic, numeric)
{
name,
max_diff,
tol,
worst_index: worst,
passed: max_diff <= tol,
}
}
///|
/// Check the analytic input gradient of a tagged network against
/// central differences. `eps` and `tol` default to the module-level
/// values; pass explicit ones to tighten or loosen a single check.
pub fn gradcheck_input(
name : String,
net : TrainableGraphNet,
g : Graph,
h : Array[Float],
labels : Array[Int],
num_classes : Int,
eps~ : Float = gradcheck_eps(),
tol~ : Float = gradcheck_tol(),
) -> GradCheckReport {
let (analytic, _) = trainable_net_input_grad(net, g, h, labels, num_classes)
let numeric = numerical_gradient(h, eps, fn(x) {
trainable_net_loss(net, g, x, labels, num_classes)
})
gradcheck_make_report(name, analytic, numeric, tol)
}
///|
/// Count how many components exceed the tolerance. A report whose max
/// is 10x the tolerance at 2 of 12 indices is a ReLU kink; one where
/// every index exceeds it is a real error. The count is what tells
/// those two apart, so it is part of the report rather than something
/// a reader has to infer from the max alone.
pub fn gradcheck_over_count(
analytic : Array[Float],
numeric : Array[Float],
tol : Float,
) -> Int {
let n = if analytic.length() < numeric.length() {
analytic.length()
} else {
numeric.length()
}
let mut count = 0
for i in 0.. tol {
count = count + 1
}
}
count
}
///|
/// A BOUNDED quadratic objective: L = 0.5 * sum_i out_i^2 over the
/// network's node embeddings. Its gradient w.r.t. the output is
/// exactly `out`, so no softmax / one-hot layer sits between the check
/// and the parameter gradients being tested.
///
/// This replaced node-level cross-entropy as the checked objective.
/// CE is the wrong choice for a gradient check here: the GIN forward
/// has no adjacency normalisation, so its logits run to ~5e3 on the
/// toy graph and the mean CE to ~1.6e3. An ABSOLUTE tolerance is
/// meaningless against gradients that scale with the logit magnitude
/// (it made GIN look 10x worse than GCN when GIN's RELATIVE error was
/// in fact the smallest of all six configs), and the Float32
/// central-difference noise floor grows with the loss, so the noise
/// band grows too. A quadratic objective is O(|out|^2) instead of
/// O(|out|) and needs no logits to be in a sane range.
pub fn trainable_net_l2_loss(
net : TrainableGraphNet,
g : Graph,
h : Array[Float],
) -> Float {
let out = trainable_net_forward(net, g, h)
let mut acc = 0.0F
for i in 0.. (Array[Float], Float) {
let out = trainable_net_forward(net, g, h)
let mut acc = 0.0F
for i in 0.. GradCheckReport {
let (analytic, _) = trainable_net_l2_input_grad(net, g, h)
let numeric = numerical_gradient(h, eps, fn(x) {
trainable_net_l2_loss(net, g, x)
})
let (max_diff, worst) = max_abs_diff(analytic, numeric)
let mut scale = 0.0F
for i in 0.. scale {
scale = a
}
}
// When the gradient is essentially zero the absolute criterion is the
// only meaningful one, so report against `tol` directly.
let eff = if scale < 0.000001F { tol } else { tol * scale }
{
name,
max_diff,
tol: eff,
worst_index: worst,
passed: max_diff <= eff,
}
}
///|
/// THE kink-vs-bug discriminator: what fraction of components deviate
/// by more than `frac` of the gradient scale.
///
/// This is the number that actually settles "is this a ReLU kink or a
/// real error", and it is cheap. Measured on the toy graph:
///
/// MPNN 2 layers 9/12 over 1% -> REAL ERROR
/// GIN 2 layers 0/12 over 1% -> clean (residual 1.5e-3 is Float32
/// GCN 2 layers 0/12 over 1% precision on a 2.3e7 gradient)
///
/// A ReLU kink perturbs the central difference at the ONE index whose
/// pre-activation sits within eps of zero, so it shows up as 1-2 of N.
/// A wrong routing shows up nearly everywhere. `max_diff` alone cannot
/// tell these apart -- it is the same 4e-2 whether one index is off by
/// a factor of two or nine are off by 4% -- and `max_diff` is also
/// scale-invariant for BOTH, so a scale sweep cannot tell them apart
/// either. Only the count can.
pub fn gradcheck_over_fraction(
analytic : Array[Float],
numeric : Array[Float],
frac : Float,
) -> Int {
let n = if analytic.length() < numeric.length() {
analytic.length()
} else {
numeric.length()
}
let mut scale = 0.0F
for i in 0.. scale {
scale = a
}
}
let thr = frac * scale
let mut count = 0
for i in 0.. thr {
count = count + 1
}
}
count
}
///|
/// Negative control: scale a gradient and confirm the checker notices.
///
/// The point is that `gradcheck_input` must FAIL on a deliberately
/// wrong gradient. If this returns `passed == true`, the checker is
/// vacuous -- e.g. if both arrays are all-zero, or the comparison is
/// against a copy of the analytic value -- and every other PASS from
/// it is meaningless.
pub fn gradcheck_negative_control(
name : String,
net : TrainableGraphNet,
g : Graph,
h : Array[Float],
labels : Array[Int],
num_classes : Int,
tol~ : Float = gradcheck_tol(),
) -> GradCheckReport {
let (analytic, _) = trainable_net_input_grad(net, g, h, labels, num_classes)
let numeric = numerical_gradient(h, gradcheck_eps(), fn(x) {
trainable_net_loss(net, g, x, labels, num_classes)
})
// Corrupt the analytic side by 50 percent.
let corrupt : Array[Float] = Array::make(analytic.length(), 0.0F)
for i in 0.. GradCheckReport {
let analytic = graph_cross_entropy_grad(logits, target)
let numeric = numerical_gradient(logits, eps, fn(x) {
graph_cross_entropy(x, target)
})
let (max_diff, worst) = max_abs_diff(analytic, numeric)
let mut scale = 0.0F
for i in 0.. scale {
scale = a
}
}
let eff = if scale < 0.000001F { tol } else { tol * scale }
{
name,
max_diff,
tol: eff,
worst_index: worst,
passed: max_diff <= eff,
}
}
///|
/// A set of logit vectors that between them reach both regimes the
/// inverted sign broke. The bug was largest where the prediction is
/// CONFIDENT AND CORRECT (`p` is the argmax, so `sum_exp -> 1` and
/// `logf(sum_exp) -> 0`), which is why a demo reporting a negative loss
/// was the thing that exposed it. `nearly_saturated_wrong` is the
/// mirror image: very confident and very wrong, where the true loss is
/// huge and the broken one is `2 * m` -- off by orders of magnitude, not
/// by a sign.
///
/// Hand-picked rather than random because the sizes are the point: a
/// random `N(0,1)` logit vector is nearly uniform, `logf(sum_exp)` is
/// ~0.7, and the check would pass on a sign-flipped function for a
/// large fraction of draws.
pub fn gradcheck_ce_probe_logits() -> Array[Array[Float]] {
let out : Array[Array[Float]] = Array::make(4, [])
let confident_right = Array::make(3, 0.0F)
confident_right[0] = 9.0F
confident_right[1] = 0.0F
confident_right[2] = 0.0F
let confident_wrong = Array::make(3, 0.0F)
confident_wrong[0] = 9.0F
confident_wrong[1] = 0.0F
confident_wrong[2] = 0.0F
let mild = Array::make(3, 0.0F)
mild[0] = 0.4F
mild[1] = -0.2F
mild[2] = 0.1F
let tied = Array::make(4, 0.0F)
tied[0] = 1.0F
tied[1] = 1.0F
tied[2] = 1.0F
tied[3] = 1.0F
out[0] = confident_right
out[1] = confident_wrong
out[2] = mild
out[3] = tied
out
}
///|
/// The target paired with each row of `gradcheck_ce_probe_logits`, and
/// the loss that row must produce. The expected values are the
/// textbook ones, computed by hand, and they are what makes this a
/// check of the FUNCTION rather than only of the loss/gradient pair --
/// a pair can be consistently wrong together, and here they were not
/// (the gradient was right), so the absolute values are pinned too.
///
/// row 0 [9,0,0] y=0 : confident correct -> -log(1 + 2e-4) ~ 2e-4
/// row 1 [9,0,0] y=1 : confident wrong -> 9 + 4.5e-5
/// row 2 [0.4,-0.2,0.1] y=1 : 0.95
/// row 3 [1,1,1,1] y=2 : uniform -> log 4
pub fn gradcheck_ce_probe_expected() -> Array[Float] {
let out : Array[Float] = Array::make(4, 0.0F)
// -log(exp(0)/(1 + 2*exp(-9))) = log(1 + 2*exp(-9))
out[0] = logf(1.0F + 2.0F * expf(0.0F - 9.0F))
// (9 - 0) + log(1 + exp(-9) + exp(-9))
out[1] = 9.0F + logf(1.0F + 2.0F * expf(0.0F - 9.0F))
// m = 0.4; sum = 1 + exp(-0.6) + exp(-0.3)
out[2] = 0.4F - (0.0F - 0.2F) + logf(
1.0F + expf(0.0F - 0.6F) + expf(0.0F - 0.3F),
)
// all equal -> uniform softmax -> -log(1/4) = log 4
out[3] = logf(4.0F)
out
}
///|
/// The target paired with each probe row.
pub fn gradcheck_ce_probe_targets() -> Array[Int] {
let out : Array[Int] = Array::make(4, 0)
out[0] = 0
out[1] = 1
out[2] = 1
out[3] = 2
out
}
///|
/// Which probe rows can be checked by FINITE DIFFERENCE, and why not the
/// others.
///
/// Row 0 (`[9,0,0]`, target 0) is the row that exposed the inverted
/// sign, and it is excluded here for a specific numerical reason rather
/// than a convenience one. It sits at the MINIMUM of the loss, where
/// the max-subtraction makes the value nearly independent of the largest
/// logit: `L = m - logits[p] + log(sum_exp)` with `p` the argmax has
/// `m - logits[p] == 0`, and perturbing `logits[0]` by +/-eps moves
/// `sum_exp` by about `2*e^-9` -- a relative change of 2.5e-4 on a
/// quantity `logf` then has to resolve near 1, where Float32's spacing
/// is 1.2e-7. The two evaluations are the same number after rounding, so
/// the central difference returns ~0 for a true derivative of -2.5e-4 and
/// the check cannot distinguish "correct" from "wrong" at all.
///
/// The VALUE check still runs on row 0, and that is the check that
/// matters for it: the inverted sign turned its expected 2.47e-4 into
/// -2.47e-4, which is unmissable. So the row is not dropped, it is
/// checked by the instrument that can actually see there.
pub fn gradcheck_ce_fd_rows() -> Array[Int] {
let out : Array[Int] = Array::make(3, 0)
out[0] = 1
out[1] = 2
out[2] = 3
out
}
///|
/// helper: a small deterministic graph for the checks. A ring of 6
/// nodes with 2 input features and an extra chord, so every node has
/// an in-degree of 1 or 2 -- enough for the degree scalers, the
/// softmax normalisers and the argmax routers all to be exercised,
/// without the cost of a large graph.
pub fn gradcheck_toy_graph() -> Graph {
let n = 6
// A 6-ring plus three chords, so in-degrees are 1 or 2: enough to
// exercise the degree scalers, the softmax normalisers and the
// argmax routers without the cost of a large graph.
let total = n + 3
let s : Array[Int] = Array::make(total, 0)
let d : Array[Int] = Array::make(total, 0)
let mut k = 0
for i in 0.. Array[Int] {
let labels : Array[Int] = Array::make(6, 0)
for v in 0..<6 {
labels[v] = if v % 3 == 0 { 0 } else { 1 }
}
labels
}
///|
/// helper: the toy graph with the symmetric GCN normalisation applied.
/// Use this for the GCN check specifically -- `gcn_sgd_step` reads
/// `edge_weight`, and a raw unweighted edge list is a legal but
/// different model.
pub fn gradcheck_toy_graph_gcn() -> Graph {
normalise_adjacency(gradcheck_toy_graph())
}