//! Noise scheduler for diffusion process. use worldgen_shared::SchedulerConfig; /// Noise scheduler for the diffusion denoising process. #[derive(Debug)] pub struct NoiseScheduler { /// Number of timesteps. num_timesteps: usize, /// Beta values for each timestep. #[allow(dead_code)] betas: Vec, /// Alpha values (1 - beta). alphas: Vec, /// Cumulative product of alphas. alphas_cumprod: Vec, /// Square root of alphas_cumprod. sqrt_alphas_cumprod: Vec, /// Square root of (1 - alphas_cumprod). sqrt_one_minus_alphas_cumprod: Vec, } impl NoiseScheduler { /// Create a new noise scheduler. pub fn new(config: &SchedulerConfig) -> Self { let num_timesteps = config.num_timesteps; let beta_start = config.beta_start as f64; let beta_end = config.beta_end as f64; // Generate beta schedule let betas: Vec = match config.beta_schedule.as_str() { "linear" => Self::linear_schedule(num_timesteps, beta_start, beta_end), "cosine" => Self::cosine_schedule(num_timesteps), "quadratic" => Self::quadratic_schedule(num_timesteps, beta_start, beta_end), _ => Self::linear_schedule(num_timesteps, beta_start, beta_end), }; // Compute alpha values let alphas: Vec = betas.iter().map(|b| 1.0 - b).collect(); // Compute cumulative product of alphas let mut alphas_cumprod = Vec::with_capacity(num_timesteps); let mut cumprod = 1.0; for &alpha in &alphas { cumprod *= alpha; alphas_cumprod.push(cumprod); } // Compute derived 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(); Self { num_timesteps, betas, alphas, alphas_cumprod, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod, } } /// Linear beta schedule. fn linear_schedule(num_timesteps: usize, beta_start: f64, beta_end: f64) -> Vec { (0..num_timesteps) .map(|i| { let t = i as f64 / (num_timesteps - 1).max(1) as f64; beta_start + t * (beta_end - beta_start) }) .collect() } /// Cosine beta schedule (improved noise schedule). fn cosine_schedule(num_timesteps: usize) -> Vec { let s = 0.008; // Small offset to prevent singularity let max_beta = 0.999; let alpha_bar = |t: f64| -> f64 { let f = ((t + s) / (1.0 + s) * std::f64::consts::FRAC_PI_2).cos(); f * f }; let mut betas = Vec::with_capacity(num_timesteps); for i in 0..num_timesteps { let t1 = i as f64 / num_timesteps as f64; let t2 = (i + 1) as f64 / num_timesteps as f64; let beta = 1.0 - alpha_bar(t2) / alpha_bar(t1); betas.push(beta.min(max_beta)); } betas } /// Quadratic beta schedule. fn quadratic_schedule(num_timesteps: usize, beta_start: f64, beta_end: f64) -> Vec { let sqrt_start = beta_start.sqrt(); let sqrt_end = beta_end.sqrt(); (0..num_timesteps) .map(|i| { let t = i as f64 / (num_timesteps - 1).max(1) as f64; let sqrt_beta = sqrt_start + t * (sqrt_end - sqrt_start); sqrt_beta * sqrt_beta }) .collect() } /// Get timesteps for inference. pub fn get_timesteps(&self, num_inference_steps: usize) -> Vec { if num_inference_steps >= self.num_timesteps { return (0..self.num_timesteps) .map(|i| i as f64 / self.num_timesteps as f64) .rev() .collect(); } // Uniform spacing for fewer steps let step = self.num_timesteps as f64 / num_inference_steps as f64; (0..num_inference_steps) .map(|i| { let idx = ((num_inference_steps - 1 - i) as f64 * step) as usize; idx as f64 / self.num_timesteps as f64 }) .collect() } /// Perform one denoising step. pub fn step( &self, latents: &[f64], noise_pred: &[f64], timestep: f64, physics_guidance: Option<&[f64]>, ) -> Vec { let t_idx = ((timestep * (self.num_timesteps - 1) as f64) as usize).min(self.num_timesteps - 1); let _alpha = self.alphas[t_idx]; let sqrt_one_minus_alpha_cumprod = self.sqrt_one_minus_alphas_cumprod[t_idx]; // DDIM-style update let mut output = vec![0.0; latents.len()]; for i in 0..latents.len() { // Predict x0 from noise let pred_x0 = (latents[i] - sqrt_one_minus_alpha_cumprod * noise_pred[i]) / self.sqrt_alphas_cumprod[t_idx].max(1e-10); // Compute previous sample let prev_alpha_cumprod = if t_idx > 0 { self.alphas_cumprod[t_idx - 1] } else { 1.0 }; let pred_sample = prev_alpha_cumprod.sqrt() * pred_x0 + (1.0 - prev_alpha_cumprod).sqrt() * noise_pred[i]; output[i] = pred_sample; // Apply physics guidance if provided if let Some(guidance) = physics_guidance && i < guidance.len() { output[i] += guidance[i] * 0.1; // Scale guidance } } output } /// Add noise to samples (forward process). pub fn add_noise(&self, samples: &[f64], noise: &[f64], timestep: f64) -> Vec { let t_idx = ((timestep * (self.num_timesteps - 1) as f64) as usize).min(self.num_timesteps - 1); let sqrt_alpha_cumprod = self.sqrt_alphas_cumprod[t_idx]; let sqrt_one_minus_alpha_cumprod = self.sqrt_one_minus_alphas_cumprod[t_idx]; samples .iter() .zip(noise.iter()) .map(|(&x, &n)| sqrt_alpha_cumprod * x + sqrt_one_minus_alpha_cumprod * n) .collect() } /// Get signal-to-noise ratio at timestep. pub fn get_snr(&self, timestep: f64) -> f64 { let t_idx = ((timestep * (self.num_timesteps - 1) as f64) as usize).min(self.num_timesteps - 1); let alpha_cumprod = self.alphas_cumprod[t_idx]; alpha_cumprod / (1.0 - alpha_cumprod).max(1e-10) } /// Get number of timesteps. #[must_use] pub fn num_timesteps(&self) -> usize { self.num_timesteps } } #[cfg(test)] mod tests { use super::*; use worldgen_shared::sample_scheduler_config; #[test] fn test_scheduler_creation() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); assert_eq!(scheduler.num_timesteps, config.num_timesteps); assert_eq!(scheduler.betas.len(), config.num_timesteps); assert_eq!(scheduler.alphas_cumprod.len(), config.num_timesteps); } #[test] fn test_linear_schedule() { let betas = NoiseScheduler::linear_schedule(10, 0.0001, 0.02); assert_eq!(betas.len(), 10); assert!((betas[0] - 0.0001).abs() < 1e-6); assert!((betas[9] - 0.02).abs() < 1e-6); // Check monotonically increasing for i in 1..betas.len() { assert!(betas[i] >= betas[i - 1]); } } #[test] fn test_cosine_schedule() { let betas = NoiseScheduler::cosine_schedule(100); assert_eq!(betas.len(), 100); // All betas should be in valid range for &beta in &betas { assert!(beta > 0.0); assert!(beta <= 0.999); } } #[test] fn test_get_timesteps() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); let timesteps = scheduler.get_timesteps(10); assert_eq!(timesteps.len(), 10); // Timesteps should be decreasing (reverse order) for i in 1..timesteps.len() { assert!(timesteps[i] <= timesteps[i - 1]); } } #[test] fn test_step() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); let latents = vec![1.0; 64]; let noise_pred = vec![0.1; 64]; let timestep = 0.5; let output = scheduler.step(&latents, &noise_pred, timestep, None); assert_eq!(output.len(), latents.len()); } #[test] fn test_step_with_physics() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); let latents = vec![1.0; 64]; let noise_pred = vec![0.1; 64]; let physics_guidance = vec![0.5; 64]; let timestep = 0.5; let output = scheduler.step(&latents, &noise_pred, timestep, Some(&physics_guidance)); assert_eq!(output.len(), latents.len()); } #[test] fn test_add_noise() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); let samples = vec![1.0; 32]; let noise = vec![0.5; 32]; let timestep = 0.5; let noisy = scheduler.add_noise(&samples, &noise, timestep); assert_eq!(noisy.len(), samples.len()); } #[test] fn test_snr() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); // SNR should decrease over time let snr_early = scheduler.get_snr(0.1); let snr_late = scheduler.get_snr(0.9); assert!(snr_early > snr_late); } #[test] fn test_alphas_cumprod_decreasing() { let config = sample_scheduler_config(); let scheduler = NoiseScheduler::new(&config); // Cumulative product of alphas should be decreasing for i in 1..scheduler.alphas_cumprod.len() { assert!(scheduler.alphas_cumprod[i] <= scheduler.alphas_cumprod[i - 1]); } } }