//! Denoiser network for the diffusion model. //! //! Uses a fully-connected architecture with sinusoidal time embedding. //! This is more memory-efficient than Conv1D U-Net while still effective //! for price sequence denoising at small sequence lengths (64-128). use ml_core::MLError; use candle_core::{DType, Device, Tensor}; use candle_nn::{linear, Linear, Module, VarBuilder}; /// Sinusoidal time embedding for diffusion timestep conditioning. /// /// Maps scalar timestep t to a fixed-dimension vector using /// sin/cos positional encoding (same idea as Transformer PE). pub struct TimeEmbedding { proj: Linear, embed_dim: usize, } impl std::fmt::Debug for TimeEmbedding { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("TimeEmbedding").finish_non_exhaustive() } } impl TimeEmbedding { pub fn new(embed_dim: usize, hidden_dim: usize, vb: VarBuilder<'_>) -> Result { let proj = linear(embed_dim, hidden_dim, vb.pp("time_proj")) .map_err(|e| MLError::ModelError(e.to_string()))?; Ok(Self { proj, embed_dim }) } /// Compute sinusoidal embedding for timestep indices. /// Input: (batch,) u32 timestep indices → Output: (batch, hidden_dim) pub fn forward(&self, t: &Tensor, device: &Device) -> Result { let half_dim = self.embed_dim / 2; let t_f32 = t.to_dtype(DType::F32) .map_err(|e| MLError::ModelError(e.to_string()))?; // freq[i] = exp(-ln(10000) * i / half_dim) let mut freq_vals = Vec::with_capacity(half_dim); for i in 0..half_dim { let freq = (-(10000.0_f64.ln()) * i as f64 / half_dim.max(1) as f64).exp(); freq_vals.push(freq as f32); } let freqs = Tensor::new(freq_vals, device) .map_err(|e| MLError::ModelError(e.to_string()))?; // (batch, 1) * (1, half_dim) → (batch, half_dim) let t_expanded = t_f32.unsqueeze(1) .map_err(|e| MLError::ModelError(e.to_string()))?; let freqs_expanded = freqs.unsqueeze(0) .map_err(|e| MLError::ModelError(e.to_string()))?; let angles = t_expanded.broadcast_mul(&freqs_expanded) .map_err(|e| MLError::ModelError(e.to_string()))?; let sin_emb = angles.sin() .map_err(|e| MLError::ModelError(e.to_string()))?; let cos_emb = angles.cos() .map_err(|e| MLError::ModelError(e.to_string()))?; // Concat sin and cos: (batch, embed_dim) let emb = Tensor::cat(&[&sin_emb, &cos_emb], 1) .map_err(|e| MLError::ModelError(e.to_string()))?; // Cast to training dtype before projection through BF16 weights let emb = ml_core::mixed_precision::ensure_training_dtype(&emb) .map_err(|e| MLError::ModelError(e.to_string()))?; // Project to hidden_dim self.proj.forward(&emb) .map_err(|e| MLError::ModelError(e.to_string())) } } /// A single denoiser block: linear → SiLU → linear + time conditioning + residual. struct DenoiserBlock { fc1: Linear, fc2: Linear, time_proj: Linear, has_residual: bool, } impl std::fmt::Debug for DenoiserBlock { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("DenoiserBlock").finish_non_exhaustive() } } impl DenoiserBlock { fn new( input_dim: usize, hidden_dim: usize, time_dim: usize, vb: VarBuilder<'_>, ) -> Result { let fc1 = linear(input_dim, hidden_dim, vb.pp("fc1")) .map_err(|e| MLError::ModelError(e.to_string()))?; let fc2 = linear(hidden_dim, hidden_dim, vb.pp("fc2")) .map_err(|e| MLError::ModelError(e.to_string()))?; let time_proj = linear(time_dim, hidden_dim, vb.pp("time")) .map_err(|e| MLError::ModelError(e.to_string()))?; Ok(Self { fc1, fc2, time_proj, has_residual: input_dim == hidden_dim, }) } fn forward(&self, x: &Tensor, t_emb: &Tensor) -> Result { let map_err = |e: candle_core::Error| MLError::ModelError(e.to_string()); // fc1 → SiLU (x * sigmoid(x)) let h = self.fc1.forward(x).map_err(map_err)?; let h_sig = candle_nn::ops::sigmoid(&h).map_err(map_err)?; let h = h.mul(&h_sig).map_err(map_err)?; // Add time embedding let t_proj = self.time_proj.forward(t_emb).map_err(map_err)?; let h = h.add(&t_proj).map_err(map_err)?; // fc2 → SiLU let h = self.fc2.forward(&h).map_err(map_err)?; let h_sig2 = candle_nn::ops::sigmoid(&h).map_err(map_err)?; let h = h.mul(&h_sig2).map_err(map_err)?; // Residual connection if self.has_residual { h.add(x).map_err(map_err) } else { Ok(h) } } } /// Fully-connected denoiser network for diffusion. /// /// Architecture: input projection → N denoiser blocks → output projection. /// Each block is time-conditioned via additive time embedding. /// /// Memory usage at batch_size=32, data_dim=64, hidden=128: /// ~32 * 128 * num_layers * 4 bytes per layer ≈ 48KB -- very safe for 4GB GPU. pub struct Denoiser { input_proj: Linear, blocks: Vec, output_proj: Linear, time_embed: TimeEmbedding, device: Device, } impl std::fmt::Debug for Denoiser { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("Denoiser").finish_non_exhaustive() } } impl Denoiser { pub fn new( data_dim: usize, hidden_dim: usize, num_layers: usize, time_embed_dim: usize, vb: VarBuilder<'_>, device: &Device, ) -> Result { if data_dim == 0 || hidden_dim == 0 { return Err(MLError::ConfigError(format!( "Diffusion denoiser requires data_dim > 0 and hidden_dim > 0 (got {}x{})", data_dim, hidden_dim ))); } let time_embed = TimeEmbedding::new(time_embed_dim, hidden_dim, vb.pp("time_embed"))?; let input_proj = linear(data_dim, hidden_dim, vb.pp("input_proj")) .map_err(|e| MLError::ModelError(e.to_string()))?; let mut blocks = Vec::with_capacity(num_layers); for i in 0..num_layers { let block = DenoiserBlock::new( hidden_dim, hidden_dim, hidden_dim, vb.pp(format!("block_{i}")), )?; blocks.push(block); } let output_proj = linear(hidden_dim, data_dim, vb.pp("output_proj")) .map_err(|e| MLError::ModelError(e.to_string()))?; Ok(Self { input_proj, blocks, output_proj, time_embed, device: device.clone(), }) } /// Predict noise epsilon given noisy input x_t and timestep t. /// /// Input x: (batch, data_dim), t: (batch,) → Output: (batch, data_dim) pub fn forward(&self, x: &Tensor, t: &Tensor) -> Result { let x = ml_core::mixed_precision::ensure_training_dtype(x) .map_err(|e| MLError::ModelError(e.to_string()))?; let map_err = |e: candle_core::Error| MLError::ModelError(e.to_string()); // Time embedding let t_emb = self.time_embed.forward(t, &self.device)?; // Input projection → SiLU let mut h = self.input_proj.forward(&x).map_err(map_err)?; let h_sig = candle_nn::ops::sigmoid(&h).map_err(map_err)?; h = h.mul(&h_sig).map_err(map_err)?; // Denoiser blocks for block in &self.blocks { h = block.forward(&h, &t_emb)?; } // Output projection → predicted noise let output = self.output_proj.forward(&h).map_err(map_err)?; // Cast output back to F32 for API compatibility output.to_dtype(candle_core::DType::F32).map_err(map_err) } } #[cfg(test)] mod tests { use super::*; use candle_nn::VarMap; use ml_core::mixed_precision::training_dtype; #[test] fn test_time_embedding_shape() { let dev = Device::Cpu; let var_map = VarMap::new(); let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); let te = TimeEmbedding::new(32, 64, vb).unwrap(); let t = Tensor::new(&[0_u32, 100, 500, 999], &dev).unwrap(); let emb = te.forward(&t, &dev).unwrap(); assert_eq!(emb.dims(), &[4, 64]); } #[test] fn test_time_embedding_different_timesteps_differ() { let dev = Device::Cpu; let var_map = VarMap::new(); let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); let te = TimeEmbedding::new(32, 64, vb).unwrap(); let t1 = Tensor::new(&[0_u32], &dev).unwrap(); let t2 = Tensor::new(&[500_u32], &dev).unwrap(); let e1 = te.forward(&t1, &dev).unwrap(); let e2 = te.forward(&t2, &dev).unwrap(); let diff: f32 = e1.sub(&e2).unwrap().abs().unwrap().sum_all().unwrap().to_scalar().unwrap(); assert!(diff > 0.0, "Different timesteps should produce different embeddings"); } #[test] fn test_denoiser_output_shape() { let dev = Device::Cpu; let var_map = VarMap::new(); let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); let denoiser = Denoiser::new(64, 128, 3, 32, vb, &dev).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[4, 64], &dev).unwrap(); let t = Tensor::new(&[100_u32, 200, 300, 400], &dev).unwrap(); let out = denoiser.forward(&x, &t).unwrap(); assert_eq!(out.dims(), &[4, 64], "Output should match input shape"); } #[test] fn test_denoiser_produces_gradients() { let dev = Device::Cpu; let var_map = VarMap::new(); let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); let denoiser = Denoiser::new(64, 128, 2, 32, vb, &dev).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[4, 64], &dev).unwrap(); let t = Tensor::new(&[50_u32, 100, 200, 300], &dev).unwrap(); let out = denoiser.forward(&x, &t).unwrap(); let loss = out.sqr().unwrap().mean_all().unwrap(); let grads = loss.backward().unwrap(); let has_grads = var_map.all_vars().iter().any(|v| grads.get(v.as_tensor()).is_some()); assert!(has_grads, "Should produce gradients for training"); } #[test] fn test_denoiser_single_layer() { let dev = Device::Cpu; let var_map = VarMap::new(); let vb = VarBuilder::from_varmap(&var_map, training_dtype(&dev), &dev); let denoiser = Denoiser::new(32, 64, 1, 16, vb, &dev).unwrap(); let x = Tensor::randn(0_f32, 1.0, &[2, 32], &dev).unwrap(); let t = Tensor::new(&[0_u32, 999], &dev).unwrap(); let out = denoiser.forward(&x, &t).unwrap(); assert_eq!(out.dims(), &[2, 32]); } }