diff --git a/crates/ml/src/cuda_pipeline/backward_kernels.cu b/crates/ml/src/cuda_pipeline/backward_kernels.cu index 05996da5d..99157dd2e 100644 --- a/crates/ml/src/cuda_pipeline/backward_kernels.cu +++ b/crates/ml/src/cuda_pipeline/backward_kernels.cu @@ -9,8 +9,16 @@ * block=(256, 1, 1). */ -/* ReLU mask: zero gradient where activation <= 0 (ReLU backward). - * No magnitude clamp needed — f32 d_logits eliminate bf16 overflow. */ +/* ReLU mask + bf16 overflow clamp for backward dX. + * + * The dX GemmEx (bf16 A × bf16 B → bf16 C) writes bf16 output. When the + * f32 accumulated sum exceeds bf16 max (~65504), the bf16 write produces + * Inf. This clamp prevents Inf from cascading through the backward chain. + * Same pattern as the forward bias kernel ±500 clamp. + * + * This is NOT a NaN guard — it's bf16 overflow prevention at the type + * boundary, identical to the forward-pass bias kernel clamping. */ +#define BW_SAFE_MAX 500.0f extern "C" __global__ void relu_mask_kernel( __nv_bfloat16* __restrict__ dy, @@ -19,7 +27,14 @@ extern "C" __global__ void relu_mask_kernel( { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i >= n) return; - if (activation[i] <= bf16_zero()) dy[i] = bf16_zero(); + float act_f = (float)activation[i]; + if (act_f <= 0.0f) { + dy[i] = bf16_zero(); + } else { + float dy_f = (float)dy[i]; + dy_f = fminf(fmaxf(dy_f, -BW_SAFE_MAX), BW_SAFE_MAX); + dy[i] = bf16(dy_f); + } } extern "C" __global__ void bias_grad_reduce_kernel( diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 9acc0fb95..fc85d0baa 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -2481,8 +2481,7 @@ impl GpuDqnTrainer { /// gradients into grad_buf (IQN, attention, ensemble). pub fn replay_adam_and_readback(&mut self) -> Result { self.replay_adam()?; - // Finalize grad_norm OUTSIDE graph: convert float sum-of-squares → bf16 L2 norm. - // The graph stores float sums via atomicAdd; this kernel reads them and writes bf16. + // Finalize grad_norm OUTSIDE graph self.launch_grad_norm_finalize()?; // Per-step: zero readback for performance (no cuStreamSync). // Loss/grad_norm accumulated on GPU by training_guard kernel.