// dynamic_routing.mbt -- Dynamic routing (v0.135.0).
//
// Dynamic routing is the mechanism that replaces pooling in a capsule
// network (Sabour et al. 2017). Instead of averaging features from
// nearby patches, the network *learns* which low-level capsule votes
// for which high-level capsule, by iteratively adjusting agreement
// coefficients `c_ij`:
//
//   1. initialise routing logits b_ij = 0
//   2. repeat r times:
//      a. w_ij = c_ij * u_ij              (weighted votes)
//      b. s_j  = squash( sum_i w_ij )     (high-level capsule output)
//      c. d_ij = dot(s_j, u_ij)           (agreement)
//      d. c_ij = softmax_i(b_ij + d_ij)   (routing coefficients)
//      e. b_ij = b_ij + d_ij               (logit update, in practice)
//
// The agreement `dot(s_j, u_ij)` is what makes the routing
// *dynamic*: a low-level capsule that predicts a pose matching the
// current high-level prediction gets its coefficient increased, so on
// the next iteration it votes more strongly. Capsules for absent
// entities converge to p ~ 0 because their votes disagree.
//
// Scope of v0.135.0:
//   - dynamic_routing: one full routing loop over low-level capsules.
//   - routing_coefficients: the final c_ij for a given input (useful
//     for visualisation / interpretability).
//
// Reference: Sabour et al. 2017, Algorithm 1.

///|
/// Dynamic routing from low-level capsules to a bank of high-level
/// capsules.
///
/// `low_caps` is the flat [n_low x dim] stack of squashed
/// low-level capsule outputs. `w` is the [n_high x n_low * dim]
/// transformation matrix (the "prediction weights" that turn a vote
/// into a predicted pose for a particular high-level capsule).
///
/// Returns the flat [n_high x dim] stack of high-level capsule
/// outputs.
pub fn dynamic_routing(
  low_caps : Array[Float],
  n_low : Int,
  w : Array[Array[Float]],
  n_high : Int,
  dim : Int,
  n_iter : Int,
) -> Array[Float] {
  // Routing logits b_ij, flat [n_high x n_low].
  let logits : Array[Float] = Array::make(n_high * n_low, 0.0F)
  // Predicted votes u_ij, flat [n_high x n_low * dim].
  let u : Array[Float] = Array::make(n_high * n_low * dim, 0.0F)
  // High-level outputs s_j, flat [n_high x dim].
  let s : Array[Float] = Array::make(n_high * dim, 0.0F)
  let votes : Array[Float] = Array::make(dim, 0.0F)
  let c : Array[Float] = Array::make(n_low, 0.0F)
  for _ in 0.. m {
          m = logits[lbase + i]
        }
      }
      let mut sum_exp = 0.0F
      for i in 0.. sum_i w_ij * u_ij, then squash.
      for d in 0.. Array[Float] {
  let out = Array::make(n_high * n_low, 0.0F)
  // Re-run routing, recording the coefficients on the last iteration.
  let logits : Array[Float] = Array::make(n_high * n_low, 0.0F)
  let s : Array[Float] = Array::make(n_high * dim, 0.0F)
  let votes : Array[Float] = Array::make(dim, 0.0F)
  let c : Array[Float] = Array::make(n_low, 0.0F)
  let u : Array[Float] = Array::make(n_low * dim, 0.0F)
  for it in 0.. m {
          m = logits[lbase + i]
        }
      }
      let mut sum_exp = 0.0F
      for i in 0..