// diffusion_forward.mbt -- DDPM forward process primitives (v0.113.0).
//
// The forward diffusion process (Ho et al. 2020, "Denoising Diffusion
// Probabilistic Models") is a fixed Markov chain that gradually adds
// Gaussian noise to a clean sample x_0 over T steps:
//
//   q(x_t | x_0) = N(x_t; sqrt(alpha_bar[t]) * x_0, (1 - alpha_bar[t]) I)
//
// where alpha_bar[t] = prod_{s=0..t} (1 - beta[s]). The betas come from
// a chosen schedule (linear or cosine).
//
// Scope of v0.113.0:
//   - DiffusionSchedule struct (betas / alphas / alpha_bars / T).
//   - linear_schedule / cosine_schedule constructors.
//   - q_sample: closed-form forward step.
//   - q_sample_pair: convenience that draws noise and returns (x_t, noise)
//     for training the reverse network.
//
// Reference: Ho et al. 2020 (DDPM); Nichol & Dhariwal 2021 (cosine).

///|
/// DiffusionSchedule: holds the T betas, the derived alphas
/// (alpha_t = 1 - beta_t) and cumulative products alpha_bars, plus a
/// few useful pre-computed square-roots used by the closed-form forward
/// step.
pub struct DiffusionSchedule {
  num_steps : Int
  betas : Array[Float]
  alphas : Array[Float]
  alpha_bars : Array[Float]
  sqrt_alpha_bars : Array[Float]
  sqrt_one_minus_alpha_bars : Array[Float]
}

///|
/// Linear beta schedule from `beta_start` to `beta_end` over T steps
/// (Ho et al. 2020 default: T=1000, beta_start=1e-4, beta_end=0.02).
pub fn linear_schedule(
  num_steps : Int,
  beta_start : Float,
  beta_end : Float,
) -> DiffusionSchedule {
  let betas : Array[Float] = Array::make(num_steps, 0.0F)
  for t in 0.. DiffusionSchedule {
  let betas : Array[Float] = Array::make(num_steps, 0.0F)
  let steps : Float = Float::from_int(num_steps)
  for t in 0.. Float {
  let c = cosf(x)
  let abs_c = if c < 0.0F { -c } else { c }
  expf(logf(abs_c) * p)
}

///|
/// helper: build the full DiffusionSchedule from a flat betas array.
fn build_schedule(
  num_steps : Int,
  betas : Array[Float],
) -> DiffusionSchedule {
  let alphas : Array[Float] = Array::make(num_steps, 0.0F)
  let alpha_bars : Array[Float] = Array::make(num_steps, 0.0F)
  let sqrt_alpha_bars : Array[Float] = Array::make(num_steps, 0.0F)
  let sqrt_one_minus_alpha_bars : Array[Float] = Array::make(num_steps, 0.0F)
  let mut ab = 1.0F
  for t in 0.. Array[Float] {
  let n = x0.length()
  let out : Array[Float] = Array::make(n, 0.0F)
  let sa = sched.sqrt_alpha_bars[t]
  let sb = sched.sqrt_one_minus_alpha_bars[t]
  for i in 0.. (Array[Float], Array[Float]) {
  let n = x0.length()
  let noise : Array[Float] = Array::make(n, 0.0F)
  let pairs = n / 2
  let mut i = 0
  while i < pairs * 2 {
    // Box-Muller gives two standard normals at a time (as Double).
    let (z1d, z2d) = box_muller(rng)
    noise[i] = Float::from_double(z1d)
    if i + 1 < n {
      noise[i + 1] = Float::from_double(z2d)
    }
    i = i + 2
  }
  // If n is odd, draw one extra.
  if n % 2 == 1 {
    let (z1d, _) = box_muller(rng)
    noise[n - 1] = Float::from_double(z1d)
  }
  let xt = q_sample(x0, noise, t, sched)
  (xt, noise)
}