Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
+375
View File
@@ -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 &timesteps {
// 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());
}
}