Initial commit
This commit is contained in:
@@ -0,0 +1,481 @@
|
||||
//! 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<B: Backend<FloatElem = f32>, const D: usize>(
|
||||
&self,
|
||||
x0: &GenericTensor<B, D>,
|
||||
noise: &GenericTensor<B, D>,
|
||||
t: usize,
|
||||
) -> GenericTensor<B, D> {
|
||||
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<B: Backend<FloatElem = f32>, const D: usize>(
|
||||
&self,
|
||||
xt: &GenericTensor<B, D>,
|
||||
pred_noise: &GenericTensor<B, D>,
|
||||
t: usize,
|
||||
) -> GenericTensor<B, D> {
|
||||
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<f32>,
|
||||
/// Alpha values (1 - beta)
|
||||
alphas: Vec<f32>,
|
||||
/// Cumulative product of alphas
|
||||
alphas_cumprod: Vec<f32>,
|
||||
/// √(ᾱ_t)
|
||||
sqrt_alphas_cumprod: Vec<f32>,
|
||||
/// √(1 - ᾱ_t)
|
||||
sqrt_one_minus_alphas_cumprod: Vec<f32>,
|
||||
/// Posterior variance σ²_t
|
||||
posterior_variance: Vec<f32>,
|
||||
}
|
||||
|
||||
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<f32> = 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<f32> = alphas_cumprod.iter().map(|a| a.sqrt()).collect();
|
||||
let sqrt_one_minus_alphas_cumprod: Vec<f32> =
|
||||
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<f32> {
|
||||
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<B: Backend<FloatElem = f32>, const D: usize>(
|
||||
&self,
|
||||
xt: &GenericTensor<B, D>,
|
||||
pred_noise: &GenericTensor<B, D>,
|
||||
t: usize,
|
||||
noise: Option<&GenericTensor<B, D>>,
|
||||
) -> GenericTensor<B, D> {
|
||||
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<usize>,
|
||||
}
|
||||
|
||||
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<usize> = (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<B: Backend<FloatElem = f32>, const D: usize>(
|
||||
&self,
|
||||
xt: &GenericTensor<B, D>,
|
||||
pred_noise: &GenericTensor<B, D>,
|
||||
t: usize,
|
||||
t_prev: usize,
|
||||
noise: Option<&GenericTensor<B, D>>,
|
||||
) -> GenericTensor<B, D> {
|
||||
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::<CpuBackend, 2>::from_slice(&[1.0, 2.0, 3.0, 4.0], [2, 2], &device);
|
||||
let noise =
|
||||
GenericTensor::<CpuBackend, 2>::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]);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user