Mechanical auto-fixes: redundant borrows, clone on Copy, or_insert_with, single-char push_str, get(0) → first(), needless borrow, let_and_return. 150 files, no behavior changes. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
300 lines
11 KiB
Rust
300 lines
11 KiB
Rust
//! 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<Self, MLError> {
|
|
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<Tensor, MLError> {
|
|
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<Self, MLError> {
|
|
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<Tensor, MLError> {
|
|
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<DenoiserBlock>,
|
|
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<Self, MLError> {
|
|
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<Tensor, MLError> {
|
|
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]);
|
|
}
|
|
}
|