fix: two-phase vaccine dot+norm reduction (3rd atomicAdd hang on H100)

gradient_dot_and_norm kernel used atomicAdd with 1264 blocks to compute
dot(g_train, g_val) and |g_val|² — same H100 L2 contention pattern as
the grad_norm hang. Replaced with two-phase: per-block partials + single-
block finalize (gradient_dot_norm_finalize). Zero atomicAdd.

Also added step-0-only per-phase GPU syncs to catch any remaining hangs
immediately (forward, aux, conditional, Adam, PER — only on step 0).

TODO: Causal intervention runs 14 full cuBLAS forward passes per step.
Should be optimized (batched or reduced frequency).

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-06 21:47:25 +02:00
parent 1cafce985a
commit adcbd6f0df
2 changed files with 75 additions and 31 deletions

View File

@@ -969,46 +969,72 @@ extern "C" __global__ void sharpness_perturb_weights(
* and read by the projection kernel — no CPU roundtrip.
* ══════════════════════════════════════════════════════════════════════ */
/* Phase 1: each block writes partial dot + norm to block_results.
* block_results[blockIdx.x * 2 + 0] = partial dot(g_train, g_val)
* block_results[blockIdx.x * 2 + 1] = partial |g_val|²
* NO atomicAdd. Launch: grid=(num_blocks), block=(256). */
extern "C" __global__ void gradient_dot_and_norm(
const float* __restrict__ g_train, /* [total_params] f32 */
const float* __restrict__ g_val, /* [total_params] f32 */
float* __restrict__ result, /* [2]: result[0]=dot, result[1]=norm_sq */
const float* __restrict__ g_train, /* [total_params] f32 */
const float* __restrict__ g_val, /* [total_params] f32 */
float* __restrict__ block_results, /* [num_blocks * 2] */
int n
) {
/* Parallel reduction: each block reduces a tile, atomicAdd into result. */
__shared__ float s_dot[256];
__shared__ float s_norm[256];
int tid = threadIdx.x;
int idx = blockIdx.x * blockDim.x + threadIdx.x;
float local_dot = 0.0f, local_norm = 0.0f;
float local_dot = 0.0f;
float local_norm = 0.0f;
/* Grid-stride loop for large n */
for (int i = idx; i < n; i += blockDim.x * gridDim.x) {
float t = g_train[i];
float v = g_val[i];
local_dot += t * v;
float t = g_train[i], v = g_val[i];
local_dot += t * v;
local_norm += v * v;
}
s_dot[tid] = local_dot;
s_dot[tid] = local_dot;
s_norm[tid] = local_norm;
__syncthreads();
/* Block reduction */
for (int s = 128; s > 0; s >>= 1) {
if (tid < s) {
s_dot[tid] += s_dot[tid + s];
s_dot[tid] += s_dot[tid + s];
s_norm[tid] += s_norm[tid + s];
}
__syncthreads();
}
if (tid == 0) {
atomicAdd(&result[0], s_dot[0]);
atomicAdd(&result[1], s_norm[0]);
block_results[blockIdx.x * 2 + 0] = s_dot[0];
block_results[blockIdx.x * 2 + 1] = s_norm[0];
}
}
/* Phase 2: reduce block partials → result[0]=dot, result[1]=norm_sq.
* Launch: grid=(1), block=(256). */
extern "C" __global__ void gradient_dot_norm_finalize(
const float* __restrict__ block_results, /* [num_blocks * 2] */
float* __restrict__ result, /* [2]: dot, norm_sq */
int num_blocks
) {
int tid = threadIdx.x;
float sum_dot = 0.0f, sum_norm = 0.0f;
for (int i = tid; i < num_blocks; i += blockDim.x) {
sum_dot += block_results[i * 2 + 0];
sum_norm += block_results[i * 2 + 1];
}
__shared__ float s_dot[256];
__shared__ float s_norm[256];
s_dot[tid] = sum_dot;
s_norm[tid] = sum_norm;
__syncthreads();
for (int s = 128; s > 0; s >>= 1) {
if (tid < s) {
s_dot[tid] += s_dot[tid + s];
s_norm[tid] += s_norm[tid + s];
}
__syncthreads();
}
if (tid == 0) {
result[0] = s_dot[0];
result[1] = s_norm[0];
}
}

View File

@@ -695,8 +695,10 @@ pub struct GpuDqnTrainer {
vaccine_grad_save: CudaSlice<f32>,
/// #32 Gradient vaccine: dot(g_train, g_val) and |g_val|^2 result [2] f32.
vaccine_dot_norm_buf: CudaSlice<f32>,
vaccine_block_results: CudaSlice<f32>, // [grad_norm_blocks * 2] partials for dot+norm
/// #32 Gradient vaccine kernels.
vaccine_dot_kernel: CudaFunction,
vaccine_dot_finalize_kernel: CudaFunction,
vaccine_project_kernel: CudaFunction,
/// #32 Whether gradient vaccine is enabled.
enable_gradient_vaccine: bool,
@@ -2141,7 +2143,7 @@ impl GpuDqnTrainer {
// per array. Stack is set once in DQNTrainer::new() (64KB for all kernels).
// ── Compile 4 utility kernels (grad_norm, adam_update, BF16 converters) ─
let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, scale_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, causal_intervene_kernel_fn, causal_reduce_kernel_fn, causal_mean_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel, nan_check_f32_kernel, nan_check_bf16_kernel) =
let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel, scale_f32_kernel, stochastic_depth_kernel, stochastic_depth_rng_kernel, bn_tanh_backward_kernel, bn_bias_grad_kernel, bn_tanh_concat_kernel_fn, vaccine_dot_kernel, vaccine_project_kernel, vaccine_dot_finalize, causal_intervene_kernel_fn, causal_reduce_kernel_fn, causal_mean_reduce_kernel_fn, pruning_mask_kernel, pruning_compute_kernel, her_inplace_kernel, nan_check_f32_kernel, nan_check_bf16_kernel) =
compile_training_kernels(&stream, &config)?;
// Separate grad_norm instance for non-graph launches (clip_grad_buf_inplace).
@@ -2664,6 +2666,8 @@ impl GpuDqnTrainer {
// #32 Gradient vaccine buffers
let vaccine_grad_save = stream.alloc_zeros::<f32>(total_params)
.map_err(|e| MLError::ModelError(format!("alloc vaccine_grad_save: {e}")))?;
let vaccine_block_results = stream.alloc_zeros::<f32>(grad_norm_blocks * 2)
.map_err(|e| MLError::ModelError(format!("alloc vaccine_block_results: {e}")))?;
let vaccine_dot_norm_buf = stream.alloc_zeros::<f32>(2)
.map_err(|e| MLError::ModelError(format!("alloc vaccine_dot_norm: {e}")))?;
@@ -2867,6 +2871,8 @@ impl GpuDqnTrainer {
vaccine_grad_save,
vaccine_dot_norm_buf,
vaccine_dot_kernel,
vaccine_dot_finalize_kernel: vaccine_dot_finalize,
vaccine_block_results,
vaccine_project_kernel,
enable_gradient_vaccine: vaccine_enabled,
bn_hidden_buf,
@@ -3466,23 +3472,33 @@ impl GpuDqnTrainer {
self.launch_cublas_backward()?;
// Now grad_buf = g_val (vaccine gradient)
// Step 4: Dot product and norm²
self.stream.memset_zeros(&mut self.vaccine_dot_norm_buf)
.map_err(|e| MLError::ModelError(format!("vaccine zero dot: {e}")))?;
// Step 4: Two-phase dot product and norm² — no atomicAdd, no memset
let tp_i32 = tp as i32;
let blocks = ((tp as u32 + 255) / 256) as u32;
let cfg = LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
let blocks = self.grad_norm_blocks as u32;
let nb = blocks as i32;
let block_res_ptr = self.vaccine_block_results.raw_ptr();
let dot_norm_ptr = self.vaccine_dot_norm_buf.raw_ptr();
let g_val_ptr = self.grad_buf.raw_ptr();
// Phase 1: per-block partials
unsafe {
self.stream
.launch_builder(&self.vaccine_dot_kernel)
.arg(&save_ptr)
.arg(&g_val_ptr)
.arg(&dot_norm_ptr)
.arg(&block_res_ptr)
.arg(&tp_i32)
.launch(cfg)
.map_err(|e| MLError::ModelError(format!("vaccine dot: {e}")))?;
.launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("vaccine dot phase1: {e}")))?;
}
// Phase 2: reduce partials
unsafe {
self.stream
.launch_builder(&self.vaccine_dot_finalize_kernel)
.arg(&block_res_ptr)
.arg(&dot_norm_ptr)
.arg(&nb)
.launch(LaunchConfig { grid_dim: (1, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("vaccine dot phase2: {e}")))?;
}
// Step 5: Conditional projection (only when dot < 0)
@@ -3493,7 +3509,7 @@ impl GpuDqnTrainer {
.arg(&g_val_ptr)
.arg(&dot_norm_ptr)
.arg(&tp_i32)
.launch(cfg)
.launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("vaccine project: {e}")))?;
}
@@ -5843,7 +5859,7 @@ impl GpuDqnTrainer {
fn compile_training_kernels(
stream: &Arc<CudaStream>,
config: &GpuDqnTrainConfig,
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
info!(
state_dim = config.state_dim,
total_params = compute_total_params(config),
@@ -5903,6 +5919,8 @@ fn compile_training_kernels(
.map_err(|e| MLError::ModelError(format!("causal_mean_reduce load: {e}")))?;
let vaccine_dot = module.load_function("gradient_dot_and_norm")
.map_err(|e| MLError::ModelError(format!("gradient_dot_and_norm load: {e}")))?;
let vaccine_dot_finalize = module.load_function("gradient_dot_norm_finalize")
.map_err(|e| MLError::ModelError(format!("gradient_dot_norm_finalize load: {e}")))?;
let vaccine_project = module.load_function("gradient_project")
.map_err(|e| MLError::ModelError(format!("gradient_project load: {e}")))?;
let bn_tanh_bw = module.load_function("bn_tanh_backward_kernel")
@@ -5920,7 +5938,7 @@ fn compile_training_kernels(
.map_err(|e| MLError::ModelError(format!("dqn_nan_check_bf16 load: {e}")))?;
info!("GpuDqnTrainer: 31 utility kernels loaded from precompiled cubin");
Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32, scale_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, causal_intervene, causal_reduce, causal_mean_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace, nan_check_f32, nan_check_bf16))
Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32, scale_f32, stochastic_depth, stochastic_depth_rng, bn_tanh_bw, bn_bias_grad, bn_tanh_concat, vaccine_dot, vaccine_project, vaccine_dot_finalize, causal_intervene, causal_reduce, causal_mean_reduce, pruning_mask_fn, pruning_compute_fn, her_inplace, nan_check_f32, nan_check_bf16))
}
/// Load the standalone Polyak EMA kernel from precompiled cubin.