Files
rustytorch/crates/models/rtx-diffuse/src/scheduler.rs
T
2026-03-04 00:08:42 +00:00

314 lines
11 KiB
Rust

use crate::error::{DiffusionError, Result};
use crate::noise::NoiseGenerator;
use rtx_tensor::Tensor;
/// Different sampling schedulers for the reverse diffusion process
#[derive(Debug, Clone)]
pub enum SchedulerType {
/// Denoising Diffusion Implicit Models (DDIM)
DDIM { eta: f32 },
/// DPM-Solver++
DPMPlusPlus,
/// Euler Ancestral
EulerAncestral,
/// DDPM (original)
DDPM,
/// UniPC (Unified Predictor-Corrector)
UniPC {
predictor_order: u8,
corrector_order: u8,
use_corrector: bool,
},
}
/// Manages the reverse diffusion sampling process
#[derive(Debug, Clone)]
pub struct DiffusionScheduler {
scheduler_type: SchedulerType,
noise_generator: NoiseGenerator,
num_inference_steps: u32,
timesteps: Vec<u32>,
}
impl DiffusionScheduler {
/// Create a new diffusion scheduler
pub fn new(
scheduler_type: SchedulerType,
noise_generator: NoiseGenerator,
num_inference_steps: u32,
) -> Result<Self> {
if num_inference_steps == 0 {
return Err(DiffusionError::Scheduler {
message: "Number of inference steps must be greater than 0".to_string(),
});
}
let timesteps = Self::compute_timesteps(&noise_generator, num_inference_steps)?;
Ok(Self {
scheduler_type,
noise_generator,
num_inference_steps,
timesteps,
})
}
/// Perform one step of the reverse diffusion process
pub fn step(
&self,
model_output: &Tensor,
timestep: u32,
sample: &Tensor,
generator: Option<&mut rand::rngs::StdRng>,
) -> Result<Tensor> {
match &self.scheduler_type {
SchedulerType::DDIM { eta } => self.ddim_step(model_output, timestep, sample, *eta),
SchedulerType::DPMPlusPlus => self.dpm_plusplus_step(model_output, timestep, sample),
SchedulerType::EulerAncestral => {
self.euler_ancestral_step(model_output, timestep, sample, generator)
}
SchedulerType::DDPM => self.ddpm_step(model_output, timestep, sample, generator),
SchedulerType::UniPC {
predictor_order,
corrector_order,
use_corrector,
} => self.unipc_step(
model_output,
timestep,
sample,
*predictor_order,
*corrector_order,
*use_corrector,
),
}
}
/// Get the timesteps for inference
pub fn timesteps(&self) -> &[u32] {
&self.timesteps
}
/// Add noise to the initial sample (for training)
pub fn add_noise(&self, original: &Tensor, noise: &Tensor, timestep: u32) -> Result<Tensor> {
self.noise_generator.add_noise(original, noise, timestep)
}
/// Get variance at timestep
pub fn get_variance(&self, timestep: u32) -> Result<f32> {
self.noise_generator.get_variance(timestep)
}
/// Scale model input according to scheduler requirements
pub fn scale_model_input(&self, sample: &Tensor, timestep: u32) -> Result<Tensor> {
// For most schedulers, no scaling is needed
// Some schedulers like UniPC may require scaling
match &self.scheduler_type {
SchedulerType::DDIM { .. }
| SchedulerType::DDPM
| SchedulerType::EulerAncestral
| SchedulerType::DPMPlusPlus
| SchedulerType::UniPC { .. } => Ok(sample.clone()),
}
}
fn ddim_step(
&self,
model_output: &Tensor,
timestep: u32,
sample: &Tensor,
eta: f32,
) -> Result<Tensor> {
let (sqrt_alpha_cumprod, sqrt_one_minus_alpha_cumprod, alpha_cumprod, alpha_cumprod_prev) =
self.noise_generator.get_schedule_params(timestep)?;
// Predict x0 from model output (assuming model predicts noise)
// x0 = (x_t - sqrt(1-alpha_cumprod) * noise) / sqrt(alpha_cumprod)
let scaled_noise = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod)?;
let x_minus_noise = sample.subtract(&scaled_noise)?;
let pred_x0 = x_minus_noise.scalar_mul(1.0 / sqrt_alpha_cumprod)?;
// Compute variance
let variance = if eta == 0.0 {
0.0
} else {
let beta_t = 1.0 - alpha_cumprod / alpha_cumprod_prev;
eta * eta * beta_t
};
let std_dev = variance.sqrt();
// Compute direction pointing to x_t
let sqrt_one_minus_alpha_cumprod_prev = (1.0 - alpha_cumprod_prev).sqrt();
let direction = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod_prev - std_dev)?;
// Compute x_{t-1}
let sqrt_alpha_cumprod_prev = alpha_cumprod_prev.sqrt();
let prev_sample = pred_x0.scalar_mul(sqrt_alpha_cumprod_prev)?;
let prev_sample = prev_sample.add(&direction)?;
if std_dev > 0.0 {
// Add noise if eta > 0
let noise_shape = sample.shape();
// For now, we'll assume zero noise - in practice, you'd generate random noise
// This is a simplification for the initial implementation
}
Ok(prev_sample)
}
fn dpm_plusplus_step(
&self,
model_output: &Tensor,
timestep: u32,
sample: &Tensor,
) -> Result<Tensor> {
// Simplified DPM++ implementation
// In practice, this would involve more complex integration schemes
let (sqrt_alpha_cumprod, sqrt_one_minus_alpha_cumprod, alpha_cumprod, alpha_cumprod_prev) =
self.noise_generator.get_schedule_params(timestep)?;
// Predict x0
let scaled_noise = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod)?;
let x_minus_noise = sample.subtract(&scaled_noise)?;
let pred_x0 = x_minus_noise.scalar_mul(1.0 / sqrt_alpha_cumprod)?;
// DPM++ uses a different integration scheme
let sqrt_alpha_cumprod_prev = alpha_cumprod_prev.sqrt();
let sqrt_one_minus_alpha_cumprod_prev = (1.0 - alpha_cumprod_prev).sqrt();
let prev_sample = pred_x0.scalar_mul(sqrt_alpha_cumprod_prev)?;
let noise_component = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod_prev)?;
prev_sample
.add(&noise_component)
.map_err(DiffusionError::Tensor)
}
fn euler_ancestral_step(
&self,
model_output: &Tensor,
timestep: u32,
sample: &Tensor,
_generator: Option<&mut rand::rngs::StdRng>,
) -> Result<Tensor> {
let (sqrt_alpha_cumprod, sqrt_one_minus_alpha_cumprod, alpha_cumprod, alpha_cumprod_prev) =
self.noise_generator.get_schedule_params(timestep)?;
// Euler ancestral method
let sigma = sqrt_one_minus_alpha_cumprod / sqrt_alpha_cumprod;
let sigma_prev = ((1.0 - alpha_cumprod_prev) / alpha_cumprod_prev).sqrt();
// Compute derivative
let derivative = sample.add(&model_output.scalar_mul(sigma)?)?;
let derivative = derivative.scalar_mul(-1.0 / sigma)?;
// Euler step
let dt = sigma_prev - sigma;
let prev_sample = sample.add(&derivative.scalar_mul(dt)?)?;
Ok(prev_sample)
}
fn ddpm_step(
&self,
model_output: &Tensor,
timestep: u32,
sample: &Tensor,
_generator: Option<&mut rand::rngs::StdRng>,
) -> Result<Tensor> {
let (sqrt_alpha_cumprod, sqrt_one_minus_alpha_cumprod, alpha_cumprod, alpha_cumprod_prev) =
self.noise_generator.get_schedule_params(timestep)?;
// DDPM step: predict x0 and sample from posterior
let scaled_noise = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod)?;
let x_minus_noise = sample.subtract(&scaled_noise)?;
let pred_x0 = x_minus_noise.scalar_mul(1.0 / sqrt_alpha_cumprod)?;
// Posterior mean
let coeff1 = (alpha_cumprod_prev.sqrt() * (1.0 - alpha_cumprod / alpha_cumprod_prev))
/ (1.0 - alpha_cumprod);
let coeff2 = ((alpha_cumprod / alpha_cumprod_prev).sqrt() * (1.0 - alpha_cumprod_prev))
/ (1.0 - alpha_cumprod);
let mean = pred_x0.scalar_mul(coeff1)?;
let mean = mean.add(&sample.scalar_mul(coeff2)?)?;
// For now, return mean without noise (deterministic)
// In practice, you'd add scaled random noise
Ok(mean)
}
fn unipc_step(
&self,
model_output: &Tensor,
timestep: u32,
sample: &Tensor,
predictor_order: u8,
corrector_order: u8,
use_corrector: bool,
) -> Result<Tensor> {
// Simplified UniPC implementation for scheduler integration
// This delegates to a basic predictor-corrector approach
// First-order predictor step (similar to DDIM)
let (sqrt_alpha_cumprod, sqrt_one_minus_alpha_cumprod, alpha_cumprod, alpha_cumprod_prev) =
self.noise_generator.get_schedule_params(timestep)?;
// Predict x0 from model output (assuming epsilon prediction)
let scaled_noise = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod)?;
let x_minus_noise = sample.subtract(&scaled_noise)?;
let pred_x0 = x_minus_noise.scalar_mul(1.0 / sqrt_alpha_cumprod)?;
// Compute predictor step
let sqrt_alpha_cumprod_prev = alpha_cumprod_prev.sqrt();
let sqrt_one_minus_alpha_cumprod_prev = (1.0 - alpha_cumprod_prev).sqrt();
let predictor_result = pred_x0.scalar_mul(sqrt_alpha_cumprod_prev)?;
let direction = model_output.scalar_mul(sqrt_one_minus_alpha_cumprod_prev)?;
let predictor_result = predictor_result.add(&direction)?;
// Apply corrector step if enabled
if use_corrector {
// Simple corrector: blend with original prediction
let correction_weight = match corrector_order {
1 => 0.1,
2 => 0.05,
_ => 0.02,
};
let corrected = predictor_result.scalar_mul(1.0 - correction_weight)?;
let correction = model_output.scalar_mul(correction_weight)?;
corrected.add(&correction).map_err(DiffusionError::Tensor)
} else {
Ok(predictor_result)
}
}
fn compute_timesteps(
noise_generator: &NoiseGenerator,
num_inference_steps: u32,
) -> Result<Vec<u32>> {
let num_train_timesteps = noise_generator.num_timesteps();
if num_inference_steps > num_train_timesteps {
return Err(DiffusionError::Scheduler {
message: format!(
"Inference steps ({}) cannot exceed training timesteps ({})",
num_inference_steps, num_train_timesteps
),
});
}
let step_size = num_train_timesteps / num_inference_steps;
let mut timesteps = Vec::with_capacity(num_inference_steps as usize);
for i in 0..num_inference_steps {
let timestep = num_train_timesteps - 1 - (i * step_size);
timesteps.push(timestep);
}
timesteps.reverse(); // Start from highest timestep and go to lowest
Ok(timesteps)
}
}