376 lines
12 KiB
Rust
376 lines
12 KiB
Rust
//! 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<GenerationResult, WorldGenError> {
|
|
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<Box<dyn Fn(TrainingProgress) + Send>>,
|
|
) -> 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<f64> {
|
|
// 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::<f64>().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<f64> {
|
|
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<VideoFrame> = (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<GenerationResult, WorldGenError> {
|
|
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::<f64>().sqrt();
|
|
assert!((norm - 1.0).abs() < 0.01);
|
|
}
|
|
|
|
#[test]
|
|
fn test_run_demo() {
|
|
let result = run_demo();
|
|
assert!(result.is_ok());
|
|
}
|
|
}
|