1use 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#[derive(Debug, Clone, Copy, PartialEq)]
23#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
24pub enum NoiseSchedule {
25 Linear,
27 Cosine,
29 Quadratic,
31 Sigmoid,
33}
34
35#[derive(Debug, Clone, Copy, PartialEq)]
37#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
38pub enum DiffusionType {
39 VariancePreserving,
41 VarianceExploding,
43 SubVariancePreserving,
45}
46
47#[derive(Debug, Clone)]
49#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
50pub struct DiffusionConfig {
51 pub num_timesteps: usize,
53 pub beta_start: f64,
55 pub beta_end: f64,
57 pub schedule: NoiseSchedule,
59 pub diffusion_type: DiffusionType,
61 pub clip_denoised: bool,
63 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#[derive(Debug)]
83#[allow(dead_code)] pub struct NoiseScheduler<T: FloatBounds> {
85 betas: Array1<T>,
87 alphas: Array1<T>,
89 alphas_cumprod: Array1<T>,
91 alphas_cumprod_prev: Array1<T>,
93 sqrt_alphas_cumprod: Array1<T>,
95 sqrt_one_minus_alphas_cumprod: Array1<T>,
97 sqrt_recip_alphas_cumprod: Array1<T>,
99 sqrt_recipm1_alphas_cumprod: Array1<T>,
101 posterior_variance: Array1<T>,
103 posterior_log_variance_clipped: Array1<T>,
105 num_timesteps: usize,
107}
108
109impl<T: FloatBounds> NoiseScheduler<T> {
110 pub fn new(config: &DiffusionConfig) -> Self {
112 let num_timesteps = config.num_timesteps;
113
114 let betas = Self::compute_beta_schedule(config);
116
117 let alphas = betas.mapv(|b| T::one() - b);
119
120 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 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 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 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 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 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 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 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 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 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 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 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 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 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
288pub trait DenoisingNetwork<T: FloatBounds> {
292 fn predict(&mut self, x_t: &Array2<T>, t: usize) -> NeuralResult<Array2<T>>;
294
295 fn update(&mut self, gradients: &Array2<T>, learning_rate: T) -> NeuralResult<()>;
297
298 fn num_parameters(&self) -> usize;
300}
301
302#[allow(dead_code)] pub struct MLPDenoiser<T: FloatBounds> {
305 weights: Vec<Array2<T>>,
307 biases: Vec<Array1<T>>,
309 input_dim: usize,
311 hidden_dims: Vec<usize>,
313 num_timesteps: usize,
315 cached_activations: Vec<Array2<T>>,
317}
318
319impl<T: FloatBounds + ScalarOperand> MLPDenoiser<T> {
320 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 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 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 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 let time_embed = self.time_embedding(t, batch_size);
393
394 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 self.cached_activations.clear();
407 self.cached_activations.push(h.clone());
408
409 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 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 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
438pub struct DDPM<T: FloatBounds, N: DenoisingNetwork<T>> {
440 scheduler: NoiseScheduler<T>,
442 network: N,
444 config: DiffusionConfig,
446 _phantom: std::marker::PhantomData<T>,
448}
449
450impl<T: FloatBounds + ScalarOperand, N: DenoisingNetwork<T>> DDPM<T, N> {
451 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 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 let t = rng.random_range(0..self.config.num_timesteps);
470
471 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 let x_t = self.scheduler.add_noise(x0, &noise, t)?;
479
480 let noise_pred = self.network.predict(&x_t, t)?;
482
483 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 pub fn sample(&mut self, n_samples: usize, input_dim: usize) -> NeuralResult<Array2<T>> {
495 let mut rng = thread_rng();
496
497 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 for t in (0..self.config.num_timesteps).rev() {
505 let noise_pred = self.network.predict(&x, t)?;
507
508 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 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 let (mean, variance) = self.scheduler.get_posterior(&x, &x0_pred, t)?;
527
528 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 pub fn config(&self) -> &DiffusionConfig {
547 &self.config
548 }
549
550 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 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 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 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, ..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 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 for beta in scheduler.betas.iter() {
742 assert!(*beta > 0.0);
743 assert!(*beta < 1.0);
744 }
745 }
746}