Initial commit
This commit is contained in:
@@ -0,0 +1,375 @@
|
||||
//! 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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user