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::() / noise_data.len() as f32; let variance: f32 = noise_data.iter().map(|x| (x - mean).powi(2)).sum::() / 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(()) }