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