//! WorldGen - Video Diffusion World Simulator. //! //! This demo showcases DiT-based video generation with physics grounding, //! enabling generation of physically plausible video from text prompts. pub mod diffusion; pub mod physics; pub mod sample_data; pub mod scheduler; use thiserror::Error; use worldgen_shared::{ GenerationRequest, GenerationResult, ModelConfig, SchedulerConfig, TrainingConfig, TrainingProgress, VideoClip, }; /// Errors that can occur in WorldGen. #[derive(Debug, Error)] pub enum WorldGenError { /// Invalid generation request. #[error("Invalid generation request: {0}")] InvalidRequest(String), /// Generation failed. #[error("Generation failed: {0}")] GenerationFailed(String), /// Model not initialized. #[error("Model not initialized")] ModelNotInitialized, /// Physics simulation failed. #[error("Physics simulation failed: {0}")] PhysicsFailed(String), } /// Main WorldGen system. #[derive(Debug)] pub struct WorldGen { /// Model configuration. model_config: ModelConfig, /// Scheduler configuration. scheduler_config: SchedulerConfig, /// Diffusion model. diffusion_model: diffusion::DiffusionTransformer, /// Physics engine. physics_engine: physics::PhysicsEngine, /// Is initialized. initialized: bool, } impl Default for WorldGen { fn default() -> Self { Self::new(ModelConfig::default(), SchedulerConfig::default()) } } impl WorldGen { /// Create a new WorldGen instance. #[must_use] pub fn new(model_config: ModelConfig, scheduler_config: SchedulerConfig) -> Self { Self { diffusion_model: diffusion::DiffusionTransformer::new(&model_config), physics_engine: physics::PhysicsEngine::new(), model_config, scheduler_config, initialized: true, } } /// Generate a video from a text prompt. pub fn generate( &mut self, request: &GenerationRequest, ) -> Result { if !self.initialized { return Err(WorldGenError::ModelNotInitialized); } if request.prompt.is_empty() { return Err(WorldGenError::InvalidRequest("Prompt is empty".to_string())); } if request.num_frames == 0 { return Err(WorldGenError::InvalidRequest( "num_frames must be > 0".to_string(), )); } let start_time = std::time::Instant::now(); // Create scheduler let scheduler = scheduler::NoiseScheduler::new(&self.scheduler_config); // Encode text prompt let text_embedding = self.encode_prompt(&request.prompt); // Initialize latent noise let seed = request.seed.unwrap_or(42); let latent_shape = ( request.num_frames, self.model_config.latent_channels, request.height / 8, request.width / 8, ); let mut latents = self.initialize_latents(latent_shape, seed); // Denoising loop let timesteps = scheduler.get_timesteps(request.num_inference_steps); for &t in ×teps { // Predict noise let noise_pred = self.diffusion_model.forward(&latents, t, &text_embedding); // Apply physics constraints let physics_guidance = if !request.physics_constraints.is_empty() { Some( self.physics_engine .compute_guidance(&latents, &request.physics_constraints), ) } else { None }; // Update latents latents = scheduler.step(&latents, &noise_pred, t, physics_guidance.as_deref()); } // Decode latents to video frames let video = self.decode_latents(&latents, request); // Calculate physics score let physics_score = if !request.physics_constraints.is_empty() { Some( self.physics_engine .evaluate_physics(&video, &request.physics_constraints), ) } else { None }; Ok(GenerationResult { video, generation_time: start_time.elapsed().as_secs_f64(), physics_score, prompt: request.prompt.clone(), }) } /// Train the model. pub fn train( &mut self, _config: &TrainingConfig, progress_callback: Option>, ) -> Result<(), WorldGenError> { let start_time = std::time::Instant::now(); let total_steps = 100; // Demo steps for step in 0..total_steps { // Simulate training step let loss = 0.5 * (-0.01 * step as f64).exp(); if let Some(ref callback) = progress_callback { callback(TrainingProgress { epoch: step / 10 + 1, total_epochs: 10, step: step + 1, total_steps, diffusion_loss: loss, physics_loss: Some(loss * 0.1), total_loss: loss * 1.1, learning_rate: 1e-4, elapsed_seconds: start_time.elapsed().as_secs_f64(), }); } } Ok(()) } /// Encode text prompt to embedding. fn encode_prompt(&self, prompt: &str) -> Vec { // Simplified text encoding (in reality, would use CLIP or similar) let mut embedding = vec![0.0; self.model_config.hidden_dim]; // Simple hash-based encoding for demo for (i, c) in prompt.chars().enumerate() { let idx = i % embedding.len(); embedding[idx] += (c as u32) as f64 / 1000.0; } // Normalize let norm: f64 = embedding.iter().map(|x| x * x).sum::().sqrt(); if norm > 1e-10 { for e in &mut embedding { *e /= norm; } } embedding } /// Initialize random latent noise. fn initialize_latents(&self, shape: (usize, usize, usize, usize), seed: u64) -> Vec { let size = shape.0 * shape.1 * shape.2 * shape.3; let mut latents = vec![0.0; size]; let mut rng_state = seed; for latent in &mut latents { // LCG random rng_state = rng_state .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); let u1 = (rng_state >> 11) as f64 / (1u64 << 53) as f64 + 1e-10; rng_state = rng_state .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); let u2 = (rng_state >> 11) as f64 / (1u64 << 53) as f64; // Box-Muller *latent = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos(); } latents } /// Decode latents to video frames. fn decode_latents(&self, latents: &[f64], request: &GenerationRequest) -> VideoClip { use worldgen_shared::VideoFrame; let latent_h = request.height / 8; let latent_w = request.width / 8; let latent_c = self.model_config.latent_channels; let frame_size = latent_c * latent_h * latent_w; let frames: Vec = (0..request.num_frames) .map(|frame_idx| { let frame_latent_start = frame_idx * frame_size; let frame_latent = &latents[frame_latent_start ..frame_latent_start + frame_size.min(latents.len() - frame_latent_start)]; // Decode latent to pixels (simplified upsampling) let mut pixels = vec![0u8; request.width * request.height * 3]; for y in 0..request.height { for x in 0..request.width { let latent_x = x / 8; let latent_y = y / 8; for c in 0..3 { let latent_idx = (c % latent_c) * latent_h * latent_w + latent_y * latent_w + latent_x; let value = if latent_idx < frame_latent.len() { frame_latent[latent_idx] } else { 0.0 }; // Map to [0, 255] let pixel_value = ((value.tanh() + 1.0) * 127.5).clamp(0.0, 255.0) as u8; pixels[(y * request.width + x) * 3 + c] = pixel_value; } } } VideoFrame { index: frame_idx, width: request.width, height: request.height, data: pixels, timestamp: frame_idx as f64 / request.fps as f64, } }) .collect(); VideoClip { frames, fps: request.fps, duration: request.num_frames as f32 / request.fps, width: request.width, height: request.height, } } /// Get model config. #[must_use] pub fn model_config(&self) -> &ModelConfig { &self.model_config } /// Get scheduler config. #[must_use] pub fn scheduler_config(&self) -> &SchedulerConfig { &self.scheduler_config } } /// Run the demo. pub fn run_demo() -> Result { let mut worldgen = WorldGen::default(); let request = worldgen_shared::sample_generation_request(); worldgen.generate(&request) } #[cfg(test)] mod tests { use super::*; #[test] fn test_worldgen_creation() { let worldgen = WorldGen::default(); assert!(worldgen.initialized); } #[test] fn test_generate_video() { let mut worldgen = WorldGen::default(); let mut request = worldgen_shared::sample_generation_request(); request.num_frames = 4; request.width = 64; request.height = 64; request.num_inference_steps = 5; let result = worldgen.generate(&request); assert!(result.is_ok()); let result = result.unwrap(); assert_eq!(result.video.frames.len(), 4); assert!(!result.prompt.is_empty()); } #[test] fn test_empty_prompt() { let mut worldgen = WorldGen::default(); let mut request = worldgen_shared::sample_generation_request(); request.prompt = String::new(); let result = worldgen.generate(&request); assert!(matches!(result, Err(WorldGenError::InvalidRequest(_)))); } #[test] fn test_physics_constraints() { let mut worldgen = WorldGen::default(); let mut request = worldgen_shared::sample_generation_request(); request.num_frames = 4; request.width = 64; request.height = 64; request.num_inference_steps = 3; let result = worldgen.generate(&request).unwrap(); assert!(result.physics_score.is_some()); } #[test] fn test_encode_prompt() { let worldgen = WorldGen::default(); let embedding = worldgen.encode_prompt("test prompt"); assert_eq!(embedding.len(), worldgen.model_config.hidden_dim); // Check normalization let norm: f64 = embedding.iter().map(|x| x * x).sum::().sqrt(); assert!((norm - 1.0).abs() < 0.01); } #[test] fn test_run_demo() { let result = run_demo(); assert!(result.is_ok()); } }