feat(bf16): f32 d_logits buffers — native atomicAdd, zero NaN from gradients
d_value_logits, d_adv_logits (+ MSE/CQL scratch): bf16 → f32 - Native atomicAdd(float*) replaces atomicAddBF16 CAS loop - Eliminates bf16 accumulation overflow in gradient kernels - Gradient value clamping ±100 removed (unnecessary with f32) - NaN guards removed from loss kernels Architecture: - f32 d_logits for gradient accumulation (atomicAdd-safe) - bf16 staging buffers (d_value_logits_bf16, d_adv_logits_bf16) cast via f32_to_bf16_kernel before backward dW GemmEx - dqn_saxpy_f32_kernel for gradient blending (MSE+C51 alpha) - CQL backward uses bf16 staging after f32→bf16 cast Remaining intermittent NaN (~1/2000 steps on long runs): - Source: bf16 params_buf weight precision loss → forward pass - Fix: f32 master weights (next commit) 895/895 unit + 359/359 ml-dqn tests pass. 9-11/11 smoke tests (intermittent NaN on 50-epoch runs). Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -1,8 +1,8 @@
|
||||
/**
|
||||
* C51 distributional RL loss gradient kernel.
|
||||
*
|
||||
* Mixed-precision: reads BF16, computes in float, writes BF16.
|
||||
* Prevents NaN from bf16 exp() overflow and intermediate product overflow.
|
||||
* Mixed-precision: reads BF16 inputs, computes in float, writes f32 d_logits.
|
||||
* f32 atomicAdd eliminates bf16 overflow that caused NaN.
|
||||
*
|
||||
* dL/d_combined[b,d,j] = is_weights[b] * (exp(current_lp[b,d,j]) - projected[b,d,j])
|
||||
* d_value[b,j] = sum_d dL/d_combined[b,d,j]
|
||||
@@ -16,8 +16,8 @@ extern "C" __global__ void c51_grad_kernel(
|
||||
const __nv_bfloat16* __restrict__ projected, // [B, 3, NA]
|
||||
const __nv_bfloat16* __restrict__ is_weights, // [B] bf16
|
||||
const int* __restrict__ actions, // [B] factored
|
||||
__nv_bfloat16* __restrict__ d_value_logits, // [B, NA]
|
||||
__nv_bfloat16* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA]
|
||||
float* __restrict__ d_value_logits, // [B, NA] f32 (native atomicAdd, no overflow)
|
||||
float* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] f32
|
||||
int batch_size,
|
||||
int num_atoms,
|
||||
int b0_size, int b1_size, int b2_size,
|
||||
@@ -47,10 +47,8 @@ extern "C" __global__ void c51_grad_kernel(
|
||||
d_combined += entropy_coeff * (1.0f + lp_clamped);
|
||||
}
|
||||
|
||||
/* Clamp: 3 branches × batch atomicAdds per element in bf16 d_value_logits.
|
||||
* max accumulated: 3 * 100 = 300 → bf16 safe. */
|
||||
d_combined = fminf(fmaxf(d_combined, -100.0f), 100.0f);
|
||||
atomicAddBF16(&d_value_logits[b * num_atoms + j], d_combined);
|
||||
/* d_value_logits is f32 — native atomicAdd, no overflow risk. */
|
||||
atomicAdd(&d_value_logits[b * num_atoms + j], d_combined);
|
||||
|
||||
/* Factored action decode */
|
||||
int factored = actions[b];
|
||||
@@ -78,6 +76,6 @@ extern "C" __global__ void c51_grad_kernel(
|
||||
float dueling_grad = (a == a_d) ? (1.0f - inv_A) : (-inv_A);
|
||||
float grad_val = d_combined * dueling_grad;
|
||||
int adv_idx = b * total_branch_atoms + branch_offset + a * num_atoms + j;
|
||||
atomicAddBF16(&d_adv_logits[adv_idx], grad_val);
|
||||
atomicAdd(&d_adv_logits[adv_idx], grad_val);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -410,8 +410,7 @@ extern "C" __global__ void c51_loss_batched(
|
||||
if (tid == 0) {
|
||||
float clamped_ce = fminf(avg_ce, MAX_PER_SAMPLE_CE);
|
||||
float weighted_loss = clamped_ce * is_weight;
|
||||
if (!fast_isfinite(weighted_loss)) weighted_loss = 0.0f;
|
||||
if (!fast_isfinite(clamped_ce)) clamped_ce = 0.0f;
|
||||
/* f32 d_logits: no NaN risk from atomicAdd overflow. */
|
||||
per_sample_loss[sample_id] = bf16(weighted_loss);
|
||||
td_errors[sample_id] = bf16(clamped_ce);
|
||||
atomicAdd(total_loss, weighted_loss / (float)batch_size);
|
||||
|
||||
@@ -18,8 +18,8 @@ extern "C" __global__ void cql_logit_grad_kernel(
|
||||
const float* __restrict__ v_logits, // [N, num_atoms] f32
|
||||
const float* __restrict__ adv_logits, // [N, total_actions * num_atoms] f32
|
||||
const int* __restrict__ actions, // [N] factored action indices (0-44)
|
||||
__nv_bfloat16* __restrict__ d_v_logits, // [N, num_atoms] output (bf16 grad)
|
||||
__nv_bfloat16* __restrict__ d_adv_logits,// [N, total_actions * num_atoms] output (bf16 grad)
|
||||
float* __restrict__ d_v_logits, // [N, num_atoms] output (f32 grad, no overflow)
|
||||
float* __restrict__ d_adv_logits,// [N, total_actions * num_atoms] output (f32 grad)
|
||||
float cql_alpha,
|
||||
int N, int num_atoms,
|
||||
int b0_size, int b1_size, int b2_size,
|
||||
@@ -120,7 +120,7 @@ extern "C" __global__ void cql_logit_grad_kernel(
|
||||
for (int a = 0; a < bd; a++) {
|
||||
const float* adv = adv_logits + (long long)i * total_actions * num_atoms
|
||||
+ (long long)(adv_offset + a) * num_atoms;
|
||||
__nv_bfloat16* d_adv = d_adv_logits + (long long)i * total_actions * num_atoms
|
||||
float* d_adv = d_adv_logits + (long long)i * total_actions * num_atoms
|
||||
+ (long long)(adv_offset + a) * num_atoms;
|
||||
|
||||
// Recompute p[j] for this action
|
||||
@@ -143,8 +143,8 @@ extern "C" __global__ void cql_logit_grad_kernel(
|
||||
// d_combined_logit[j] = d_cql_dq[a] * p * (z - Q)
|
||||
float d_combined = d_cql_dq[a] * p * (z - eq);
|
||||
|
||||
// Split combined gradient to adv and val (bf16 output)
|
||||
d_adv[j] = bf16(d_combined);
|
||||
// Split combined gradient to adv and val (f32 output)
|
||||
d_adv[j] = d_combined;
|
||||
if (j < 256) d_val_accum[j] += d_combined;
|
||||
}
|
||||
}
|
||||
@@ -152,8 +152,8 @@ extern "C" __global__ void cql_logit_grad_kernel(
|
||||
}
|
||||
|
||||
// Write accumulated value logit gradient (summed across all branches and actions)
|
||||
__nv_bfloat16* d_val = d_v_logits + (long long)i * num_atoms;
|
||||
float* d_val = d_v_logits + (long long)i * num_atoms;
|
||||
for (int j = 0; j < num_atoms && j < 256; j++) {
|
||||
d_val[j] = bf16(d_val_accum[j]);
|
||||
d_val[j] = d_val_accum[j];
|
||||
}
|
||||
}
|
||||
|
||||
@@ -155,6 +155,26 @@ extern "C" __global__ void dqn_saxpy_kernel(
|
||||
if (i < n) y[i] = y[i] + bf16(alpha) * x[i];
|
||||
}
|
||||
|
||||
/* ══════════════════════════════════════════════════════════════════════
|
||||
* F32 SAXPY KERNEL
|
||||
*
|
||||
* y[i] += alpha * x[i] for i = 0..n-1
|
||||
*
|
||||
* Float variant for f32 d_logits blending (MSE+C51 gradient mix).
|
||||
*
|
||||
* Launch config: grid=(ceil(n/256), 1, 1), block=(256, 1, 1).
|
||||
* ══════════════════════════════════════════════════════════════════════ */
|
||||
|
||||
extern "C" __global__ void dqn_saxpy_f32_kernel(
|
||||
float* __restrict__ y,
|
||||
const float* __restrict__ x,
|
||||
float alpha,
|
||||
int n
|
||||
) {
|
||||
int i = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (i < n) y[i] = y[i] + alpha * x[i];
|
||||
}
|
||||
|
||||
/* ══════════════════════════════════════════════════════════════════════
|
||||
* CLIPPED SAXPY KERNEL
|
||||
*
|
||||
|
||||
@@ -392,6 +392,7 @@ pub struct GpuDqnTrainer {
|
||||
f32_to_bf16_kernel: CudaFunction,
|
||||
bf16_to_f32_kernel: CudaFunction,
|
||||
saxpy_kernel: CudaFunction,
|
||||
saxpy_f32_kernel: CudaFunction,
|
||||
zero_kernel: CudaFunction,
|
||||
regime_scale_kernel: CudaFunction,
|
||||
shrink_perturb_kernel: CudaFunction,
|
||||
@@ -570,12 +571,16 @@ pub struct GpuDqnTrainer {
|
||||
/// When this differs from `loss_mode`, the graph must be recaptured.
|
||||
last_captured_loss_mode: Option<LossMode>,
|
||||
/// Scratch buffers for blended loss (MSE grad stored here, then blended into d_value/d_adv)
|
||||
d_value_logits_mse: CudaSlice<half::bf16>,
|
||||
d_adv_logits_mse: CudaSlice<half::bf16>,
|
||||
/// Gradient w.r.t. value logits: [B, NA]
|
||||
d_value_logits_buf: CudaSlice<half::bf16>,
|
||||
/// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA]
|
||||
d_adv_logits_buf: CudaSlice<half::bf16>,
|
||||
d_value_logits_mse: CudaSlice<f32>,
|
||||
d_adv_logits_mse: CudaSlice<f32>,
|
||||
/// Gradient w.r.t. value logits: [B, NA] — f32 for native atomicAdd (no bf16 overflow)
|
||||
d_value_logits_buf: CudaSlice<f32>,
|
||||
/// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA] — f32 for native atomicAdd
|
||||
d_adv_logits_buf: CudaSlice<f32>,
|
||||
/// BF16 staging for backward pass: cast from f32 d_value_logits before cuBLAS GEMM
|
||||
d_value_logits_bf16: CudaSlice<half::bf16>,
|
||||
/// BF16 staging for backward pass: cast from f32 d_adv_logits before cuBLAS GEMM
|
||||
d_adv_logits_bf16: CudaSlice<half::bf16>,
|
||||
|
||||
|
||||
// ── cuBLAS batched backward (Phase 2 Task 2) ──────────────────────
|
||||
@@ -609,10 +614,10 @@ pub struct GpuDqnTrainer {
|
||||
/// Computes CQL logit gradients: dCQL/d_value_logits and dCQL/d_adv_logits.
|
||||
/// Only used when `config.use_cql == true && config.cql_alpha > 0`.
|
||||
cql_logit_grad_kernel: Option<CudaFunction>,
|
||||
/// CQL scratch: value logit gradients [B, NA]
|
||||
cql_d_value_logits: CudaSlice<half::bf16>,
|
||||
/// CQL scratch: advantage logit gradients [B, (B0+B1+B2)*NA]
|
||||
cql_d_adv_logits: CudaSlice<half::bf16>,
|
||||
/// CQL scratch: value logit gradients [B, NA] — f32 for native atomicAdd
|
||||
cql_d_value_logits: CudaSlice<f32>,
|
||||
/// CQL scratch: advantage logit gradients [B, (B0+B1+B2)*NA] — f32
|
||||
cql_d_adv_logits: CudaSlice<f32>,
|
||||
}
|
||||
|
||||
impl Drop for GpuDqnTrainer {
|
||||
@@ -1312,13 +1317,38 @@ impl GpuDqnTrainer {
|
||||
let param_sizes = compute_param_sizes(&self.config);
|
||||
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes);
|
||||
|
||||
// Construct d_adv_logits pointers per branch
|
||||
let f32_size = std::mem::size_of::<half::bf16>();
|
||||
let d_adv_base = d_adv_ptr;
|
||||
// Cast f32 CQL d_logits → bf16 staging for cuBLAS backward GEMM.
|
||||
// Reuse d_value_logits_bf16 / d_adv_logits_bf16 staging buffers (main backward
|
||||
// has already consumed them by the time CQL runs between graph phases).
|
||||
{
|
||||
let total_actions = b0 + b1 + b2;
|
||||
let n_val = (b * na) as i32;
|
||||
let n_adv = (b * total_actions * na) as i32;
|
||||
let val_blocks = ((n_val as u32 + 255) / 256) as u32;
|
||||
let adv_blocks = ((n_adv as u32 + 255) / 256) as u32;
|
||||
let cfg = |blocks: u32| LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
|
||||
let val_dst = self.d_value_logits_bf16.raw_ptr();
|
||||
let adv_dst = self.d_adv_logits_bf16.raw_ptr();
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&d_v_ptr).arg(&val_dst).arg(&n_val)
|
||||
.launch(cfg(val_blocks))
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 cql_d_value: {e}")))?;
|
||||
self.stream.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&d_adv_ptr).arg(&adv_dst).arg(&n_adv)
|
||||
.launch(cfg(adv_blocks))
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 cql_d_adv: {e}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
// Construct d_adv_logits pointers per branch (bf16 staging)
|
||||
let bf16_size = std::mem::size_of::<half::bf16>();
|
||||
let d_val_bf16 = self.d_value_logits_bf16.raw_ptr();
|
||||
let d_adv_bf16_base = self.d_adv_logits_bf16.raw_ptr();
|
||||
let d_adv_ptrs = [
|
||||
d_adv_base,
|
||||
d_adv_base + (b0 * na * f32_size) as u64,
|
||||
d_adv_base + ((b0 + b1) * na * f32_size) as u64,
|
||||
d_adv_bf16_base,
|
||||
d_adv_bf16_base + (b0 * na * bf16_size) as u64,
|
||||
d_adv_bf16_base + ((b0 + b1) * na * bf16_size) as u64,
|
||||
];
|
||||
|
||||
// Saved activations from the forward pass (still valid)
|
||||
@@ -1342,11 +1372,11 @@ impl GpuDqnTrainer {
|
||||
self.stream.memset_zeros(&mut self.cql_grad_scratch)
|
||||
.map_err(|e| MLError::ModelError(format!("zero cql_grad_scratch: {e}")))?;
|
||||
|
||||
// Run full backward pass with CQL logit gradients into ISOLATED scratch buffer.
|
||||
// Run full backward pass with CQL logit gradients (bf16 staging) into ISOLATED scratch buffer.
|
||||
// Produces CQL parameter gradients WITHOUT mixing with C51's grad_buf.
|
||||
self.cublas_backward.backward_full(
|
||||
&self.stream,
|
||||
d_v_ptr,
|
||||
d_val_bf16,
|
||||
&d_adv_ptrs,
|
||||
states_ptr_fw,
|
||||
h_s1_ptr, h_s2_ptr, h_v_ptr,
|
||||
@@ -1759,7 +1789,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) =
|
||||
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) =
|
||||
compile_training_kernels(&stream, &config)?;
|
||||
|
||||
// Separate grad_norm instance for non-graph launches (clip_grad_buf_inplace).
|
||||
@@ -1934,15 +1964,24 @@ impl GpuDqnTrainer {
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let cql_d_value_logits = alloc_bf16(&stream, b * pad32(config.num_atoms), "cql_d_value_logits")?;
|
||||
let cql_d_adv_logits = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "cql_d_adv_logits")?;
|
||||
let cql_d_value_logits = stream.alloc_zeros::<f32>(b * pad32(config.num_atoms))
|
||||
.map_err(|e| MLError::ModelError(format!("alloc cql_d_value_logits f32: {e}")))?;
|
||||
let cql_d_adv_logits = stream.alloc_zeros::<f32>(b * total_branch_atoms + 32 * 3)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc cql_d_adv_logits f32: {e}")))?;
|
||||
|
||||
// ── Gradient output buffers for cuBLAS backward ──────────────
|
||||
let d_value_logits_buf = alloc_bf16(&stream, b * pad32(config.num_atoms), "d_value_logits")?;
|
||||
let d_adv_logits_buf = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "d_adv_logits")?;
|
||||
// ── Gradient output buffers (f32 for native atomicAdd — eliminates bf16 overflow NaN) ─
|
||||
let d_value_logits_buf = stream.alloc_zeros::<f32>(b * pad32(config.num_atoms))
|
||||
.map_err(|e| MLError::ModelError(format!("alloc d_value_logits f32: {e}")))?;
|
||||
let d_adv_logits_buf = stream.alloc_zeros::<f32>(b * total_branch_atoms + 32 * 3)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc d_adv_logits f32: {e}")))?;
|
||||
// Scratch buffers for blended MSE+C51 loss (MSE grad stored here, then blended)
|
||||
let d_value_logits_mse = alloc_bf16(&stream, b * pad32(config.num_atoms), "d_value_logits_mse")?;
|
||||
let d_adv_logits_mse = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "d_adv_logits_mse")?;
|
||||
let d_value_logits_mse = stream.alloc_zeros::<f32>(b * pad32(config.num_atoms))
|
||||
.map_err(|e| MLError::ModelError(format!("alloc d_value_logits_mse f32: {e}")))?;
|
||||
let d_adv_logits_mse = stream.alloc_zeros::<f32>(b * total_branch_atoms + 32 * 3)
|
||||
.map_err(|e| MLError::ModelError(format!("alloc d_adv_logits_mse f32: {e}")))?;
|
||||
// BF16 staging buffers — cast from f32 before cuBLAS backward GEMM
|
||||
let d_value_logits_bf16 = alloc_bf16(&stream, b * pad32(config.num_atoms), "d_value_logits_bf16")?;
|
||||
let d_adv_logits_bf16 = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "d_adv_logits_bf16")?;
|
||||
|
||||
// ── Spectral normalization singular vectors ─────────────────
|
||||
// Initialize with random unit vectors for proper power iteration convergence.
|
||||
@@ -2136,6 +2175,7 @@ impl GpuDqnTrainer {
|
||||
f32_to_bf16_kernel,
|
||||
bf16_to_f32_kernel,
|
||||
saxpy_kernel,
|
||||
saxpy_f32_kernel,
|
||||
zero_kernel,
|
||||
regime_scale_kernel,
|
||||
shrink_perturb_kernel: shrink_perturb,
|
||||
@@ -2229,6 +2269,8 @@ impl GpuDqnTrainer {
|
||||
last_captured_loss_mode: None,
|
||||
d_value_logits_buf,
|
||||
d_adv_logits_buf,
|
||||
d_value_logits_bf16,
|
||||
d_adv_logits_bf16,
|
||||
d_value_logits_mse,
|
||||
d_adv_logits_mse,
|
||||
cublas_backward,
|
||||
@@ -3052,29 +3094,32 @@ impl GpuDqnTrainer {
|
||||
let adv_mse_ptr = self.d_adv_logits_mse.raw_ptr();
|
||||
|
||||
unsafe {
|
||||
// d_value += (α-1) * d_value → d_value *= α
|
||||
self.stream.launch_builder(&self.saxpy_kernel)
|
||||
// d_value += (α-1) * d_value → d_value *= α (f32 SAXPY)
|
||||
self.stream.launch_builder(&self.saxpy_f32_kernel)
|
||||
.arg(&val_ptr).arg(&val_ptr)
|
||||
.arg(&scale_c51).arg(&n_val)
|
||||
.launch(cfg_val).map_err(|e| MLError::ModelError(format!("blend c51 val: {e}")))?;
|
||||
// d_value += (1-α) * mse_scratch
|
||||
self.stream.launch_builder(&self.saxpy_kernel)
|
||||
self.stream.launch_builder(&self.saxpy_f32_kernel)
|
||||
.arg(&val_ptr).arg(&val_mse_ptr)
|
||||
.arg(&scale_mse).arg(&n_val)
|
||||
.launch(cfg_val).map_err(|e| MLError::ModelError(format!("blend mse val: {e}")))?;
|
||||
// d_adv += (α-1) * d_adv → d_adv *= α
|
||||
self.stream.launch_builder(&self.saxpy_kernel)
|
||||
self.stream.launch_builder(&self.saxpy_f32_kernel)
|
||||
.arg(&adv_ptr).arg(&adv_ptr)
|
||||
.arg(&scale_c51).arg(&n_adv)
|
||||
.launch(cfg_adv).map_err(|e| MLError::ModelError(format!("blend c51 adv: {e}")))?;
|
||||
// d_adv += (1-α) * mse_scratch
|
||||
self.stream.launch_builder(&self.saxpy_kernel)
|
||||
self.stream.launch_builder(&self.saxpy_f32_kernel)
|
||||
.arg(&adv_ptr).arg(&adv_mse_ptr)
|
||||
.arg(&scale_mse).arg(&n_adv)
|
||||
.launch(cfg_adv).map_err(|e| MLError::ModelError(format!("blend mse adv: {e}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
// ── 3.5. Cast f32 d_logits → bf16 staging for cuBLAS backward ─
|
||||
self.cast_d_logits_to_bf16()?;
|
||||
|
||||
// ── 4. Backward (cuBLAS SGEMM, chain rule through layers) ─
|
||||
self.launch_cublas_backward()?;
|
||||
|
||||
@@ -3636,8 +3681,8 @@ impl GpuDqnTrainer {
|
||||
/// Writes gradient outputs to the provided destination buffers.
|
||||
fn launch_mse_grad_inner(
|
||||
&self,
|
||||
d_value_dst: &CudaSlice<half::bf16>,
|
||||
d_adv_dst: &CudaSlice<half::bf16>,
|
||||
d_value_dst: &CudaSlice<f32>,
|
||||
d_adv_dst: &CudaSlice<f32>,
|
||||
) -> Result<(), MLError> {
|
||||
let b = self.config.batch_size;
|
||||
let na = self.config.num_atoms;
|
||||
@@ -3684,11 +3729,50 @@ impl GpuDqnTrainer {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Cast f32 d_logits → bf16 staging buffers for cuBLAS backward GEMM.
|
||||
///
|
||||
/// The gradient kernels write to f32 buffers (native atomicAdd, no overflow).
|
||||
/// cuBLAS backward expects bf16 dY inputs for tensor core GEMM. This method
|
||||
/// converts the f32 gradients to bf16 in the staging buffers.
|
||||
fn cast_d_logits_to_bf16(&self) -> Result<(), MLError> {
|
||||
let na = self.config.num_atoms;
|
||||
let b = self.config.batch_size;
|
||||
let b0 = self.config.branch_0_size;
|
||||
let b1 = self.config.branch_1_size;
|
||||
let b2 = self.config.branch_2_size;
|
||||
|
||||
let n_val = (b * pad32(na)) as i32;
|
||||
let n_adv = (b * (b0 + b1 + b2) * na + 32 * 3) as i32;
|
||||
|
||||
let val_src = self.d_value_logits_buf.raw_ptr();
|
||||
let val_dst = self.d_value_logits_bf16.raw_ptr();
|
||||
let adv_src = self.d_adv_logits_buf.raw_ptr();
|
||||
let adv_dst = self.d_adv_logits_bf16.raw_ptr();
|
||||
|
||||
let val_blocks = ((n_val as u32 + 255) / 256) as u32;
|
||||
let adv_blocks = ((n_adv as u32 + 255) / 256) as u32;
|
||||
let cfg_val = LaunchConfig { grid_dim: (val_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
|
||||
let cfg_adv = LaunchConfig { grid_dim: (adv_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
|
||||
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&val_src).arg(&val_dst).arg(&n_val)
|
||||
.launch(cfg_val)
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 d_value_logits: {e}")))?;
|
||||
self.stream.launch_builder(&self.f32_to_bf16_kernel)
|
||||
.arg(&adv_src).arg(&adv_dst).arg(&n_adv)
|
||||
.launch(cfg_adv)
|
||||
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 d_adv_logits: {e}")))?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// cuBLAS SGEMM backward pass: chain rule through all layers.
|
||||
///
|
||||
/// Reads dL/d_logits from `d_value_logits_buf` and `d_adv_logits_buf`
|
||||
/// (populated by `launch_c51_grad` or `launch_mse_grad`), propagates gradients
|
||||
/// through all layers using cuBLAS GEMM, and accumulates into `grad_buf`.
|
||||
/// Reads dL/d_logits from bf16 staging buffers (`d_value_logits_bf16` and
|
||||
/// `d_adv_logits_bf16`, cast from f32 by `cast_d_logits_to_bf16`),
|
||||
/// propagates gradients through all layers using cuBLAS GEMM, and
|
||||
/// accumulates into `grad_buf`.
|
||||
fn launch_cublas_backward(&self) -> Result<(), MLError> {
|
||||
let bw = &self.cublas_backward;
|
||||
|
||||
@@ -3711,15 +3795,15 @@ impl GpuDqnTrainer {
|
||||
let d_h_b1_ptr = bw_raw_ptr(&self.bw_d_h_b1, &self.stream);
|
||||
let d_h_b2_ptr = bw_raw_ptr(&self.bw_d_h_b2, &self.stream);
|
||||
|
||||
// dL/d_logits from c51_grad_kernel
|
||||
let d_value_logits_ptr = bw_raw_ptr(&self.d_value_logits_buf, &self.stream);
|
||||
let d_adv_logits_ptr = bw_raw_ptr(&self.d_adv_logits_buf, &self.stream);
|
||||
// dL/d_logits from bf16 staging (cast from f32 by cast_d_logits_to_bf16)
|
||||
let d_value_logits_ptr = bw_raw_ptr(&self.d_value_logits_bf16, &self.stream);
|
||||
let d_adv_logits_ptr = bw_raw_ptr(&self.d_adv_logits_bf16, &self.stream);
|
||||
|
||||
let na = self.config.num_atoms;
|
||||
let f32_size = std::mem::size_of::<half::bf16>() as u64;
|
||||
let bf16_size = std::mem::size_of::<half::bf16>() as u64;
|
||||
let d_adv0 = d_adv_logits_ptr;
|
||||
let d_adv1 = d_adv0 + (self.config.batch_size * self.config.branch_0_size * na) as u64 * f32_size;
|
||||
let d_adv2 = d_adv1 + (self.config.batch_size * self.config.branch_1_size * na) as u64 * f32_size;
|
||||
let d_adv1 = d_adv0 + (self.config.batch_size * self.config.branch_0_size * na) as u64 * bf16_size;
|
||||
let d_adv2 = d_adv1 + (self.config.batch_size * self.config.branch_1_size * na) as u64 * bf16_size;
|
||||
|
||||
bw.backward_full(
|
||||
&self.stream,
|
||||
@@ -4101,7 +4185,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), MLError> {
|
||||
) -> Result<(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),
|
||||
@@ -4125,6 +4209,8 @@ fn compile_training_kernels(
|
||||
.map_err(|e| MLError::ModelError(format!("bf16_to_f32_kernel load: {e}")))?;
|
||||
let saxpy = module.load_function("dqn_saxpy_kernel")
|
||||
.map_err(|e| MLError::ModelError(format!("dqn_saxpy_kernel load: {e}")))?;
|
||||
let saxpy_f32 = module.load_function("dqn_saxpy_f32_kernel")
|
||||
.map_err(|e| MLError::ModelError(format!("dqn_saxpy_f32_kernel load: {e}")))?;
|
||||
let zero = module.load_function("dqn_zero_kernel")
|
||||
.map_err(|e| MLError::ModelError(format!("dqn_zero_kernel load: {e}")))?;
|
||||
let regime_scale = module.load_function("dqn_regime_scale_kernel")
|
||||
@@ -4142,8 +4228,8 @@ fn compile_training_kernels(
|
||||
let pad_states = module.load_function("pad_states_kernel")
|
||||
.map_err(|e| MLError::ModelError(format!("pad_states_kernel load: {e}")))?;
|
||||
|
||||
info!("GpuDqnTrainer: 13 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))
|
||||
info!("GpuDqnTrainer: 14 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))
|
||||
}
|
||||
|
||||
/// Load the standalone Polyak EMA kernel from precompiled cubin.
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
/**
|
||||
* MSE loss gradient kernel through softmax expectation.
|
||||
*
|
||||
* Mixed-precision: reads BF16, computes in float, writes BF16.
|
||||
* Prevents NaN from bf16 intermediate product overflow.
|
||||
* Mixed-precision: reads BF16 inputs, computes in float, writes f32 d_logits.
|
||||
* f32 atomicAdd eliminates bf16 overflow that caused NaN.
|
||||
*
|
||||
* For each sample [b], branch [d], atom [j]:
|
||||
* d_logit_j = td_error * is_weight * p_j * (z_j - E[Q])
|
||||
@@ -15,8 +15,8 @@ extern "C" __global__ void mse_grad_kernel(
|
||||
const __nv_bfloat16* __restrict__ save_eq_td, // [B, 3, NA] layout: [td_error, E_Q, 0, ...]
|
||||
const __nv_bfloat16* __restrict__ is_weights, // [B] bf16
|
||||
const int* __restrict__ actions, // [B] factored
|
||||
__nv_bfloat16* __restrict__ d_value_logits, // [B, NA]
|
||||
__nv_bfloat16* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA]
|
||||
float* __restrict__ d_value_logits, // [B, NA] f32 (native atomicAdd, no overflow)
|
||||
float* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] f32
|
||||
int batch_size,
|
||||
int num_atoms,
|
||||
int b0_size, int b1_size, int b2_size,
|
||||
@@ -46,10 +46,8 @@ extern "C" __global__ void mse_grad_kernel(
|
||||
float d_combined = isw * td_error * p_j * (z_j - e_q);
|
||||
|
||||
/* Route through dueling: d_value[b,j] += d_combined.
|
||||
* d_value_logits is bf16 — atomicAddBF16 accumulates. Clamp d_combined
|
||||
* to prevent bf16 overflow (3 branches × batch atomicAdds per element). */
|
||||
d_combined = fminf(fmaxf(d_combined, -100.0f), 100.0f);
|
||||
atomicAddBF16(&d_value_logits[b * num_atoms + j], d_combined);
|
||||
* d_value_logits is f32 — native atomicAdd, no overflow risk. */
|
||||
atomicAdd(&d_value_logits[b * num_atoms + j], d_combined);
|
||||
|
||||
/* Factored action decode */
|
||||
int factored = actions[b];
|
||||
@@ -77,6 +75,6 @@ extern "C" __global__ void mse_grad_kernel(
|
||||
float dueling_grad = (a == a_d) ? (1.0f - inv_A) : (-inv_A);
|
||||
float grad_val = d_combined * dueling_grad;
|
||||
int adv_idx = b * total_branch_atoms + branch_offset + a * num_atoms + j;
|
||||
atomicAddBF16(&d_adv_logits[adv_idx], grad_val);
|
||||
atomicAdd(&d_adv_logits[adv_idx], grad_val);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -354,10 +354,7 @@ extern "C" __global__ void mse_loss_batched(
|
||||
|
||||
if (tid == 0) {
|
||||
float weighted_loss = avg_mse * is_weight;
|
||||
/* Guard: bf16 d_logits atomicAddBF16 can overflow despite per-thread
|
||||
* clamping. Root fix: convert d_value_logits/d_adv_logits to f32. */
|
||||
if (!fast_isfinite(weighted_loss)) weighted_loss = 0.0f;
|
||||
if (!fast_isfinite(avg_td)) avg_td = 0.0f;
|
||||
/* f32 d_logits: no NaN risk from atomicAdd overflow. */
|
||||
per_sample_loss[sample_id] = bf16(weighted_loss);
|
||||
td_errors[sample_id] = bf16(avg_td);
|
||||
atomicAdd(total_loss, weighted_loss / (float)batch_size);
|
||||
|
||||
Reference in New Issue
Block a user