Initial commit
This commit is contained in:
@@ -0,0 +1,317 @@
|
||||
//! 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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user