// 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..