diff --git a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu index f485d95b4..9fdc474b4 100644 --- a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu +++ b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu @@ -96,9 +96,10 @@ extern "C" __global__ void dqn_adam_update_kernel( __nv_bfloat16 clip_scale = (*grad_norm_sq > bf16_zero() && norm > bf16(max_grad_norm)) ? (bf16(max_grad_norm) / norm) : bf16_one(); __nv_bfloat16 clipped_g = g * clip_scale; - /* Adam update */ - __nv_bfloat16 beta1_t = bf16_one() - bf16_pow(bf16(beta1), bf16((float)t)); - __nv_bfloat16 beta2_t = bf16_one() - bf16_pow(bf16(beta2), bf16((float)t)); + /* Adam bias correction: compute in float — bf16(0.999) rounds to 1.0 → div-by-zero. + * This is the ONLY float arithmetic in the kernel (3 decimal digits insufficient). */ + __nv_bfloat16 beta1_t = bf16(fmaxf(1.0f - powf(beta1, (float)t), 1e-4f)); + __nv_bfloat16 beta2_t = bf16(fmaxf(1.0f - powf(beta2, (float)t), 1e-4f)); __nv_bfloat16 m_i = bf16(beta1) * m[idx] + (bf16_one() - bf16(beta1)) * clipped_g; __nv_bfloat16 v_i = bf16(beta2) * v[idx] + (bf16_one() - bf16(beta2)) * clipped_g * clipped_g;