// graph_attention.mbt -- Graph Attention Network (v0.139.0).
//
// GAT (Velickovic et al. 2018) replaces the fixed mean/sum aggregation
// of GCN with a learned attention coefficient per edge:
//
//   e_uv = a_L^T [W h_u || W h_v]
//   alpha_uv = softmax_v(e_uv)        (softmax over the IN-neighbours)
//   h'_v = sigma( sum_u alpha_uv · W h_u )
//
// where `||` is concatenation. The attention head `a` is itself
// learnable, which is what lets GAT focus on important neighbours
// instead of averaging all of them.
//
// Multi-head attention concatenates K independent heads (rather than
// averaging them) to keep the output dimension stable.
//
// Scope of v0.139.0:
//   - GATLayer (single head) + gat_layer_forward.
//   - GraphAttention (multi-head stack) + graph_attention_forward.
//   - LeakyReLU (slope 0.2, the GAT default).
//
// Reference: Velickovic et al. 2018.

///|
/// LeakyReLU with slope 0.2 (the GAT default negative slope).
pub fn gat_leaky_relu(x : Float) -> Float {
  if x > 0.0F { x } else { 0.2F * x }
}

///|
/// GATLayer: one attention head mapping in_dim -> out_dim.
pub struct GATLayer {
  in_dim : Int
  out_dim : Int
  // Shared linear projection W: (out_dim x in_dim).
  w : Array[Array[Float]]
  // Attention vector a, split into a_src and a_dst halves so the
  // score e_uv = a_src . Wh_u + a_dst . Wh_v.
  a_src : Array[Float]
  a_dst : Array[Float]
}

///|
/// Build a GATLayer with xavier-normal init.
pub fn GATLayer::new(
  in_dim : Int,
  out_dim : Int,
  seed : UInt64,
) -> GATLayer {
  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 a_src : Array[Float] = Array::make(out_dim, 0.0F)
  let a_dst : Array[Float] = Array::make(out_dim, 0.0F)
  for i in 0.. Array[Float] {
  let n = g.n_nodes
  let d = layer.out_dim
  // Project every node: p = W h, flat [n x d].
  let p : Array[Float] = Array::make(n * d, 0.0F)
  for v in 0..= n || dst < 0 || dst >= n {
      logits[e] = 0.0F
    } else {
      logits[e] = t_src[s] + t_dst[dst]
    }
  }
  let deg = graph_in_degree(g)
  let alpha : Array[Float] = Array::make(g.n_edges, 0.0F)
  for v in 0.. m {
        m = logits[e]
      }
    }
    let mut sum_exp = 0.0F
    for e in 0..= n || dst < 0 || dst >= n {
      continue
    }
    let a = alpha[e]
    let src_off = s * d
    let dst_off = dst * d
    for o in 0..