Files
foxhunt/crates/ml-supervised/src/diffusion/denoiser.rs
jgrusewski 7ef92983f9 fix(clippy): apply cargo clippy --fix across workspace
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>
2026-03-10 11:17:51 +01:00

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]);
}
}