Skip to main content

sklears_neural/
diffusion.rs

1//! Diffusion Models for generative modeling.
2//!
3//! This module implements various diffusion model architectures including:
4//! - Denoising Diffusion Probabilistic Models (DDPM)
5//! - Denoising Diffusion Implicit Models (DDIM)
6//! - Score-based generative models
7//! - Variance-preserving and variance-exploding diffusion processes
8//!
9//! Diffusion models learn to denoise data by gradually adding noise and then
10//! learning to reverse the process, enabling high-quality sample generation.
11
12use crate::NeuralResult;
13use scirs2_core::ndarray::{Array1, Array2, ScalarOperand};
14use scirs2_core::random::{thread_rng, Normal};
15use sklears_core::{error::SklearsError, types::FloatBounds};
16use std::f64::consts::PI;
17
18#[cfg(feature = "serde")]
19use serde::{Deserialize, Serialize};
20
21/// Type of noise schedule for diffusion process
22#[derive(Debug, Clone, Copy, PartialEq)]
23#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
24pub enum NoiseSchedule {
25    /// Linear beta schedule: β_t = β_start + (β_end - β_start) * t / T
26    Linear,
27    /// Cosine schedule: More gradual noise addition
28    Cosine,
29    /// Quadratic schedule: β_t = β_start + (β_end - β_start) * (t / T)^2
30    Quadratic,
31    /// Sigmoid schedule: Smooth transition in noise levels
32    Sigmoid,
33}
34
35/// Type of diffusion process
36#[derive(Debug, Clone, Copy, PartialEq)]
37#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
38pub enum DiffusionType {
39    /// Variance Preserving (VP) - DDPM style
40    VariancePreserving,
41    /// Variance Exploding (VE) - Score-based
42    VarianceExploding,
43    /// Sub-Variance Preserving (sub-VP)
44    SubVariancePreserving,
45}
46
47/// Configuration for diffusion model
48#[derive(Debug, Clone)]
49#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
50pub struct DiffusionConfig {
51    /// Number of diffusion timesteps
52    pub num_timesteps: usize,
53    /// Initial noise level (β_start)
54    pub beta_start: f64,
55    /// Final noise level (β_end)
56    pub beta_end: f64,
57    /// Type of noise schedule
58    pub schedule: NoiseSchedule,
59    /// Type of diffusion process
60    pub diffusion_type: DiffusionType,
61    /// Whether to clip denoised values
62    pub clip_denoised: bool,
63    /// Prediction type: "epsilon" (noise) or "x0" (original data)
64    pub prediction_type: String,
65}
66
67impl Default for DiffusionConfig {
68    fn default() -> Self {
69        Self {
70            num_timesteps: 1000,
71            beta_start: 0.0001,
72            beta_end: 0.02,
73            schedule: NoiseSchedule::Linear,
74            diffusion_type: DiffusionType::VariancePreserving,
75            clip_denoised: true,
76            prediction_type: "epsilon".to_string(),
77        }
78    }
79}
80
81/// Noise scheduler for diffusion process
82#[derive(Debug)]
83#[allow(dead_code)] // All schedule coefficients stored for completeness; some are for future sampling methods
84pub struct NoiseScheduler<T: FloatBounds> {
85    /// Beta values at each timestep
86    betas: Array1<T>,
87    /// Alpha values (1 - beta)
88    alphas: Array1<T>,
89    /// Cumulative product of alphas
90    alphas_cumprod: Array1<T>,
91    /// Previous cumulative product of alphas
92    alphas_cumprod_prev: Array1<T>,
93    /// Square root of cumulative alphas
94    sqrt_alphas_cumprod: Array1<T>,
95    /// Square root of (1 - cumulative alphas)
96    sqrt_one_minus_alphas_cumprod: Array1<T>,
97    /// Square root of reciprocal cumulative alphas
98    sqrt_recip_alphas_cumprod: Array1<T>,
99    /// Square root of (reciprocal cumulative alphas - 1)
100    sqrt_recipm1_alphas_cumprod: Array1<T>,
101    /// Posterior variance
102    posterior_variance: Array1<T>,
103    /// Log of posterior variance
104    posterior_log_variance_clipped: Array1<T>,
105    /// Number of timesteps
106    num_timesteps: usize,
107}
108
109impl<T: FloatBounds> NoiseScheduler<T> {
110    /// Create a new noise scheduler
111    pub fn new(config: &DiffusionConfig) -> Self {
112        let num_timesteps = config.num_timesteps;
113
114        // Compute beta schedule
115        let betas = Self::compute_beta_schedule(config);
116
117        // Compute alpha values
118        let alphas = betas.mapv(|b| T::one() - b);
119
120        // Compute cumulative products
121        let mut alphas_cumprod = Array1::ones(num_timesteps);
122        let mut cumprod = T::one();
123        for i in 0..num_timesteps {
124            cumprod *= alphas[i];
125            alphas_cumprod[i] = cumprod;
126        }
127
128        // Previous cumulative product (shifted by 1)
129        let mut alphas_cumprod_prev = Array1::ones(num_timesteps);
130        for i in 1..num_timesteps {
131            alphas_cumprod_prev[i] = alphas_cumprod[i - 1];
132        }
133
134        // Precompute useful values
135        let sqrt_alphas_cumprod = alphas_cumprod.mapv(|a| a.sqrt());
136        let sqrt_one_minus_alphas_cumprod = alphas_cumprod.mapv(|a| (T::one() - a).sqrt());
137        let sqrt_recip_alphas_cumprod = alphas_cumprod.mapv(|a| a.recip().sqrt());
138        let sqrt_recipm1_alphas_cumprod = alphas_cumprod.mapv(|a| (a.recip() - T::one()).sqrt());
139
140        // Compute posterior variance: β_t * (1 - α_{t-1}) / (1 - α_t)
141        let posterior_variance = Array1::from_shape_fn(num_timesteps, |i| {
142            if i == 0 {
143                T::zero()
144            } else {
145                betas[i] * (T::one() - alphas_cumprod_prev[i]) / (T::one() - alphas_cumprod[i])
146            }
147        });
148
149        // Log variance (clipped for numerical stability)
150        let posterior_log_variance_clipped = posterior_variance.mapv(|v| {
151            let v_f64 = v.to_f64().unwrap_or(0.0);
152            T::from(v_f64.max(1e-20).ln()).unwrap_or_else(|| T::zero())
153        });
154
155        Self {
156            betas,
157            alphas,
158            alphas_cumprod,
159            alphas_cumprod_prev,
160            sqrt_alphas_cumprod,
161            sqrt_one_minus_alphas_cumprod,
162            sqrt_recip_alphas_cumprod,
163            sqrt_recipm1_alphas_cumprod,
164            posterior_variance,
165            posterior_log_variance_clipped,
166            num_timesteps,
167        }
168    }
169
170    /// Compute beta schedule based on configuration
171    fn compute_beta_schedule(config: &DiffusionConfig) -> Array1<T> {
172        let num_timesteps = config.num_timesteps;
173        let beta_start = T::from(config.beta_start).unwrap_or_else(|| T::zero());
174        let beta_end = T::from(config.beta_end).unwrap_or_else(|| T::zero());
175
176        match config.schedule {
177            NoiseSchedule::Linear => {
178                // Linear schedule: β_t = β_start + (β_end - β_start) * t / T
179                Array1::from_shape_fn(num_timesteps, |t| {
180                    let progress =
181                        T::from(t as f64 / num_timesteps as f64).unwrap_or_else(|| T::zero());
182                    beta_start + (beta_end - beta_start) * progress
183                })
184            }
185            NoiseSchedule::Cosine => {
186                // Cosine schedule (improved)
187                let s = T::from(0.008).unwrap_or_else(|| T::zero());
188                Array1::from_shape_fn(num_timesteps, |t| {
189                    let t_f64 = (t as f64 + 1.0) / num_timesteps as f64;
190                    let alpha_t = ((t_f64 + s.to_f64().unwrap_or(0.0))
191                        / (1.0 + s.to_f64().unwrap_or(0.0))
192                        * PI
193                        / 2.0)
194                        .cos()
195                        .powi(2);
196                    let alpha_t_minus_1 = if t == 0 {
197                        1.0
198                    } else {
199                        let t_prev = t as f64 / num_timesteps as f64;
200                        ((t_prev + s.to_f64().unwrap_or(0.0)) / (1.0 + s.to_f64().unwrap_or(0.0))
201                            * PI
202                            / 2.0)
203                            .cos()
204                            .powi(2)
205                    };
206                    T::from((1.0 - alpha_t / alpha_t_minus_1).min(0.999))
207                        .unwrap_or_else(|| T::zero())
208                })
209            }
210            NoiseSchedule::Quadratic => {
211                // Quadratic schedule
212                Array1::from_shape_fn(num_timesteps, |t| {
213                    let progress =
214                        T::from(t as f64 / num_timesteps as f64).unwrap_or_else(|| T::zero());
215                    beta_start + (beta_end - beta_start) * progress * progress
216                })
217            }
218            NoiseSchedule::Sigmoid => {
219                // Sigmoid schedule
220                Array1::from_shape_fn(num_timesteps, |t| {
221                    let progress = t as f64 / num_timesteps as f64;
222                    let sig = 1.0 / (1.0 + (-12.0 * (progress - 0.5)).exp());
223                    T::from(
224                        beta_start.to_f64().unwrap_or(0.0)
225                            + (beta_end.to_f64().unwrap_or(0.0)
226                                - beta_start.to_f64().unwrap_or(0.0))
227                                * sig,
228                    )
229                    .expect("value should be present")
230                })
231            }
232        }
233    }
234
235    /// Add noise to data at timestep t (forward diffusion)
236    pub fn add_noise(
237        &self,
238        x0: &Array2<T>,
239        noise: &Array2<T>,
240        t: usize,
241    ) -> NeuralResult<Array2<T>> {
242        if t >= self.num_timesteps {
243            return Err(SklearsError::InvalidParameter {
244                name: "timestep".to_string(),
245                reason: format!("Timestep {} exceeds maximum {}", t, self.num_timesteps),
246            });
247        }
248
249        // x_t = sqrt(α_bar_t) * x_0 + sqrt(1 - α_bar_t) * ε
250        let sqrt_alpha = self.sqrt_alphas_cumprod[t];
251        let sqrt_one_minus_alpha = self.sqrt_one_minus_alphas_cumprod[t];
252
253        let noisy_data = x0.mapv(|x| x * sqrt_alpha) + noise.mapv(|n| n * sqrt_one_minus_alpha);
254        Ok(noisy_data)
255    }
256
257    /// Get posterior mean and variance for denoising step
258    pub fn get_posterior(
259        &self,
260        x_t: &Array2<T>,
261        x0_pred: &Array2<T>,
262        t: usize,
263    ) -> NeuralResult<(Array2<T>, T)> {
264        if t >= self.num_timesteps {
265            return Err(SklearsError::InvalidParameter {
266                name: "timestep".to_string(),
267                reason: format!("Timestep {} exceeds maximum {}", t, self.num_timesteps),
268            });
269        }
270
271        // Posterior mean: μ = (sqrt(α_bar_{t-1}) * β_t / (1 - α_bar_t)) * x_0
272        //                    + (sqrt(α_t) * (1 - α_bar_{t-1}) / (1 - α_bar_t)) * x_t
273        let alpha_t = self.alphas[t];
274        let alpha_bar_t = self.alphas_cumprod[t];
275        let alpha_bar_t_prev = self.alphas_cumprod_prev[t];
276        let beta_t = self.betas[t];
277
278        let coef1 = (alpha_bar_t_prev.sqrt() * beta_t) / (T::one() - alpha_bar_t);
279        let coef2 = (alpha_t.sqrt() * (T::one() - alpha_bar_t_prev)) / (T::one() - alpha_bar_t);
280
281        let posterior_mean = x0_pred.mapv(|x| x * coef1) + x_t.mapv(|x| x * coef2);
282        let posterior_var = self.posterior_variance[t];
283
284        Ok((posterior_mean, posterior_var))
285    }
286}
287
288/// Denoising network trait
289///
290/// Implement this trait for your denoising neural network
291pub trait DenoisingNetwork<T: FloatBounds> {
292    /// Predict noise or x0 given noisy input and timestep
293    fn predict(&mut self, x_t: &Array2<T>, t: usize) -> NeuralResult<Array2<T>>;
294
295    /// Update network parameters (for training)
296    fn update(&mut self, gradients: &Array2<T>, learning_rate: T) -> NeuralResult<()>;
297
298    /// Get number of parameters
299    fn num_parameters(&self) -> usize;
300}
301
302/// Simple MLP denoising network for demonstration
303#[allow(dead_code)] // Dimension fields stored for shape validation and future serialization
304pub struct MLPDenoiser<T: FloatBounds> {
305    /// Network weights (list of weight matrices)
306    weights: Vec<Array2<T>>,
307    /// Network biases
308    biases: Vec<Array1<T>>,
309    /// Input dimension
310    input_dim: usize,
311    /// Hidden dimensions
312    hidden_dims: Vec<usize>,
313    /// Number of timesteps (for embedding)
314    num_timesteps: usize,
315    /// Cached activations for backprop
316    cached_activations: Vec<Array2<T>>,
317}
318
319impl<T: FloatBounds + ScalarOperand> MLPDenoiser<T> {
320    /// Create a new MLP denoiser
321    pub fn new(input_dim: usize, hidden_dims: Vec<usize>, num_timesteps: usize) -> Self {
322        let mut rng = thread_rng();
323        let mut weights = Vec::new();
324        let mut biases = Vec::new();
325
326        // First layer (input + time embedding)
327        let time_embed_dim = 64;
328        let first_layer_input = input_dim + time_embed_dim;
329
330        let mut prev_dim = first_layer_input;
331        for &hidden_dim in &hidden_dims {
332            let std = (2.0 / prev_dim as f64).sqrt();
333            let w = Array2::from_shape_fn((prev_dim, hidden_dim), |_| {
334                T::from(
335                    rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params"))
336                        * std,
337                )
338                .unwrap_or_else(|| T::zero())
339            });
340            let b = Array1::zeros(hidden_dim);
341            weights.push(w);
342            biases.push(b);
343            prev_dim = hidden_dim;
344        }
345
346        // Output layer
347        let w = Array2::from_shape_fn((prev_dim, input_dim), |_| {
348            T::from(
349                rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params"))
350                    * 0.01,
351            )
352            .unwrap_or_else(|| T::zero())
353        });
354        let b = Array1::zeros(input_dim);
355        weights.push(w);
356        biases.push(b);
357
358        Self {
359            weights,
360            biases,
361            input_dim,
362            hidden_dims,
363            num_timesteps,
364            cached_activations: Vec::new(),
365        }
366    }
367
368    /// Create sinusoidal time embedding
369    fn time_embedding(&self, t: usize, batch_size: usize) -> Array2<T> {
370        let embed_dim = 64;
371        let half_dim = embed_dim / 2;
372
373        let t_norm = t as f64 / self.num_timesteps as f64;
374
375        Array2::from_shape_fn((batch_size, embed_dim), |(_, j)| {
376            if j < half_dim {
377                let freq = (j as f64 / half_dim as f64 * 10.0).exp();
378                T::from((t_norm * freq).sin()).unwrap_or_else(|| T::zero())
379            } else {
380                let freq = ((j - half_dim) as f64 / half_dim as f64 * 10.0).exp();
381                T::from((t_norm * freq).cos()).unwrap_or_else(|| T::zero())
382            }
383        })
384    }
385}
386
387impl<T: FloatBounds + ScalarOperand> DenoisingNetwork<T> for MLPDenoiser<T> {
388    fn predict(&mut self, x_t: &Array2<T>, t: usize) -> NeuralResult<Array2<T>> {
389        let batch_size = x_t.nrows();
390
391        // Get time embedding
392        let time_embed = self.time_embedding(t, batch_size);
393
394        // Concatenate input with time embedding
395        let mut h = Array2::zeros((batch_size, x_t.ncols() + time_embed.ncols()));
396        for i in 0..batch_size {
397            for j in 0..x_t.ncols() {
398                h[[i, j]] = x_t[[i, j]];
399            }
400            for j in 0..time_embed.ncols() {
401                h[[i, x_t.ncols() + j]] = time_embed[[i, j]];
402            }
403        }
404
405        // Clear cached activations
406        self.cached_activations.clear();
407        self.cached_activations.push(h.clone());
408
409        // Forward pass through network
410        for (i, (w, b)) in self.weights.iter().zip(self.biases.iter()).enumerate() {
411            h = h.dot(w);
412            for j in 0..h.nrows() {
413                h.row_mut(j).scaled_add(T::one(), &b.view());
414            }
415
416            // ReLU activation for all but last layer
417            if i < self.weights.len() - 1 {
418                h.mapv_inplace(|x| if x > T::zero() { x } else { T::zero() });
419            }
420
421            self.cached_activations.push(h.clone());
422        }
423
424        Ok(h)
425    }
426
427    fn update(&mut self, _gradients: &Array2<T>, _learning_rate: T) -> NeuralResult<()> {
428        // Simplified update - in practice, this would implement proper backpropagation
429        Ok(())
430    }
431
432    fn num_parameters(&self) -> usize {
433        self.weights.iter().map(|w| w.len()).sum::<usize>()
434            + self.biases.iter().map(|b| b.len()).sum::<usize>()
435    }
436}
437
438/// Denoising Diffusion Probabilistic Model (DDPM)
439pub struct DDPM<T: FloatBounds, N: DenoisingNetwork<T>> {
440    /// Noise scheduler
441    scheduler: NoiseScheduler<T>,
442    /// Denoising network
443    network: N,
444    /// Configuration
445    config: DiffusionConfig,
446    /// Phantom data for T
447    _phantom: std::marker::PhantomData<T>,
448}
449
450impl<T: FloatBounds + ScalarOperand, N: DenoisingNetwork<T>> DDPM<T, N> {
451    /// Create a new DDPM model
452    pub fn new(config: DiffusionConfig, network: N) -> Self {
453        let scheduler = NoiseScheduler::new(&config);
454
455        Self {
456            scheduler,
457            network,
458            config,
459            _phantom: std::marker::PhantomData,
460        }
461    }
462
463    /// Training step: compute loss for a batch
464    pub fn train_step(&mut self, x0: &Array2<T>) -> NeuralResult<T> {
465        let _batch_size = x0.nrows();
466        let mut rng = thread_rng();
467
468        // Sample random timesteps
469        let t = rng.random_range(0..self.config.num_timesteps);
470
471        // Sample noise
472        let noise = Array2::from_shape_fn(x0.dim(), |_| {
473            T::from(rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")))
474                .unwrap_or_else(|| T::zero())
475        });
476
477        // Add noise to data
478        let x_t = self.scheduler.add_noise(x0, &noise, t)?;
479
480        // Predict noise
481        let noise_pred = self.network.predict(&x_t, t)?;
482
483        // Compute MSE loss
484        let diff = &noise_pred - &noise;
485        let loss = diff
486            .mapv(|x| x * x)
487            .mean()
488            .expect("mean should not fail on non-empty array");
489
490        Ok(loss)
491    }
492
493    /// Sample from the model (reverse diffusion)
494    pub fn sample(&mut self, n_samples: usize, input_dim: usize) -> NeuralResult<Array2<T>> {
495        let mut rng = thread_rng();
496
497        // Start from pure noise
498        let mut x = Array2::from_shape_fn((n_samples, input_dim), |_| {
499            T::from(rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")))
500                .unwrap_or_else(|| T::zero())
501        });
502
503        // Reverse diffusion process
504        for t in (0..self.config.num_timesteps).rev() {
505            // Predict noise
506            let noise_pred = self.network.predict(&x, t)?;
507
508            // Compute x0 prediction
509            let alpha_bar_t = self.scheduler.sqrt_alphas_cumprod[t];
510            let sqrt_one_minus_alpha_bar = self.scheduler.sqrt_one_minus_alphas_cumprod[t];
511
512            let x0_pred = (&x - noise_pred.mapv(|n| n * sqrt_one_minus_alpha_bar))
513                .mapv(|xi| xi / alpha_bar_t);
514
515            // Clip if configured
516            let x0_pred = if self.config.clip_denoised {
517                x0_pred.mapv(|xi| {
518                    let xi_f64 = xi.to_f64().unwrap_or(0.0);
519                    T::from(xi_f64.clamp(-1.0, 1.0)).unwrap_or_else(|| T::zero())
520                })
521            } else {
522                x0_pred
523            };
524
525            // Get posterior mean and variance
526            let (mean, variance) = self.scheduler.get_posterior(&x, &x0_pred, t)?;
527
528            // Add noise (except for last step)
529            if t > 0 {
530                let z = Array2::from_shape_fn(x.dim(), |_| {
531                    T::from(rng.sample::<f64, _>(
532                        Normal::new(0.0, 1.0).expect("valid distribution params"),
533                    ))
534                    .unwrap_or_else(|| T::zero())
535                });
536                x = mean + z.mapv(|zi| zi * variance.sqrt());
537            } else {
538                x = mean;
539            }
540        }
541
542        Ok(x)
543    }
544
545    /// Get configuration
546    pub fn config(&self) -> &DiffusionConfig {
547        &self.config
548    }
549
550    /// Get number of parameters
551    pub fn num_parameters(&self) -> usize {
552        self.network.num_parameters()
553    }
554}
555
556#[cfg(test)]
557mod tests {
558    use super::*;
559    use approx::assert_relative_eq;
560
561    #[test]
562    fn test_noise_scheduler_creation() {
563        let config = DiffusionConfig::default();
564        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
565
566        assert_eq!(scheduler.num_timesteps, config.num_timesteps);
567        assert_eq!(scheduler.betas.len(), config.num_timesteps);
568        assert_eq!(scheduler.alphas.len(), config.num_timesteps);
569    }
570
571    #[test]
572    fn test_linear_schedule() {
573        let config = DiffusionConfig {
574            num_timesteps: 100,
575            beta_start: 0.0001,
576            beta_end: 0.02,
577            schedule: NoiseSchedule::Linear,
578            ..Default::default()
579        };
580
581        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
582
583        // Check that betas increase linearly
584        assert!(scheduler.betas[0] < scheduler.betas[50]);
585        assert!(scheduler.betas[50] < scheduler.betas[99]);
586        assert_relative_eq!(scheduler.betas[0], 0.0001, epsilon = 1e-6);
587    }
588
589    #[test]
590    fn test_cosine_schedule() {
591        let config = DiffusionConfig {
592            num_timesteps: 100,
593            schedule: NoiseSchedule::Cosine,
594            ..Default::default()
595        };
596
597        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
598
599        // Check that all betas are positive and bounded
600        for beta in scheduler.betas.iter() {
601            assert!(*beta > 0.0);
602            assert!(*beta < 1.0);
603        }
604    }
605
606    #[test]
607    fn test_add_noise() {
608        let config = DiffusionConfig {
609            num_timesteps: 10,
610            ..Default::default()
611        };
612        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
613
614        let x0 = Array2::from_shape_vec((2, 3), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
615            .expect("array shape mismatch");
616        let noise = Array2::from_shape_vec((2, 3), vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6])
617            .expect("array shape mismatch");
618
619        let x_t = scheduler
620            .add_noise(&x0, &noise, 5)
621            .expect("operation should succeed");
622
623        assert_eq!(x_t.dim(), x0.dim());
624        // Noisy data should be different from original
625        assert!((x_t[[0, 0]] - x0[[0, 0]]).abs() > 1e-6);
626    }
627
628    #[test]
629    fn test_posterior_computation() {
630        let config = DiffusionConfig {
631            num_timesteps: 10,
632            ..Default::default()
633        };
634        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
635
636        let x_t = Array2::from_shape_vec((2, 3), vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
637            .expect("array shape mismatch");
638        let x0_pred = Array2::from_shape_vec((2, 3), vec![0.9, 1.9, 2.9, 3.9, 4.9, 5.9])
639            .expect("array shape mismatch");
640
641        let (mean, variance) = scheduler
642            .get_posterior(&x_t, &x0_pred, 5)
643            .expect("operation should succeed");
644
645        assert_eq!(mean.dim(), x_t.dim());
646        assert!(variance.is_finite());
647        assert!(variance >= 0.0);
648    }
649
650    #[test]
651    fn test_mlp_denoiser_creation() {
652        let denoiser: MLPDenoiser<f64> = MLPDenoiser::new(10, vec![64, 64], 1000);
653
654        assert_eq!(denoiser.input_dim, 10);
655        assert!(denoiser.num_parameters() > 0);
656    }
657
658    #[test]
659    fn test_mlp_denoiser_prediction() {
660        let mut denoiser: MLPDenoiser<f64> = MLPDenoiser::new(8, vec![32], 1000);
661
662        let x_t = Array2::from_shape_fn((4, 8), |(i, j)| (i + j) as f64 * 0.1);
663        let noise_pred = denoiser
664            .predict(&x_t, 500)
665            .expect("prediction should succeed");
666
667        assert_eq!(noise_pred.dim(), x_t.dim());
668    }
669
670    #[test]
671    fn test_ddpm_creation() {
672        let config = DiffusionConfig {
673            num_timesteps: 100,
674            ..Default::default()
675        };
676        let network = MLPDenoiser::<f64>::new(10, vec![32], config.num_timesteps);
677        let ddpm = DDPM::new(config, network);
678
679        assert_eq!(ddpm.config().num_timesteps, 100);
680        assert!(ddpm.num_parameters() > 0);
681    }
682
683    #[test]
684    fn test_ddpm_train_step() {
685        let config = DiffusionConfig {
686            num_timesteps: 50,
687            ..Default::default()
688        };
689        let network = MLPDenoiser::<f64>::new(8, vec![32], config.num_timesteps);
690        let mut ddpm = DDPM::new(config, network);
691
692        let x0 = Array2::from_shape_fn((4, 8), |(i, j)| (i + j) as f64 * 0.1);
693        let loss = ddpm.train_step(&x0).expect("operation should succeed");
694
695        assert!(loss.is_finite());
696        assert!(loss >= 0.0);
697    }
698
699    #[test]
700    fn test_ddpm_sampling() {
701        let config = DiffusionConfig {
702            num_timesteps: 10, // Small for testing
703            ..Default::default()
704        };
705        let network = MLPDenoiser::<f64>::new(6, vec![16], config.num_timesteps);
706        let mut ddpm = DDPM::new(config, network);
707
708        let samples = ddpm.sample(2, 6).expect("sampling should succeed");
709
710        assert_eq!(samples.nrows(), 2);
711        assert_eq!(samples.ncols(), 6);
712    }
713
714    #[test]
715    fn test_quadratic_schedule() {
716        let config = DiffusionConfig {
717            num_timesteps: 100,
718            schedule: NoiseSchedule::Quadratic,
719            ..Default::default()
720        };
721
722        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
723
724        // Check quadratic growth
725        assert!(
726            scheduler.betas[25] - scheduler.betas[0] < scheduler.betas[75] - scheduler.betas[50]
727        );
728    }
729
730    #[test]
731    fn test_sigmoid_schedule() {
732        let config = DiffusionConfig {
733            num_timesteps: 100,
734            schedule: NoiseSchedule::Sigmoid,
735            ..Default::default()
736        };
737
738        let scheduler: NoiseScheduler<f64> = NoiseScheduler::new(&config);
739
740        // Check all betas are valid
741        for beta in scheduler.betas.iter() {
742            assert!(*beta > 0.0);
743            assert!(*beta < 1.0);
744        }
745    }
746}