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

256 lines
6.7 KiB
Rust

use rtx_diffuse::{NoiseGenerator, NoiseSchedule, Result};
use rtx_tensor::Tensor;
#[test]
fn test_noise_generator_creation() -> Result<()> {
// Test linear schedule
let linear_gen = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(42),
)?;
assert_eq!(linear_gen.num_timesteps(), 1000);
// Test cosine schedule
let cosine_gen = NoiseGenerator::new(NoiseSchedule::Cosine { s: 0.008 }, 1000, Some(42))?;
assert_eq!(cosine_gen.num_timesteps(), 1000);
// Test scaled linear schedule
let scaled_gen = NoiseGenerator::new(
NoiseSchedule::ScaledLinear {
beta_start: 0.00085,
beta_end: 0.012,
},
1000,
Some(42),
)?;
assert_eq!(scaled_gen.num_timesteps(), 1000);
Ok(())
}
#[test]
fn test_noise_generation() -> Result<()> {
let mut noise_gen = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(42),
)?;
// Generate noise for different shapes
let noise1 = noise_gen.generate_noise(&[2, 3, 32, 32])?;
assert_eq!(noise1.shape().dims(), &[2, 3, 32, 32]);
let noise2 = noise_gen.generate_noise(&[1, 4, 64, 64])?;
assert_eq!(noise2.shape().dims(), &[1, 4, 64, 64]);
// Verify noise has reasonable statistics (approximately N(0,1))
let noise_data = noise1.data()?;
let mean: f32 = noise_data.iter().sum::<f32>() / noise_data.len() as f32;
let variance: f32 =
noise_data.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / noise_data.len() as f32;
// For standard normal, mean ≈ 0, variance ≈ 1
assert!(
(mean.abs() < 0.1),
"Mean should be close to 0, got {}",
mean
);
assert!(
(variance - 1.0).abs() < 0.2,
"Variance should be close to 1, got {}",
variance
);
Ok(())
}
#[test]
fn test_add_noise_forward_process() -> Result<()> {
let noise_gen = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(42),
)?;
// Create clean image and noise
let clean_data = vec![1.0; 2 * 3 * 8 * 8];
let x0 = Tensor::new(clean_data, vec![2, 3, 8, 8])?;
let noise_data = vec![0.5; 2 * 3 * 8 * 8];
let noise = Tensor::new(noise_data, vec![2, 3, 8, 8])?;
// Test different timesteps
let x_t_0 = noise_gen.add_noise(&x0, &noise, 0)?;
let x_t_500 = noise_gen.add_noise(&x0, &noise, 500)?;
let x_t_999 = noise_gen.add_noise(&x0, &noise, 999)?;
// At t=0, should be mostly clean image
let x_t_0_data = x_t_0.data()?;
assert!(x_t_0_data[0] > 0.9, "At t=0, should be close to original");
// At high timesteps, should be mostly noise
let x_t_999_data = x_t_999.data()?;
assert!(x_t_999_data[0] < 0.7, "At t=999, should be mostly noise");
Ok(())
}
#[test]
fn test_variance_schedule() -> Result<()> {
let noise_gen = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(42),
)?;
// Variance at t=0 should be 0
let var_0 = noise_gen.get_variance(0)?;
assert_eq!(var_0, 0.0);
// Variance should increase with timestep
let var_100 = noise_gen.get_variance(100)?;
let var_500 = noise_gen.get_variance(500)?;
let var_900 = noise_gen.get_variance(900)?;
assert!(var_100 > 0.0);
assert!(var_500 > var_100);
assert!(var_900 > var_500);
Ok(())
}
#[test]
fn test_schedule_parameters() -> Result<()> {
let noise_gen = NoiseGenerator::new(NoiseSchedule::Cosine { s: 0.008 }, 1000, Some(42))?;
// Test parameter extraction
let (sqrt_alpha_cumprod, sqrt_one_minus_alpha_cumprod, alpha_cumprod, alpha_cumprod_prev) =
noise_gen.get_schedule_params(500)?;
// Basic sanity checks
assert!(sqrt_alpha_cumprod > 0.0 && sqrt_alpha_cumprod <= 1.0);
assert!(sqrt_one_minus_alpha_cumprod >= 0.0 && sqrt_one_minus_alpha_cumprod <= 1.0);
assert!(alpha_cumprod > 0.0 && alpha_cumprod <= 1.0);
assert!(alpha_cumprod_prev > 0.0 && alpha_cumprod_prev <= 1.0);
// Square relationship
let expected_sum = sqrt_alpha_cumprod.powi(2) + sqrt_one_minus_alpha_cumprod.powi(2);
assert!(
(expected_sum - 1.0).abs() < 1e-6,
"sqrt terms should sum to 1"
);
Ok(())
}
#[test]
fn test_invalid_noise_schedules() {
// Test invalid linear schedule
let result = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.02,
beta_end: 0.0001,
}, // start > end
1000,
Some(42),
);
assert!(result.is_err());
// Test invalid cosine schedule
let result = NoiseGenerator::new(
NoiseSchedule::Cosine { s: -0.1 }, // negative s
1000,
Some(42),
);
assert!(result.is_err());
// Test zero timesteps
let result = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
0, // zero timesteps
Some(42),
);
assert!(result.is_err());
}
#[test]
fn test_timestep_bounds() -> Result<()> {
let noise_gen = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(42),
)?;
let x = Tensor::new(vec![1.0; 4], vec![1, 1, 2, 2])?;
let noise = Tensor::new(vec![0.5; 4], vec![1, 1, 2, 2])?;
// Test invalid timestep
let result = noise_gen.add_noise(&x, &noise, 1000); // >= num_timesteps
assert!(result.is_err());
let result = noise_gen.get_variance(1000);
assert!(result.is_err());
let result = noise_gen.get_schedule_params(1000);
assert!(result.is_err());
Ok(())
}
#[test]
fn test_reproducibility_with_seed() -> Result<()> {
let mut gen1 = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(12345), // Same seed
)?;
let mut gen2 = NoiseGenerator::new(
NoiseSchedule::Linear {
beta_start: 0.0001,
beta_end: 0.02,
},
1000,
Some(12345), // Same seed
)?;
let noise1 = gen1.generate_noise(&[2, 2, 4, 4])?;
let noise2 = gen2.generate_noise(&[2, 2, 4, 4])?;
// Should generate identical noise with same seed
let data1 = noise1.data()?;
let data2 = noise2.data()?;
assert_eq!(data1.len(), data2.len());
for (a, b) in data1.iter().zip(data2.iter()) {
assert!(
(a - b).abs() < 1e-6,
"Values should be identical with same seed"
);
}
Ok(())
}