394 lines
11 KiB
Rust
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, ×teps, 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, ×teps, 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, ×teps1, 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, ×teps2, 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, ×teps1, 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, ×teps3, 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(×teps)?;
|
|
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, ×teps, 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, ×teps, 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, ×teps, None)?;
|
|
let output2 = dit.forward(&input, ×teps, 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, ×teps, Some(&labels))?;
|
|
assert_eq!(output.shape().dims(), &[1, 1, 8, 8]);
|
|
}
|
|
|
|
Ok(())
|
|
}
|