Files
rustytorch/demos/rtx-worldgen-demo/src/diffusion.rs
T
2026-03-04 00:08:42 +00:00

198 lines
5.8 KiB
Rust

//! Diffusion Transformer (DiT) implementation for video generation.
use worldgen_shared::ModelConfig;
/// Diffusion Transformer.
#[derive(Debug)]
pub struct DiffusionTransformer {
/// Hidden dimension.
hidden_dim: usize,
/// Number of layers.
num_layers: usize,
/// Number of attention heads.
#[allow(dead_code)]
num_heads: usize,
/// Patch size.
#[allow(dead_code)]
patch_size: usize,
/// Weights (simplified).
weights: Vec<f64>,
/// Use temporal attention.
temporal_attention: bool,
/// RNG state.
rng_state: u64,
}
impl DiffusionTransformer {
/// Create a new DiT model.
pub fn new(config: &ModelConfig) -> Self {
let num_params = Self::calculate_params(config);
let mut model = Self {
hidden_dim: config.hidden_dim,
num_layers: config.num_layers,
num_heads: config.num_heads,
patch_size: config.patch_size,
weights: vec![0.0; num_params],
temporal_attention: config.temporal_attention,
rng_state: 42,
};
model.initialize_weights();
model
}
/// Calculate number of parameters.
fn calculate_params(config: &ModelConfig) -> usize {
let d = config.hidden_dim;
let l = config.num_layers;
// Embedding + transformer layers + output projection
let embedding = d * 1024; // Patch embedding
let transformer = l * (4 * d * d + 2 * d); // MLP + attention
let output = d * 4; // Output projection
embedding + transformer + output
}
/// Initialize weights.
fn initialize_weights(&mut self) {
let scale = (self.hidden_dim as f64).sqrt().recip();
let num_weights = self.weights.len();
let init_values: Vec<f64> = (0..num_weights)
.map(|_| self.random_normal() * scale)
.collect();
for (weight, value) in self.weights.iter_mut().zip(init_values) {
*weight = value;
}
}
/// Random number.
fn random(&mut self) -> f64 {
self.rng_state = self
.rng_state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(self.rng_state >> 11) as f64 / (1u64 << 53) as f64
}
/// Random normal.
fn random_normal(&mut self) -> f64 {
let u1 = self.random() + 1e-10;
let u2 = self.random();
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
}
/// Forward pass: predict noise given noisy latents, timestep, and text embedding.
pub fn forward(&self, latents: &[f64], timestep: f64, text_embedding: &[f64]) -> Vec<f64> {
let mut output = vec![0.0; latents.len()];
// Timestep embedding
let t_embed = self.timestep_embedding(timestep);
// Process each latent
for (i, &latent) in latents.iter().enumerate() {
// Combine latent, timestep, and text embedding
let t_idx = i % t_embed.len();
let txt_idx = i % text_embedding.len();
let combined = latent + t_embed[t_idx] * 0.1 + text_embedding[txt_idx] * 0.1;
// Apply transformer layers (simplified)
let mut hidden = combined;
for layer in 0..self.num_layers {
let weight_idx = (i + layer) % self.weights.len();
// Attention (simplified as weighted sum)
hidden = hidden * self.weights[weight_idx] + hidden;
// MLP
hidden = hidden.tanh() * self.weights[(weight_idx + 1) % self.weights.len()];
// Temporal attention (if enabled)
if self.temporal_attention {
hidden *= 1.0 + 0.01 * (i as f64).sin();
}
}
// Output projection
output[i] = hidden.tanh();
}
output
}
/// Create timestep embedding.
fn timestep_embedding(&self, timestep: f64) -> Vec<f64> {
let dim = self.hidden_dim;
let mut embedding = vec![0.0; dim];
for i in 0..dim / 2 {
let freq = (-((i as f64) / (dim as f64 / 2.0) * 10.0)).exp();
embedding[i] = (timestep * freq).sin();
embedding[i + dim / 2] = (timestep * freq).cos();
}
embedding
}
/// Get hidden dimension.
#[must_use]
pub fn hidden_dim(&self) -> usize {
self.hidden_dim
}
/// Get number of parameters.
#[must_use]
pub fn num_parameters(&self) -> usize {
self.weights.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
use worldgen_shared::sample_model_config;
#[test]
fn test_dit_creation() {
let config = sample_model_config();
let dit = DiffusionTransformer::new(&config);
assert_eq!(dit.hidden_dim, config.hidden_dim);
assert_eq!(dit.num_layers, config.num_layers);
}
#[test]
fn test_forward() {
let config = sample_model_config();
let dit = DiffusionTransformer::new(&config);
let latents = vec![0.5; 64];
let text_embedding = vec![0.1; config.hidden_dim];
let timestep = 0.5;
let output = dit.forward(&latents, timestep, &text_embedding);
assert_eq!(output.len(), latents.len());
// Check output is bounded
for &val in &output {
assert!(val >= -1.0 && val <= 1.0);
}
}
#[test]
fn test_timestep_embedding() {
let config = sample_model_config();
let dit = DiffusionTransformer::new(&config);
let embedding = dit.timestep_embedding(0.5);
assert_eq!(embedding.len(), config.hidden_dim);
}
#[test]
fn test_num_parameters() {
let config = sample_model_config();
let dit = DiffusionTransformer::new(&config);
assert!(dit.num_parameters() > 0);
}
}