//! Shared types for WorldGen - Video Diffusion World Simulator. //! //! This crate defines the IPC types for video generation and world simulation. use serde::{Deserialize, Serialize}; // ============================================================================ // Video Types // ============================================================================ /// Video frame. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct VideoFrame { /// Frame index. pub index: usize, /// Width in pixels. pub width: usize, /// Height in pixels. pub height: usize, /// RGB pixel data (flattened, row-major). pub data: Vec, /// Timestamp in seconds. pub timestamp: f64, } /// Video clip. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct VideoClip { /// Frames. pub frames: Vec, /// Frames per second. pub fps: f32, /// Duration in seconds. pub duration: f32, /// Width. pub width: usize, /// Height. pub height: usize, } /// Video generation request. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct GenerationRequest { /// Text prompt describing the video. pub prompt: String, /// Negative prompt (what to avoid). pub negative_prompt: Option, /// Number of frames to generate. pub num_frames: usize, /// Frame width. pub width: usize, /// Frame height. pub height: usize, /// Frames per second. pub fps: f32, /// Physics constraints to apply. pub physics_constraints: Vec, /// Random seed. pub seed: Option, /// Guidance scale (CFG). pub guidance_scale: f32, /// Number of inference steps. pub num_inference_steps: usize, } impl Default for GenerationRequest { fn default() -> Self { Self { prompt: String::new(), negative_prompt: None, num_frames: 16, width: 512, height: 512, fps: 8.0, physics_constraints: vec![], seed: Some(42), guidance_scale: 7.5, num_inference_steps: 50, } } } /// Generation result. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct GenerationResult { /// Generated video. pub video: VideoClip, /// Generation time in seconds. pub generation_time: f64, /// Physics consistency score. pub physics_score: Option, /// Prompt used. pub prompt: String, } // ============================================================================ // Physics Types // ============================================================================ /// Physics constraint type. #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Hash)] pub enum PhysicsConstraintType { /// Gravity (objects fall). Gravity, /// Collision (objects don't pass through each other). Collision, /// Conservation of momentum. Momentum, /// Fluid dynamics. Fluid, /// Rigid body dynamics. RigidBody, /// Soft body deformation. SoftBody, /// Lighting consistency. Lighting, /// Shadow consistency. Shadows, /// Temporal coherence. TemporalCoherence, } /// Physics constraint. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct PhysicsConstraint { /// Constraint type. pub constraint_type: PhysicsConstraintType, /// Weight (0.0 to 1.0). pub weight: f32, /// Parameters. pub parameters: PhysicsParameters, } /// Physics parameters. #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct PhysicsParameters { /// Gravity vector (m/s²). pub gravity: Option<[f32; 3]>, /// Friction coefficient. pub friction: Option, /// Restitution (bounciness). pub restitution: Option, /// Fluid viscosity. pub viscosity: Option, /// Light direction. pub light_direction: Option<[f32; 3]>, } // ============================================================================ // Model Types // ============================================================================ /// Diffusion model type. #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Hash)] pub enum ModelType { /// Diffusion Transformer (DiT). DiT, /// U-Net based diffusion. UNet, /// Video Diffusion Model. VDM, /// Latent Video Diffusion. LVD, } /// Model configuration. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ModelConfig { /// Model type. pub model_type: ModelType, /// Hidden dimension. pub hidden_dim: usize, /// Number of layers. pub num_layers: usize, /// Number of attention heads. pub num_heads: usize, /// Patch size for DiT. pub patch_size: usize, /// Latent channels. pub latent_channels: usize, /// Use temporal attention. pub temporal_attention: bool, /// Use cross-attention for text conditioning. pub cross_attention: bool, } impl Default for ModelConfig { fn default() -> Self { Self { model_type: ModelType::DiT, hidden_dim: 768, num_layers: 12, num_heads: 12, patch_size: 2, latent_channels: 4, temporal_attention: true, cross_attention: true, } } } /// Noise scheduler type. #[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)] pub enum SchedulerType { /// DDPM scheduler. DDPM, /// DDIM scheduler. DDIM, /// Euler scheduler. Euler, /// DPM++ scheduler. DPMPlusPlus, /// Flow matching. FlowMatching, } /// Scheduler configuration. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SchedulerConfig { /// Scheduler type. pub scheduler_type: SchedulerType, /// Number of timesteps. pub num_timesteps: usize, /// Beta start. pub beta_start: f32, /// Beta end. pub beta_end: f32, /// Beta schedule. pub beta_schedule: String, } impl Default for SchedulerConfig { fn default() -> Self { Self { scheduler_type: SchedulerType::DDIM, num_timesteps: 1000, beta_start: 0.0001, beta_end: 0.02, beta_schedule: "linear".to_string(), } } } // ============================================================================ // Training Types // ============================================================================ /// Training configuration. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TrainingConfig { /// Number of epochs. pub epochs: usize, /// Batch size. pub batch_size: usize, /// Learning rate. pub learning_rate: f64, /// Weight decay. pub weight_decay: f64, /// Gradient accumulation steps. pub gradient_accumulation: usize, /// Mixed precision training. pub mixed_precision: bool, /// Physics loss weight. pub physics_weight: f32, /// Random seed. pub seed: Option, } impl Default for TrainingConfig { fn default() -> Self { Self { epochs: 100, batch_size: 4, learning_rate: 1e-4, weight_decay: 0.01, gradient_accumulation: 4, mixed_precision: true, physics_weight: 0.1, seed: Some(42), } } } /// Training progress. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TrainingProgress { /// Current epoch. pub epoch: usize, /// Total epochs. pub total_epochs: usize, /// Current step. pub step: usize, /// Total steps. pub total_steps: usize, /// Diffusion loss. pub diffusion_loss: f64, /// Physics loss. pub physics_loss: Option, /// Total loss. pub total_loss: f64, /// Learning rate. pub learning_rate: f64, /// Elapsed time. pub elapsed_seconds: f64, } // ============================================================================ // World State Types // ============================================================================ /// World state for simulation. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WorldState { /// Objects in the world. pub objects: Vec, /// Camera position. pub camera_position: [f32; 3], /// Camera rotation (euler angles). pub camera_rotation: [f32; 3], /// Lighting. pub lighting: LightingState, /// Time in simulation. pub time: f32, } /// World object. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WorldObject { /// Object ID. pub id: String, /// Object type. pub object_type: String, /// Position. pub position: [f32; 3], /// Velocity. pub velocity: [f32; 3], /// Rotation (euler angles). pub rotation: [f32; 3], /// Scale. pub scale: [f32; 3], /// Mass. pub mass: f32, } /// Lighting state. #[derive(Debug, Clone, Serialize, Deserialize)] pub struct LightingState { /// Ambient light color. pub ambient: [f32; 3], /// Directional light direction. pub sun_direction: [f32; 3], /// Directional light color. pub sun_color: [f32; 3], /// Time of day (0-24). pub time_of_day: f32, } impl Default for LightingState { fn default() -> Self { Self { ambient: [0.1, 0.1, 0.1], sun_direction: [0.5, -1.0, 0.5], sun_color: [1.0, 0.95, 0.9], time_of_day: 12.0, } } } // ============================================================================ // Sample Functions // ============================================================================ /// Create a sample generation request. #[must_use] pub fn sample_generation_request() -> GenerationRequest { GenerationRequest { prompt: "A ball bouncing on a wooden floor in a sunlit room".to_string(), negative_prompt: Some("blurry, distorted, unrealistic".to_string()), num_frames: 16, width: 256, height: 256, fps: 8.0, physics_constraints: vec![ PhysicsConstraint { constraint_type: PhysicsConstraintType::Gravity, weight: 1.0, parameters: PhysicsParameters { gravity: Some([0.0, -9.81, 0.0]), ..Default::default() }, }, PhysicsConstraint { constraint_type: PhysicsConstraintType::Collision, weight: 0.8, parameters: PhysicsParameters { restitution: Some(0.7), ..Default::default() }, }, ], seed: Some(42), guidance_scale: 7.5, num_inference_steps: 50, } } /// Create a sample model config. #[must_use] pub fn sample_model_config() -> ModelConfig { ModelConfig { model_type: ModelType::DiT, hidden_dim: 384, num_layers: 6, num_heads: 6, patch_size: 4, latent_channels: 4, temporal_attention: true, cross_attention: true, } } /// Create a sample scheduler config. #[must_use] pub fn sample_scheduler_config() -> SchedulerConfig { SchedulerConfig { scheduler_type: SchedulerType::DDIM, num_timesteps: 1000, beta_start: 0.0001, beta_end: 0.02, beta_schedule: "linear".to_string(), } } /// Create a sample world state. #[must_use] pub fn sample_world_state() -> WorldState { WorldState { objects: vec![ WorldObject { id: "ball".to_string(), object_type: "sphere".to_string(), position: [0.0, 2.0, 0.0], velocity: [0.0, 0.0, 0.0], rotation: [0.0, 0.0, 0.0], scale: [0.5, 0.5, 0.5], mass: 1.0, }, WorldObject { id: "floor".to_string(), object_type: "plane".to_string(), position: [0.0, 0.0, 0.0], velocity: [0.0, 0.0, 0.0], rotation: [0.0, 0.0, 0.0], scale: [10.0, 1.0, 10.0], mass: f32::INFINITY, }, ], camera_position: [5.0, 3.0, 5.0], camera_rotation: [-0.5, 0.8, 0.0], lighting: LightingState::default(), time: 0.0, } } // ============================================================================ // Tests // ============================================================================ #[cfg(test)] mod tests { use super::*; #[test] fn test_generation_request() { let req = sample_generation_request(); assert!(!req.prompt.is_empty()); assert_eq!(req.num_frames, 16); assert!(!req.physics_constraints.is_empty()); } #[test] fn test_model_config() { let config = sample_model_config(); assert_eq!(config.model_type, ModelType::DiT); assert!(config.temporal_attention); } #[test] fn test_scheduler_config() { let config = sample_scheduler_config(); assert_eq!(config.scheduler_type, SchedulerType::DDIM); assert_eq!(config.num_timesteps, 1000); } #[test] fn test_world_state() { let state = sample_world_state(); assert_eq!(state.objects.len(), 2); assert_eq!(state.objects[0].id, "ball"); } #[test] fn test_video_frame() { let frame = VideoFrame { index: 0, width: 64, height: 64, data: vec![0; 64 * 64 * 3], timestamp: 0.0, }; assert_eq!(frame.data.len(), 64 * 64 * 3); } #[test] fn test_physics_constraint() { let constraint = PhysicsConstraint { constraint_type: PhysicsConstraintType::Gravity, weight: 1.0, parameters: PhysicsParameters { gravity: Some([0.0, -9.81, 0.0]), ..Default::default() }, }; assert_eq!(constraint.constraint_type, PhysicsConstraintType::Gravity); } #[test] fn test_serialization() { let req = sample_generation_request(); let json = serde_json::to_string(&req).unwrap(); assert!(json.contains("bouncing")); let parsed: GenerationRequest = serde_json::from_str(&json).unwrap(); assert_eq!(parsed.prompt, req.prompt); } }