From e12cdfa0050dea066102e0484fbf410b66e65d13 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 2 Mar 2026 17:30:19 +0100 Subject: [PATCH] feat(ml): diffusion trainer GPU-accumulated loss + batched grad norm MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replace per-param gradient norm extraction with batched Tensor::stack pattern (N GPU syncs → 1). Remove per-batch loss scalar extraction from backward(), defer to epoch-level accumulation. Apply GPU-accumulated validation loss (stack + mean_all instead of per-batch to_scalar). Co-Authored-By: Claude Opus 4.6 --- crates/ml/src/diffusion/trainable.rs | 50 +++++++++++++++++++--------- 1 file changed, 34 insertions(+), 16 deletions(-) diff --git a/crates/ml/src/diffusion/trainable.rs b/crates/ml/src/diffusion/trainable.rs index 6e486f438..2c13da98c 100644 --- a/crates/ml/src/diffusion/trainable.rs +++ b/crates/ml/src/diffusion/trainable.rs @@ -145,24 +145,35 @@ impl UnifiedTrainable for DiffusionTrainableAdapter { let grads = loss.backward() .map_err(|e| MLError::ModelError(e.to_string()))?; - let mut total_norm: f64 = 0.0; - for var in self.var_map.all_vars() { + // Collect all per-parameter squared norms on GPU, then stack+sum once + // to avoid per-parameter GPU sync (to_scalar) which serializes the pipeline. + let mut norm_parts = Vec::new(); + let vars_lock = self.var_map.data().lock() + .map_err(|e| MLError::ModelError(format!("diffusion var_map lock: {e}")))?; + + for (_name, var) in vars_lock.iter() { if let Some(grad) = grads.get(var.as_tensor()) { - let norm: f64 = grad.sqr() - .and_then(|s| s.sum_all()) - .and_then(|s| s.to_scalar::()) - .unwrap_or(0.0) as f64; - total_norm += norm; + if let Ok(norm_sq) = grad.sqr().and_then(|s| s.sum_all()) { + norm_parts.push(norm_sq); + } } } + drop(vars_lock); - let loss_val = loss.to_scalar::().unwrap_or(f32::NAN) as f64; - self.loss_history.push(loss_val); + let grad_norm_sq = if norm_parts.is_empty() { + 0.0_f64 + } else { + let stacked = Tensor::stack(&norm_parts, 0) + .map_err(|e| MLError::ModelError(format!("diffusion grad norm stack: {e}")))?; + stacked.sum_all() + .and_then(|s| s.to_scalar::()) + .map_err(|e| MLError::ModelError(format!("diffusion grad norm: {e}")))? as f64 + }; // Store gradients for optimizer_step to consume self.last_grads = Some(grads); - Ok(total_norm.sqrt()) + Ok(grad_norm_sq.sqrt()) } fn optimizer_step(&mut self) -> Result<(), MLError> { @@ -256,16 +267,23 @@ impl UnifiedTrainable for DiffusionTrainableAdapter { }); } - let mut total_loss = 0.0; - let mut count = 0; + // Accumulate per-batch loss tensors on GPU, then stack+mean once + // to avoid per-batch GPU sync (to_scalar) which serializes the pipeline. + let mut loss_tensors = Vec::with_capacity(val_data.len()); for (input, target) in val_data { let output = self.forward(input)?; let loss = self.compute_loss(&output, target)?; - total_loss += loss.to_scalar::() - .map_err(|e| MLError::ModelError(e.to_string()))? as f64; - count += 1; + loss_tensors.push(loss); } - Ok(total_loss / count as f64) + + let stacked = Tensor::stack(&loss_tensors, 0) + .map_err(|e| MLError::ModelError(format!("diffusion val loss stack: {e}")))?; + let avg_loss = stacked + .mean_all() + .and_then(|t| t.to_scalar::()) + .map_err(|e| MLError::ModelError(format!("diffusion val loss mean: {e}")))? as f64; + + Ok(avg_loss) } }