Consistent formatting pass: line wrapping, import sorting, trailing whitespace removal, let-chain indentation, merged derive attributes, and unsafe block reformatting. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
319 lines
10 KiB
Rust
319 lines
10 KiB
Rust
//! 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<f64>,
|
|
/// Alpha values (1 - beta).
|
|
alphas: Vec<f64>,
|
|
/// Cumulative product of alphas.
|
|
alphas_cumprod: Vec<f64>,
|
|
/// Square root of alphas_cumprod.
|
|
sqrt_alphas_cumprod: Vec<f64>,
|
|
/// Square root of (1 - alphas_cumprod).
|
|
sqrt_one_minus_alphas_cumprod: Vec<f64>,
|
|
}
|
|
|
|
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<f64> = 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<f64> = 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<f64> = alphas_cumprod.iter().map(|a| a.sqrt()).collect();
|
|
let sqrt_one_minus_alphas_cumprod: Vec<f64> =
|
|
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<f64> {
|
|
(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<f64> {
|
|
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<f64> {
|
|
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<f64> {
|
|
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<f64> {
|
|
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<f64> {
|
|
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]);
|
|
}
|
|
}
|
|
}
|