Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
+317
View File
@@ -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]);
}
}
}