// pna.mbt -- Principal Neighbourhood Aggregation (v0.142.0).
//
// Reference: Corso, Cavalleri, Beaini, Liò, Veličković 2020,
// "Principal Neighbourhood Aggregation for Graph Neural Networks",
// NeurIPS.
//
// PNA attacks the weakness every aggregator in v0.137.0 / v0.141.0
// shares: a single reduction (mean, sum, max) discards information
// about the *shape* of a neighbourhood. A node with neighbours
// {0, 0} and a node with neighbours {0, 1} have the same max, and
// a node with neighbours {1, 1, 1} and {-1, -1, -1} differ in sign
// but not in max. PNA concatenates four complementary reductions —
//
// mean = (1/d) sum_u h_u order-sensitive magnitude
// max = max_u h_u upper envelope
// min = min_u h_u lower envelope
// std = sqrt(E[x^2] - E[x]^2) spread
//
// — so a downstream Linear sees the whole first two moments of the
// neighbour multiset instead of one summary of it.
//
// On top of that, PNA rescales each aggregate by the node's degree.
// High-degree hubs and low-degree leaves have wildly different
// neighbour-count statistics, and a fixed aggregation weight makes
// the two populations incomparable. Three scalers are provided:
// identity, amplification (log(d+1) up-weighting hubs), and
// attenuation (the inverse).
//
// Scope of v0.142.0:
// - PNAAggregator / PNAScaler enums.
// - pna_aggregate: mean / max / min / std over the in-neighbour
// multiset, all UNWEIGHTED (the paper aggregates a multiset, not
// a weighted adjacency).
// - pna_degree_scale: the degree scaler.
// - PNALayer: concat([h, mean, max, min, std]) -> 2-layer MLP.
// - PNA: an L-layer stack.
// - pna_forward + pna_num_params + pna_mean_in_degree.
//
// Convention: the four aggregators are always concatenated in the
// fixed order mean, max, min, std, so the first Linear's input width
// is always 5 * in_dim. A layer is not built from a configurable
// subset, because varying that width per layer is the fastest way to
// produce a silent shape mismatch that `moon check` cannot see.
///|
/// The neighbourhood reductions PNA concatenates.
pub(all) enum PNAAggregator {
Mean
Max
Min
Std
} derive(Eq, Debug)
///|
/// How a per-node aggregate is rescaled by the node's degree.
pub(all) enum PNAScaler {
/// No degree dependence.
Identity
/// scale = log(d + 1) / mean(d): up-weights high-degree hubs.
Amplification
/// scale = mean(d) / log(d + 1): down-weights high-degree hubs.
Attenuation
} derive(Eq, Debug)
///|
/// helper: number of distinct aggregators concatenated by a layer.
pub fn pna_num_aggregators() -> Int {
4
}
///|
/// helper: raw (unweighted) sums of the in-neighbour features, plus
/// the per-node in-degree. Shared by the mean and std reducers.
fn pna_sum_and_deg_w(
g : Graph,
h : Array[Float],
width : Int,
) -> (Array[Float], Array[Int]) {
let sum : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
let deg : Array[Int] = Array::make(g.n_nodes, 0)
for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
continue
}
let src_off = s * width
let dst_off = d * width
for k in 0.. (Array[Float], Array[Int]) {
pna_sum_and_deg_w(g, h, g.feat_dim)
}
///|
/// Mean over the in-neighbour multiset. Isolated nodes get a zero row.
pub fn pna_mean(g : Graph, h : Array[Float]) -> Array[Float] {
pna_mean_w(g, h, g.feat_dim)
}
///|
/// `pna_mean` with an EXPLICIT message width. A PNA layer l > 0
/// receives `dims[l]`-wide activations, so the layer forward must pass
/// `layer.in_dim`, never `g.feat_dim`.
pub fn pna_mean_w(g : Graph, h : Array[Float], width : Int) -> Array[Float] {
let (sum, deg) = pna_sum_and_deg_w(g, h, width)
let deg = graph_in_degree(g)
for v in 0.. Array[Float] {
pna_max_w(g, h, g.feat_dim)
}
///|
/// `pna_max` at an EXPLICIT message width. Layer l > 0 receives
/// `dims[l]`-wide activations, so the layer forward passes
/// `layer.in_dim` here, never `g.feat_dim`.
pub fn pna_max_w(
g : Graph,
h : Array[Float],
width : Int,
) -> Array[Float] {
let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
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 src_off = s * width
let dst_off = d * width
let first = seeded[d] == 0
seeded[d] = 1
for k in 0.. out[dst_off + k] {
out[dst_off + k] = m
}
}
}
out
}
///|
/// Min over the in-neighbour multiset. Isolated nodes get a zero row.
pub fn pna_min(g : Graph, h : Array[Float]) -> Array[Float] {
pna_min_w(g, h, g.feat_dim)
}
///|
/// `pna_min` at an EXPLICIT message width.
pub fn pna_min_w(
g : Graph,
h : Array[Float],
width : Int,
) -> Array[Float] {
let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
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 src_off = s * width
let dst_off = d * width
let first = seeded[d] == 0
seeded[d] = 1
for k in 0.. Array[Float] {
pna_std_w(g, h, g.feat_dim)
}
///|
/// `pna_std` at an EXPLICIT message width.
pub fn pna_std_w(
g : Graph,
h : Array[Float],
width : Int,
) -> Array[Float] {
let (sum, deg) = pna_sum_and_deg_w(g, h, width)
let sq : 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 src_off = s * width
let dst_off = d * width
for k in 0.. 0.0F {
out[off + k] = sqrtf(variance)
}
}
}
out
}
///|
/// pna_max at an EXPLICIT message width, ALSO recording which incoming
/// edge won per (node, feature). The backward needs this index and
/// cannot recover it from the forward's value alone.
///
/// Ties go to the FIRST edge in list order, matching the `first ||
/// m > out[...]` tie-break, so the recorded index and the returned
/// value always agree.
pub fn pna_max_forward_with_idx_w(
g : Graph,
h : Array[Float],
width : Int,
) -> (Array[Float], Array[Int]) {
let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
let argmax : Array[Int] = Array::make(g.n_nodes * width, -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 src_off = s * width
let dst_off = d * width
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)
}
///|
/// pna_min at an EXPLICIT message width, recording the argmin.
pub fn pna_min_forward_with_idx_w(
g : Graph,
h : Array[Float],
width : Int,
) -> (Array[Float], Array[Int]) {
let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
let argmin : Array[Int] = Array::make(g.n_nodes * width, -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 src_off = s * width
let dst_off = d * width
let first = seeded[d] == 0
seeded[d] = 1
for k in 0.. (Array[Float], Array[Int]) {
pna_max_forward_with_idx_w(g, h, g.feat_dim)
}
///|
/// pna_min at `g.feat_dim` width, also recording the argmin.
pub fn pna_min_forward_with_idx(
g : Graph,
h : Array[Float],
) -> (Array[Float], Array[Int]) {
pna_min_forward_with_idx_w(g, h, g.feat_dim)
}
///|
/// Reduce the in-neighbour multiset with the chosen aggregator.
pub fn pna_aggregate(
g : Graph,
h : Array[Float],
agg : PNAAggregator,
) -> Array[Float] {
match agg {
Mean => pna_mean(g, h)
Max => pna_max(g, h)
Min => pna_min(g, h)
Std => pna_std(g, h)
}
}
///|
/// Max over the in-neighbour multiset, ALSO recording which incoming
/// edge won per (node, feature). The backward (pna_backward.mbt,
/// v0.148.0) needs this index and cannot recover it from the forward's
/// value alone -- two different neighbourhoods can share a maximum.
///
/// Ties go to the FIRST edge in list order, matching the `first ||
/// m > out[...]` tie-break above, so the recorded index and the
/// returned value always agree.
///|
/// Mean in-degree over the graph, floored at `delta`. The paper uses
/// a small delta (0.001 by default) to stop the normalisation from
/// blowing up on a near-empty graph.
pub fn pna_mean_in_degree(g : Graph, delta : Float) -> Float {
let deg = graph_in_degree(g)
let mut total = 0
for i in 0.. Float {
match scaler {
Identity => 1.0F
Amplification => {
let d = Float::from_int(deg)
logf(d + 1.0F) / mean_deg
}
Attenuation => {
if deg == 0 {
1.0F
} else {
let d = Float::from_int(deg)
mean_deg / logf(d + 1.0F)
}
}
}
}
///|
/// A single PNA layer: five scaled blocks (self + 4 aggregations)
/// concatenated and fed to a 2-layer MLP.
pub struct PNALayer {
in_dim : Int
hidden_dim : Int
out_dim : Int
delta : Float
scaler : PNAScaler
pre_lin : GraphLinear
post_lin : GraphLinear
}
///|
/// Build a PNALayer.
pub fn PNALayer::new(
in_dim : Int,
hidden_dim : Int,
out_dim : Int,
delta : Float,
scaler : PNAScaler,
seed : UInt64,
) -> PNALayer {
let n_agg = pna_num_aggregators()
let concat_dim = in_dim * (n_agg + 1)
{
in_dim,
hidden_dim,
out_dim,
delta,
scaler,
pre_lin: GraphLinear::new(concat_dim, hidden_dim, seed),
post_lin: GraphLinear::new(hidden_dim, out_dim, seed + 10UL),
}
}
///|
/// Forward one PNA layer:
/// z_v = concat( scale_v * h_v, scale_v * mean, scale_v * max,
/// scale_v * min, scale_v * std )
/// h_v' = act( W_2 · relu( W_1 · z_v + b_1 ) + b_2 )
/// The SAME per-node scale multiplies all five blocks, matching the
/// paper (the scaler is a property of the node's degree, not of the
/// individual reduction).
pub fn pna_layer_forward(
layer : PNALayer,
g : Graph,
h : Array[Float],
) -> Array[Float] {
// The reductions run at the layer's INPUT width, not g.feat_dim --
// the two differ for any hidden layer, and using feat_dim here
// truncates the concat so it no longer matches pre_lin's in_dim.
let dim = layer.in_dim
let n_agg = pna_num_aggregators()
let aggs = [
pna_mean_w(g, h, dim),
pna_max_w(g, h, dim),
pna_min_w(g, h, dim),
pna_std_w(g, h, dim),
]
let in_deg = graph_in_degree(g)
let mean_deg = pna_mean_in_degree(g, layer.delta)
let concat_dim = dim * (n_agg + 1)
let z : Array[Float] = Array::make(g.n_nodes * concat_dim, 0.0F)
for v in 0.. Int {
layer.pre_lin.in_dim * layer.pre_lin.out_dim + layer.pre_lin.out_dim +
layer.post_lin.in_dim * layer.post_lin.out_dim + layer.post_lin.out_dim
}
///|
/// Stack of L PNA layers.
pub struct PNA {
layers : Array[PNALayer]
num_layers : Int
scaler : PNAScaler
delta : Float
out_dim : Int
}
///|
/// Build an L-layer PNA. `dims` must have `num_layers + 1` entries.
/// The default scaler is Amplification, which is the paper's
/// recommended setting for most benchmarks.
pub fn PNA::new(
dims : Array[Int],
hidden_dim : Int,
delta : Float,
scaler : PNAScaler,
seed : UInt64,
) -> PNA {
let num_layers = dims.length() - 1
let layers : Array[PNALayer] = Array::make(
num_layers,
PNALayer::new(dims[0], hidden_dim, dims[1], delta, scaler, seed),
)
for l in 0.. Array[Float] {
let mut acc = pna_layer_forward(net.layers[0], g, h)
for l in 1.. Int {
let mut total = 0
for l in 0..