perf(cuda): batch DQN metrics readbacks — 15 DtoH transfers → 4

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-18 08:29:26 +01:00
parent b53df16b5f
commit d572dc466a

View File

@@ -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<f64, crate::MLError> {
Ok(t.to_dtype(ml_core::native_types::NativeDType::F32)
.and_then(|t| t.to_scalar::<f32>())
.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<f32> = 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<f64, MLError> {
Ok(t.to_dtype(ml_core::native_types::NativeDType::F32)?
.to_scalar::<f32>()? 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::<f32>())
.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<f32> = 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<Vec<f64>> {
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::<f32>()? as f64);
}
Ok(q_values)
// Single bulk download instead of per-element get(i).to_scalar() loop
let q_f32: Vec<f32> = 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::<f32>())
.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::<f32>())
.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<f32> = 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 {