// ddpm.mbt -- DDPM composite: schedule + score network (v0.115.0).
//
// A Denoising Diffusion Probabilistic Model (Ho et al. 2020) pairs the
// fixed forward process (DiffusionSchedule) with a learned reverse
// process parameterized by a ScoreNetwork. This file ships:
//
// - DDPM struct (DiffusionSchedule + ScoreNetwork + input_dim).
// - ddpm_loss: simple MSE between predicted and true noise.
// - ddpm_sample_step: one reverse step x_t -> x_{t-1}.
// - ddpm_sample: full reverse chain starting from pure Gaussian noise.
// - ddpm_predict_x0: derive x_0 prediction from epsilon prediction.
//
// The reverse step is the DDPM parameterisation:
//
// x_{t-1} = (1/sqrt(alpha[t])) * (x_t - (1 - alpha[t]) / sqrt(1 - alpha_bar[t]) * eps_theta(x_t, t))
// + sigma[t] * z, z ~ N(0, I), sigma[t] = sqrt(beta[t])
//
// (with sigma[t] = 0 at t = 0).
//
// Reference: Ho et al. 2020.
///|
/// DDPM: composite of a forward schedule and a learned score network.
pub struct DDPM {
input_dim : Int
schedule : DiffusionSchedule
net : ScoreNetwork
}
///|
/// Build a fresh DDPM. `schedule` is created by `linear_schedule` or
/// `cosine_schedule`; `net` is a ScoreNetwork matching `input_dim`.
pub fn DDPM::new(
input_dim : Int,
schedule : DiffusionSchedule,
net : ScoreNetwork,
) -> DDPM {
{ input_dim, schedule, net }
}
///|
/// MSE loss between predicted noise (from the score network) and the
/// true noise. Returns the mean squared error.
pub fn ddpm_loss(
ddpm : DDPM,
x0 : Array[Float],
t : Int,
noise : Array[Float],
) -> Float {
let (xt, _) = q_sample_pair(x0, t, ddpm.schedule, Xoshiro::new(0UL))
// Use the supplied noise instead of regenerating it.
let sa = ddpm.schedule.sqrt_alpha_bars[t]
let sb = ddpm.schedule.sqrt_one_minus_alpha_bars[t]
let n = ddpm.input_dim
let xt_supplied : Array[Float] = Array::make(n, 0.0F)
for i in 0.. Array[Float] {
let sa = ddpm.schedule.sqrt_alpha_bars[t]
let sb = ddpm.schedule.sqrt_one_minus_alpha_bars[t]
let n = ddpm.input_dim
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0.. x_{t-1} (Ho et al. 2020 parameterisation).
/// At t = 0, no extra noise is added.
pub fn ddpm_sample_step(
ddpm : DDPM,
xt : Array[Float],
t : Int,
rng : Xoshiro,
) -> Array[Float] {
let n = ddpm.input_dim
let eps = score_network_predict(ddpm.net, xt, t)
let sa = ddpm.schedule.alphas[t]
let sa_bar = ddpm.schedule.alpha_bars[t]
let sa_bar_prev = if t == 0 { 1.0F } else { ddpm.schedule.alpha_bars[t - 1] }
let beta_t = ddpm.schedule.betas[t]
// mean = (1/sqrt(alpha_t)) * (x_t - (1 - alpha_t)/sqrt(1 - alpha_bar_t) * eps)
let coef = (1.0F - sa) / sqrtf(1.0F - sa_bar)
let mean : Array[Float] = Array::make(n, 0.0F)
let inv_sqrt_alpha = 1.0F / sqrtf(sa)
for i in 0.. x_0. Returns x_0.
pub fn ddpm_sample(
ddpm : DDPM,
rng : Xoshiro,
) -> Array[Float] {
let n = ddpm.input_dim
let mut xt : Array[Float] = Array::make(n, 0.0F)
let pairs = n / 2
let mut i = 0
while i < pairs * 2 {
let (z1d, z2d) = box_muller(rng)
xt[i] = Float::from_double(z1d)
if i + 1 < n {
xt[i + 1] = Float::from_double(z2d)
}
i = i + 2
}
if n % 2 == 1 {
let (z1d, _) = box_muller(rng)
xt[n - 1] = Float::from_double(z1d)
}
let mut t = ddpm.schedule.num_steps - 1
while t >= 0 {
xt = ddpm_sample_step(ddpm, xt, t, rng)
t = t - 1
}
xt
}