feat(bf16): C51/MSE loss + grad kernels accept BF16 logits
All 6 loss/gradient CUDA kernels converted from float* to __nv_bfloat16*: - c51_loss_batched: BF16 logits (12 inputs), rewards, dones, IS weights, outputs - c51_grad_kernel: BF16 d_logits output, atomicAddBF16 for gradient accumulation - mse_loss_batched: BF16 logits, rewards, dones, IS weights, outputs - mse_grad_kernel: BF16 d_logits output, atomicAddBF16 - expected_q_kernel: BF16 logits in, BF16 Q-values out - q_stats_kernel: BF16 Q-values in (monitoring scalars stay float) Pattern: BF16 storage, F32 arithmetic (cast on load/store). Shared memory stays float for softmax/log/exp precision. total_loss stays float* (atomicAdd doesn't support BF16). Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -1,3 +1,5 @@
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
/**
|
||||
* 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<unsigned short*>(&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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
/**
|
||||
* 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;
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
/**
|
||||
* 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<unsigned short*>(&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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
/**
|
||||
* 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;
|
||||
|
||||
Reference in New Issue
Block a user