//! Noise schedulers for diffusion models. //! //! This module implements DDPM and DDIM schedulers that control the noise //! addition and removal process in diffusion models. use rtx_backend::Backend; use rtx_tensor::GenericTensor; /// Configuration for noise schedulers. #[derive(Debug, Clone)] pub struct SchedulerConfig { /// Number of diffusion timesteps pub num_timesteps: usize, /// Starting value for beta schedule pub beta_start: f32, /// Ending value for beta schedule pub beta_end: f32, /// Type of beta schedule pub schedule_type: BetaSchedule, } impl Default for SchedulerConfig { fn default() -> Self { Self { num_timesteps: 1000, beta_start: 1e-4, beta_end: 0.02, schedule_type: BetaSchedule::Linear, } } } /// Beta schedule types for noise variance. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum BetaSchedule { /// Linear schedule: β_t = β_start + t/T * (β_end - β_start) Linear, /// Cosine schedule (from "Improved Denoising Diffusion Probabilistic Models") Cosine, /// Scaled linear schedule ScaledLinear, } /// Trait for noise schedulers. pub trait NoiseScheduler { /// Get the number of timesteps. fn num_timesteps(&self) -> usize; /// Get alpha_cumprod (ᾱ_t) at timestep t. fn alpha_cumprod(&self, t: usize) -> f32; /// Get beta (β_t) at timestep t. fn beta(&self, t: usize) -> f32; /// Add noise to clean data: x_t = √(ᾱ_t) * x_0 + √(1 - ᾱ_t) * ε fn add_noise, const D: usize>( &self, x0: &GenericTensor, noise: &GenericTensor, t: usize, ) -> GenericTensor { let alpha_cumprod = self.alpha_cumprod(t); let sqrt_alpha_cumprod = alpha_cumprod.sqrt(); let sqrt_one_minus_alpha_cumprod = (1.0 - alpha_cumprod).sqrt(); // x_t = √(ᾱ_t) * x_0 + √(1 - ᾱ_t) * ε x0.mul_scalar(sqrt_alpha_cumprod) .add(&noise.mul_scalar(sqrt_one_minus_alpha_cumprod)) } /// Predict x_0 from x_t and predicted noise. /// x_0 = (x_t - √(1 - ᾱ_t) * ε) / √(ᾱ_t) fn predict_x0, const D: usize>( &self, xt: &GenericTensor, pred_noise: &GenericTensor, t: usize, ) -> GenericTensor { let alpha_cumprod = self.alpha_cumprod(t); let sqrt_alpha_cumprod = alpha_cumprod.sqrt(); let sqrt_one_minus_alpha_cumprod = (1.0 - alpha_cumprod).sqrt(); // x_0 = (x_t - √(1 - ᾱ_t) * ε) / √(ᾱ_t) xt.sub(&pred_noise.mul_scalar(sqrt_one_minus_alpha_cumprod)) .mul_scalar(1.0 / sqrt_alpha_cumprod) } } /// DDPM (Denoising Diffusion Probabilistic Models) scheduler. /// /// Implements the forward and reverse diffusion process from /// Ho et al., "Denoising Diffusion Probabilistic Models" (2020). #[derive(Debug, Clone)] pub struct DDPMScheduler { /// Configuration config: SchedulerConfig, /// Beta values for each timestep betas: Vec, /// Alpha values (1 - beta) alphas: Vec, /// Cumulative product of alphas alphas_cumprod: Vec, /// √(ᾱ_t) sqrt_alphas_cumprod: Vec, /// √(1 - ᾱ_t) sqrt_one_minus_alphas_cumprod: Vec, /// Posterior variance σ²_t posterior_variance: Vec, } impl DDPMScheduler { /// Create a new DDPM scheduler. pub fn new(num_timesteps: usize, beta_start: f32, beta_end: f32) -> Self { let config = SchedulerConfig { num_timesteps, beta_start, beta_end, schedule_type: BetaSchedule::Linear, }; Self::from_config(config) } /// Create from configuration. pub fn from_config(config: SchedulerConfig) -> Self { let betas = Self::compute_betas(&config); let alphas: Vec = betas.iter().map(|b| 1.0 - b).collect(); // Compute cumulative products let mut alphas_cumprod = Vec::with_capacity(config.num_timesteps); let mut cumprod = 1.0f32; for &alpha in &alphas { cumprod *= alpha; alphas_cumprod.push(cumprod); } // Precompute sqrt values let sqrt_alphas_cumprod: Vec = alphas_cumprod.iter().map(|a| a.sqrt()).collect(); let sqrt_one_minus_alphas_cumprod: Vec = alphas_cumprod.iter().map(|a| (1.0 - a).sqrt()).collect(); // Posterior variance: β̃_t = β_t * (1 - ᾱ_{t-1}) / (1 - ᾱ_t) let mut posterior_variance = Vec::with_capacity(config.num_timesteps); posterior_variance.push(betas[0]); // t=0 for t in 1..config.num_timesteps { let var = betas[t] * (1.0 - alphas_cumprod[t - 1]) / (1.0 - alphas_cumprod[t]); posterior_variance.push(var.max(1e-20)); // Clip for numerical stability } Self { config, betas, alphas, alphas_cumprod, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod, posterior_variance, } } /// Compute beta schedule. fn compute_betas(config: &SchedulerConfig) -> Vec { let t = config.num_timesteps; match config.schedule_type { BetaSchedule::Linear => { // Linear interpolation from beta_start to beta_end (0..t) .map(|i| { config.beta_start + (i as f32 / (t - 1) as f32) * (config.beta_end - config.beta_start) }) .collect() } BetaSchedule::Cosine => { // Cosine schedule from "Improved DDPM" let s = 0.008f32; // Small offset to prevent β from being too small let max_beta = 0.999f32; let f = |t: f32| -> f32 { let angle = (t + s) / (1.0 + s) * std::f32::consts::FRAC_PI_2; angle.cos().powi(2) }; let mut betas = Vec::with_capacity(t); for i in 0..t { let t1 = i as f32 / t as f32; let t2 = (i + 1) as f32 / t as f32; let beta = 1.0 - f(t2) / f(t1); betas.push(beta.min(max_beta)); } betas } BetaSchedule::ScaledLinear => { // Scaled linear: β varies from sqrt(β_start) to sqrt(β_end), then squared let sqrt_start = config.beta_start.sqrt(); let sqrt_end = config.beta_end.sqrt(); (0..t) .map(|i| { let sqrt_beta = sqrt_start + (i as f32 / (t - 1) as f32) * (sqrt_end - sqrt_start); sqrt_beta * sqrt_beta }) .collect() } } } /// Perform one reverse diffusion step (DDPM sampling). /// /// x_{t-1} = μ_θ(x_t, t) + σ_t * z, where z ~ N(0, 1) pub fn step, const D: usize>( &self, xt: &GenericTensor, pred_noise: &GenericTensor, t: usize, noise: Option<&GenericTensor>, ) -> GenericTensor { let alpha = self.alphas[t]; let alpha_cumprod = self.alphas_cumprod[t]; let beta = self.betas[t]; // Predict x_0 let pred_x0 = self.predict_x0(xt, pred_noise, t); // Compute mean for posterior: μ_θ(x_t, t) = √(ᾱ_{t-1}) * β_t / (1 - ᾱ_t) * x_0 // + √(α_t) * (1 - ᾱ_{t-1}) / (1 - ᾱ_t) * x_t let (coef_x0, coef_xt) = if t == 0 { (1.0, 0.0) } else { let alpha_cumprod_prev = self.alphas_cumprod[t - 1]; let one_minus_alpha_cumprod = 1.0 - alpha_cumprod; let coef_x0 = alpha_cumprod_prev.sqrt() * beta / one_minus_alpha_cumprod; let coef_xt = alpha.sqrt() * (1.0 - alpha_cumprod_prev) / one_minus_alpha_cumprod; (coef_x0, coef_xt) }; let mean = pred_x0.mul_scalar(coef_x0).add(&xt.mul_scalar(coef_xt)); // Add noise if not at t=0 if t > 0 { let std = self.posterior_variance[t].sqrt(); if let Some(z) = noise { mean.add(&z.mul_scalar(std)) } else { mean } } else { mean } } /// Get posterior variance at timestep t. pub fn posterior_variance(&self, t: usize) -> f32 { self.posterior_variance[t] } } impl NoiseScheduler for DDPMScheduler { fn num_timesteps(&self) -> usize { self.config.num_timesteps } fn alpha_cumprod(&self, t: usize) -> f32 { self.alphas_cumprod[t] } fn beta(&self, t: usize) -> f32 { self.betas[t] } } /// DDIM (Denoising Diffusion Implicit Models) scheduler. /// /// Implements deterministic sampling from Song et al., /// "Denoising Diffusion Implicit Models" (2021). #[derive(Debug, Clone)] pub struct DDIMScheduler { /// Inner DDPM scheduler (shares noise schedule) inner: DDPMScheduler, /// DDIM eta parameter (0 = deterministic, 1 = DDPM) eta: f32, /// Timesteps to use for sampling (can be subset) timesteps: Vec, } impl DDIMScheduler { /// Create a new DDIM scheduler. /// /// # Arguments /// * `ddpm` - Base DDPM scheduler /// * `eta` - Stochasticity parameter (0 = deterministic, 1 = DDPM) /// * `num_inference_steps` - Number of steps for sampling (can be < training steps) pub fn new(ddpm: DDPMScheduler, eta: f32, num_inference_steps: usize) -> Self { let num_timesteps = ddpm.num_timesteps(); // Create evenly spaced timesteps let step_ratio = num_timesteps / num_inference_steps; let timesteps: Vec = (0..num_inference_steps) .map(|i| (num_timesteps - 1 - i * step_ratio).min(num_timesteps - 1)) .collect(); Self { inner: ddpm, eta, timesteps, } } /// Get the timesteps for inference. pub fn timesteps(&self) -> &[usize] { &self.timesteps } /// Perform one DDIM sampling step. /// /// DDIM update: /// x_{t-1} = √(ᾱ_{t-1}) * x_0 + √(1 - ᾱ_{t-1} - σ²) * ε_θ + σ * z /// /// where σ² = η² * β̃_t (posterior variance) pub fn step, const D: usize>( &self, xt: &GenericTensor, pred_noise: &GenericTensor, t: usize, t_prev: usize, noise: Option<&GenericTensor>, ) -> GenericTensor { let alpha_cumprod = self.inner.alphas_cumprod[t]; let alpha_cumprod_prev = if t_prev == 0 { 1.0 // Convention: ᾱ_0 = 1 } else { self.inner.alphas_cumprod[t_prev] }; // Predict x_0 let pred_x0 = self.inner.predict_x0(xt, pred_noise, t); // Compute σ for DDIM let sigma = if t_prev > 0 { let sigma_sq = self.eta * self.eta * self.inner.posterior_variance[t]; sigma_sq.sqrt() } else { 0.0 }; // Direction pointing to x_t let dir_xt = (1.0 - alpha_cumprod_prev - sigma * sigma).max(0.0).sqrt(); // x_{t-1} = √(ᾱ_{t-1}) * x_0 + dir_xt * ε_θ + σ * z let x_prev = pred_x0 .mul_scalar(alpha_cumprod_prev.sqrt()) .add(&pred_noise.mul_scalar(dir_xt)); if sigma > 0.0 { if let Some(z) = noise { x_prev.add(&z.mul_scalar(sigma)) } else { x_prev } } else { x_prev } } } impl NoiseScheduler for DDIMScheduler { fn num_timesteps(&self) -> usize { self.inner.num_timesteps() } fn alpha_cumprod(&self, t: usize) -> f32 { self.inner.alpha_cumprod(t) } fn beta(&self, t: usize) -> f32 { self.inner.beta(t) } } #[cfg(test)] mod tests { use super::*; use rtx_backend_cpu::{CpuBackend, CpuDevice}; #[test] fn test_ddpm_scheduler_creation() { let scheduler = DDPMScheduler::new(1000, 1e-4, 0.02); assert_eq!(scheduler.num_timesteps(), 1000); assert!(scheduler.alpha_cumprod(0) > 0.99); // First alpha should be close to 1 assert!(scheduler.alpha_cumprod(999) < 0.1); // Last alpha should be small } #[test] fn test_beta_schedule_monotonic() { let scheduler = DDPMScheduler::new(100, 1e-4, 0.02); // Betas should be monotonically increasing (for linear schedule) for t in 1..100 { assert!( scheduler.beta(t) >= scheduler.beta(t - 1), "Beta should be monotonic at t={}", t ); } } #[test] fn test_alpha_cumprod_decreasing() { let scheduler = DDPMScheduler::new(100, 1e-4, 0.02); // Alpha cumprod should be monotonically decreasing for t in 1..100 { assert!( scheduler.alpha_cumprod(t) < scheduler.alpha_cumprod(t - 1), "Alpha cumprod should decrease at t={}", t ); } } #[test] fn test_add_noise_and_predict_x0() { let scheduler = DDPMScheduler::new(100, 1e-4, 0.02); let device = CpuDevice::default(); // Create test tensors let x0 = GenericTensor::::from_slice(&[1.0, 2.0, 3.0, 4.0], [2, 2], &device); let noise = GenericTensor::::from_slice(&[0.1, 0.2, 0.3, 0.4], [2, 2], &device); let t = 50; // Add noise let xt = scheduler.add_noise(&x0, &noise, t); // Predict x0 back let pred_x0 = scheduler.predict_x0(&xt, &noise, t); // Should recover x0 approximately (using known noise) let x0_vec = x0.to_vec(); let pred_vec = pred_x0.to_vec(); for (orig, pred) in x0_vec.iter().zip(pred_vec.iter()) { assert!( (orig - pred).abs() < 1e-5, "x0 recovery failed: {} vs {}", orig, pred ); } } #[test] fn test_cosine_schedule() { let config = SchedulerConfig { num_timesteps: 100, beta_start: 1e-4, beta_end: 0.02, schedule_type: BetaSchedule::Cosine, }; let scheduler = DDPMScheduler::from_config(config); // Cosine schedule should have smaller betas at start assert!(scheduler.beta(0) < 0.01); // And should never exceed max_beta for t in 0..100 { assert!(scheduler.beta(t) < 1.0); } } #[test] fn test_ddim_scheduler() { let ddpm = DDPMScheduler::new(1000, 1e-4, 0.02); let ddim = DDIMScheduler::new(ddpm, 0.0, 50); // 50 steps, deterministic assert_eq!(ddim.timesteps().len(), 50); // Timesteps should be evenly spaced from high to low assert!(ddim.timesteps()[0] > ddim.timesteps()[1]); } }