314 lines
11 KiB
Rust
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)
|
|
}
|
|
}
|