// capsule.mbt -- Capsule primitives (v0.133.0).
//
// A capsule (Sabour et al. 2017 "Dynamic Routing Between Capsules")
// is a group of neurons that outputs a vector, not a scalar. The
// vector encodes a "pose" (for an image) of an entity. The norm of
// the output vector is the existence probability p, and the direction
// carries the pose.
//
// Capsules must be *squashed* so that p is in [0, 1):
//
//   v_hat = v / ||v||_2
//   p    = 1 / (1 + exp(-||v||_2))
//   s(v) = (p / ||v||_2) * v_hat
//
// so s(v) = v / (1 + ||v||^2), which is the numerically convenient
// form used here.
//
// Scope of v0.133.0:
//   - Capsule struct: a vector-valued neuron group.
//   - capsule_squash: the squash activation.
//   - capsule_norm: L2 norm of a capsule vector.
//   - capsule_dot: dot product of two capsule vectors.
//   - capsule_num_params: counting helper.
//
// Reference: Sabour et al. 2017; Hinton et al. 2018 "The Matrix Capsule
// Networks".

///|
/// Capsule: a vector-valued output with a squashing activation.
/// `dim` is the length of the pose vector.
pub struct Capsule {
  dim : Int
  // Squashed output vector of length `dim`.
  v : Array[Float]
  // Existence probability (the norm of the squashed output).
  prob : Float
}

///|
/// The squash activation: s(v) = v / (1 + ||v||^2).
///
/// This is the closed-form simplification of the two-step
/// (normalise -> scale by p) definition, and it is the version used
/// throughout the capsule literature because it avoids two
/// intermediate arrays and a division.
pub fn capsule_squash(v : Array[Float]) -> Capsule {
  let n = v.length()
  let mut sum = 0.0F
  for i in 0.. Float {
  let n = v.length()
  let mut sum = 0.0F
  for i in 0.. Float {
  let n = a.length()
  let mut sum = 0.0F
  for i in 0.. Array[Float] {
  let n = v.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Array[Float] {
  let n = v.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  for i in 0.. Float {
  let mut total = 0.0F
  for k in 0.. 0.0F {
        total = total + diff * diff
      }
    } else {
      // Absent class: push the length below m_neg.
      let diff = len - m_neg
      if diff > 0.0F {
        total = total + absent_weight * diff * diff
      }
    }
  }
  total
}