// affine_coupling_layer.mbt — AffineCouplingLayer (v0.101.0), the core
// building block of normalizing flows.
//
// An affine coupling layer (Dinh et al. 2016 "Real NVP") splits the
// input x ∈ R^d in two halves x = (x₁, x₂) and applies:
//   y₂ = x₂ ⊙ exp(s(x₁)) + t(x₁)
//   y = (y₁, y₂) where y₁ = x₁
// where s, t : R^{d/2} → R^{d/2} are learnable scale and translation
// functions.
//
// The Jacobian of this transformation is triangular:
//   J = [ I    0  ]
//       [ *  diag(exp(s(x₁))) ]
// so log |det J| = Σ_i s(x₁)[i].
//
// The inverse is straightforward:
//   x₂ = (y₂ - t(x₁)) ⊙ exp(-s(x₁))
//   x = (x₁, x₂)
//
// For v0.101.0 we use simple LINEAR projections for s and t (no NN):
//   s(x₁) = scale_w · x₁ + scale_b
//   t(x₁) = trans_w · x₁ + trans_b
// This keeps the layer cheap and exposes the canonical coupling
// structure. A neural-network s/t is straightforward to add later.
//
// Reference: Dinh et al. 2016 "Density Estimation with Real NVP";
// Dinh et al. 2018 "Glow" Section 3.

///|
/// Affine coupling layer. Holds linear projection weights for scale and
/// translation (both d/2 × d/2). The input is split into two halves of
/// length `half_dim` (must be even input dim).
pub struct AffineCouplingLayer {
  half_dim : Int
  // scale_w: (half_dim × half_dim), scale_b: (half_dim)
  scale_w : Array[Array[Float]]
  scale_b : Array[Float]
  // trans_w: (half_dim × half_dim), trans_b: (half_dim)
  trans_w : Array[Array[Float]]
  trans_b : Array[Float]
}

///|
/// Build a fresh AffineCouplingLayer with xavier-normal weights.
pub fn AffineCouplingLayer::new(
  half_dim : Int,
  seed : UInt64,
) -> AffineCouplingLayer {
  let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std = sqrtf(2.0F / Float::from_int(half_dim))
  let scale_w = xavier_normal(half_dim, half_dim, std, rng1)
  let scale_b : Array[Float] = Array::make(half_dim, 0.0F)
  let rng2 = Xoshiro::from_state(seed + 4UL, seed + 5UL, seed + 6UL, seed + 7UL)
  let trans_w = xavier_normal(half_dim, half_dim, std, rng2)
  let trans_b : Array[Float] = Array::make(half_dim, 0.0F)
  { half_dim, scale_w, scale_b, trans_w, trans_b }
}

///|
/// Apply the linear projection: out = w · x + b.
fn linear_project(
  w : Array[Array[Float]],
  b : Array[Float],
  x : Array[Float],
) -> Array[Float] {
  let out : Array[Float] = Array::make(b.length(), 0.0F)
  for i in 0.. (Array[Float], Float) {
  let half = layer.half_dim
  let x1 : Array[Float] = Array::make(half, 0.0F)
  let x2 : Array[Float] = Array::make(half, 0.0F)
  for i in 0.. Array[Float] {
  let half = layer.half_dim
  let y1 : Array[Float] = Array::make(half, 0.0F)
  let y2 : Array[Float] = Array::make(half, 0.0F)
  for i in 0.. Float {
  let half = layer.half_dim
  let x1 : Array[Float] = Array::make(half, 0.0F)
  for i in 0..