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, } impl DiffusionScheduler { /// Create a new diffusion scheduler pub fn new( scheduler_type: SchedulerType, noise_generator: NoiseGenerator, num_inference_steps: u32, ) -> Result { 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 { 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 { self.noise_generator.add_noise(original, noise, timestep) } /// Get variance at timestep pub fn get_variance(&self, timestep: u32) -> Result { self.noise_generator.get_variance(timestep) } /// Scale model input according to scheduler requirements pub fn scale_model_input(&self, sample: &Tensor, timestep: u32) -> Result { // 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 { 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 { // 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 { 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 { 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 { // 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> { 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) } }