// glow.mbt 鈥?Glow normalizing flow (v0.103.0): RealNVP + ActNorm +
// 1脳1 convolution permutation.
//
// Glow (Dinh et al. 2018) extends Real NVP with two additional
// layers:
//   1. ActNorm: a learnable per-dimension affine
//        y = scale 鈯?x + bias
//      where scale and bias are initialized from the first batch's
//      per-dim std/mean (data-dependent init).
//   2. 1脳1 convolution permutation: replaces the fixed reverse
//      permutation with a learnable permutation matrix W. The log
//      |det W| contributes to the total log det.
// A Glow "step" chains: ActNorm 鈫?1脳1 conv 鈫?AffineCoupling.
//
// Reference: Dinh et al. 2018 "Glow: Generative Flow with Invertible
// 1脳1 Convolutions".

///|
/// ActNorm layer: per-dim learnable affine y = scale 路 x + bias.
/// Initialized from data statistics (mean, std of the first batch).
pub struct ActNorm {
  dim : Int
  scale : Array[Float]
  bias : Array[Float]
}

///|
/// Build ActNorm with zero scale + zero bias (init from data later).
pub fn ActNorm::new(dim : Int) -> ActNorm {
  let scale : Array[Float] = Array::make(dim, 1.0F)
  let bias : Array[Float] = Array::make(dim, 0.0F)
  { dim, scale, bias }
}

///|
/// Initialize ActNorm's scale + bias from data (per-dim std + mean).
pub fn actnorm_init_from_data(
  layer : ActNorm,
  data : Array[Float],
  n : Int,
) -> ActNorm {
  if n <= 0 || layer.dim == 0 {
    return layer
  }
  let scale : Array[Float] = Array::make(layer.dim, 0.0F)
  let bias : Array[Float] = Array::make(layer.dim, 0.0F)
  for k in 0.. Array[Float] {
  let y : Array[Float] = Array::make(layer.dim, 0.0F)
  for i in 0.. Array[Float] {
  let x : Array[Float] = Array::make(layer.dim, 0.0F)
  for i in 0.. Float {
  let mut ld = 0.0F
  for i in 0.. OneByOneConv {
  let w : Array[Array[Float]] = Array::make(
    dim, Array::make(dim, 0.0F),
  )
  // Initialize as identity + small noise.
  let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  for i in 0.. Float {
  let (z1, _) = box_muller(r)
  Float::from_double(z1)
}

///|
/// Compute log |det A| for a square matrix A via Gaussian elimination
/// (LU decomposition without pivoting). For v0.100.0 we just track
/// the diagonal of the upper-triangular U.
fn matrix_log_det(a : Array[Array[Float]], n : Int) -> Float {
  let lu : Array[Array[Float]] = Array::make(
    n, Array::make(n, 0.0F),
  )
  for i in 0.. pivot_val {
        pivot_val = v
        pivot_row = i
      }
    }
    if pivot_val < 1.0e-12F {
      return 0.0F  // singular
    }
    if pivot_row != k {
      // Swap rows.
      let tmp = lu[k]
      lu[k] = lu[pivot_row]
      lu[pivot_row] = tmp
      sign = -sign
    }
    log_det = log_det + logf(if lu[k][k] < 0.0F { -lu[k][k] } else { lu[k][k] })
    // Eliminate below.
    for i in (k + 1).. Array[Float] {
  let y : Array[Float] = Array::make(conv.dim, 0.0F)
  for i in 0.. Array[Float] {
  // Build augmented matrix [W | y] and solve via Gaussian elimination.
  let n = conv.dim
  let aug : Array[Array[Float]] = Array::make(
    n, Array::make(n + 1, 0.0F),
  )
  for i in 0.. pivot_val {
        pivot_val = v
        pivot_row = i
      }
    }
    if pivot_row != k {
      let tmp = aug[k]
      aug[k] = aug[pivot_row]
      aug[pivot_row] = tmp
    }
    for i in (k + 1).. GlowStep {
  let half = dim / 2
  let actnorm = ActNorm::new(dim)
  let conv = OneByOneConv::new(dim, seed)
  let coupling = AffineCouplingLayer::new(half, seed + 100UL)
  { actnorm, conv, coupling }
}

///|
/// Initialize ActNorm scale + bias from data (caller must call this
/// before the first forward pass for proper init).
pub fn glow_step_init_from_data(
  step : GlowStep,
  data : Array[Float],
  n : Int,
) -> GlowStep {
  let new_actnorm = actnorm_init_from_data(step.actnorm, data, n)
  { ..step, actnorm: new_actnorm }
}

///|
/// Glow step forward: y = coupling(conv(actnorm(x))).
/// Returns `(y, log_det)` where log_det = log_det_actnorm +
/// log_det_W + log_det_coupling.
pub fn glow_step_forward(
  step : GlowStep,
  x : Array[Float],
) -> (Array[Float], Float) {
  let y_actnorm = actnorm_forward(step.actnorm, x)
  let ld_actnorm = actnorm_log_det(step.actnorm)
  let y_conv = one_by_one_conv_forward(step.conv, y_actnorm)
  let ld_conv = step.conv.log_det_w
  let (y_coupling, ld_coupling) = coupling_forward(step.coupling, y_conv)
  let total_ld = ld_actnorm + ld_conv + ld_coupling
  (y_coupling, total_ld)
}

///|
/// Glow step inverse: x = actnorm_inverse(conv_inverse(coupling_inverse(y))).
pub fn glow_step_inverse(
  step : GlowStep,
  y : Array[Float],
) -> Array[Float] {
  let x_coupling = coupling_inverse(step.coupling, y)
  let x_conv = one_by_one_conv_inverse(step.conv, x_coupling)
  let x = actnorm_inverse(step.actnorm, x_conv)
  x
}

///|
/// Glow flow: N stacked GlowSteps.
pub struct Glow {
  n_steps : Int
  steps : Array[GlowStep]
  dim : Int
}

///|
/// Build a fresh Glow flow.
pub fn Glow::new(dim : Int, n_steps : Int, seed : UInt64) -> Glow {
  let steps : Array[GlowStep] = Array::make(
    n_steps, GlowStep::new(dim, seed),
  )
  for i in 0.. Glow {
  for i in 0.. (Array[Float], Float) {
  let mut cur = x.copy()
  let mut total_log_det = 0.0F
  for i in 0.. Array[Float] {
  let mut cur = y.copy()
  let mut i = flow.n_steps - 1
  while i >= 0 {
    cur = glow_step_inverse(flow.steps[i], cur)
    i = i - 1
  }
  cur
}