75 lines
1.6 KiB
Rust
75 lines
1.6 KiB
Rust
use rtx_diffuse::{NoiseGenerator, NoiseSchedule, Result};
|
|
|
|
#[test]
|
|
fn test_basic_integration() -> Result<()> {
|
|
// This is our first passing test in TDD green phase
|
|
let noise_gen = NoiseGenerator::new(
|
|
NoiseSchedule::Linear {
|
|
beta_start: 0.0001,
|
|
beta_end: 0.02,
|
|
},
|
|
1000,
|
|
Some(42),
|
|
);
|
|
|
|
// Just verify creation works
|
|
assert!(noise_gen.is_ok());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_noise_schedule_types() -> Result<()> {
|
|
// Test that we can create different schedule types
|
|
let linear = NoiseGenerator::new(
|
|
NoiseSchedule::Linear {
|
|
beta_start: 0.0001,
|
|
beta_end: 0.02,
|
|
},
|
|
100,
|
|
Some(42),
|
|
);
|
|
assert!(linear.is_ok());
|
|
|
|
let cosine = NoiseGenerator::new(NoiseSchedule::Cosine { s: 0.008 }, 100, Some(42));
|
|
assert!(cosine.is_ok());
|
|
|
|
let scaled = NoiseGenerator::new(
|
|
NoiseSchedule::ScaledLinear {
|
|
beta_start: 0.00085,
|
|
beta_end: 0.012,
|
|
},
|
|
100,
|
|
Some(42),
|
|
);
|
|
assert!(scaled.is_ok());
|
|
|
|
Ok(())
|
|
}
|
|
|
|
#[test]
|
|
fn test_models_creation() -> Result<()> {
|
|
use rtx_diffuse::{DiT, DiTConfig, UNet, UNetConfig};
|
|
|
|
// Test UNet creation
|
|
let unet = UNet::new(UNetConfig::default());
|
|
assert!(unet.is_ok());
|
|
|
|
// Test DiT creation
|
|
let dit_config = DiTConfig {
|
|
input_size: 8,
|
|
patch_size: 2,
|
|
in_channels: 3,
|
|
hidden_size: 64,
|
|
depth: 1,
|
|
num_heads: 4,
|
|
mlp_ratio: 2.0,
|
|
num_classes: None,
|
|
learn_sigma: false,
|
|
};
|
|
let dit = DiT::new(dit_config);
|
|
assert!(dit.is_ok());
|
|
|
|
Ok(())
|
|
}
|