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:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user