// pna_backward.mbt -- PNA backward pass and the shared GNN training
// loop (v0.148.0).
//
// PNA is the fourth and last of the Batch X architectures to get a
// backward. It is the interesting one, because it is the only
// aggregation here whose derivative is neither a scatter (sum, mean)
// nor a routing (max, min): the standard deviation's derivative is
//
//   d sigma_i / d x_j = (x_j - mean_i) / (deg_i * sigma_i)
//
// for every neighbour j of node i. Every incoming edge therefore
// contributes a share of the destination's std gradient, but the
// share depends on that edge's own message, not just on a degree
// count. Where sigma_i == 0 the derivative does not exist (a constant
// neighbourhood is a flat function); the backward returns 0 there,
// which is the conventional choice and keeps the training step from
// producing a NaN.
//
// PNA's aggregators are UNWEIGHTED (they reduce a neighbour multiset,
// not a normalised adjacency), so this file ships its own mean /
// max / min backwards rather than the `scatter_*_backward` family in
// graph_backward.mbt, which multiplies by `edge_weight[e]`. Using the
// weighted versions here would be a silent per-edge rescale by the
// adjacency normalisation -- invisible on an unweighted graph, wrong
// on a normalised one.
//
// The second half of the file is the payoff: a tagged `TrainableGraphNet`
// and a `graph_net_train_step` that runs loss -> backward -> SGD for
// GCN / GIN / PNA / MPNN behind one signature. Node-level
// cross-entropy is the objective, which is the standard node
// classification task (Cora / Citeseer) that GCN was invented for.

///|
/// Backward of `pna_mean`: the unweighted analogue of
/// `scatter_mean_backward`, with no `edge_weight` factor.
pub fn pna_mean_backward(g : Graph, d_out : Array[Float]) -> Array[Float] {
  pna_mean_backward_w(g, d_out, g.feat_dim)
}

///|
/// `pna_mean_backward` at an EXPLICIT message width.
pub fn pna_mean_backward_w(
  g : Graph,
  d_out : Array[Float],
  width : Int,
) -> Array[Float] {
  let deg = graph_in_degree(g)
  let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
  for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
      continue
    }
    if deg[d] == 0 {
      continue
    }
    let scale = 1.0F / Float::from_int(deg[d])
    let src_off = s * width
    let dst_off = d * width
    for k in 0.. Array[Float] {
  pna_max_backward_w(g, d_out, argmax, g.feat_dim)
}

///|
/// `pna_max_backward` at an EXPLICIT message width.
pub fn pna_max_backward_w(
  g : Graph,
  d_out : Array[Float],
  argmax : Array[Int],
  width : Int,
) -> Array[Float] {
  let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
  for d in 0..= g.n_edges {
        continue
      }
      let s = g.edge_src[e]
      if s < 0 || s >= g.n_nodes {
        continue
      }
      let src_off = s * width
      out[src_off + k] = out[src_off + k] + d_out[dst_off + k]
    }
  }
  out
}

///|
/// Backward of `pna_min`. Structurally identical to
/// `pna_max_backward` -- only the recorded argmin differs.
pub fn pna_min_backward(
  g : Graph,
  d_out : Array[Float],
  argmin : Array[Int],
) -> Array[Float] {
  pna_max_backward_w(g, d_out, argmin, g.feat_dim)
}

///|
/// `pna_min_backward` at an EXPLICIT message width.
pub fn pna_min_backward_w(
  g : Graph,
  d_out : Array[Float],
  argmin : Array[Int],
  width : Int,
) -> Array[Float] {
  pna_max_backward_w(g, d_out, argmin, width)
}

///|
/// Backward of `pna_std`.
///
///   d sigma_i / d x_j = (x_j - mean_i) / (deg_i * sigma_i)
///
/// `mean` and `sigma` are the per-node values from the forward pass;
/// they can be recomputed with `pna_mean` / `pna_std` on the same `h`.
/// Nodes with deg == 0 or sigma == 0 get a zero contribution, because
/// the derivative is undefined there (a constant neighbourhood is a
/// flat function) and dividing by it would emit a NaN.
pub fn pna_std_backward(
  g : Graph,
  h : Array[Float],
  d_out : Array[Float],
  mean : Array[Float],
  sigma : Array[Float],
) -> Array[Float] {
  let width = if g.n_nodes == 0 {
    0
  } else {
    mean.length() / g.n_nodes
  }
  pna_std_backward_w(g, h, d_out, mean, sigma, width)
}

///|
/// `pna_std_backward` at an EXPLICIT message width.
pub fn pna_std_backward_w(
  g : Graph,
  h : Array[Float],
  d_out : Array[Float],
  mean : Array[Float],
  sigma : Array[Float],
  width : Int,
) -> Array[Float] {
  let deg = graph_in_degree(g)
  let out : Array[Float] = Array::make(g.n_nodes * width, 0.0F)
  for e in 0..= g.n_nodes || d < 0 || d >= g.n_nodes {
      continue
    }
    if deg[d] == 0 {
      continue
    }
    let src_off = s * width
    let dst_off = d * width
    let inv = 1.0F / Float::from_int(deg[d])
    for k in 0.. PNAGrad {
  let layers : Array[PNALayerGrad] = Array::make(
    net.num_layers,
    PNALayerGrad::{
      pre_lin: GraphLinearGrad::zero(net.layers[0].pre_lin),
      post_lin: GraphLinearGrad::zero(net.layers[0].post_lin),
    },
  )
  for l in 0.. (Array[Float], Array[Float], Array[Int], Array[Int], Array[Float], Array[Float]) {
  let dim = layer.in_dim
  let n_agg = pna_num_aggregators()
  let mean = pna_mean_w(g, h, dim)
  let (mx, argmax) = pna_max_forward_with_idx_w(g, h, dim)
  let (mn, argmin) = pna_min_forward_with_idx_w(g, h, dim)
  let sigma = pna_std_w(g, h, dim)
  let aggs = [mean, mx, mn, sigma]
  let in_deg = graph_in_degree(g)
  let mean_deg = pna_mean_in_degree(g, layer.delta)
  let scale_size = g.n_nodes
  let scales : Array[Float] = Array::make(scale_size, 1.0F)
  for v in 0.. (Array[Float], PNALayerGrad) {
  // MUST be the layer's input width, matching
  // `pna_layer_backward_state`. Leaving this as `g.feat_dim` sizes the
  // d_input buffer and the reducer-gradient stride to the NODE feature
  // width, which for a hidden layer under-allocates it; the previous
  // layer then receives a gradient that is too narrow and reads out of
  // bounds one layer up. That is exactly the PanicError this fixed.
  let dim = layer.in_dim
  let n_agg = pna_num_aggregators()
  let (z, scales, argmax, argmin, mean, sigma) =
    pna_layer_backward_state(layer, g, h)
  // Recompute both pre-activations for the ReLU masks.
  let z1_pre = graph_linear_forward(layer.pre_lin, z, g.n_nodes)
  let mid : Array[Float] = Array::make(z1_pre.length(), 0.0F)
  for i in 0.. 0.0F { d_out[i] } else { 0.0F }
  }
  // Backward through the post Linear.
  let (d_mid, g_post) = graph_linear_backward(
    layer.post_lin, d_pre1, g.n_nodes, mid,
  )
  // Backward through the pre Linear + its ReLU.
  let d_z1 : Array[Float] = Array::make(d_mid.length(), 0.0F)
  for i in 0.. 0.0F { d_mid[i] } else { 0.0F }
  }
  let (d_z, g_pre) = graph_linear_backward(
    layer.pre_lin, d_z1, g.n_nodes, z,
  )
  // Split d_z into its five blocks, undo the degree scaling, and
  // route each reducer's share back to the input embeddings.
  let concat_dim = dim * (n_agg + 1)
  // Each reducer's gradient is a FULL [n_nodes x dim] tensor, because
  // the forward aggregate `aggs[a]` is per-node. Writing a single
  // node's row into a `dim`-wide buffer and overwriting it per v
  // silently keeps only the last node's share.
  let per_node = g.n_nodes * dim
  let d_mean : Array[Float] = Array::make(per_node, 0.0F)
  let d_max : Array[Float] = Array::make(per_node, 0.0F)
  let d_min : Array[Float] = Array::make(per_node, 0.0F)
  let d_std : Array[Float] = Array::make(per_node, 0.0F)
  for v in 0.. (Array[Float], Array[Array[Float]]) {
  let inputs : Array[Array[Float]] = Array::make(net.num_layers, [])
  let mut acc = h
  for l in 0.. (Array[Float], PNAGrad) {
  let grads = PNAGrad::zero(net)
  let mut d = d_out
  for i in 0.. (PNA, Array[Float]) {
  let (out, inputs) = pna_forward_with_inputs(net, g, h)
  let (d_input, grads) = pna_backward(net, g, inputs, d_out)
  ignore(out)
  let layers : Array[PNALayer] = Array::make(net.num_layers, net.layers[0])
  for l in 0.. Array[Float] {
  match net {
    Gcn(m) => gcn_forward(m, g, h)
    Gin(m) => gin_forward(m, g, h)
    Pna(m) => pna_forward(m, g, h)
    Mpnn(m) => mpnn_forward(m, g, h)
    Gat(m) => graph_attention_forward(m, g, h)
  }
}

///|
/// Node-embedding width of a tagged network.
pub fn trainable_net_node_dim(net : TrainableGraphNet) -> Int {
  match net {
    Gcn(m) => m.out_dim
    Gin(m) => m.out_dim
    Pna(m) => m.out_dim
    Mpnn(m) => m.out_dim
    // GAT CONCATENATES its heads, so the node row is heads times the
    // last layer's per-head width. Reading `m.out_dim` here would tell
    // every downstream caller (readouts, pooled readouts, the node
    // cross-entropy) that the embedding is `num_heads` times too
    // narrow -- and every one of those would then read the wrong
    // elements rather than fail.
    Gat(m) => gat_stack_node_dim(m)
  }
}

///|
/// Parameter count of a tagged network.
pub fn trainable_net_num_params(net : TrainableGraphNet) -> Int {
  match net {
    Gcn(m) => {
      let mut t = 0
      for l in 0.. gin_num_params(m)
    Pna(m) => pna_num_params(m)
    Mpnn(m) => {
      let mut t = 0
      for l in 0.. graph_attention_num_params(m)
  }
}

///|
/// Forward + backward + SGD for a tagged network. Returns
/// `(new_net, d_input)`.
pub fn trainable_net_sgd_step(
  net : TrainableGraphNet,
  g : Graph,
  h : Array[Float],
  d_out : Array[Float],
  lr : Float,
) -> (TrainableGraphNet, Array[Float]) {
  match net {
    Gcn(m) => {
      let (m2, d) = gcn_sgd_step(m, g, h, d_out, lr)
      (Gcn(m2), d)
    }
    Gin(m) => {
      let (m2, d) = gin_sgd_step(m, g, h, d_out, lr)
      (Gin(m2), d)
    }
    Pna(m) => {
      let (m2, d) = pna_sgd_step(m, g, h, d_out, lr)
      (Pna(m2), d)
    }
    Mpnn(m) => {
      let (m2, d) = mpnn_sgd_step(m, g, h, d_out, lr)
      (Mpnn(m2), d)
    }
    Gat(m) => {
      let (m2, d) = graph_attention_sgd_step(m, g, h, d_out, lr)
      (Gat(m2), d)
    }
  }
}

///|
/// Mean cross-entropy over nodes, where every node carries its own
/// per-node logit row. `logits` is [n_nodes x num_classes].
pub fn node_ce_loss(
  logits : Array[Float],
  labels : Array[Int],
  n_nodes : Int,
  num_classes : Int,
) -> Float {
  if n_nodes == 0 {
    return 0.0F
  }
  let mut total = 0.0F
  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 inv = 1.0F / Float::from_int(n_nodes)
  for v in 0.. (TrainableGraphNet, Float, Array[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 (new_net, d_input) = trainable_net_sgd_step(net, g, h, d_logits, lr)
  (new_net, loss, d_input)
}

///|
/// One graph-classification training step: forward, mean readout, mean
/// cross-entropy on the pooled graph embedding, backward, SGD.
///
/// The readout is a mean pool over node embeddings, so the gradient
/// reaching the backbone is a scatter of the graph-level gradient back
/// to every node -- which is why the loss over one pooled vector
/// trains every node in the graph.
pub fn graph_net_graph_train_step(
  net : TrainableGraphNet,
  g : Graph,
  h : Array[Float],
  target : Int,
  num_classes : Int,
  lr : Float,
) -> (TrainableGraphNet, Float, Array[Float]) {
  let node_emb = trainable_net_forward(net, g, h)
  let node_dim = trainable_net_node_dim(net)
  let graph_emb = mean_pool(node_emb, g.n_nodes, node_dim)
  let loss = graph_cross_entropy(graph_emb, target)
  let d_graph = graph_cross_entropy_grad(graph_emb, target)
  // Mean-pool adjoint: every node gets the same share, scaled by 1/n.
  let d_node : Array[Float] = Array::make(g.n_nodes * node_dim, 0.0F)
  if g.n_nodes > 0 {
    let inv = 1.0F / Float::from_int(g.n_nodes)
    for v in 0..