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:
@@ -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];
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user