perf(dqn): wire GpuTrainingGuard into train_step_single_batch (zero-sync loss/grad)

Replace the GPU->CPU sync barrier (to_vec1 readback) in the DQN training
hot path with the GpuTrainingGuard CUDA kernel that performs NaN detection,
loss clipping, and gradient collapse checks entirely on-device. The guard
is lazy-initialized on first training step and falls back to the original
CPU readback path if CUDA kernel compilation fails.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-10 14:21:49 +01:00
parent 0aaa53bd18
commit 38f0bdec00

View File

@@ -238,6 +238,10 @@ pub struct DQNTrainer {
#[cfg(feature = "cuda")]
gpu_action_selector: Option<crate::cuda_pipeline::gpu_action_selector::GpuActionSelector>,
/// GPU training guard for zero-sync safety checks (loss clip, NaN, grad collapse)
#[cfg(feature = "cuda")]
training_guard: Option<crate::cuda_pipeline::gpu_training_guard::GpuTrainingGuard>,
/// Reusable GPU staging buffers for zero-alloc fold transitions
buffer_pool: Option<crate::cuda_pipeline::GpuBufferPool>,
@@ -902,6 +906,8 @@ impl DQNTrainer {
gpu_experience_collector: None,
#[cfg(feature = "cuda")]
gpu_action_selector: None,
#[cfg(feature = "cuda")]
training_guard: None,
// GPU pipeline: staging buffer pool (auto-initialized on CUDA devices)
buffer_pool,
@@ -4520,38 +4526,135 @@ impl DQNTrainer {
.train_step(explicit_batch)
.map_err(|e| anyhow::anyhow!("Training step failed: {}", e))?;
// Batch both GPU scalars into a single [2]-element tensor for ONE DMA
// readback instead of two separate GPU sync barriers. Each to_scalar()
// triggers a cudaMemcpy + stream synchronize; by stacking first we pay
// that cost only once.
let loss_unsqueezed = gpu_result.loss_gpu.unsqueeze(0)
.map_err(|e| anyhow::anyhow!("Loss unsqueeze: {e}"))?;
let grad_unsqueezed = gpu_result.grad_norm_gpu.unsqueeze(0)
.map_err(|e| anyhow::anyhow!("Grad norm unsqueeze: {e}"))?;
let stacked = candle_core::Tensor::cat(&[&loss_unsqueezed, &grad_unsqueezed], 0)
.map_err(|e| anyhow::anyhow!("Loss+grad stack: {e}"))?;
let readback = stacked.to_vec1::<f32>()
.map_err(|e| anyhow::anyhow!("Batched readback: {e}"))?;
let loss_f32 = readback.first().copied().unwrap_or(0.0);
let grad_norm_f32 = readback.get(1).copied().unwrap_or(0.0);
// GPU training guard: on-device NaN/loss-clip/grad-collapse checks.
// Zero cudaStreamSynchronize — kernel writes halt flags to pinned host memory.
#[cfg(feature = "cuda")]
let (loss_clipped, grad_norm) = {
// Lazy-init training guard on first call
if self.training_guard.is_none() && self.device.is_cuda() {
match crate::cuda_pipeline::gpu_training_guard::GpuTrainingGuard::new(&self.device) {
Ok(guard) => {
info!("GPU training guard initialized");
self.training_guard = Some(guard);
}
Err(e) => {
warn!("GPU training guard init failed, using CPU fallback: {e}");
}
}
}
// Check for gradient collapse (early stopping)
agent.log_diagnostics(grad_norm_f32)
.map_err(|e| {
if let Some(ref mut guard) = self.training_guard {
let grad_collapse_threshold =
self.hyperparams.learning_rate as f32
* self.hyperparams.gradient_collapse_multiplier as f32;
let warmup_steps = (self.hyperparams.buffer_size as f64 * 0.2) as u64;
let past_warmup = self.gradient_logging_step as u64 > warmup_steps;
let result = guard
.check_and_accumulate(
&gpu_result.loss_gpu,
&gpu_result.grad_norm_gpu,
1e6_f32, // loss clip threshold
grad_collapse_threshold,
!past_warmup,
)
.map_err(|e| anyhow::anyhow!("GPU guard check: {e}"))?;
// Handle halt conditions
if result.halt_nan {
return Err(anyhow::anyhow!(
"NaN/Inf detected in loss ({}) or grad_norm ({})",
result.raw_loss,
result.raw_grad_norm
));
}
if result.halt_loss_clip {
warn!(
"Loss clipped from {:.2e} to 1.0e6 (TD error explosion, epoch {})",
result.raw_loss,
self.loss_history.len() + 1
);
}
// Always call log_diagnostics for dead neuron detection + collapse counter
agent.log_diagnostics(result.raw_grad_norm).map_err(|e| {
tracing::info!("Early stopping triggered (gradient collapse): {}", e);
anyhow::anyhow!("Early stopping: {}", e)
})?;
(result.clipped_loss as f64, result.raw_grad_norm as f64)
} else {
// CPU fallback (guard init failed or non-CUDA device)
let loss_unsqueezed = gpu_result
.loss_gpu
.unsqueeze(0)
.map_err(|e| anyhow::anyhow!("Loss unsqueeze: {e}"))?;
let grad_unsqueezed = gpu_result
.grad_norm_gpu
.unsqueeze(0)
.map_err(|e| anyhow::anyhow!("Grad norm unsqueeze: {e}"))?;
let stacked =
candle_core::Tensor::cat(&[&loss_unsqueezed, &grad_unsqueezed], 0)
.map_err(|e| anyhow::anyhow!("Loss+grad stack: {e}"))?;
let readback = stacked
.to_vec1::<f32>()
.map_err(|e| anyhow::anyhow!("Batched readback: {e}"))?;
let loss_f32 = readback.first().copied().unwrap_or(0.0);
let grad_norm_f32 = readback.get(1).copied().unwrap_or(0.0);
agent.log_diagnostics(grad_norm_f32).map_err(|e| {
tracing::info!("Early stopping triggered (gradient collapse): {}", e);
anyhow::anyhow!("Early stopping: {}", e)
})?;
let loss_clipped_val = if loss_f32 > 1e6 {
warn!(
"Loss clipped from {:.2e} to 1.0e6 (TD error explosion, epoch {})",
loss_f32,
self.loss_history.len() + 1
);
1e6_f32
} else {
loss_f32
};
(loss_clipped_val as f64, grad_norm_f32 as f64)
}
};
#[cfg(not(feature = "cuda"))]
let (loss_clipped, grad_norm) = {
let loss_unsqueezed = gpu_result
.loss_gpu
.unsqueeze(0)
.map_err(|e| anyhow::anyhow!("Loss unsqueeze: {e}"))?;
let grad_unsqueezed = gpu_result
.grad_norm_gpu
.unsqueeze(0)
.map_err(|e| anyhow::anyhow!("Grad norm unsqueeze: {e}"))?;
let stacked =
candle_core::Tensor::cat(&[&loss_unsqueezed, &grad_unsqueezed], 0)
.map_err(|e| anyhow::anyhow!("Loss+grad stack: {e}"))?;
let readback = stacked
.to_vec1::<f32>()
.map_err(|e| anyhow::anyhow!("Batched readback: {e}"))?;
let loss_f32 = readback.first().copied().unwrap_or(0.0);
let grad_norm_f32 = readback.get(1).copied().unwrap_or(0.0);
agent.log_diagnostics(grad_norm_f32).map_err(|e| {
tracing::info!("Early stopping triggered (gradient collapse): {}", e);
anyhow::anyhow!("Early stopping: {}", e)
})?;
// Emergency brake — clip extreme losses
let loss_clipped = if loss_f32 > 1e6 {
warn!(
"Loss clipped from {:.2e} to 1.0e6 (TD error explosion, epoch {})",
loss_f32,
self.loss_history.len() + 1
);
1e6
} else {
loss_f32
let loss_clipped_val = if loss_f32 > 1e6 {
warn!(
"Loss clipped from {:.2e} to 1.0e6 (TD error explosion, epoch {})",
loss_f32,
self.loss_history.len() + 1
);
1e6_f32
} else {
loss_f32
};
(loss_clipped_val as f64, grad_norm_f32 as f64)
};
// Q-value estimation: periodic (every 50 steps)
@@ -4561,8 +4664,6 @@ impl DQNTrainer {
}
let avg_q_value = self.cached_avg_q;
let grad_norm = grad_norm_f32 as f64;
debug!("Gradient norm after clip (actual): {:.4}", grad_norm);
self.gradient_logging_step += 1;
@@ -4573,7 +4674,7 @@ impl DQNTrainer {
);
}
Ok((loss_clipped as f64, avg_q_value, grad_norm))
Ok((loss_clipped, avg_q_value, grad_norm))
}
/// Training step with true gradient accumulation across N mini-batches.