feat(ml): diffusion trainer GPU-accumulated loss + batched grad norm

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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-02 17:30:19 +01:00
parent 4cf181c8fb
commit e12cdfa005

View File

@@ -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::<f32>())
.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::<f32>().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::<f32>())
.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::<f32>()
.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::<f32>())
.map_err(|e| MLError::ModelError(format!("diffusion val loss mean: {e}")))? as f64;
Ok(avg_loss)
}
}