From 4cf181c8fb45cdf2865aabe753569b95ca4165f3 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 2 Mar 2026 17:30:10 +0100 Subject: [PATCH] =?UTF-8?q?feat(ml):=20TFT=20batched=20gradient=20norm=20?= =?UTF-8?q?=E2=80=94=20N=20GPU=20syncs=20=E2=86=92=201?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replace per-param gradient norm extraction loop with batched Tensor::stack pattern. Single GPU→CPU sync instead of one per parameter tensor. Co-Authored-By: Claude Opus 4.6 --- crates/ml/src/tft/training.rs | 34 ++++++++++++++++++++++------------ 1 file changed, 22 insertions(+), 12 deletions(-) diff --git a/crates/ml/src/tft/training.rs b/crates/ml/src/tft/training.rs index 56b745a04..96704dc6e 100644 --- a/crates/ml/src/tft/training.rs +++ b/crates/ml/src/tft/training.rs @@ -479,18 +479,28 @@ impl TFTTrainer { let varmap_data = self.model.varmap().data().lock().map_err(|e| { MLError::TrainingError(format!("Failed to lock VarMap: {}", e)) })?; - let mut total_norm_sq = 0.0_f64; - for (_name, var) in varmap_data.iter() { - if let Some(grad) = grads.get(var.as_tensor()) { - let norm_sq = grad - .sqr() - .and_then(|t| t.sum_all()) - .and_then(|t| t.to_dtype(DType::F64)) - .and_then(|t| t.to_scalar::()) - .unwrap_or(0.0); - total_norm_sq += norm_sq; - } - } + let norm_parts: Vec = varmap_data + .iter() + .filter_map(|(_name, var)| { + grads + .get(var.as_tensor()) + .and_then(|g| g.sqr().and_then(|s| s.sum_all()).ok()) + }) + .collect(); + + let total_norm_sq: f64 = if !norm_parts.is_empty() { + let stacked = Tensor::stack(&norm_parts, 0).map_err(|e| { + MLError::TrainingError(format!("TFT grad norm stack: {e}")) + })?; + stacked + .sum_all() + .and_then(|s| s.to_scalar::()) + .map_err(|e| { + MLError::TrainingError(format!("TFT grad norm: {e}")) + })? + } else { + 0.0 + }; drop(varmap_data); let grad_norm = total_norm_sq.sqrt(); self.last_gradient_norm = grad_norm;