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

394 lines
11 KiB
Rust

use rtx_diffuse::{DiT, DiTConfig, Result};
use rtx_tensor::Tensor;
#[test]
fn test_dit_creation() -> Result<()> {
let dit = DiT::new(DiTConfig::default())?;
let config = dit.config();
assert_eq!(config.input_size, 32);
assert_eq!(config.patch_size, 2);
assert_eq!(config.in_channels, 4);
assert_eq!(config.hidden_size, 1152);
assert_eq!(config.depth, 28);
assert_eq!(config.num_heads, 16);
Ok(())
}
#[test]
fn test_dit_custom_config() -> Result<()> {
let config = DiTConfig {
input_size: 16,
patch_size: 4,
in_channels: 3,
hidden_size: 384,
depth: 12,
num_heads: 6,
mlp_ratio: 4.0,
num_classes: Some(100),
learn_sigma: false,
};
let dit = DiT::new(config.clone())?;
let stored_config = dit.config();
assert_eq!(stored_config.input_size, 16);
assert_eq!(stored_config.patch_size, 4);
assert_eq!(stored_config.in_channels, 3);
assert_eq!(stored_config.hidden_size, 384);
assert_eq!(stored_config.depth, 12);
assert_eq!(stored_config.num_heads, 6);
assert_eq!(stored_config.learn_sigma, false);
Ok(())
}
#[test]
fn test_dit_forward_pass() -> Result<()> {
let config = DiTConfig {
input_size: 16,
patch_size: 2,
in_channels: 4,
hidden_size: 192,
depth: 4, // Smaller for testing
num_heads: 8,
mlp_ratio: 2.0,
num_classes: Some(10),
learn_sigma: true,
};
let dit = DiT::new(config)?;
// Create input tensor [batch_size, channels, height, width]
let input_data = vec![0.5; 2 * 4 * 16 * 16];
let input = Tensor::new(input_data, vec![2, 4, 16, 16])?;
// Create timesteps tensor [batch_size]
let timesteps_data = vec![500.0, 300.0];
let timesteps = Tensor::new(timesteps_data, vec![2])?;
// Create class labels [batch_size]
let labels_data = vec![3.0, 7.0];
let labels = Tensor::new(labels_data, vec![2])?;
// Forward pass with class conditioning
let output = dit.forward(&input, &timesteps, Some(&labels))?;
// With learn_sigma=true, output channels should be 2x input channels
assert_eq!(output.shape().dims(), &[2, 8, 16, 16]); // 4 * 2 = 8 channels
// Verify output is finite
let output_data = output.data()?;
for val in &output_data {
assert!(val.is_finite(), "Output should contain finite values");
}
Ok(())
}
#[test]
fn test_dit_without_class_conditioning() -> Result<()> {
let config = DiTConfig {
input_size: 8,
patch_size: 2,
in_channels: 3,
hidden_size: 96,
depth: 2,
num_heads: 4,
mlp_ratio: 2.0,
num_classes: None, // No class conditioning
learn_sigma: false,
};
let dit = DiT::new(config)?;
let input = Tensor::new(vec![0.3; 1 * 3 * 8 * 8], vec![1, 3, 8, 8])?;
let timesteps = Tensor::new(vec![750.0], vec![1])?;
// Forward pass without class labels
let output = dit.forward(&input, &timesteps, None)?;
// Without learn_sigma, output channels should match input channels
assert_eq!(output.shape().dims(), &[1, 3, 8, 8]);
Ok(())
}
#[test]
fn test_dit_different_patch_sizes() -> Result<()> {
// Test patch size 1
let config1 = DiTConfig {
input_size: 8,
patch_size: 1,
in_channels: 1,
hidden_size: 64,
depth: 1,
num_heads: 4,
mlp_ratio: 2.0,
num_classes: None,
learn_sigma: false,
};
let dit1 = DiT::new(config1)?;
let input1 = Tensor::new(vec![0.1; 1 * 1 * 8 * 8], vec![1, 1, 8, 8])?;
let timesteps1 = Tensor::new(vec![100.0], vec![1])?;
let output1 = dit1.forward(&input1, &timesteps1, None)?;
assert_eq!(output1.shape().dims(), &[1, 1, 8, 8]);
// Test patch size 2
let config2 = DiTConfig {
input_size: 8,
patch_size: 2,
in_channels: 1,
hidden_size: 64,
depth: 1,
num_heads: 4,
mlp_ratio: 2.0,
num_classes: None,
learn_sigma: false,
};
let dit2 = DiT::new(config2)?;
let input2 = Tensor::new(vec![0.1; 1 * 1 * 8 * 8], vec![1, 1, 8, 8])?;
let timesteps2 = Tensor::new(vec![100.0], vec![1])?;
let output2 = dit2.forward(&input2, &timesteps2, None)?;
assert_eq!(output2.shape().dims(), &[1, 1, 8, 8]);
Ok(())
}
#[test]
fn test_dit_invalid_patch_size() {
// Patch size that doesn't divide image size evenly
let config = DiTConfig {
input_size: 15, // Not divisible by patch_size=4
patch_size: 4,
in_channels: 3,
hidden_size: 96,
depth: 2,
num_heads: 4,
mlp_ratio: 2.0,
num_classes: None,
learn_sigma: false,
};
let result = DiT::new(config);
assert!(result.is_err());
}
#[test]
fn test_dit_different_batch_sizes() -> Result<()> {
let config = DiTConfig {
input_size: 8,
patch_size: 2,
in_channels: 2,
hidden_size: 64,
depth: 1,
num_heads: 4,
mlp_ratio: 2.0,
num_classes: Some(5),
learn_sigma: false,
};
let dit = DiT::new(config)?;
// Test batch size 1
let input1 = Tensor::new(vec![0.2; 1 * 2 * 8 * 8], vec![1, 2, 8, 8])?;
let timesteps1 = Tensor::new(vec![200.0], vec![1])?;
let labels1 = Tensor::new(vec![2.0], vec![1])?;
let output1 = dit.forward(&input1, &timesteps1, Some(&labels1))?;
assert_eq!(output1.shape().dims(), &[1, 2, 8, 8]);
// Test batch size 3
let input3 = Tensor::new(vec![0.2; 3 * 2 * 8 * 8], vec![3, 2, 8, 8])?;
let timesteps3 = Tensor::new(vec![100.0, 300.0, 500.0], vec![3])?;
let labels3 = Tensor::new(vec![0.0, 1.0, 4.0], vec![3])?;
let output3 = dit.forward(&input3, &timesteps3, Some(&labels3))?;
assert_eq!(output3.shape().dims(), &[3, 2, 8, 8]);
Ok(())
}
#[test]
fn test_patch_embed() -> Result<()> {
use rtx_diffuse::models::dit::PatchEmbed;
let patch_embed = PatchEmbed::new(16, 4, 3, 192)?;
assert_eq!(patch_embed.num_patches(), 16); // (16/4)^2 = 16
let x = Tensor::new(vec![0.1; 2 * 3 * 16 * 16], vec![2, 3, 16, 16])?;
let patches = patch_embed.forward(&x)?;
assert_eq!(patches.shape().dims(), &[2, 16, 192]); // [batch, num_patches, embed_dim]
Ok(())
}
#[test]
fn test_patch_embed_invalid_size() {
use rtx_diffuse::models::dit::PatchEmbed;
// Image size not divisible by patch size
let result = PatchEmbed::new(15, 4, 3, 192);
assert!(result.is_err());
}
#[test]
fn test_timestep_embedder() -> Result<()> {
use rtx_diffuse::models::dit::TimestepEmbedder;
let embedder = TimestepEmbedder::new(256, 128)?;
let timesteps = Tensor::new(vec![50.0, 500.0, 950.0], vec![3])?;
let embeddings = embedder.forward(&timesteps)?;
assert_eq!(embeddings.shape().dims(), &[3, 256]);
// Check embeddings are finite
let emb_data = embeddings.data()?;
for val in &emb_data {
assert!(val.is_finite(), "Timestep embeddings should be finite");
}
Ok(())
}
#[test]
fn test_label_embedder() -> Result<()> {
use rtx_diffuse::models::dit::LabelEmbedder;
let embedder = LabelEmbedder::new(100, 512)?;
let labels = Tensor::new(vec![5.0, 23.0, 99.0], vec![3])?;
let embeddings = embedder.forward(&labels)?;
assert_eq!(embeddings.shape().dims(), &[3, 512]);
Ok(())
}
#[test]
fn test_dit_block() -> Result<()> {
use rtx_diffuse::models::dit::DiTBlock;
let block = DiTBlock::new(192, 8, 4.0)?;
let x = Tensor::new(vec![0.1; 2 * 16 * 192], vec![2, 16, 192])?; // [batch, seq, hidden]
let c = Tensor::new(vec![0.2; 2 * 192], vec![2, 192])?; // [batch, hidden]
let output = block.forward(&x, &c)?;
assert_eq!(output.shape().dims(), &[2, 16, 192]);
Ok(())
}
#[test]
fn test_dit_block_invalid_heads() {
use rtx_diffuse::models::dit::DiTBlock;
// Hidden size not divisible by num_heads
let result = DiTBlock::new(193, 8, 4.0); // 193 not divisible by 8
assert!(result.is_err());
// Valid configuration
let result = DiTBlock::new(192, 8, 4.0); // 192 divisible by 8
assert!(result.is_ok());
}
#[test]
fn test_dit_learn_sigma_variants() -> Result<()> {
// Test with learn_sigma = false
let config_no_sigma = 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_no_sigma = DiT::new(config_no_sigma)?;
let input = Tensor::new(vec![0.1; 1 * 3 * 8 * 8], vec![1, 3, 8, 8])?;
let timesteps = Tensor::new(vec![400.0], vec![1])?;
let output_no_sigma = dit_no_sigma.forward(&input, &timesteps, None)?;
assert_eq!(output_no_sigma.shape().dims(), &[1, 3, 8, 8]); // Same as input channels
// Test with learn_sigma = true
let config_with_sigma = 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: true,
};
let dit_with_sigma = DiT::new(config_with_sigma)?;
let output_with_sigma = dit_with_sigma.forward(&input, &timesteps, None)?;
assert_eq!(output_with_sigma.shape().dims(), &[1, 6, 8, 8]); // 2x input channels
Ok(())
}
#[test]
fn test_dit_deterministic_output() -> Result<()> {
let config = DiTConfig {
input_size: 8,
patch_size: 2,
in_channels: 2,
hidden_size: 64,
depth: 1,
num_heads: 4,
mlp_ratio: 2.0,
num_classes: None,
learn_sigma: false,
};
let dit = DiT::new(config)?;
let input = Tensor::new(vec![0.3; 1 * 2 * 8 * 8], vec![1, 2, 8, 8])?;
let timesteps = Tensor::new(vec![600.0], vec![1])?;
// Multiple forward passes should be deterministic
let output1 = dit.forward(&input, &timesteps, None)?;
let output2 = dit.forward(&input, &timesteps, None)?;
let data1 = output1.data()?;
let data2 = output2.data()?;
for (a, b) in data1.iter().zip(data2.iter()) {
assert!((a - b).abs() < 1e-6, "DiT should be deterministic");
}
Ok(())
}
#[test]
fn test_dit_memory_efficiency() -> Result<()> {
// Test creating multiple DiTs with small configurations
for i in 0..3 {
let config = DiTConfig {
input_size: 8,
patch_size: 2,
in_channels: 1,
hidden_size: 32,
depth: 1,
num_heads: 2,
mlp_ratio: 2.0,
num_classes: Some(5),
learn_sigma: false,
};
let dit = DiT::new(config)?;
let input = Tensor::new(vec![0.1; 1 * 1 * 8 * 8], vec![1, 1, 8, 8])?;
let timesteps = Tensor::new(vec![i as f32 * 100.0], vec![1])?;
let labels = Tensor::new(vec![i as f32], vec![1])?;
let output = dit.forward(&input, &timesteps, Some(&labels))?;
assert_eq!(output.shape().dims(), &[1, 1, 8, 8]);
}
Ok(())
}