diff --git a/crates/ml/src/trainers/tlob.rs b/crates/ml/src/trainers/tlob.rs index 33a3eb01f..09eb43568 100644 --- a/crates/ml/src/trainers/tlob.rs +++ b/crates/ml/src/trainers/tlob.rs @@ -416,9 +416,13 @@ impl TLOBTrainer { } /// Train one epoch + /// + /// Loss is accumulated on-device as a GPU tensor to avoid per-batch + /// `to_scalar()` synchronisation. A single scalar extraction happens + /// after the loop completes, with a NaN guard every 100 batches. async fn train_epoch(&mut self, sequences: &[OrderBookSequence]) -> Result { - let mut total_loss = 0.0; - let mut num_batches = 0; + let mut loss_accum = Tensor::zeros((), DType::F32, &self.device)?; + let mut num_batches: usize = 0; // Process in batches for batch_sequences in sequences.chunks(self.hyperparams.batch_size) { @@ -437,18 +441,37 @@ impl TLOBTrainer { // Gradient clipping self.clip_gradients()?; - total_loss += loss.to_scalar::()? as f64; + // Accumulate loss on-device (detach to prevent graph growth) + loss_accum = loss_accum.add(&loss.detach())?; num_batches += 1; + + // NaN guard every 100 batches — single cheap sync + if num_batches % 100 == 0 { + let check = loss_accum.to_scalar::()?; + if check.is_nan() || check.is_infinite() { + warn!(batch = num_batches, "TLOB train loss diverged (NaN/Inf)"); + return Ok(f64::NAN); + } + } } - Ok(total_loss / num_batches as f64) + if num_batches == 0 { + return Ok(0.0); + } + + // Single scalar extraction after all batches + let total = loss_accum.to_scalar::()? as f64; + Ok(total / num_batches as f64) } /// Validate one epoch + /// + /// Loss is accumulated on-device to avoid per-batch `to_scalar()`. + /// MAE already returns `f64` (CPU scalar) so it stays as f64 sum. async fn validate_epoch(&self, sequences: &[OrderBookSequence]) -> Result<(f64, f64)> { - let mut total_loss = 0.0; + let mut loss_accum = Tensor::zeros((), DType::F32, &self.device)?; let mut total_mae = 0.0; - let mut num_batches = 0; + let mut num_batches: usize = 0; // Process in batches (no gradient computation) for batch_sequences in sequences.chunks(self.hyperparams.batch_size) { @@ -461,11 +484,27 @@ impl TLOBTrainer { let loss = self.calculate_mse_loss(&predictions, &target_tensor)?.detach(); let mae = self.calculate_mae(&predictions, &target_tensor)?; - total_loss += loss.to_scalar::()? as f64; + // Accumulate loss on-device + loss_accum = loss_accum.add(&loss)?; total_mae += mae; num_batches += 1; + + // NaN guard every 100 batches + if num_batches % 100 == 0 { + let check = loss_accum.to_scalar::()?; + if check.is_nan() || check.is_infinite() { + warn!(batch = num_batches, "TLOB val loss diverged (NaN/Inf)"); + return Ok((f64::NAN, f64::NAN)); + } + } } + if num_batches == 0 { + return Ok((0.0, 0.0)); + } + + // Single scalar extraction after all batches + let total_loss = loss_accum.to_scalar::()? as f64; Ok(( total_loss / num_batches as f64, total_mae / num_batches as f64,