// message_passing.mbt -- Message passing layers (v0.138.0).
//
// Message Passing Neural Networks (Gilmer et al. 2017) generalise
// every GNN as three steps applied at every node:
//
//   m_v   = AGGREGATE({ MESSAGE(h_u, h_v, e_uv) : u in N(v) })
//   h_v'  = UPDATE(h_v, m_v)
//
// The MPNN formulation factors a GNN into a message function, an
// aggregation function, and an update function, which is why it
// covers GCN, GAT, GraphSAGE and most others as special cases.
//
// Two concrete layers ship here:
//
//   GCNLayer  (Kipf & Welling 2017): h' = act(W · D^-1/2(A+I)D^-1/2 · h)
//   MPnnLayer (Gilmer et al. 2017):  h' = act(W_self · h
//                                               + W_neigh · mean_{u in N(v)} h_u)
//
// Scope of v0.138.0:
//   - GCNLayer struct + gcn_layer_forward (stacked L layers).
//   - MPnnLayer struct + mpnn_layer_forward (stacked L layers).
//   - Linear + ReLU helpers used by both.
//
// Reference: Gilmer et al. 2017; Kipf & Welling 2017.

///|
/// A single Linear layer stored as (out_dim x in_dim) weights.
pub struct GraphLinear {
  in_dim : Int
  out_dim : Int
  w : Array[Array[Float]]
  b : Array[Float]
}

///|
/// Build a Linear layer with xavier-normal initialisation.
pub fn GraphLinear::new(
  in_dim : Int,
  out_dim : Int,
  seed : UInt64,
) -> GraphLinear {
  let std = sqrtf(2.0F / Float::from_int(in_dim))
  let rng = Xoshiro::from_state(
    seed + 10UL, seed + 11UL, seed + 12UL, seed + 13UL,
  )
  let w = xavier_normal(out_dim, in_dim, std, rng)
  let b : Array[Float] = Array::make(out_dim, 0.0F)
  { in_dim, out_dim, w, b }
}

///|
/// Forward: h [n_nodes x in_dim] -> [n_nodes x out_dim].
pub fn graph_linear_forward(
  lin : GraphLinear,
  h : Array[Float],
  n_nodes : Int,
) -> Array[Float] {
  let out : Array[Float] = Array::make(n_nodes * lin.out_dim, 0.0F)
  for n in 0.. Array[Float] {
  let out : Array[Float] = Array::make(x.length(), 0.0F)
  for i in 0.. 0.0F { x[i] } else { 0.0F }
  }
  out
}

///|
/// GCNLayer: one graph convolution with the symmetric normalised
/// adjacency. `alpha` is the GCN self-loop weight (1.0 in the paper).
pub struct GCNLayer {
  in_dim : Int
  out_dim : Int
  alpha : Float
  w : GraphLinear
}

///|
/// Build a GCNLayer.
pub fn GCNLayer::new(
  in_dim : Int,
  out_dim : Int,
  alpha : Float,
  seed : UInt64,
) -> GCNLayer {
  { in_dim, out_dim, alpha, w: GraphLinear::new(in_dim, out_dim, seed), }
}

///|
/// Forward one GCN layer:
///   h' = act( W · ( alpha * h + D^-1/2 A D^-1/2 h ) )
/// The adjacency is normalised once by the caller via
/// `normalise_adjacency`, so this function reads the pre-baked weights
/// from `g.edge_weight`.
pub fn gcn_layer_forward(
  layer : GCNLayer,
  g : Graph,
  h : Array[Float],
) -> Array[Float] {
  // Self contribution (alpha * h) plus the weighted neighbour sum.
  // `gcn_layer_support` (gcn_backward.mbt, v0.147.0) owns that
  // computation so the forward and the backward cannot drift apart.
  let support = gcn_layer_support(layer, g, h)
  let projected = graph_linear_forward(layer.w, support, g.n_nodes)
  graph_relu(projected)
}

///|
/// Stack of L GCN layers. The last layer is applied without a ReLU so
/// the output is a usable logit (Kipf & Welling apply the activation
/// to every *hidden* layer only).
pub struct GCN {
  layers : Array[GCNLayer]
  num_layers : Int
  hidden_dim : Int
  out_dim : Int
}

///|
/// Build an L-layer GCN. `dims` must have `num_layers + 1` entries.
pub fn GCN::new(
  dims : Array[Int],
  alpha : Float,
  seed : UInt64,
) -> GCN {
  let num_layers = dims.length() - 1
  let layers : Array[GCNLayer] = Array::make(
    num_layers, GCNLayer::new(dims[0], dims[1], alpha, seed),
  )
  for l in 0.. Array[Float] {
  let out = gcn_layer_forward(net.layers[0], g, h)
  let mut acc = out
  for l in 1.. MPnnLayer {
  {
    in_dim,
    out_dim,
    w_self: GraphLinear::new(in_dim, out_dim, seed),
    w_neigh: GraphLinear::new(in_dim, out_dim, seed + 50UL),
  }
}

///|
/// Forward one MPNN layer:
///   h' = act( W_self · h + W_neigh · mean_{u in N(v)} h_u )
/// The UPDATE function is an element-wise sum of two linear maps,
/// which is the GraphSAGE-style update (Hamilton et al. 2017).
pub fn mpnn_layer_forward(
  layer : MPnnLayer,
  g : Graph,
  h : Array[Float],
) -> Array[Float] {
  let self_out = graph_linear_forward(layer.w_self, h, g.n_nodes)
  // The mean is over the layer's INPUT width, not g.feat_dim -- the
  // two differ for any hidden layer.
  let neigh_mean = scatter_mean_w(g, h, layer.in_dim)
  let neigh_out = graph_linear_forward(layer.w_neigh, neigh_mean, g.n_nodes)
  let out : Array[Float] = Array::make(g.n_nodes * layer.out_dim, 0.0F)
  for i in 0.. Float {
  if x > 0.0F { x } else { 0.0F }
}

///|
/// Stack of L MPNN layers.
pub struct MPnn {
  layers : Array[MPnnLayer]
  num_layers : Int
  out_dim : Int
}

///|
/// Build an L-layer MPNN stack. `dims` must have `num_layers + 1`
/// entries.
pub fn MPnn::new(dims : Array[Int], seed : UInt64) -> MPnn {
  let num_layers = dims.length() - 1
  let layers : Array[MPnnLayer] = Array::make(
    num_layers, MPnnLayer::new(dims[0], dims[1], seed),
  )
  for l in 0.. Array[Float] {
  let mut acc = mpnn_layer_forward(net.layers[0], g, h)
  for l in 1..