// set2set.mbt -- Set2Set attention readout + attention-pooling
// graph classifier (v0.143.0).
//
// Reference: Vinyals, Keras, Bengio 2016, "Set2Set: Learning from
// Sets", ICLR 2017; Lee et al. 2019, "A Set-to-Set Approach to
// Multi-Class Molecular Graph Classification" (PNA paper's
// companion result for readout).
//
// The readouts in v0.140.0 (mean / sum / max) are PERMUTATION
// INVARIANT: they collapse the whole node set to one vector with a
// fixed reduction, so any reordering of the nodes is invisible. That
// is a virtue for bags of tokens and a limitation for molecular
// graphs, where the *identity* of a specific substructure matters --
// "there exists a nitrogen bonded to a carbon bonded to oxygen" is a
// statement about one element of the set, and mean-pooling cannot
// express it.
//
// Set2Set replaces the single reduction with `processing_steps` rounds
// of soft attention. Each round lets the model ask one question of the
// node set ("which node has the most mass in direction k?") and fold
// the answer into a running query state, so T rounds can extract T
// components of information instead of one summary:
//
//   q*_0 = 0
//   for t in 1..T:
//     q*_t   = LSTM(q*_{t-1}, q*_{t-1})          running query
//     e_i,t  =                attention logits
//     a_i,t  = softmax_i(e_i,t)
//     c*_t   = sum_i a_i,t H_i                   attended read-out
//   return c*_T
//
// The LSTM cell is the one from lstm_cell.mbt (v0.29.0), so Set2Set
// inherits its tested gates and its BPTT cache.
//
// Scope of v0.143.0:
//   - Set2Set parameter bundle + Set2SetState (q, c).
//   - set2set_forward: run T attention rounds, return the final read
//     -out and the advanced state (so a caller can continue a read
//     -out across chunks without restarting the query).
//   - GraphBackbone enum: a tagged GIN / PNA node encoder.
//   - Set2SetClassifier: backbone + Set2Set + MLP head, with
//     forward / predict / loss / accuracy.
//
// Deliberate simplification vs the paper: the paper scores
// e_i = a^T tanh(W_s H_i + W_q q). This uses the two projections
// inside a plain inner product, dropping the tanh and the learned
// output vector. It is the common implementation (and the paper's
// own ablation notes the tanh contributes little); keeping it means
// two Linear maps instead of three.

///|
/// Set2Set parameters: the query LSTM plus the two attention
/// projections. The LSTM's input and hidden widths are both
/// `in_dim`, matching the paper (the previous query is fed as both
/// the input and the hidden state).
pub struct Set2Set {
  in_dim : Int
  processing_steps : Int
  lstm : LstmCellParam
  w_source : GraphLinear
  w_query : GraphLinear
}

///|
/// Build a Set2Set readout for `in_dim`-wide node embeddings.
pub fn Set2Set::new(
  in_dim : Int,
  processing_steps : Int,
  seed : UInt64,
) -> Set2Set {
  {
    in_dim,
    processing_steps,
    lstm: LstmCellParam::new(in_dim, in_dim, seed),
    w_source: GraphLinear::new(in_dim, in_dim, seed + 20UL),
    w_query: GraphLinear::new(in_dim, in_dim, seed + 30UL),
  }
}

///|
/// The carried state of an (optionally interrupted) read-out: the
/// running query `q` and the last attended vector `c`. Both have
/// length `in_dim`.
pub struct Set2SetState {
  q : Array[Float]
  c : Array[Float]
}

///|
/// A zero state, i.e. the state before the first attention round.
pub fn Set2SetState::new(in_dim : Int) -> Set2SetState {
  {
    q: Array::make(in_dim, 0.0F),
    c: Array::make(in_dim, 0.0F),
  }
}

///|
/// Number of parameters in the readout.
pub fn set2set_num_params(m : Set2Set) -> Int {
  let d_h = m.lstm.d_h
  let d_x = m.lstm.d_x
  let in_dim = d_h + d_x
  // 4 gates, each a d_h x in_dim matrix plus a d_h bias.
  let lstm_params = 4 * (d_h * in_dim + d_h)
  let source = m.w_source.in_dim * m.w_source.out_dim + m.w_source.out_dim
  let query = m.w_query.in_dim * m.w_query.out_dim + m.w_query.out_dim
  lstm_params + source + query
}

///|
/// Run `processing_steps` attention rounds over the node embeddings.
///
/// Returns (graph_embedding, advanced_state). The returned embedding
/// has length `in_dim` and equals the last round's attended read-out.
/// Because the state is returned rather than hidden, a caller may
/// split a large node set into chunks, run the read-out per chunk,
/// and continue from the returned state.
pub fn set2set_forward(
  m : Set2Set,
  state : Set2SetState,
  node_emb : Array[Float],
  n_nodes : Int,
  feat_dim : Int,
) -> (Array[Float], Set2SetState) {
  let dim = m.in_dim
  let mut q = state.q
  let mut c = state.c
  if n_nodes == 0 {
    return (c, { q, c, })
  }
  // The source projection W_s H_i does not depend on the query, so it
  // is computed once and reused across all T rounds.
  let src = graph_linear_forward(m.w_source, node_emb, n_nodes)
  for _ in 0..