//! Diffusion Transformer (DiT) implementation for video generation. use worldgen_shared::ModelConfig; /// Diffusion Transformer. #[derive(Debug)] pub struct DiffusionTransformer { /// Hidden dimension. hidden_dim: usize, /// Number of layers. num_layers: usize, /// Number of attention heads. #[allow(dead_code)] num_heads: usize, /// Patch size. #[allow(dead_code)] patch_size: usize, /// Weights (simplified). weights: Vec, /// Use temporal attention. temporal_attention: bool, /// RNG state. rng_state: u64, } impl DiffusionTransformer { /// Create a new DiT model. pub fn new(config: &ModelConfig) -> Self { let num_params = Self::calculate_params(config); let mut model = Self { hidden_dim: config.hidden_dim, num_layers: config.num_layers, num_heads: config.num_heads, patch_size: config.patch_size, weights: vec![0.0; num_params], temporal_attention: config.temporal_attention, rng_state: 42, }; model.initialize_weights(); model } /// Calculate number of parameters. fn calculate_params(config: &ModelConfig) -> usize { let d = config.hidden_dim; let l = config.num_layers; // Embedding + transformer layers + output projection let embedding = d * 1024; // Patch embedding let transformer = l * (4 * d * d + 2 * d); // MLP + attention let output = d * 4; // Output projection embedding + transformer + output } /// Initialize weights. fn initialize_weights(&mut self) { let scale = (self.hidden_dim as f64).sqrt().recip(); let num_weights = self.weights.len(); let init_values: Vec = (0..num_weights) .map(|_| self.random_normal() * scale) .collect(); for (weight, value) in self.weights.iter_mut().zip(init_values) { *weight = value; } } /// Random number. fn random(&mut self) -> f64 { self.rng_state = self .rng_state .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (self.rng_state >> 11) as f64 / (1u64 << 53) as f64 } /// Random normal. fn random_normal(&mut self) -> f64 { let u1 = self.random() + 1e-10; let u2 = self.random(); (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos() } /// Forward pass: predict noise given noisy latents, timestep, and text embedding. pub fn forward(&self, latents: &[f64], timestep: f64, text_embedding: &[f64]) -> Vec { let mut output = vec![0.0; latents.len()]; // Timestep embedding let t_embed = self.timestep_embedding(timestep); // Process each latent for (i, &latent) in latents.iter().enumerate() { // Combine latent, timestep, and text embedding let t_idx = i % t_embed.len(); let txt_idx = i % text_embedding.len(); let combined = latent + t_embed[t_idx] * 0.1 + text_embedding[txt_idx] * 0.1; // Apply transformer layers (simplified) let mut hidden = combined; for layer in 0..self.num_layers { let weight_idx = (i + layer) % self.weights.len(); // Attention (simplified as weighted sum) hidden = hidden * self.weights[weight_idx] + hidden; // MLP hidden = hidden.tanh() * self.weights[(weight_idx + 1) % self.weights.len()]; // Temporal attention (if enabled) if self.temporal_attention { hidden *= 1.0 + 0.01 * (i as f64).sin(); } } // Output projection output[i] = hidden.tanh(); } output } /// Create timestep embedding. fn timestep_embedding(&self, timestep: f64) -> Vec { let dim = self.hidden_dim; let mut embedding = vec![0.0; dim]; for i in 0..dim / 2 { let freq = (-((i as f64) / (dim as f64 / 2.0) * 10.0)).exp(); embedding[i] = (timestep * freq).sin(); embedding[i + dim / 2] = (timestep * freq).cos(); } embedding } /// Get hidden dimension. #[must_use] pub fn hidden_dim(&self) -> usize { self.hidden_dim } /// Get number of parameters. #[must_use] pub fn num_parameters(&self) -> usize { self.weights.len() } } #[cfg(test)] mod tests { use super::*; use worldgen_shared::sample_model_config; #[test] fn test_dit_creation() { let config = sample_model_config(); let dit = DiffusionTransformer::new(&config); assert_eq!(dit.hidden_dim, config.hidden_dim); assert_eq!(dit.num_layers, config.num_layers); } #[test] fn test_forward() { let config = sample_model_config(); let dit = DiffusionTransformer::new(&config); let latents = vec![0.5; 64]; let text_embedding = vec![0.1; config.hidden_dim]; let timestep = 0.5; let output = dit.forward(&latents, timestep, &text_embedding); assert_eq!(output.len(), latents.len()); // Check output is bounded for &val in &output { assert!(val >= -1.0 && val <= 1.0); } } #[test] fn test_timestep_embedding() { let config = sample_model_config(); let dit = DiffusionTransformer::new(&config); let embedding = dit.timestep_embedding(0.5); assert_eq!(embedding.len(), config.hidden_dim); } #[test] fn test_num_parameters() { let config = sample_model_config(); let dit = DiffusionTransformer::new(&config); assert!(dit.num_parameters() > 0); } }