198 lines
5.8 KiB
Rust
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);
|
|
}
|
|
}
|