256 lines
6.7 KiB
Rust
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(())
|
|
}
|