fix(critical): recursive_confidence_backward — per-weight-element, zero syncthreads

ROOT CAUSE of graph_mega hang: the warp+block reduction in
recursive_confidence_backward had 256 iterations of __syncthreads
inside a loop (2 syncs × SH2=256 = 512 barriers per thread).
This overwhelmed the CUDA graph node scheduler on H100.

Fix: converted to per-weight-element pattern (same as other
deterministic backward kernels). One thread per weight (SH2+1=257),
loops over B samples. Zero __syncthreads, zero atomicAdd,
zero shared memory. Graph-mega safe.

Grid: ceil(257/256)=2 blocks (was ceil(8192/256)=32 blocks).
Massive resource reduction + deterministic + graph compatible.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-17 10:54:50 +02:00
parent 341346af09
commit 5ce0454ceb
2 changed files with 29 additions and 51 deletions

View File

@@ -5354,74 +5354,50 @@ extern "C" __global__ void recursive_confidence_forward(
/* Kernel: recursive_confidence_backward — MSE grad into trunk */
/* ================================================================== */
/**
* Phase 1: Per-sample trunk gradient + warp-reduced weight gradient.
* d_h_s2[i, k] += d_sigmoid * w_conf[k] (plain write — single writer per (i,k)).
* d_w_conf and d_b_conf use warp+block reduction → one atomicAdd per BLOCK.
* Deterministic per-weight-element accumulation. One thread per weight
* element (SH2 + 1 for bias), each loops over B samples.
* Zero atomicAdd, zero __syncthreads, graph_mega safe.
*
* Grid: ceil(B/256), Block: 256.
* Grid: ceil((SH2+1)/256), Block: 256.
*/
extern "C" __global__ void recursive_confidence_backward(
const float* __restrict__ h_s2, /* [B, SH2] */
const float* __restrict__ predicted_error, /* [B] */
const float* __restrict__ lagged_td_error, /* [1] pinned — target */
const float* __restrict__ w_conf, /* [SH2] */
float* __restrict__ d_w_conf, /* [SH2] gradient accumulator */
float* __restrict__ d_b_conf, /* [1] gradient accumulator */
float* __restrict__ d_w_conf, /* [SH2] gradient output */
float* __restrict__ d_b_conf, /* [1] gradient output */
float* __restrict__ d_h_s2, /* [B, SH2] trunk gradient (accumulate) */
int B, int SH2,
float loss_weight /* 0.01 */
) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
int widx = blockIdx.x * blockDim.x + threadIdx.x;
int total_weights = SH2 + 1; /* SH2 weights + 1 bias */
if (widx >= total_weights) return;
float d_sigmoid = 0.0f;
if (i < B) {
float target = lagged_td_error[0];
float sum = 0.0f;
for (int i = 0; i < B; i++) {
float pred = predicted_error[i];
float target = lagged_td_error[0];
float d_loss = loss_weight * 2.0f * (pred - target) / (float)B;
d_sigmoid = d_loss * pred * (1.0f - pred);
float d_sigmoid = d_loss * pred * (1.0f - pred);
/* Per-sample trunk gradient — no cross-sample contention, plain write */
for (int k = 0; k < SH2; k++) {
d_h_s2[(long long)i * SH2 + k] += d_sigmoid * w_conf[k];
if (widx < SH2) {
/* Weight gradient: sum over batch */
sum += d_sigmoid * h_s2[(long long)i * SH2 + widx];
/* Trunk gradient: accumulate (each (i, widx) has unique writer) */
d_h_s2[(long long)i * SH2 + widx] += d_sigmoid * w_conf[widx];
} else {
/* Bias gradient: sum over batch */
sum += d_sigmoid;
}
}
/* Weight gradient reduction: d_b_conf = sum_b(d_sigmoid_b) */
float my_db = d_sigmoid;
for (int offset = 16; offset > 0; offset >>= 1)
my_db += __shfl_xor_sync(0xFFFFFFFF, my_db, offset);
__shared__ float ws_db[8];
int warp_id = threadIdx.x / 32;
int lane = threadIdx.x % 32;
if (lane == 0) ws_db[warp_id] = my_db;
__syncthreads();
if (warp_id == 0) {
float val = (lane < blockDim.x / 32) ? ws_db[lane] : 0.0f;
for (int off = 16; off > 0; off >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, off);
if (lane == 0) atomicAdd(d_b_conf, val);
}
/* Weight gradient reduction: d_w_conf[k] = sum_b(d_sigmoid_b * h_s2[b, k]) */
for (int k = 0; k < SH2; k++) {
float my_dw = (i < B) ? d_sigmoid * h_s2[(long long)i * SH2 + k] : 0.0f;
for (int offset = 16; offset > 0; offset >>= 1)
my_dw += __shfl_xor_sync(0xFFFFFFFF, my_dw, offset);
__shared__ float ws_dw[8];
if (lane == 0) ws_dw[warp_id] = my_dw;
__syncthreads();
if (warp_id == 0) {
float val = (lane < blockDim.x / 32) ? ws_dw[lane] : 0.0f;
for (int off = 16; off > 0; off >>= 1)
val += __shfl_xor_sync(0xFFFFFFFF, val, off);
if (lane == 0) atomicAdd(&d_w_conf[k], val);
}
__syncthreads();
if (widx < SH2) {
d_w_conf[widx] = sum;
} else {
d_b_conf[0] = sum;
}
}

View File

@@ -2267,7 +2267,9 @@ impl GpuDqnTrainer {
let d_w_conf = self.ptrs.grad_buf + padded_byte_offset(&param_sizes, 76);
let d_b_conf = self.ptrs.grad_buf + padded_byte_offset(&param_sizes, 77);
let blocks = ((batch_size as u32 + 255) / 256).max(1);
// Per-weight-element grid: SH2+1 threads (one per weight + bias)
let total_weights = (self.config.shared_h2 + 1) as u32;
let blocks = ((total_weights + 255) / 256).max(1);
let b_i32 = batch_size as i32;
let sh2 = self.config.shared_h2 as i32;
let loss_weight = 0.01_f32;