diff --git a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu index 76ef5e111..bdbc11b55 100644 --- a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu @@ -1,3 +1,5 @@ +#include + /** * C51 distributional RL loss gradient kernel. * @@ -13,13 +15,36 @@ * Launch config: grid=(ceil(batch_size*3*num_atoms/256), 1, 1), block=(256, 1, 1). */ +/* BF16 atomicAdd via 16-bit CAS loop — no native BF16 atomicAdd on any SM */ +__device__ __forceinline__ void atomicAddBF16(__nv_bfloat16* addr, float val) { + /* Pack the BF16 into the low or high half of a 32-bit word and use + * 32-bit atomicCAS. This handles alignment correctly. */ + unsigned int* base = (unsigned int*)((size_t)addr & ~(size_t)3); + unsigned int shift = ((unsigned int)((size_t)addr & 2)) << 3; /* 0 or 16 */ + unsigned int mask = 0x0000FFFFu << shift; + + unsigned int old_val = *base; + unsigned int assumed; + do { + assumed = old_val; + /* Extract the BF16 bits, convert to float, add, convert back */ + unsigned short old_bits = (unsigned short)((assumed >> shift) & 0xFFFF); + __nv_bfloat16 old_bf16 = *reinterpret_cast<__nv_bfloat16*>(&old_bits); + float old_f = __bfloat162float(old_bf16); + __nv_bfloat16 new_bf16 = __float2bfloat16(old_f + val); + unsigned short new_bits = *reinterpret_cast(&new_bf16); + unsigned int new_word = (assumed & ~mask) | ((unsigned int)new_bits << shift); + old_val = atomicCAS(base, assumed, new_word); + } while (old_val != assumed); +} + extern "C" __global__ void c51_grad_kernel( - const float* __restrict__ current_lp, // [B, 3, NA] - const float* __restrict__ projected, // [B, 3, NA] - const float* __restrict__ is_weights, // [B] + const __nv_bfloat16* __restrict__ current_lp, // [B, 3, NA] + const __nv_bfloat16* __restrict__ projected, // [B, 3, NA] + const __nv_bfloat16* __restrict__ is_weights, // [B] const int* __restrict__ actions, // [B] factored - float* __restrict__ d_value_logits, // [B, NA] - float* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] + __nv_bfloat16* __restrict__ d_value_logits, // [B, NA] + __nv_bfloat16* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] int batch_size, int num_atoms, int b0_size, int b1_size, int b2_size, @@ -34,9 +59,9 @@ extern "C" __global__ void c51_grad_kernel( int d = (tid / num_atoms) % 3; int b = tid / (3 * num_atoms); - float isw = is_weights[b]; - float lp = current_lp[tid]; - float proj = projected[tid]; + float isw = __bfloat162float(is_weights[b]); + float lp = __bfloat162float(current_lp[tid]); + float proj = __bfloat162float(projected[tid]); /* Cross-entropy gradient: d/d_logits(-Sigma proj * lp) = exp(lp) - proj */ float d_combined = isw * (expf(lp) - proj); @@ -52,7 +77,7 @@ extern "C" __global__ void c51_grad_kernel( } // Route through dueling: d_value[b,j] += d_combined - atomicAdd(&d_value_logits[b * num_atoms + j], d_combined); + atomicAddBF16(&d_value_logits[b * num_atoms + j], d_combined); // Factored action decode int factored = actions[b]; @@ -80,6 +105,6 @@ extern "C" __global__ void c51_grad_kernel( for (int a = 0; a < A_d; a++) { float dueling_grad = (a == a_d) ? (1.0f - inv_A) : (-inv_A); int adv_idx = b * total_branch_atoms + branch_offset + a * num_atoms + j; - atomicAdd(&d_adv_logits[adv_idx], d_combined * dueling_grad); + atomicAddBF16(&d_adv_logits[adv_idx], d_combined * dueling_grad); } } diff --git a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu index ca7b70af5..ec92e814c 100644 --- a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu @@ -265,37 +265,37 @@ __device__ void block_bellman_project( extern "C" __global__ void c51_loss_batched( /* ── Online network outputs on current STATES ─────────────────── */ - const float* __restrict__ on_value_logits, /* [B, num_atoms] */ - const float* __restrict__ on_adv_logits_b0, /* [B, b0_size * num_atoms] */ - const float* __restrict__ on_adv_logits_b1, /* [B, b1_size * num_atoms] */ - const float* __restrict__ on_adv_logits_b2, /* [B, b2_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_value_logits, /* [B, num_atoms] */ + const __nv_bfloat16* __restrict__ on_adv_logits_b0, /* [B, b0_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_adv_logits_b1, /* [B, b1_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_adv_logits_b2, /* [B, b2_size * num_atoms] */ /* ── Target network outputs on NEXT_STATES ────────────────────── */ - const float* __restrict__ tg_value_logits, /* [B, num_atoms] */ - const float* __restrict__ tg_adv_logits_b0, /* [B, b0_size * num_atoms] */ - const float* __restrict__ tg_adv_logits_b1, /* [B, b1_size * num_atoms] */ - const float* __restrict__ tg_adv_logits_b2, /* [B, b2_size * num_atoms] */ + const __nv_bfloat16* __restrict__ tg_value_logits, /* [B, num_atoms] */ + const __nv_bfloat16* __restrict__ tg_adv_logits_b0, /* [B, b0_size * num_atoms] */ + const __nv_bfloat16* __restrict__ tg_adv_logits_b1, /* [B, b1_size * num_atoms] */ + const __nv_bfloat16* __restrict__ tg_adv_logits_b2, /* [B, b2_size * num_atoms] */ /* ── Online network outputs on NEXT_STATES (Double DQN selector) */ - const float* __restrict__ on_next_value_logits, /* [B, num_atoms] */ - const float* __restrict__ on_next_adv_logits_b0, /* [B, b0_size * num_atoms] */ - const float* __restrict__ on_next_adv_logits_b1, /* [B, b1_size * num_atoms] */ - const float* __restrict__ on_next_adv_logits_b2, /* [B, b2_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_value_logits, /* [B, num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_adv_logits_b0, /* [B, b0_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_adv_logits_b1, /* [B, b1_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_adv_logits_b2, /* [B, b2_size * num_atoms] */ /* ── Batch data ───────────────────────────────────────────────── */ const int* __restrict__ actions, /* [B] factored action indices 0-44 */ - const float* __restrict__ rewards, /* [B] */ - const float* __restrict__ dones, /* [B] */ - const float* __restrict__ is_weights, /* [B] PER importance-sampling weights */ + const __nv_bfloat16* __restrict__ rewards, /* [B] */ + const __nv_bfloat16* __restrict__ dones, /* [B] */ + const __nv_bfloat16* __restrict__ is_weights, /* [B] PER importance-sampling weights */ /* ── Outputs ──────────────────────────────────────────────────── */ - float* __restrict__ per_sample_loss, /* [B] IS-weighted loss per sample */ - float* __restrict__ td_errors, /* [B] unweighted, for PER priority update */ - float* __restrict__ total_loss, /* [1] batch mean loss (atomicAdd) */ + __nv_bfloat16* __restrict__ per_sample_loss, /* [B] IS-weighted loss per sample */ + __nv_bfloat16* __restrict__ td_errors, /* [B] unweighted, for PER priority update */ + float* __restrict__ total_loss, /* [1] batch mean loss (atomicAdd) — stays F32, no BF16 atomicAdd */ /* ── Saved tensors for backward pass ─────────────────────────── */ - float* __restrict__ save_current_lp, /* [B, NUM_BRANCHES, num_atoms] */ - float* __restrict__ save_projected, /* [B, NUM_BRANCHES, num_atoms] */ + __nv_bfloat16* __restrict__ save_current_lp, /* [B, NUM_BRANCHES, num_atoms] */ + __nv_bfloat16* __restrict__ save_projected, /* [B, NUM_BRANCHES, num_atoms] */ /* ── Config ───────────────────────────────────────────────────── */ float gamma, @@ -360,30 +360,30 @@ extern "C" __global__ void c51_loss_batched( const int branch_sizes[NUM_BRANCHES] = { b0_size, b1_size, b2_size }; /* Per-sample row pointers into each branch's advantage logit tensors */ - const float* on_adv_row[NUM_BRANCHES] = { + const __nv_bfloat16* on_adv_row[NUM_BRANCHES] = { on_adv_logits_b0 + (long long)sample_id * b0_atoms, on_adv_logits_b1 + (long long)sample_id * b1_atoms, on_adv_logits_b2 + (long long)sample_id * b2_atoms }; - const float* tg_adv_row[NUM_BRANCHES] = { + const __nv_bfloat16* tg_adv_row[NUM_BRANCHES] = { tg_adv_logits_b0 + (long long)sample_id * b0_atoms, tg_adv_logits_b1 + (long long)sample_id * b1_atoms, tg_adv_logits_b2 + (long long)sample_id * b2_atoms }; - const float* on_next_adv_row[NUM_BRANCHES] = { + const __nv_bfloat16* on_next_adv_row[NUM_BRANCHES] = { on_next_adv_logits_b0 + (long long)sample_id * b0_atoms, on_next_adv_logits_b1 + (long long)sample_id * b1_atoms, on_next_adv_logits_b2 + (long long)sample_id * b2_atoms }; /* Value logit row pointers (same num_atoms for all branches) */ - const float* on_val_row = on_value_logits + (long long)sample_id * num_atoms; - const float* tg_val_row = tg_value_logits + (long long)sample_id * num_atoms; - const float* on_next_val_row = on_next_value_logits + (long long)sample_id * num_atoms; + const __nv_bfloat16* on_val_row = on_value_logits + (long long)sample_id * num_atoms; + const __nv_bfloat16* tg_val_row = tg_value_logits + (long long)sample_id * num_atoms; + const __nv_bfloat16* on_next_val_row = on_next_value_logits + (long long)sample_id * num_atoms; - float reward = rewards[sample_id]; - float done = dones[sample_id]; - float is_weight = is_weights[sample_id]; + float reward = __bfloat162float(rewards[sample_id]); + float done = __bfloat162float(dones[sample_id]); + float is_weight = __bfloat162float(is_weights[sample_id]); float total_ce = 0.0f; @@ -411,14 +411,14 @@ extern "C" __global__ void c51_loss_batched( * (online network, current states) * ═══════════════════════════════════════════════════════════ */ - /* Load value logits */ + /* Load value logits (BF16 → F32) */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - shmem_val[j] = on_val_row[j]; + shmem_val[j] = __bfloat162float(on_val_row[j]); __syncthreads(); - /* Load advantage logits for all n_d actions */ + /* Load advantage logits for all n_d actions (BF16 → F32) */ for (int i = tid; i < n_atoms; i += BLOCK_THREADS) - shmem_adv[i] = on_adv_row[d][i]; + shmem_adv[i] = __bfloat162float(on_adv_row[d][i]); __syncthreads(); /* Dueling for action a_d: Q[j] = V[j] + A[a_d,j] - mean_a(A[*,j]) */ @@ -439,9 +439,9 @@ extern "C" __global__ void c51_loss_batched( for (int j = tid; j < num_atoms; j += BLOCK_THREADS) shmem_current_lp[j] = shmem_lp[j]; - /* Save for backward kernel */ + /* Save for backward kernel (F32 → BF16) */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - save_current_lp[save_off + j] = shmem_lp[j]; + save_current_lp[save_off + j] = __float2bfloat16(shmem_lp[j]); __syncthreads(); /* ═══════════════════════════════════════════════════════════ @@ -451,13 +451,13 @@ extern "C" __global__ void c51_loss_batched( * then expected Q. Select action with highest expected Q. * ═══════════════════════════════════════════════════════════ */ - /* Load online_next value and advantage logits */ + /* Load online_next value and advantage logits (BF16 → F32) */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - shmem_val[j] = on_next_val_row[j]; + shmem_val[j] = __bfloat162float(on_next_val_row[j]); __syncthreads(); for (int i = tid; i < n_atoms; i += BLOCK_THREADS) - shmem_adv[i] = on_next_adv_row[d][i]; + shmem_adv[i] = __bfloat162float(on_next_adv_row[d][i]); __syncthreads(); /* Compute average advantage across actions for dueling (same for all a) */ @@ -507,13 +507,13 @@ extern "C" __global__ void c51_loss_batched( * (target network, next states) * ═══════════════════════════════════════════════════════════ */ - /* Load target value and advantage logits */ + /* Load target value and advantage logits (BF16 → F32) */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - shmem_val[j] = tg_val_row[j]; + shmem_val[j] = __bfloat162float(tg_val_row[j]); __syncthreads(); for (int i = tid; i < n_atoms; i += BLOCK_THREADS) - shmem_adv[i] = tg_adv_row[d][i]; + shmem_adv[i] = __bfloat162float(tg_adv_row[d][i]); __syncthreads(); /* Per-atom mean advantage for target network */ @@ -566,9 +566,9 @@ extern "C" __global__ void c51_loss_batched( /* shmem_lp[0..num_atoms] = smoothed projected target distribution */ - /* Save projected for backward kernel */ + /* Save projected for backward kernel (F32 → BF16) */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - save_projected[save_off + j] = shmem_lp[j]; + save_projected[save_off + j] = __float2bfloat16(shmem_lp[j]); __syncthreads(); /* ═══════════════════════════════════════════════════════════ @@ -598,8 +598,8 @@ extern "C" __global__ void c51_loss_batched( * so clamping the LOSS only affects PER priority, not gradient direction. */ float clamped_ce = fminf(avg_ce, MAX_PER_SAMPLE_CE); float weighted_loss = clamped_ce * is_weight; - per_sample_loss[sample_id] = weighted_loss; - td_errors[sample_id] = clamped_ce; /* PER sees clamped too */ + per_sample_loss[sample_id] = __float2bfloat16(weighted_loss); + td_errors[sample_id] = __float2bfloat16(clamped_ce); /* PER sees clamped too */ atomicAdd(total_loss, weighted_loss / (float)batch_size); } } diff --git a/crates/ml/src/cuda_pipeline/expected_q_kernel.cu b/crates/ml/src/cuda_pipeline/expected_q_kernel.cu index 7ee0168b1..daa173f9e 100644 --- a/crates/ml/src/cuda_pipeline/expected_q_kernel.cu +++ b/crates/ml/src/cuda_pipeline/expected_q_kernel.cu @@ -1,3 +1,5 @@ +#include + /** * Expected Q-value kernel for ad-hoc validation forward pass. * @@ -8,9 +10,9 @@ */ extern "C" __global__ void compute_expected_q( - const float* __restrict__ v_logits, // [N, num_atoms] - const float* __restrict__ b_logits, // [N, (b0+b1+b2)*num_atoms] - float* __restrict__ q_values, // [N, b0+b1+b2] + const __nv_bfloat16* __restrict__ v_logits, // [N, num_atoms] + const __nv_bfloat16* __restrict__ b_logits, // [N, (b0+b1+b2)*num_atoms] + __nv_bfloat16* __restrict__ q_values, // [N, b0+b1+b2] int N, int num_atoms, int b0_size, int b1_size, int b2_size, float v_min, float v_max) @@ -22,7 +24,7 @@ extern "C" __global__ void compute_expected_q( float dz = (num_atoms > 1) ? (v_max - v_min) / (float)(num_atoms - 1) : 0.0f; // Value logits for this sample: [num_atoms] - const float* val = v_logits + (long long)i * num_atoms; + const __nv_bfloat16* val = v_logits + (long long)i * num_atoms; int branch_sizes[3]; branch_sizes[0] = b0_size; @@ -34,16 +36,16 @@ extern "C" __global__ void compute_expected_q( for (int d = 0; d < 3; d++) { int bd = branch_sizes[d]; for (int a = 0; a < bd; a++) { - const float* adv = b_logits + (long long)i * total_actions * num_atoms + const __nv_bfloat16* adv = b_logits + (long long)i * total_actions * num_atoms + (long long)(adv_offset + a) * num_atoms; // Compute mean advantage logit sum for this branch (for dueling centering) float mean_adv_sum = 0.0f; for (int aa = 0; aa < bd; aa++) { - const float* adv_aa = b_logits + (long long)i * total_actions * num_atoms + const __nv_bfloat16* adv_aa = b_logits + (long long)i * total_actions * num_atoms + (long long)(adv_offset + aa) * num_atoms; for (int j = 0; j < num_atoms; j++) { - mean_adv_sum += adv_aa[j]; + mean_adv_sum += __bfloat162float(adv_aa[j]); } } float mean_adv_per_atom = mean_adv_sum / (float)(bd * num_atoms); @@ -51,12 +53,12 @@ extern "C" __global__ void compute_expected_q( // Numerically stable log_softmax over combined = val[j] + adv[j] - mean_adv_per_atom float max_logit = -1e30f; for (int j = 0; j < num_atoms; j++) { - float combined = val[j] + adv[j] - mean_adv_per_atom; + float combined = __bfloat162float(val[j]) + __bfloat162float(adv[j]) - mean_adv_per_atom; if (combined > max_logit) max_logit = combined; } float sum_exp = 0.0f; for (int j = 0; j < num_atoms; j++) { - float combined = val[j] + adv[j] - mean_adv_per_atom; + float combined = __bfloat162float(val[j]) + __bfloat162float(adv[j]) - mean_adv_per_atom; sum_exp += expf(combined - max_logit); } float log_sum = logf(sum_exp + 1e-8f) + max_logit; @@ -64,13 +66,13 @@ extern "C" __global__ void compute_expected_q( // Expected Q = sum_j( softmax_j * z_j ) float eq = 0.0f; for (int j = 0; j < num_atoms; j++) { - float combined = val[j] + adv[j] - mean_adv_per_atom; + float combined = __bfloat162float(val[j]) + __bfloat162float(adv[j]) - mean_adv_per_atom; float p = expf(combined - log_sum); float z = v_min + (float)j * dz; eq += p * z; } - q_values[(long long)i * total_actions + q_offset + a] = eq; + q_values[(long long)i * total_actions + q_offset + a] = __float2bfloat16(eq); } adv_offset += bd; q_offset += bd; diff --git a/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu b/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu index 9b128b743..31e7ae75f 100644 --- a/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu @@ -1,3 +1,5 @@ +#include + /** * MSE loss gradient kernel through softmax expectation. * @@ -13,13 +15,33 @@ * Launch config: grid=(ceil(batch_size*3*num_atoms/256), 1, 1), block=(256, 1, 1). */ +/* BF16 atomicAdd via 16-bit CAS loop — no native BF16 atomicAdd on any SM */ +__device__ __forceinline__ void atomicAddBF16(__nv_bfloat16* addr, float val) { + unsigned int* base = (unsigned int*)((size_t)addr & ~(size_t)3); + unsigned int shift = ((unsigned int)((size_t)addr & 2)) << 3; /* 0 or 16 */ + unsigned int mask = 0x0000FFFFu << shift; + + unsigned int old_val = *base; + unsigned int assumed; + do { + assumed = old_val; + unsigned short old_bits = (unsigned short)((assumed >> shift) & 0xFFFF); + __nv_bfloat16 old_bf16 = *reinterpret_cast<__nv_bfloat16*>(&old_bits); + float old_f = __bfloat162float(old_bf16); + __nv_bfloat16 new_bf16 = __float2bfloat16(old_f + val); + unsigned short new_bits = *reinterpret_cast(&new_bf16); + unsigned int new_word = (assumed & ~mask) | ((unsigned int)new_bits << shift); + old_val = atomicCAS(base, assumed, new_word); + } while (old_val != assumed); +} + extern "C" __global__ void mse_grad_kernel( - const float* __restrict__ save_probs, // [B, 3, NA] softmax probs - const float* __restrict__ save_eq_td, // [B, 3, NA] layout: [td_error, E_Q, 0, ...] - const float* __restrict__ is_weights, // [B] + const __nv_bfloat16* __restrict__ save_probs, // [B, 3, NA] softmax probs + const __nv_bfloat16* __restrict__ save_eq_td, // [B, 3, NA] layout: [td_error, E_Q, 0, ...] + const __nv_bfloat16* __restrict__ is_weights, // [B] const int* __restrict__ actions, // [B] factored - float* __restrict__ d_value_logits, // [B, NA] - float* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] + __nv_bfloat16* __restrict__ d_value_logits, // [B, NA] + __nv_bfloat16* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] int batch_size, int num_atoms, int b0_size, int b1_size, int b2_size, @@ -34,14 +56,14 @@ extern "C" __global__ void mse_grad_kernel( int d = (tid / num_atoms) % 3; int b = tid / (3 * num_atoms); - float isw = is_weights[b]; - float p_j = save_probs[tid]; /* softmax prob for atom j */ + float isw = __bfloat162float(is_weights[b]); + float p_j = __bfloat162float(save_probs[tid]); /* softmax prob for atom j */ /* Read td_error and E[Q] from save_eq_td. * Layout: for each [b, d, :], element [0] = td_error, element [1] = E[Q] */ int base = (b * 3 + d) * num_atoms; - float td_error = save_eq_td[base + 0]; - float e_q = save_eq_td[base + 1]; + float td_error = __bfloat162float(save_eq_td[base + 0]); + float e_q = __bfloat162float(save_eq_td[base + 1]); /* C51 support: z_j = v_min + j * delta_z */ float delta_z = (num_atoms > 1) ? (v_max - v_min) / (float)(num_atoms - 1) : 0.0f; @@ -56,7 +78,7 @@ 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 - atomicAdd(&d_value_logits[b * num_atoms + j], d_combined); + atomicAddBF16(&d_value_logits[b * num_atoms + j], d_combined); // Factored action decode int factored = actions[b]; @@ -84,6 +106,6 @@ extern "C" __global__ void mse_grad_kernel( for (int a = 0; a < A_d; a++) { float dueling_grad = (a == a_d) ? (1.0f - inv_A) : (-inv_A); int adv_idx = b * total_branch_atoms + branch_offset + a * num_atoms + j; - atomicAdd(&d_adv_logits[adv_idx], d_combined * dueling_grad); + atomicAddBF16(&d_adv_logits[adv_idx], d_combined * dueling_grad); } } diff --git a/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu b/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu index ebd0e35b0..12a3e5be4 100644 --- a/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu @@ -73,7 +73,7 @@ __device__ float block_softmax_expected_q( float* __restrict__ shmem_logits, /* [num_atoms] in shmem (modified in-place) */ const float* __restrict__ shmem_support, /* [num_atoms] z_j */ float* __restrict__ shmem_reduce, /* [NUM_WARPS] scratch */ - float* __restrict__ save_probs_out, /* [num_atoms] global mem or NULL */ + __nv_bfloat16* __restrict__ save_probs_out, /* [num_atoms] global mem (BF16) or NULL */ int tid, int num_atoms ) { @@ -96,7 +96,7 @@ __device__ float block_softmax_expected_q( float lp = shmem_logits[i] - block_max - log_sum_exp; float prob = expf(lp); shmem_logits[i] = prob; /* overwrite logits with probs */ - if (save_probs_out) save_probs_out[i] = prob; + if (save_probs_out) save_probs_out[i] = __float2bfloat16(prob); /* F32 → BF16 */ } __syncthreads(); @@ -124,37 +124,37 @@ __device__ float block_softmax_expected_q( extern "C" __global__ void mse_loss_batched( /* ── Online network outputs on current STATES ─────────────────── */ - const float* __restrict__ on_value_logits, /* [B, num_atoms] */ - const float* __restrict__ on_adv_logits_b0, /* [B, b0_size * num_atoms] */ - const float* __restrict__ on_adv_logits_b1, /* [B, b1_size * num_atoms] */ - const float* __restrict__ on_adv_logits_b2, /* [B, b2_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_value_logits, /* [B, num_atoms] */ + const __nv_bfloat16* __restrict__ on_adv_logits_b0, /* [B, b0_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_adv_logits_b1, /* [B, b1_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_adv_logits_b2, /* [B, b2_size * num_atoms] */ /* ── Target network outputs on NEXT_STATES ────────────────────── */ - const float* __restrict__ tg_value_logits, /* [B, num_atoms] */ - const float* __restrict__ tg_adv_logits_b0, /* [B, b0_size * num_atoms] */ - const float* __restrict__ tg_adv_logits_b1, /* [B, b1_size * num_atoms] */ - const float* __restrict__ tg_adv_logits_b2, /* [B, b2_size * num_atoms] */ + const __nv_bfloat16* __restrict__ tg_value_logits, /* [B, num_atoms] */ + const __nv_bfloat16* __restrict__ tg_adv_logits_b0, /* [B, b0_size * num_atoms] */ + const __nv_bfloat16* __restrict__ tg_adv_logits_b1, /* [B, b1_size * num_atoms] */ + const __nv_bfloat16* __restrict__ tg_adv_logits_b2, /* [B, b2_size * num_atoms] */ /* ── Online network outputs on NEXT_STATES (Double DQN selector) */ - const float* __restrict__ on_next_value_logits, /* [B, num_atoms] */ - const float* __restrict__ on_next_adv_logits_b0, /* [B, b0_size * num_atoms] */ - const float* __restrict__ on_next_adv_logits_b1, /* [B, b1_size * num_atoms] */ - const float* __restrict__ on_next_adv_logits_b2, /* [B, b2_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_value_logits, /* [B, num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_adv_logits_b0, /* [B, b0_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_adv_logits_b1, /* [B, b1_size * num_atoms] */ + const __nv_bfloat16* __restrict__ on_next_adv_logits_b2, /* [B, b2_size * num_atoms] */ /* ── Batch data ───────────────────────────────────────────────── */ const int* __restrict__ actions, /* [B] factored action indices 0-44 */ - const float* __restrict__ rewards, /* [B] */ - const float* __restrict__ dones, /* [B] */ - const float* __restrict__ is_weights, /* [B] PER importance-sampling weights */ + const __nv_bfloat16* __restrict__ rewards, /* [B] */ + const __nv_bfloat16* __restrict__ dones, /* [B] */ + const __nv_bfloat16* __restrict__ is_weights, /* [B] PER importance-sampling weights */ /* ── Outputs ──────────────────────────────────────────────────── */ - float* __restrict__ per_sample_loss, /* [B] IS-weighted loss per sample */ - float* __restrict__ td_errors, /* [B] unweighted, for PER priority update */ - float* __restrict__ total_loss, /* [1] batch mean loss (atomicAdd) */ + __nv_bfloat16* __restrict__ per_sample_loss, /* [B] IS-weighted loss per sample */ + __nv_bfloat16* __restrict__ td_errors, /* [B] unweighted, for PER priority update */ + float* __restrict__ total_loss, /* [1] batch mean loss (atomicAdd) — stays F32, no BF16 atomicAdd */ /* ── Saved tensors for backward pass ─────────────────────────── */ - float* __restrict__ save_current_lp, /* [B, NUM_BRANCHES, num_atoms] online probs */ - float* __restrict__ save_projected, /* [B, NUM_BRANCHES, num_atoms] [td, E_Q, 0..] */ + __nv_bfloat16* __restrict__ save_current_lp, /* [B, NUM_BRANCHES, num_atoms] online probs */ + __nv_bfloat16* __restrict__ save_projected, /* [B, NUM_BRANCHES, num_atoms] [td, E_Q, 0..] */ /* ── Config ───────────────────────────────────────────────────── */ float gamma, @@ -215,29 +215,29 @@ extern "C" __global__ void mse_loss_batched( const int branch_action[NUM_BRANCHES] = { a0, a1, a2 }; const int branch_sizes[NUM_BRANCHES] = { b0_size, b1_size, b2_size }; - const float* on_adv_row[NUM_BRANCHES] = { + const __nv_bfloat16* on_adv_row[NUM_BRANCHES] = { on_adv_logits_b0 + (long long)sample_id * b0_atoms, on_adv_logits_b1 + (long long)sample_id * b1_atoms, on_adv_logits_b2 + (long long)sample_id * b2_atoms }; - const float* tg_adv_row[NUM_BRANCHES] = { + const __nv_bfloat16* tg_adv_row[NUM_BRANCHES] = { tg_adv_logits_b0 + (long long)sample_id * b0_atoms, tg_adv_logits_b1 + (long long)sample_id * b1_atoms, tg_adv_logits_b2 + (long long)sample_id * b2_atoms }; - const float* on_next_adv_row[NUM_BRANCHES] = { + const __nv_bfloat16* on_next_adv_row[NUM_BRANCHES] = { on_next_adv_logits_b0 + (long long)sample_id * b0_atoms, on_next_adv_logits_b1 + (long long)sample_id * b1_atoms, on_next_adv_logits_b2 + (long long)sample_id * b2_atoms }; - const float* on_val_row = on_value_logits + (long long)sample_id * num_atoms; - const float* tg_val_row = tg_value_logits + (long long)sample_id * num_atoms; - const float* on_next_val_row = on_next_value_logits + (long long)sample_id * num_atoms; + const __nv_bfloat16* on_val_row = on_value_logits + (long long)sample_id * num_atoms; + const __nv_bfloat16* tg_val_row = tg_value_logits + (long long)sample_id * num_atoms; + const __nv_bfloat16* on_next_val_row = on_next_value_logits + (long long)sample_id * num_atoms; - float reward = rewards[sample_id]; - float done = dones[sample_id]; - float is_weight = is_weights[sample_id]; + float reward = __bfloat162float(rewards[sample_id]); + float done = __bfloat162float(dones[sample_id]); + float is_weight = __bfloat162float(is_weights[sample_id]); float total_mse = 0.0f; float total_abs_td = 0.0f; @@ -253,14 +253,14 @@ extern "C" __global__ void mse_loss_batched( * STEP a: Online E[Q] for the taken action a_d * ═══════════════════════════════════════════════════════════ */ - /* Load online value logits */ + /* Load online value logits (BF16 → F32) */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - shmem_val[j] = on_val_row[j]; + shmem_val[j] = __bfloat162float(on_val_row[j]); __syncthreads(); - /* Load advantage logits */ + /* Load advantage logits (BF16 → F32) */ for (int i = tid; i < n_atoms; i += BLOCK_THREADS) - shmem_adv[i] = on_adv_row[d][i]; + shmem_adv[i] = __bfloat162float(on_adv_row[d][i]); __syncthreads(); /* Dueling for action a_d: Q[j] = V[j] + A[a_d,j] - mean_a(A[*,j]) */ @@ -284,11 +284,11 @@ extern "C" __global__ void mse_loss_batched( * ═══════════════════════════════════════════════════════════ */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - shmem_val[j] = on_next_val_row[j]; + shmem_val[j] = __bfloat162float(on_next_val_row[j]); __syncthreads(); for (int i = tid; i < n_atoms; i += BLOCK_THREADS) - shmem_adv[i] = on_next_adv_row[d][i]; + shmem_adv[i] = __bfloat162float(on_next_adv_row[d][i]); __syncthreads(); /* Per-atom mean advantage for dueling */ @@ -311,7 +311,7 @@ extern "C" __global__ void mse_loss_batched( float eq = block_softmax_expected_q( shmem_lp, shmem_support, shmem_reduce, - (float*)0, tid, num_atoms /* no save needed for action selection */ + (__nv_bfloat16*)0, tid, num_atoms /* no save needed for action selection */ ); if (tid == 0) eq_per_action[a] = eq; __syncthreads(); @@ -331,11 +331,11 @@ extern "C" __global__ void mse_loss_batched( * ═══════════════════════════════════════════════════════════ */ for (int j = tid; j < num_atoms; j += BLOCK_THREADS) - shmem_val[j] = tg_val_row[j]; + shmem_val[j] = __bfloat162float(tg_val_row[j]); __syncthreads(); for (int i = tid; i < n_atoms; i += BLOCK_THREADS) - shmem_adv[i] = tg_adv_row[d][i]; + shmem_adv[i] = __bfloat162float(tg_adv_row[d][i]); __syncthreads(); for (int j = tid; j < num_atoms; j += BLOCK_THREADS) { @@ -349,7 +349,7 @@ extern "C" __global__ void mse_loss_batched( float target_eq = block_softmax_expected_q( shmem_lp, shmem_support, shmem_reduce, - (float*)0, tid, num_atoms + (__nv_bfloat16*)0, tid, num_atoms ); /* ═══════════════════════════════════════════════════════════ @@ -367,12 +367,12 @@ extern "C" __global__ void mse_loss_batched( * save_projected[save_off + 1] = E[Q_online] * save_projected[save_off + 2..] = 0 (padding) */ if (tid == 0) { - save_projected[save_off + 0] = td; - save_projected[save_off + 1] = online_eq; + save_projected[save_off + 0] = __float2bfloat16(td); + save_projected[save_off + 1] = __float2bfloat16(online_eq); } /* Zero padding */ for (int j = tid + 2; j < num_atoms; j += BLOCK_THREADS) - save_projected[save_off + j] = 0.0f; + save_projected[save_off + j] = __float2bfloat16(0.0f); __syncthreads(); } /* end branch loop */ @@ -383,8 +383,8 @@ extern "C" __global__ void mse_loss_batched( if (tid == 0) { float weighted_loss = avg_mse * is_weight; - per_sample_loss[sample_id] = weighted_loss; - td_errors[sample_id] = avg_td; + per_sample_loss[sample_id] = __float2bfloat16(weighted_loss); + td_errors[sample_id] = __float2bfloat16(avg_td); atomicAdd(total_loss, weighted_loss / (float)batch_size); } } diff --git a/crates/ml/src/cuda_pipeline/q_stats_kernel.cu b/crates/ml/src/cuda_pipeline/q_stats_kernel.cu index 10a253956..fac96a783 100644 --- a/crates/ml/src/cuda_pipeline/q_stats_kernel.cu +++ b/crates/ml/src/cuda_pipeline/q_stats_kernel.cu @@ -1,3 +1,5 @@ +#include + /** * Q-value statistics reduction kernel. * @@ -8,8 +10,8 @@ */ extern "C" __global__ void q_stats_reduce( - const float* __restrict__ q_values, // [N, total_actions] - float* __restrict__ out, // [5]: avg_max_q, q_min, q_max, q_mean, q_var + const __nv_bfloat16* __restrict__ q_values, // [N, total_actions] + float* __restrict__ out, // [5]: avg_max_q, q_min, q_max, q_mean, q_var — stays F32 (monitoring scalars) int N, int total_actions) { @@ -26,7 +28,7 @@ extern "C" __global__ void q_stats_reduce( for (int i = 0; i < N; i++) { float row_max = -1e30f; for (int a = 0; a < total_actions; a++) { - float v = q_values[i * total_actions + a]; + float v = __bfloat162float(q_values[i * total_actions + a]); if (v < global_min) global_min = v; if (v > global_max) global_max = v; if (v > row_max) row_max = v; @@ -38,7 +40,7 @@ extern "C" __global__ void q_stats_reduce( float mean = (total > 0) ? global_sum / (float)total : 0.0f; float var_sum = 0.0f; for (int i = 0; i < total; i++) { - float d = q_values[i] - mean; + float d = __bfloat162float(q_values[i]) - mean; var_sum += d * d; } float variance = (total > 0) ? var_sum / (float)total : 0.0f;