// capsule_network.mbt -- CapsuleNetwork: full CapsNet (v0.136.0).
//
// The full two-layer CapsNet:
//
//   x [c_in, h, w]
//   -> ReLU convolution (c_in -> ch, KxK, stride 1)          [c1_h, c1_w]
//   -> PrimaryCapsule (ch -> c_base*cd, K, K)                [n_patches x cd]
//   -> dynamic routing over the patches                       [n_classes x cd]
//
// The class prediction is the argmax over the length (existence
// probability) of each class capsule; the pose vector is discarded
// for classification.
//
// Reference: Sabour et al. 2017; Hinton et al. 2018.

///|
/// CapsuleNetwork: a PrimaryCapsule bank + a digit-capsule layer
/// connected by dynamic routing.
pub struct CapsuleNetwork {
  c_in : Int
  in_h : Int
  in_w : Int
  ch : Int
  conv_k : Int
  c1_h : Int
  c1_w : Int
  primary : PrimaryCapsule
  n_low : Int
  n_classes : Int
  capsule_dim : Int
  n_routing : Int
  // ReLU conv before the primary capsules: (ch rows x c_in*conv_k*conv_k).
  conv_w : Array[Array[Float]]
  conv_b : Array[Float]
  // Routing transformation W: n_classes rows of (n_low * capsule_dim).
  route_w : Array[Array[Float]]
}

///|
/// Build a fresh CapsuleNetwork.
pub fn CapsuleNetwork::new(
  c_in : Int,
  in_h : Int,
  in_w : Int,
  ch : Int,
  conv_k : Int,
  c_base : Int,
  capsule_dim : Int,
  n_classes : Int,
  n_routing : Int,
  seed : UInt64,
) -> CapsuleNetwork {
  // ReLU convolution.
  let conv_in = c_in * conv_k * conv_k
  let conv_w : Array[Array[Float]] = Array::make(
    ch, Array::make(conv_in, 0.0F),
  )
  let std_c = sqrtf(2.0F / Float::from_int(conv_in))
  let rng_c = Xoshiro::from_state(
    seed + 1UL, seed + 2UL, seed + 3UL, seed + 4UL,
  )
  for o in 0.. flat row-major
/// [n_classes x capsule_dim] class-capsule poses.
pub fn capsule_network_forward(
  net : CapsuleNetwork,
  image : Array[Float],
) -> Array[Float] {
  // 1. ReLU convolution.
  let k = net.conv_k
  let chn = net.ch
  let c1 : Array[Float] = Array::make(chn * net.c1_h * net.c1_w, 0.0F)
  for o in 0.. 0.0F { acc } else { 0.0F }
      }
    }
  }
  // 2. Primary capsules: [n_patches * c_base x capsule_dim].
  let low = primary_capsule_forward(net.primary, c1)
  // 3. Dynamic routing to the class capsules.
  dynamic_routing(
    low, net.n_low, net.route_w, net.n_classes, net.capsule_dim, net.n_routing,
  )
}

///|
/// Class probabilities from the class-capsule outputs: the norm of each
/// capsule is the existence probability.
pub fn capsule_network_probs(
  net : CapsuleNetwork,
  image : Array[Float],
) -> Array[Float] {
  let caps = capsule_network_forward(net, image)
  let out : Array[Float] = Array::make(net.n_classes, 0.0F)
  for k in 0.. Int {
  let probs = capsule_network_probs(net, image)
  let mut best = 0
  let mut best_val = probs[0]
  for k in 1.. best_val {
      best_val = probs[k]
      best = k
    }
  }
  best
}

///|
/// CapsuleNetworkTrainer: pairs a CapsuleNetwork with training
/// hyperparameters.
pub struct CapsuleNetworkTrainer {
  net : CapsuleNetwork
  lr : Float
  m_pos : Float
  m_neg : Float
  absent_weight : Float
}

///|
/// Build a fresh CapsuleNetworkTrainer. The default margins are the
/// paper's m+ = 0.9, m- = 0.1, absent weight = 0.9.
pub fn CapsuleNetworkTrainer::new(
  net : CapsuleNetwork,
  lr : Float,
  m_pos : Float,
  m_neg : Float,
  absent_weight : Float,
) -> CapsuleNetworkTrainer {
  { net, lr, m_pos, m_neg, absent_weight }
}

///|
/// One training step: forward pass + margin loss. Returns
/// (loss, predicted_class). Parameter updates are deferred to a
/// follow-up batch, consistent with the forward-only pattern.
pub fn capsule_train_step(
  trainer : CapsuleNetworkTrainer,
  image : Array[Float],
  target : Int,
) -> (Float, Int) {
  let caps = capsule_network_forward(trainer.net, image)
  let loss = capsule_margin_loss(
    caps,
    trainer.net.n_classes,
    trainer.net.capsule_dim,
    target,
    trainer.m_pos,
    trainer.m_neg,
    trainer.absent_weight,
  )
  let _ = trainer.lr
  let probs = capsule_network_probs(trainer.net, image)
  let mut best = 0
  let mut best_val = probs[0]
  for k in 1.. best_val {
      best_val = probs[k]
      best = k
    }
  }
  (loss, best)
}

///|
/// Classification accuracy over a batch.
pub fn capsule_eval_accuracy(
  trainer : CapsuleNetworkTrainer,
  batch : Array[Array[Float]],
  targets : Array[Int],
) -> Float {
  let m = batch.length()
  if m == 0 {
    return 0.0F
  }
  let mut correct = 0.0F
  for i in 0.. Int {
  let mut total = 0
  total = total + net.conv_w.length() * net.conv_w[0].length()
  total = total + net.conv_b.length()
  total = total + primary_capsule_num_params(net.primary)
  total = total + net.route_w.length() * net.route_w[0].length()
  total
}