// 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..