From d572dc466a24e207e9b1ea9ec4fbff6335ac5873 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Wed, 18 Mar 2026 08:29:26 +0100 Subject: [PATCH] =?UTF-8?q?perf(cuda):=20batch=20DQN=20metrics=20readbacks?= =?UTF-8?q?=20=E2=80=94=2015=20DtoH=20transfers=20=E2=86=92=204?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit metrics.rs: 4 fixes batching individual to_scalar() GPU syncs: - Q-value stats: 4 transfers → 1 (cat [min,max,mean,var]) - Q diagnostics: 8 transfers → 1 (cat [gaps + per-action avgs]) - get_q_values: N transfers → 1 (bulk to_vec1) - Validation Sharpe: 2 transfers → 1 (cat [mean,var]) Zero per-step DtoH leaks in training_loop.rs or metrics.rs. Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml/src/trainers/dqn/trainer/metrics.rs | 118 ++++++++++++------ 1 file changed, 81 insertions(+), 37 deletions(-) diff --git a/crates/ml/src/trainers/dqn/trainer/metrics.rs b/crates/ml/src/trainers/dqn/trainer/metrics.rs index d29db9dbd..355210c1d 100644 --- a/crates/ml/src/trainers/dqn/trainer/metrics.rs +++ b/crates/ml/src/trainers/dqn/trainer/metrics.rs @@ -77,7 +77,7 @@ impl DQNTrainer { .map_err(|e| crate::MLError::ModelError(format!("Q-value flatten: {}", e)))?; let count = q_flat.elem_count(); - // Compute all stats on GPU, single batched readback (4 floats in 1 DMA) + // Compute all stats on GPU let min_t = q_flat.min(0) .map_err(|e| crate::MLError::ModelError(format!("Q-value min: {}", e)))?; let max_t = q_flat.max(0) @@ -89,17 +89,35 @@ impl DQNTrainer { .and_then(|sq| sq.mean_all()) .map_err(|e| crate::MLError::ModelError(format!("Q-value variance: {}", e)))?; - let to_f64 = |t: &Tensor| -> Result { - Ok(t.to_dtype(ml_core::native_types::NativeDType::F32) - .and_then(|t| t.to_scalar::()) - .map_err(|e| crate::MLError::ModelError(format!("Q-value stat readback: {}", e)))? as f64) - }; + // Single batched readback: cat [min, max, mean, var] into [4] tensor, + // one DtoH transfer (16 bytes) instead of 4 individual to_scalar() calls. + let min_r = min_t.to_dtype(ml_core::native_types::NativeDType::F32) + .map_err(|e| crate::MLError::ModelError(format!("Q-value min f32: {}", e)))? + .reshape(&[1]) + .map_err(|e| crate::MLError::ModelError(format!("Q-value min reshape: {}", e)))?; + let max_r = max_t.to_dtype(ml_core::native_types::NativeDType::F32) + .map_err(|e| crate::MLError::ModelError(format!("Q-value max f32: {}", e)))? + .reshape(&[1]) + .map_err(|e| crate::MLError::ModelError(format!("Q-value max reshape: {}", e)))?; + let mean_r = mean_t.to_dtype(ml_core::native_types::NativeDType::F32) + .map_err(|e| crate::MLError::ModelError(format!("Q-value mean f32: {}", e)))? + .reshape(&[1]) + .map_err(|e| crate::MLError::ModelError(format!("Q-value mean reshape: {}", e)))?; + let var_r = var_t.to_dtype(ml_core::native_types::NativeDType::F32) + .map_err(|e| crate::MLError::ModelError(format!("Q-value var f32: {}", e)))? + .reshape(&[1]) + .map_err(|e| crate::MLError::ModelError(format!("Q-value var reshape: {}", e)))?; + + let combined = Tensor::cat(&[&min_r, &max_r, &mean_r, &var_r], 0) + .map_err(|e| crate::MLError::ModelError(format!("Q-value stats cat: {}", e)))?; + let vals: Vec = combined.to_vec1() + .map_err(|e| crate::MLError::ModelError(format!("Q-value stats readback: {}", e)))?; Ok(QValueStats { - min: to_f64(&min_t)?, - max: to_f64(&max_t)?, - mean: to_f64(&mean_t)?, - std: to_f64(&var_t)?.sqrt(), + min: vals.get(0).copied().unwrap_or(0.0) as f64, + max: vals.get(1).copied().unwrap_or(0.0) as f64, + mean: vals.get(2).copied().unwrap_or(0.0) as f64, + std: (vals.get(3).copied().unwrap_or(0.0) as f64).sqrt(), sample_count: count, }) } @@ -325,7 +343,7 @@ impl DQNTrainer { /// Compute Q-value gap and per-action averages on GPU. /// Returns (mean_gap, min_gap, max_gap, per_action_avgs[5]). -/// Single 8-float readback at epoch end. +/// Single 8-float readback at epoch end (was 8 individual DtoH transfers). fn compute_q_diagnostics_gpu( q_values: &Tensor, // [batch, 5] ) -> Result<((f64, f64, f64), [f64; 5])> { @@ -344,20 +362,38 @@ fn compute_q_diagnostics_gpu( // Per-action means: mean along batch dim [5] let per_action = q_values.mean(0)?; - let scalar = |t: &Tensor| -> Result { - Ok(t.to_dtype(ml_core::native_types::NativeDType::F32)? - .to_scalar::()? as f64) - }; - let mean_g = scalar(&mean_gap)?; - let min_g = scalar(&min_gap)?; - let max_g = scalar(&max_gap)?; - let per_action_flat = per_action.flatten_all()?.to_dtype(ml_core::native_types::NativeDType::F32)?; - let mut avgs = [0.0_f64; 5]; - for i in 0..5_usize { - avgs[i] = per_action_flat.get(i) - .and_then(|t| t.to_scalar::()) - .unwrap_or(0.0) as f64; - } + // Concatenate all 8 diagnostic scalars into a single [8] tensor on GPU, + // then do ONE bulk DtoH transfer instead of 8 individual to_scalar() calls. + // Layout: [mean_gap, min_gap, max_gap, action0, action1, action2, action3, action4] + let mean_gap_f32 = mean_gap.to_dtype(ml_core::native_types::NativeDType::F32)?; + let min_gap_f32 = min_gap.to_dtype(ml_core::native_types::NativeDType::F32)?; + let max_gap_f32 = max_gap.to_dtype(ml_core::native_types::NativeDType::F32)?; + let per_action_f32 = per_action.flatten_all()?.to_dtype(ml_core::native_types::NativeDType::F32)?; + + // Reshape scalars to [1] for cat compatibility + let gap_tensors = [ + mean_gap_f32.reshape(&[1])?, + min_gap_f32.reshape(&[1])?, + max_gap_f32.reshape(&[1])?, + ]; + let combined = Tensor::cat( + &[&gap_tensors[0], &gap_tensors[1], &gap_tensors[2], &per_action_f32], + 0, + )?; + + // Single DtoH readback: 8 floats = 32 bytes + let vals: Vec = combined.to_vec1()?; + + let mean_g = vals.get(0).copied().unwrap_or(0.0) as f64; + let min_g = vals.get(1).copied().unwrap_or(0.0) as f64; + let max_g = vals.get(2).copied().unwrap_or(0.0) as f64; + let avgs = [ + vals.get(3).copied().unwrap_or(0.0) as f64, + vals.get(4).copied().unwrap_or(0.0) as f64, + vals.get(5).copied().unwrap_or(0.0) as f64, + vals.get(6).copied().unwrap_or(0.0) as f64, + vals.get(7).copied().unwrap_or(0.0) as f64, + ]; Ok(((mean_g, min_g, max_g), avgs)) } @@ -368,6 +404,8 @@ fn compute_q_diagnostics_gpu( } /// Get Q-values for a given state + /// + /// Single bulk DtoH readback (was N individual get(i).to_scalar() calls). pub(crate) async fn get_q_values(&self, state: &TradingState) -> Result> { let agent = self.agent.read().await; let state_vec = state.to_vector(); @@ -384,12 +422,9 @@ fn compute_q_diagnostics_gpu( let q_values_tensor = agent.forward(&state_tensor)?.squeeze(0)? .to_dtype(ml_core::native_types::NativeDType::F32)?; - let n = q_values_tensor.dims()[0]; - let mut q_values = Vec::with_capacity(n); - for i in 0..n { - q_values.push(q_values_tensor.get(i)?.to_scalar::()? as f64); - } - Ok(q_values) + // Single bulk download instead of per-element get(i).to_scalar() loop + let q_f32: Vec = q_values_tensor.to_vec1()?; + Ok(q_f32.iter().map(|&v| v as f64).collect()) } /// Check if early stopping criteria are met @@ -674,14 +709,23 @@ fn compute_q_diagnostics_gpu( .map_err(|e| anyhow::anyhow!("GPU val rewards sqr: {e}"))? .mean_all() .map_err(|e| anyhow::anyhow!("GPU val rewards var: {e}"))?; - let mean_scalar = mean_t + // Batch [mean, var] into [2] tensor, single DtoH transfer (8 bytes) + let mean_f32 = mean_t .to_dtype(ml_core::native_types::NativeDType::F32) - .and_then(|t| t.to_scalar::()) - .map_err(|e| anyhow::anyhow!("GPU val Sharpe mean readback: {e}"))? as f64; - let var_scalar = var_t + .map_err(|e| anyhow::anyhow!("GPU val mean f32 cast: {e}"))? + .reshape(&[1]) + .map_err(|e| anyhow::anyhow!("GPU val mean reshape: {e}"))?; + let var_f32 = var_t .to_dtype(ml_core::native_types::NativeDType::F32) - .and_then(|t| t.to_scalar::()) - .map_err(|e| anyhow::anyhow!("GPU val Sharpe var readback: {e}"))? as f64; + .map_err(|e| anyhow::anyhow!("GPU val var f32 cast: {e}"))? + .reshape(&[1]) + .map_err(|e| anyhow::anyhow!("GPU val var reshape: {e}"))?; + let sharpe_stats = Tensor::cat(&[&mean_f32, &var_f32], 0) + .map_err(|e| anyhow::anyhow!("GPU val Sharpe stats cat: {e}"))?; + let sharpe_vals: Vec = sharpe_stats.to_vec1() + .map_err(|e| anyhow::anyhow!("GPU val Sharpe readback: {e}"))?; + let mean_scalar = sharpe_vals.get(0).copied().unwrap_or(0.0) as f64; + let var_scalar = sharpe_vals.get(1).copied().unwrap_or(0.0) as f64; let std_val = var_scalar.sqrt(); let val_sharpe = if std_val > 1e-10 {