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