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

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(())
}