From dcd851a62d5d042312ef109a11ace508bdc8a291 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 26 May 2026 17:17:56 +0200 Subject: [PATCH] perf(rl): cuBLAS SGEMM for DQN/IQN + Bellman/outcome kernel fusion MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit DQN: replace hand-written fwd/bwd/grad (27.5% GPU time) with cuBLAS SGEMM tensor-core operations. Eliminates 30MiB per-batch grad scratch. IQN: replace embedding matmul + output projection with cuBLAS SGEMM. Pre-allocate scratch buffers for graph-capture safety. Fix device_ptr → raw_ptr in helpers. Convert remaining 6 launch_builder to raw_launch. Bellman: fuse select_action_atoms + bellman_target_projection into single kernel with shared-memory intermediate. Saves 424 launches. Outcome: fuse fwd + CE + bwd into single kernel with shared-memory logits/grad. Saves 848 launches. Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/build.rs | 1 + .../cuda/bellman_target_projection.cu | 121 ++++ crates/ml-alpha/cuda/rl_iqn_backward.cu | 156 +++-- crates/ml-alpha/cuda/rl_iqn_forward.cu | 227 +++---- crates/ml-alpha/cuda/rl_outcome_fused.cu | 125 ++++ crates/ml-alpha/src/rl/dqn.rs | 424 ++++++++++-- crates/ml-alpha/src/rl/iqn.rs | 614 +++++++++++++----- crates/ml-alpha/src/rl/outcome_head.rs | 70 ++ crates/ml-alpha/src/trainer/integrated.rs | 117 +--- 9 files changed, 1362 insertions(+), 493 deletions(-) create mode 100644 crates/ml-alpha/cuda/rl_outcome_fused.cu diff --git a/crates/ml-alpha/build.rs b/crates/ml-alpha/build.rs index a9e43ca65..580e2e4a7 100644 --- a/crates/ml-alpha/build.rs +++ b/crates/ml-alpha/build.rs @@ -121,6 +121,7 @@ const KERNELS: &[&str] = &[ "rl_outcome_ce", // Outcome aux: softmax CE loss + gradient (masked by sentinel -1) "rl_outcome_label", // Outcome aux: assign labels from reward/done (Profit/Timeout/Loss) "rl_outcome_bwd", // Outcome aux: backward through linear layer → grad_W/b/h_t + "rl_outcome_fused", // Outcome aux: fused fwd + CE + bwd — eliminates 2 global round-trips (logits, grad_logits kept in smem) "rl_curriculum_weights", // E8: per-segment difficulty-weighted softmax from Sharpe → PER weights "rl_adversarial_boost", // Adversarial: boost PER priority for negative-reward transitions "rl_outcome_bwd", // Outcome aux: single linear layer backward — dW (per-batch), db (per-batch), dh_t (per-batch); 1 block per batch, 128 threads diff --git a/crates/ml-alpha/cuda/bellman_target_projection.cu b/crates/ml-alpha/cuda/bellman_target_projection.cu index aba1f4896..e5db6bad0 100644 --- a/crates/ml-alpha/cuda/bellman_target_projection.cu +++ b/crates/ml-alpha/cuda/bellman_target_projection.cu @@ -89,6 +89,127 @@ extern "C" __global__ void dqn_select_action_atoms( } +// ───────────────────────────────────────────────────────────────────── +// bellman_fused_select_project — fused kernel combining +// `dqn_select_action_atoms` + `bellman_target_projection` into a +// single launch. Eliminates the intermediate `action_logits[B, Q_N_ATOMS]` +// global memory round-trip by staging the selected atom row in shared +// memory. +// +// nsys profiling showed the two kernels always launch back-to-back with +// identical grid/block dims: Grid=(B, 1, 1), Block=(Q_N_ATOMS, 1, 1). +// The intermediate `action_logits[B, Q_N_ATOMS]` exists only to shuttle +// data from the first kernel's output to the second kernel's input. +// Fusing them saves one global write + one global read of B*Q_N_ATOMS +// floats and one kernel launch overhead (~1.2μs avg). +// +// Inputs: +// full_logits [B × N_ACTIONS × Q_N_ATOMS] — target-net atom logits +// actions [B] — action indices (0..N_ACTIONS) +// rewards [B] — per-transition reward r_t +// dones [B] — 0/1 done flag +// n_step_gammas [B] — γⁿ per transition +// isv [≥ 486] — reads V_MIN/V_MAX/γ +// B int — batch size +// +// Outputs: +// target_dist [B × Q_N_ATOMS] — projected target distribution +// +// Block layout: +// grid = (B, 1, 1) +// block = (Q_N_ATOMS, 1, 1) +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void bellman_fused_select_project( + const float* __restrict__ full_logits, // [B × N_ACTIONS × Q_N_ATOMS] + const int* __restrict__ actions, // [B] + const float* __restrict__ rewards, // [B] + const float* __restrict__ dones, // [B] + const float* __restrict__ n_step_gammas, // [B] + const float* __restrict__ isv, // ≥ 486 + int B, + float* __restrict__ target_dist // [B × Q_N_ATOMS] +) { + const int batch = blockIdx.x; + const int atom = threadIdx.x; + if (batch >= B || atom >= Q_N_ATOMS) return; + + // ── Phase 1: select action atoms into shared memory ───────────── + __shared__ float s_logits[Q_N_ATOMS]; + + int a = actions[batch]; + if (a < 0) a = 0; + if (a >= N_ACTIONS) a = 0; + + const long long src_idx = + (long long)batch * N_ACTIONS * Q_N_ATOMS + + (long long)a * Q_N_ATOMS + + (long long)atom; + s_logits[atom] = full_logits[src_idx]; + __syncthreads(); + + // ── Phase 2: bellman target projection from shared memory ─────── + __shared__ float s_softmax[Q_N_ATOMS]; + __shared__ float s_max; + __shared__ float s_sumexp; + __shared__ float s_proj[Q_N_ATOMS]; + + // Softmax over target logits (numerically-stable max-subtract) + if (atom == 0) { + float m = s_logits[0]; + #pragma unroll + for (int z = 1; z < Q_N_ATOMS; ++z) + m = fmaxf(m, s_logits[z]); + s_max = m; + } + __syncthreads(); + + const float e = expf(s_logits[atom] - s_max); + s_softmax[atom] = e; + s_proj[atom] = 0.0f; + __syncthreads(); + + if (atom == 0) { + float sum = 0.0f; + #pragma unroll + for (int z = 0; z < Q_N_ATOMS; ++z) sum += s_softmax[z]; + s_sumexp = sum; + } + __syncthreads(); + + const float p = s_softmax[atom] / s_sumexp; + + const float r = rewards[batch]; + const float gamma_eff = n_step_gammas[batch]; + + const float V_MIN_eff = isv[RL_C51_V_MIN_INDEX]; + const float V_MAX_eff = isv[RL_C51_V_MAX_INDEX]; + const float DELTA_Z = (V_MAX_eff - V_MIN_eff) / (float)(Q_N_ATOMS - 1); + + const float atom_value = V_MIN_eff + (float)atom * DELTA_Z; + const float t_z = r + gamma_eff * atom_value; + const float t_z_clamp = fmaxf(V_MIN_eff, fminf(V_MAX_eff, t_z)); + const float b_frac = (t_z_clamp - V_MIN_eff) / DELTA_Z; + const int l = max(0, min(Q_N_ATOMS - 1, (int)floorf(b_frac))); + const int u = max(0, min(Q_N_ATOMS - 1, (int)ceilf(b_frac))); + const float frac = b_frac - (float)l; + + // Distribute mass — serialised walk per `feedback_no_atomicadd.md` + for (int src = 0; src < Q_N_ATOMS; ++src) { + __syncthreads(); + if (atom == src) { + if (l == u) { + s_proj[l] += p; + } else { + s_proj[l] += p * (1.0f - frac); + s_proj[u] += p * frac; + } + } + } + __syncthreads(); + target_dist[batch * Q_N_ATOMS + atom] = s_proj[atom]; +} + + extern "C" __global__ void bellman_target_projection( const float* __restrict__ target_logits, // [B × Q_N_ATOMS] const float* __restrict__ rewards, // [B] (n-step discounted R_n) diff --git a/crates/ml-alpha/cuda/rl_iqn_backward.cu b/crates/ml-alpha/cuda/rl_iqn_backward.cu index 229d8a307..463e2f679 100644 --- a/crates/ml-alpha/cuda/rl_iqn_backward.cu +++ b/crates/ml-alpha/cuda/rl_iqn_backward.cu @@ -1,49 +1,33 @@ -// rl_iqn_backward.cu — IQN backward through the forward pass. +// rl_iqn_backward.cu — IQN backward pass: hybrid cuBLAS + custom kernel. // -// Given grad_output [B, N_TAU, N_ACTIONS] (from rl_iqn_loss_fwd), backprop -// through the IQN forward to produce per-batch gradients for: -// - w_out [HIDDEN_DIM, N_ACTIONS] (per-batch scratch) -// - b_out [N_ACTIONS] (per-batch scratch) -// - w_embed [EMBED_DIM, HIDDEN_DIM] (per-batch scratch) -// - b_embed [HIDDEN_DIM] (per-batch scratch) +// The backward pass is split into: // -// Forward recap: -// phi(tau)[c] = ReLU(sum_i cos((i+1)*pi*tau) * W_embed[i,c] + b_embed[c]) -// combined[c] = h_t[c] * phi(tau)[c] -// Q[tau, a] = sum_c W_out[c, a] * combined[c] + b_out[a] +// 1. cuBLAS SGEMM (from Rust): +// grad_combined[B*N_TAU, HIDDEN_DIM] = grad_q[B*N_TAU, N_ACTIONS] +// @ W_out^T[N_ACTIONS, HIDDEN_DIM] +// Replaces the per-thread inner loop over N_ACTIONS. // -// Backward: -// dQ/dW_out[c, a] = combined[c] (for the given tau) -// dQ/db_out[a] = 1 -// dQ/d_combined[c] = sum_a grad_Q[a] * W_out[c, a] -// d_combined/d_phi[c] = h_t[c] -// d_phi/d_embed_input[c] = 1(embed_input > 0) (ReLU mask) -// d_embed_input/dW_embed[i,c] = cos((i+1)*pi*tau) -// d_embed_input/db_embed[c] = 1 +// 2. rl_iqn_backward (custom kernel, this file): +// Per-batch backward accumulation. Recomputes phi(tau) and combined +// from (h_t, tau, W_embed, b_embed) inline (same as the original +// monolithic kernel). Uses the cuBLAS-computed grad_combined instead +// of recomputing it from grad_q × W_out. // -// Block layout: -// grid = (B, N_TAU, 1) +// Produces per-batch gradients for: +// - grad_w_out_pb [B, HIDDEN_DIM * N_ACTIONS] +// - grad_b_out_pb [B, N_ACTIONS] +// - grad_w_embed_pb [B, EMBED_DIM * HIDDEN_DIM] +// - grad_b_embed_pb [B, HIDDEN_DIM] +// +// 3. rl_iqn_bwd_relu_hadamard (custom kernel, this file): +// Element-wise backward through hadamard + ReLU. Used when the +// forward cache provides embed_pre_relu (alternative to inline +// recompute). Kept for future use when forward caching is added. +// +// Block layout for rl_iqn_backward: +// grid = (B, 1, 1) // block = (HIDDEN_DIM, 1, 1) -// One block per (batch, tau) pair — mirrors the forward kernel. -// Each thread handles one hidden-dim index c. -// -// Output layout (per-batch scratch accumulated across N_TAU): -// grad_w_out_per_batch [B, HIDDEN_DIM, N_ACTIONS] -// grad_b_out_per_batch [B, N_ACTIONS] -// grad_w_embed_per_batch [B, EMBED_DIM, HIDDEN_DIM] -// grad_b_embed_per_batch [B, HIDDEN_DIM] -// -// The per-tau contributions are ACCUMULATED (+=) into the per-batch -// scratch via atomicAdd-free patterns: each (batch, tau) block writes -// to its own output offset, and a subsequent reduce_axis0 pass sums -// across batches (same as C51 head). The tau-dimension accumulation -// happens via the atomic-free pattern of writing to [B, N_TAU, ...] and -// then launching a second reduction kernel over the tau dimension. -// -// SIMPLIFIED APPROACH: since HIDDEN_DIM threads can cooperate and -// N_TAU is typically 32, we launch grid=(B, 1, 1) block=(HIDDEN_DIM, 1, 1) -// and each thread loops over all N_TAU to accumulate its contributions. -// This avoids cross-block accumulation entirely. +// One thread per hidden dim; each thread loops over N_TAU. // // Per `feedback_no_atomicadd.md`: no atomicAdd. // Per `feedback_cpu_is_read_only.md`: all compute on GPU. @@ -54,12 +38,39 @@ #define EMBED_DIM 64 #define PI_F 3.14159265f +// ───────────────────────────────────────────────────────────────────── +// rl_iqn_backward: +// Per-batch backward with cuBLAS-precomputed grad_combined. +// +// Thread c recomputes phi(tau)[c] and combined[c] from (tau, W_embed, +// b_embed, h_t) for each quantile sample, then accumulates: +// - grad_w_out[c, a] = Σ_t combined[c] * grad_q[t, a] +// - grad_phi[c] = grad_combined[c] * h_t[c] (from cuBLAS) +// - grad_embed_input[c] = grad_phi[c] * relu_mask +// - grad_w_embed[i, c] = Σ_t cos((i+1)*pi*tau[t]) * grad_embed_input[c] +// - grad_b_embed[c] = Σ_t grad_embed_input[c] +// - grad_b_out[a] = Σ_t grad_q[t, a] (thread 0) +// +// Inputs: +// h_t [B, HIDDEN_DIM] +// tau [B, N_TAU] — saved from forward +// w_embed [EMBED_DIM, HIDDEN_DIM] — online weights +// b_embed [HIDDEN_DIM] +// grad_combined [B*N_TAU, HIDDEN_DIM] — from cuBLAS (grad_q @ W_out^T) +// grad_output [B, N_TAU, N_ACTIONS] — from loss +// B, N_TAU +// Outputs: +// grad_w_out_pb [B, HIDDEN_DIM * N_ACTIONS] +// grad_b_out_pb [B, N_ACTIONS] +// grad_w_embed_pb [B, EMBED_DIM * HIDDEN_DIM] +// grad_b_embed_pb [B, HIDDEN_DIM] +// ───────────────────────────────────────────────────────────────────── extern "C" __global__ void rl_iqn_backward( const float* __restrict__ h_t, // [B, HIDDEN_DIM] const float* __restrict__ tau, // [B, N_TAU] const float* __restrict__ w_embed, // [EMBED_DIM, HIDDEN_DIM] const float* __restrict__ b_embed, // [HIDDEN_DIM] - const float* __restrict__ w_out, // [HIDDEN_DIM, N_ACTIONS] + const float* __restrict__ grad_combined, // [B*N_TAU, HIDDEN_DIM] (cuBLAS) const float* __restrict__ grad_output, // [B, N_TAU, N_ACTIONS] int B, int N_TAU, @@ -83,6 +94,7 @@ extern "C" __global__ void rl_iqn_backward( float acc_grad_b_embed = 0.0f; for (int t = 0; t < N_TAU; ++t) { + const int row = batch * N_TAU + t; const float tau_val = tau[batch * N_TAU + t]; // Recompute forward: phi(tau)[c] @@ -100,18 +112,15 @@ extern "C" __global__ void rl_iqn_backward( float combined_c = h_c * phi_c; // grad_output for this (batch, tau): [N_ACTIONS] - const int go_base = batch * N_TAU * N_ACTIONS + t * N_ACTIONS; + const int go_base = row * N_ACTIONS; // ─── Grad w.r.t. w_out: dL/dW_out[c,a] += combined_c * grad_Q[a] for (int a = 0; a < N_ACTIONS; ++a) { acc_grad_w_out[a] += combined_c * grad_output[go_base + a]; } - // ─── Grad w.r.t. combined: dL/d_combined[c] = Σ_a grad_Q[a] * W_out[c,a] - float grad_combined_c = 0.0f; - for (int a = 0; a < N_ACTIONS; ++a) { - grad_combined_c += grad_output[go_base + a] * w_out[c * N_ACTIONS + a]; - } + // ─── Grad w.r.t. combined: READ from cuBLAS output ─────────── + float grad_combined_c = grad_combined[row * HIDDEN_DIM + c]; // ─── Chain through element-wise product: d_combined/d_phi = h_t[c] float grad_phi_c = grad_combined_c * h_c; @@ -128,15 +137,13 @@ extern "C" __global__ void rl_iqn_backward( acc_grad_b_embed += grad_embed_input_c; } - // Write accumulated grad_w_out for this (batch, c) — row c of the - // per-batch grad_w_out matrix [HIDDEN_DIM, N_ACTIONS]. + // Write accumulated grad_w_out for this (batch, c). const int wo_base = batch * HIDDEN_DIM * N_ACTIONS + c * N_ACTIONS; for (int a = 0; a < N_ACTIONS; ++a) { grad_w_out_pb[wo_base + a] = acc_grad_w_out[a]; } - // Write accumulated grad_w_embed for this (batch, c) — column c of - // the per-batch grad_w_embed matrix [EMBED_DIM, HIDDEN_DIM]. + // Write accumulated grad_w_embed for this (batch, c). for (int i = 0; i < EMBED_DIM; ++i) { grad_w_embed_pb[batch * EMBED_DIM * HIDDEN_DIM + i * HIDDEN_DIM + c] = acc_grad_w_embed[i]; } @@ -144,12 +151,7 @@ extern "C" __global__ void rl_iqn_backward( // Write grad_b_embed for this (batch, c). grad_b_embed_pb[batch * HIDDEN_DIM + c] = acc_grad_b_embed; - // grad_b_out: each action gets contribution from ALL hidden dims. - // Only thread c=0 writes it to avoid races — it accumulates across - // all hidden dims by reading grad_output directly. - // Actually: dL/db_out[a] = Σ_tau grad_Q[tau, a] (since dQ/db_out = 1). - // Each thread has access to grad_output — use a shared-mem reduce. - // Simpler: thread 0 computes it by summing over tau. + // grad_b_out: thread 0 computes by summing over tau. if (c == 0) { for (int a = 0; a < N_ACTIONS; ++a) { float sum = 0.0f; @@ -160,3 +162,41 @@ extern "C" __global__ void rl_iqn_backward( } } } + +// ───────────────────────────────────────────────────────────────────── +// rl_iqn_bwd_relu_hadamard: +// Element-wise backward through hadamard + ReLU when embed_pre_relu +// is available from a cached forward pass. +// +// Grid = (B*N_TAU, ceil(HIDDEN_DIM / 256), 1) +// Block = (min(HIDDEN_DIM, 256), 1, 1) +// +// Inputs: +// grad_combined [M, HIDDEN_DIM] — from cuBLAS (grad_q @ W_out^T) +// h_t [B, HIDDEN_DIM] — encoder hidden state +// embed_pre_relu [M, HIDDEN_DIM] — saved from forward (before ReLU) +// M, B, N_TAU +// Outputs: +// grad_embed_input [M, HIDDEN_DIM] — gradient for embedding input +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void rl_iqn_bwd_relu_hadamard( + const float* __restrict__ grad_combined, // [M, HIDDEN_DIM] + const float* __restrict__ h_t, // [B, HIDDEN_DIM] + const float* __restrict__ embed_pre_relu, // [M, HIDDEN_DIM] + int M, + int B, + int N_TAU, + float* __restrict__ grad_embed_input // [M, HIDDEN_DIM] +) { + const int row = blockIdx.x; + const int c = blockIdx.y * blockDim.x + threadIdx.x; + if (row >= M) return; + if (c >= HIDDEN_DIM) return; + + const int batch_idx = row / N_TAU; + const float h_c = h_t[batch_idx * HIDDEN_DIM + c]; + + float grad_phi_c = grad_combined[row * HIDDEN_DIM + c] * h_c; + float relu_mask = (embed_pre_relu[row * HIDDEN_DIM + c] > 0.0f) ? 1.0f : 0.0f; + grad_embed_input[row * HIDDEN_DIM + c] = grad_phi_c * relu_mask; +} diff --git a/crates/ml-alpha/cuda/rl_iqn_forward.cu b/crates/ml-alpha/cuda/rl_iqn_forward.cu index 190a09c16..ee00882b6 100644 --- a/crates/ml-alpha/cuda/rl_iqn_forward.cu +++ b/crates/ml-alpha/cuda/rl_iqn_forward.cu @@ -1,31 +1,39 @@ -// rl_iqn_forward.cu — Implicit Quantile Network (IQN) forward pass. +// rl_iqn_forward.cu — IQN forward pass: split pipeline for cuBLAS SGEMM. // -// Complementary distributional Q-head running alongside C51. The IQN -// head models the full return distribution via learned quantile -// functions rather than a fixed atom support. +// The monolithic kernel is replaced by three custom kernels interleaved +// with two cuBLAS SGEMM calls driven from Rust. The pipeline: // -// Forward pass: -// 1. Inline tau sampling: thread 0 of each block generates one U(0,1) -// tau value via xorshift32, writes to tau[batch, tau_idx]. -// Eliminates the separate rl_sample_tau kernel launch (847 per step). -// 2. Quantile embedding: phi(tau) = ReLU(W_embed × cos(i×π×τ) + b_embed) -// where i=1..EMBED_DIM (64). Output: [B, N_TAU, HIDDEN_DIM] -// 3. Element-wise: combined = h_t ⊙ phi(tau) → [B, N_TAU, HIDDEN_DIM] -// 4. Action value: Q = W_out × combined + b_out → [B, N_TAU, N_ACTIONS] +// 1. rl_iqn_tau_cos_features (custom): +// - Thread 0 of each (batch, tau) block samples tau ~ U(0,1) via +// inline xorshift32 and writes to tau[batch, tau_idx]. +// - All threads compute cos_features[j] = cos((j+1) * pi * tau) +// for j = 0..EMBED_DIM-1. +// Output: tau[B, N_TAU], cos_features[B*N_TAU, EMBED_DIM]. // -// Expected Q for ensemble action selection: mean over the tau dimension. +// 2. cuBLAS SGEMM (from Rust): +// embed_out[B*N_TAU, HIDDEN_DIM] = cos_features[B*N_TAU, EMBED_DIM] +// @ W_embed[EMBED_DIM, HIDDEN_DIM] // -// Block layout: -// rl_iqn_forward: grid = (B, N_TAU, 1); block = (HIDDEN_DIM, 1, 1). -// One block per (batch, tau) pair. HIDDEN_DIM threads cooperate on -// the quantile embedding and then compute action values via a -// strided inner loop over N_ACTIONS. +// 3. rl_iqn_relu_hadamard (custom): +// embed_out += b_embed (bias add) +// phi = ReLU(embed_out) +// combined = h_t ⊙ phi (hadamard product with broadcast over tau) +// Output: combined[B*N_TAU, HIDDEN_DIM]. // -// rl_iqn_expected_q: grid = (B, 1, 1); block = (N_ACTIONS, 1, 1). -// One block per batch. Each thread (one per action) reduces across -// N_TAU quantile samples to compute E[Q(s,a)] = mean_tau Q(s,tau,a). -// Uses block tree-reduce pattern (no atomicAdd per -// `feedback_no_atomicadd.md`). +// 4. cuBLAS SGEMM (from Rust): +// q_raw[B*N_TAU, N_ACTIONS] = combined[B*N_TAU, HIDDEN_DIM] +// @ W_out[HIDDEN_DIM, N_ACTIONS] +// +// 5. rl_iqn_bias_add_q (custom): +// q_values = q_raw + b_out (broadcast bias) +// Output: q_values[B, N_TAU, N_ACTIONS]. +// +// 6. rl_iqn_expected_q (unchanged): +// E[Q(s,a)] = mean_tau Q(s, tau, a). +// +// Per `feedback_no_atomicadd.md`: no atomicAdd. +// Per `feedback_cpu_is_read_only.md`: all compute on GPU. +// Per `feedback_no_nvrtc.md`: pre-compiled cubin via build.rs. #include @@ -35,8 +43,7 @@ #define EMBED_DIM 64 #define PI_F 3.14159265f -// Inline xorshift32 PRNG — identical to rl_sample_tau.cu. Fused here -// to eliminate 847 separate kernel launches per training step. +// Inline xorshift32 PRNG — identical to rl_sample_tau.cu. __device__ static uint32_t xorshift32_iqn(uint32_t* state) { uint32_t x = *state; x ^= x << 13; @@ -47,81 +54,51 @@ __device__ static uint32_t xorshift32_iqn(uint32_t* state) { } // ───────────────────────────────────────────────────────────────────── -// rl_iqn_forward: -// Fused tau sampling + Q(s, tau, a) computation. +// rl_iqn_tau_cos_features: +// Stage 1: inline tau sampling + cosine basis computation. // -// Thread 0 of each block samples tau for its (batch, tau_idx) pair -// using xorshift32, identical to the standalone rl_sample_tau kernel. -// The sampled tau is written to the tau buffer so downstream consumers -// (loss, backward) still have access to it. -// -// The block with tau_idx == 0 also advances the persistent per-batch -// PRNG seed so the next step produces different tau values. All blocks -// for a given batch read prng_state[batch] at block start before the -// tau_idx==0 block writes the advance back at block end, so reads are -// consistent within a single kernel launch. +// Grid = (B, N_TAU, 1) +// Block = (EMBED_DIM, 1, 1) — one thread per embedding dimension. // // Inputs: -// prng_state [B] — per-batch xorshift32 seed (mutated) -// h_t [B, HIDDEN_DIM] — encoder hidden state -// w_embed [EMBED_DIM, HIDDEN_DIM] — quantile embedding weight -// b_embed [HIDDEN_DIM] — quantile embedding bias -// w_out [HIDDEN_DIM, N_ACTIONS] — output projection weight -// b_out [N_ACTIONS] — output projection bias -// B — batch size -// N_TAU — number of quantile samples +// prng_state [B] — per-batch xorshift32 seed (mutated) +// B — batch size +// N_TAU — number of quantile samples // Outputs: -// tau [B, N_TAU] — sampled U(0,1) quantile fractions -// q_values [B, N_TAU, N_ACTIONS] — quantile action values +// tau [B, N_TAU] — sampled U(0,1) quantile fractions +// cos_features [B*N_TAU, EMBED_DIM] — cos((j+1) * pi * tau) // ───────────────────────────────────────────────────────────────────── -extern "C" __global__ void rl_iqn_forward( - uint32_t* __restrict__ prng_state, // [B] per-batch seed (mutated) - const float* __restrict__ h_t, // [B, HIDDEN_DIM] - float* __restrict__ tau, // [B, N_TAU] (output) - const float* __restrict__ w_embed, // [EMBED_DIM, HIDDEN_DIM] - const float* __restrict__ b_embed, // [HIDDEN_DIM] - const float* __restrict__ w_out, // [HIDDEN_DIM, N_ACTIONS] - const float* __restrict__ b_out, // [N_ACTIONS] +extern "C" __global__ void rl_iqn_tau_cos_features( + uint32_t* __restrict__ prng_state, // [B] int B, int N_TAU, - float* __restrict__ q_values // [B, N_TAU, N_ACTIONS] + float* __restrict__ tau, // [B, N_TAU] + float* __restrict__ cos_features // [B*N_TAU, EMBED_DIM] ) { const int batch = blockIdx.x; const int tau_idx = blockIdx.y; - const int c = threadIdx.x; // hidden dim index + const int j = threadIdx.x; // embedding dim index if (batch >= B) return; if (tau_idx >= N_TAU) return; - if (c >= HIDDEN_DIM) return; + if (j >= EMBED_DIM) return; - // Inline tau sampling (fused from rl_sample_tau). Thread 0 of each - // block generates the tau value for this (batch, tau_idx) pair. - // Shared memory broadcasts the result to all HIDDEN_DIM threads. + // Thread 0 samples tau for this (batch, tau_idx) pair. __shared__ float s_tau_val; - if (c == 0) { - // Self-seed on first call (alloc_zeros → 0). Knuth golden-ratio - // hash mixes batch index + 0xBEEF into a nonzero initial state. + if (j == 0) { uint32_t seed = prng_state[batch]; if (seed == 0u) seed = (uint32_t)(batch + 1) * 2654435761u + 0xBEEFu; - // Derive a local PRNG state from the per-batch seed, mixed with - // the tau index to decorrelate across quantiles. uint32_t local_state = seed ^ (uint32_t)(tau_idx * 2654435761u); - // Warmup to dispel low-quality seeding correlations. #pragma unroll for (int w = 0; w < 4; ++w) xorshift32_iqn(&local_state); const uint32_t r = xorshift32_iqn(&local_state); - // Convert to [0, 1) with 24-bit mantissa precision. const float u = (float)(r >> 8) * (1.0f / 16777216.0f); tau[batch * N_TAU + tau_idx] = u; s_tau_val = u; - // The block with tau_idx == 0 advances the persistent per-batch - // seed so the next step produces different tau samples. Uses - // `seed` (which includes the self-seed value) not the raw - // prng_state[batch] (which may still be zero on the first call). if (tau_idx == 0) { uint32_t adv = seed; #pragma unroll @@ -133,46 +110,74 @@ extern "C" __global__ void rl_iqn_forward( const float tau_val = s_tau_val; - // Step 1: Compute quantile embedding phi(tau)[c] - // phi(tau) = ReLU( sum_{i=1..EMBED_DIM} cos(i * pi * tau) * W_embed[i, c] + b_embed[c] ) - float embed_acc = b_embed[c]; - #pragma unroll - for (int i = 0; i < EMBED_DIM; ++i) { - float cos_feat = cosf((float)(i + 1) * PI_F * tau_val); - embed_acc += cos_feat * w_embed[i * HIDDEN_DIM + c]; - } - // ReLU activation - float phi_c = fmaxf(embed_acc, 0.0f); + // cos_features[row, j] = cos((j+1) * pi * tau_val) + const int row = batch * N_TAU + tau_idx; + cos_features[row * EMBED_DIM + j] = cosf((float)(j + 1) * PI_F * tau_val); +} - // Step 2: Element-wise product with encoder hidden state - float combined_c = h_t[batch * HIDDEN_DIM + c] * phi_c; +// ───────────────────────────────────────────────────────────────────── +// rl_iqn_relu_hadamard: +// Stage 3: bias-add → ReLU → element-wise product with h_t. +// +// Grid = (B*N_TAU, ceil(HIDDEN_DIM / 256), 1) +// Block = (min(HIDDEN_DIM, 256), 1, 1) +// +// Inputs: +// embed_out [B*N_TAU, HIDDEN_DIM] — output of cuBLAS SGEMM (no bias) +// b_embed [HIDDEN_DIM] — embedding bias +// h_t [B, HIDDEN_DIM] — encoder hidden state +// M — total rows = B * N_TAU +// B — batch size (for h_t indexing) +// N_TAU — quantile count +// Outputs: +// combined [B*N_TAU, HIDDEN_DIM] — h_t ⊙ ReLU(embed_out + b_embed) +// embed_pre_relu [B*N_TAU, HIDDEN_DIM] — embed_out + b_embed (before ReLU, +// saved for backward) +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void rl_iqn_relu_hadamard( + const float* __restrict__ embed_out, // [M, HIDDEN_DIM] + const float* __restrict__ b_embed, // [HIDDEN_DIM] + const float* __restrict__ h_t, // [B, HIDDEN_DIM] + int M, // B * N_TAU + int B, + int N_TAU, + float* __restrict__ combined, // [M, HIDDEN_DIM] + float* __restrict__ embed_pre_relu // [M, HIDDEN_DIM] (saved for bwd) +) { + const int row = blockIdx.x; + const int c = blockIdx.y * blockDim.x + threadIdx.x; + if (row >= M) return; + if (c >= HIDDEN_DIM) return; - // Store combined in shared memory for the matmul reduction - __shared__ float s_combined[HIDDEN_DIM]; - s_combined[c] = combined_c; - __syncthreads(); + const int batch = row / N_TAU; - // Step 3: Compute Q(s, tau, a) = W_out^T × combined + b_out - // Each thread computes one action's contribution from its hidden dim, - // then we need a reduction across hidden dims. Instead, thread c - // contributes to all actions via the output projection. - // With HIDDEN_DIM threads and N_ACTIONS outputs, each thread iterates - // over actions and contributes its portion. - // - // Output: q_values[batch, tau_idx, a] = sum_c(W_out[c, a] * combined[c]) + b_out[a] - // Since HIDDEN_DIM=128 and N_ACTIONS=11, we assign each thread to - // compute partial sums and use warp-shuffle reduction. - // - // Strategy: thread c writes combined[c] to shared mem (done above). - // First N_ACTIONS threads each compute one full dot product. - if (c < N_ACTIONS) { - float q_acc = b_out[c]; - #pragma unroll - for (int k = 0; k < HIDDEN_DIM; ++k) { - q_acc += w_out[k * N_ACTIONS + c] * s_combined[k]; - } - q_values[batch * N_TAU * N_ACTIONS + tau_idx * N_ACTIONS + c] = q_acc; - } + float val = embed_out[row * HIDDEN_DIM + c] + b_embed[c]; + embed_pre_relu[row * HIDDEN_DIM + c] = val; + + float phi_c = fmaxf(val, 0.0f); + combined[row * HIDDEN_DIM + c] = h_t[batch * HIDDEN_DIM + c] * phi_c; +} + +// ───────────────────────────────────────────────────────────────────── +// rl_iqn_bias_add_q: +// Stage 5: add b_out bias to the cuBLAS SGEMM output. +// +// Grid = (B*N_TAU, 1, 1) +// Block = (N_ACTIONS, 1, 1) +// +// In-place: q_values[row, a] += b_out[a] +// ───────────────────────────────────────────────────────────────────── +extern "C" __global__ void rl_iqn_bias_add_q( + float* __restrict__ q_values, // [M, N_ACTIONS] (mutated in-place) + const float* __restrict__ b_out, // [N_ACTIONS] + int M // B * N_TAU +) { + const int row = blockIdx.x; + const int a = threadIdx.x; + if (row >= M) return; + if (a >= N_ACTIONS) return; + + q_values[row * N_ACTIONS + a] += b_out[a]; } // ───────────────────────────────────────────────────────────────────── @@ -180,6 +185,8 @@ extern "C" __global__ void rl_iqn_forward( // Compute E[Q(s, a)] = (1/N_TAU) × Σ_{tau} Q(s, tau, a) // for ensemble action selection. // +// Unchanged from the original monolithic kernel. +// // Inputs: // q_values [B, N_TAU, N_ACTIONS] — full quantile Q tensor // B — batch size @@ -188,8 +195,6 @@ extern "C" __global__ void rl_iqn_forward( // expected_q [B, N_ACTIONS] — mean Q per action // // Block layout: grid = (B, 1, 1); block = (N_ACTIONS, 1, 1). -// One thread per action; each thread sums over N_TAU quantiles. -// No cross-thread reduction needed (each action is independent). // ───────────────────────────────────────────────────────────────────── extern "C" __global__ void rl_iqn_expected_q( const float* __restrict__ q_values, // [B, N_TAU, N_ACTIONS] diff --git a/crates/ml-alpha/cuda/rl_outcome_fused.cu b/crates/ml-alpha/cuda/rl_outcome_fused.cu new file mode 100644 index 000000000..bbb75c670 --- /dev/null +++ b/crates/ml-alpha/cuda/rl_outcome_fused.cu @@ -0,0 +1,125 @@ +/* ===================================================================== + * rl_outcome_fused.cu -- Fused fwd + CE + bwd for K=3 outcome head. + * + * Eliminates two global-memory round-trips (logits, grad_logits) by + * keeping all intermediates in shared memory. + * + * Grid = (B, 1, 1) -- one block per batch element. + * Block = (128, 1, 1) -- matches backward's HIDDEN_DIM parallelism. + * + * Three phases separated by __syncthreads(): + * + * Phase 1 (threads 0..2): linear forward logits = h_t × W + b + * → s_logits[3] in shared memory + * + * Phase 2 (thread 0): softmax CE + gradient + * → loss_pb[batch], s_grad_logits[3] + * + * Phase 3 (all 128 threads): backward through linear layer + * → grad_w_per_batch, grad_b_per_batch, grad_h_t + * + * No atomicAdd -- sole-writer per output element. + * No nvrtc -- precompiled cubin. + * ===================================================================== */ + +#include + +#define HIDDEN_DIM 128 +#define K_CLASSES 3 + +extern "C" __global__ void rl_outcome_fused( + const float* __restrict__ h_t, /* [B, HIDDEN_DIM=128] */ + const float* __restrict__ w, /* [HIDDEN_DIM, K_CLASSES=3] col-major */ + const float* __restrict__ b, /* [K_CLASSES=3] */ + const int* __restrict__ labels, /* [B] (0..2 or -1=skip) */ + int b_size, + float* __restrict__ loss_pb, /* [B] */ + float* __restrict__ grad_w_per_batch, /* [B, HIDDEN_DIM, K_CLASSES] */ + float* __restrict__ grad_b_per_batch, /* [B, K_CLASSES] */ + float* __restrict__ grad_h_t /* [B, HIDDEN_DIM] */ +) { + const int batch = blockIdx.x; + const int tid = threadIdx.x; + + if (batch >= b_size) return; + + /* Shared intermediates -- eliminates two global round-trips. */ + __shared__ float s_logits[K_CLASSES]; + __shared__ float s_grad_logits[K_CLASSES]; + + const float* h_row = h_t + batch * HIDDEN_DIM; + + /* ── Phase 1: linear forward (threads 0..2) ────────────────────── */ + if (tid < K_CLASSES) { + float sum = b[tid]; + #pragma unroll 16 + for (int i = 0; i < HIDDEN_DIM; ++i) { + sum += h_row[i] * w[i * K_CLASSES + tid]; + } + s_logits[tid] = sum; + } + __syncthreads(); + + /* ── Phase 2: softmax CE + gradient (thread 0 only) ──────────── */ + if (tid == 0) { + int label = labels[batch]; + + if (label < 0 || label >= K_CLASSES) { + /* Masked sample: no label available. */ + loss_pb[batch] = 0.0f; + s_grad_logits[0] = 0.0f; + s_grad_logits[1] = 0.0f; + s_grad_logits[2] = 0.0f; + } else { + float l0 = s_logits[0]; + float l1 = s_logits[1]; + float l2 = s_logits[2]; + + /* Numerically stable softmax: subtract max. */ + float mx = fmaxf(l0, fmaxf(l1, l2)); + float e0 = expf(l0 - mx); + float e1 = expf(l1 - mx); + float e2 = expf(l2 - mx); + float sum_exp = e0 + e1 + e2; + + float p0 = e0 / sum_exp; + float p1 = e1 / sum_exp; + float p2 = e2 / sum_exp; + + /* CE loss = -log(p[label]). */ + float p_label = (label == 0) ? p0 : ((label == 1) ? p1 : p2); + loss_pb[batch] = -logf(fmaxf(p_label, 1e-12f)); + + /* Gradient: p[k] - 1(k == label). */ + s_grad_logits[0] = p0 - ((label == 0) ? 1.0f : 0.0f); + s_grad_logits[1] = p1 - ((label == 1) ? 1.0f : 0.0f); + s_grad_logits[2] = p2 - ((label == 2) ? 1.0f : 0.0f); + } + } + __syncthreads(); + + /* ── Phase 3: backward through linear layer (all 128 threads) ── */ + if (tid >= HIDDEN_DIM) return; + + const float h_bi = h_row[tid]; + + /* grad_w_per_batch[batch, tid, k] = h_bi × s_grad_logits[k] */ + const int row_off = (batch * HIDDEN_DIM + tid) * K_CLASSES; + #pragma unroll + for (int k = 0; k < K_CLASSES; ++k) { + grad_w_per_batch[row_off + k] = h_bi * s_grad_logits[k]; + } + + /* grad_h_t[batch, tid] = Σ_k W[tid, k] × s_grad_logits[k] */ + float acc = 0.0f; + #pragma unroll + for (int k = 0; k < K_CLASSES; ++k) { + acc += w[tid * K_CLASSES + k] * s_grad_logits[k]; + } + grad_h_t[batch * HIDDEN_DIM + tid] = acc; + + /* grad_b_per_batch[batch, k] = s_grad_logits[k] — threads 0..2 only. */ + if (tid < K_CLASSES) { + grad_b_per_batch[batch * K_CLASSES + tid] = s_grad_logits[tid]; + } +} diff --git a/crates/ml-alpha/src/rl/dqn.rs b/crates/ml-alpha/src/rl/dqn.rs index 464da837a..7b11b745b 100644 --- a/crates/ml-alpha/src/rl/dqn.rs +++ b/crates/ml-alpha/src/rl/dqn.rs @@ -41,10 +41,14 @@ use std::sync::Arc; use anyhow::{Context, Result}; +use cudarc::cublas::CudaBlas; +use cudarc::cublas::sys as cublas_sys; +use cudarc::cublas::sys::cublasOperation_t; use cudarc::driver::{ CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut, LaunchConfig, PushKernelArg, }; use ml_core::cuda_autograd::init::scoped_init_seed; +use ml_core::cuda_autograd::linear::BiasKernels; use ml_core::device::MlDevice; use rand::{Rng, SeedableRng}; use rand_chacha::ChaCha8Rng; @@ -53,6 +57,58 @@ use crate::heads::HIDDEN_DIM; use crate::pinned_mem::MappedF32Buffer; use crate::rl::common::{N_ACTIONS, Q_N_ATOMS}; +/// Output dimension of the Q-head linear layer: one logit per (action, atom). +const K_OUT: usize = N_ACTIONS * Q_N_ATOMS; + +// ── cuBLAS SGEMM helper (mirrors ml-core::cuda_autograd::linear::gemm_ex_f32) ── + +/// F32 x F32 -> F32 GEMM via `cublasGemmEx` with F32 internal accumulation. +/// +/// # Safety +/// All device pointers must be valid and dimensions must be correct. +#[allow(clippy::too_many_arguments)] +unsafe fn gemm_f32( + cublas: &CudaBlas, + transa: cublasOperation_t, + transb: cublasOperation_t, + m: i32, + n: i32, + k: i32, + alpha: f32, + a_ptr: u64, + lda: i32, + b_ptr: u64, + ldb: i32, + beta: f32, + c_ptr: u64, + ldc: i32, + label: &str, +) -> Result<()> { + cudarc::cublas::result::gemm_ex( + *cublas.handle(), + transa, + transb, + m, + n, + k, + (&alpha as *const f32).cast(), + a_ptr as *const std::ffi::c_void, + cublas_sys::cudaDataType_t::CUDA_R_32F, + lda, + b_ptr as *const std::ffi::c_void, + cublas_sys::cudaDataType_t::CUDA_R_32F, + ldb, + (&beta as *const f32).cast(), + c_ptr as *mut std::ffi::c_void, + cublas_sys::cudaDataType_t::CUDA_R_32F, + ldc, + cublas_sys::cublasComputeType_t::CUBLAS_COMPUTE_32F, + cublas_sys::cublasGemmAlgo_t::CUBLAS_GEMM_DFALT, + ) + .map_err(|e| anyhow::anyhow!("cublasGemmEx {label}: {e:?}"))?; + Ok(()) +} + const DQN_HEAD_CUBIN: &[u8] = include_bytes!(concat!( env!("OUT_DIR"), "/dqn_distributional_q.cubin" @@ -120,9 +176,16 @@ pub struct DqnHead { /// `[B × Q_N_ATOMS]` shape that `bellman_target_projection` consumes. /// See `cuda/bellman_target_projection.cu::dqn_select_action_atoms`. pub select_action_atoms_fn: CudaFunction, - /// Owns the Bellman projection cubin lifetime (both - /// `bellman_target_projection` and `dqn_select_action_atoms` live in - /// the same translation unit). + /// Fused `dqn_select_action_atoms` + `bellman_target_projection` in + /// a single launch. Eliminates the intermediate + /// `action_logits[B, Q_N_ATOMS]` global memory round-trip by staging + /// the selected atom row in shared memory. Hot-path replacement for + /// the two sequential launches; the standalone kernels are retained + /// for potential standalone use. + pub bellman_fused_fn: CudaFunction, + /// Owns the Bellman projection cubin lifetime (all three kernels — + /// `bellman_target_projection`, `dqn_select_action_atoms`, and + /// `bellman_fused_select_project` — live in the same translation unit). _bellman_module: Arc, /// Phase R5: element-wise target-net soft update kernel handle. @@ -142,6 +205,18 @@ pub struct DqnHead { pub w_target_d: CudaSlice, /// Target-network biases, same shape as `b_d`. pub b_target_d: CudaSlice, + + // ── cuBLAS infrastructure ──────────────────────────────────────── + /// cuBLAS handle for SGEMM-based forward / backward passes. + /// Replaces the hand-written `dqn_distributional_q_fwd` and + /// `dqn_grad_w_b_h_t` kernels: cuBLAS tensor-core SGEMM is 5-10x + /// faster at b=256 for the [B, 128] x [128, 231] matrix shapes. + pub cublas: CudaBlas, + /// Pre-allocated cuBLAS workspace (8 MiB). Prevents per-call + /// cudaMalloc inside cuBLAS that would break CUDA Graph capture. + _cublas_workspace: CudaSlice, + /// Bias-add and reduce-sum-axis0 kernels from ml-core. + bias_kernels: BiasKernels, } impl DqnHead { @@ -174,6 +249,9 @@ impl DqnHead { let select_action_atoms_fn = bellman_module .load_function("dqn_select_action_atoms") .context("load dqn_select_action_atoms")?; + let bellman_fused_fn = bellman_module + .load_function("bellman_fused_select_project") + .context("load bellman_fused_select_project")?; // Phase R5: target soft-update kernel. let target_soft_update_module = ctx @@ -208,6 +286,30 @@ impl DqnHead { let w_target_d = upload(&stream, &w_host)?; let b_target_d = upload(&stream, &b_host)?; + // cuBLAS handle + pre-allocated workspace (same pattern as + // Mamba2Block::new). 8 MiB workspace prevents per-call + // cudaMalloc that would break CUDA Graph stream capture. + let cublas = CudaBlas::new(Arc::clone(&stream)) + .context("DqnHead: cuBLAS init")?; + const CUBLAS_WORKSPACE_BYTES: usize = 8 * 1024 * 1024; + let cublas_workspace = stream + .alloc_zeros::(CUBLAS_WORKSPACE_BYTES) + .context("DqnHead: cuBLAS workspace alloc")?; + unsafe { + let ws_ptr = cublas_workspace.raw_ptr(); + cudarc::cublas::sys::cublasSetWorkspace_v2( + *cublas.handle(), + ws_ptr as *mut std::ffi::c_void, + CUBLAS_WORKSPACE_BYTES, + ) + .result() + .map_err(|e| anyhow::anyhow!("DqnHead: cublasSetWorkspace_v2: {e:?}"))?; + } + + // Bias-add and reduce-sum-axis0 kernel handles from ml-core. + let bias_kernels = BiasKernels::shared(&stream) + .map_err(|e| anyhow::anyhow!("DqnHead: BiasKernels init: {e}"))?; + Ok(Self { cfg, stream, @@ -217,6 +319,7 @@ impl DqnHead { grad_w_b_h_t_fn, bellman_proj_fn, select_action_atoms_fn, + bellman_fused_fn, _bellman_module: bellman_module, target_soft_update_fn, _target_soft_update_module: target_soft_update_module, @@ -224,6 +327,9 @@ impl DqnHead { b_d, w_target_d, b_target_d, + cublas, + _cublas_workspace: cublas_workspace, + bias_kernels, }) } @@ -282,10 +388,15 @@ impl DqnHead { Ok(()) } - /// Phase E.2 forward: launch `dqn_distributional_q_fwd` on `h_t`, - /// writing raw atom logits `[B × N_ACTIONS × Q_N_ATOMS]` into - /// `logits_out`. The trainer pre-allocates the output buffer so the - /// same memory is reused across step() calls. + /// Phase E.2 forward via cuBLAS SGEMM: + /// `logits[B, K_OUT] = h_t[B, HIDDEN_DIM] @ W^T[HIDDEN_DIM, K_OUT] + bias[K_OUT]` + /// + /// Replaces the hand-written `dqn_distributional_q_fwd` kernel whose + /// per-thread serial loop over HIDDEN_DIM=128 consumed 18.1% of GPU + /// time. cuBLAS tensor-core SGEMM is 5-10x faster at b=256. + /// + /// The trainer pre-allocates `logits_out` so the same memory is + /// reused across step() calls. pub fn forward( &self, h_t: &CudaSlice, @@ -293,23 +404,79 @@ impl DqnHead { logits_out: &mut CudaSlice, ) -> Result<()> { debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM); - debug_assert_eq!(logits_out.len(), b_size * N_ACTIONS * Q_N_ATOMS); + debug_assert_eq!(logits_out.len(), b_size * K_OUT); + self.forward_gemm(&self.w_d, &self.b_d, h_t, b_size, logits_out) + .context("dqn forward (online)") + } - let b_i = b_size as i32; + /// cuBLAS SGEMM forward shared by online and target network paths. + /// + /// Computes `logits[B, K_OUT] = h_t[B, HIDDEN_DIM] @ W^T + bias`: + /// - cuBLAS column-major: C_col[K_OUT, B] = W_col^T[K_OUT, HIDDEN_DIM] @ X_col[HIDDEN_DIM, B] + /// transA=T, transB=N, m=K_OUT, n=B, k=HIDDEN_DIM + /// - Then add_bias_2d kernel broadcasts bias[K_OUT] into logits[B, K_OUT]. + fn forward_gemm( + &self, + w: &CudaSlice, + b: &CudaSlice, + h_t: &CudaSlice, + b_size: usize, + logits_out: &mut CudaSlice, + ) -> Result<()> { + let w_ptr = w.raw_ptr(); + let h_ptr = h_t.raw_ptr(); + let out_ptr = logits_out.raw_ptr(); + // SGEMM: logits = h_t @ W^T (alpha=1, beta=0 overwrites logits_out) + unsafe { + gemm_f32( + &self.cublas, + cublasOperation_t::CUBLAS_OP_T, // transA: W[K_OUT, HIDDEN_DIM] stored row-major = col[HIDDEN_DIM, K_OUT] + cublasOperation_t::CUBLAS_OP_N, // transB: h_t[B, HIDDEN_DIM] stored row-major = col[HIDDEN_DIM, B] + K_OUT as i32, // m + b_size as i32, // n + HIDDEN_DIM as i32, // k + 1.0, // alpha + w_ptr, // A = W + HIDDEN_DIM as i32, // lda (W row-major [K_OUT, HIDDEN_DIM] → col lda = HIDDEN_DIM) + h_ptr, // B = h_t + HIDDEN_DIM as i32, // ldb + 0.0, // beta + out_ptr, // C = logits_out + K_OUT as i32, // ldc + "dqn_fwd_sgemm", + )?; + } + // Broadcast bias: logits_out[b, j] += bias[j] + self.add_bias(logits_out, b, b_size)?; + Ok(()) + } + + /// Launch the `add_bias_2d_kernel` from ml-core's BiasKernels. + fn add_bias( + &self, + y: &mut CudaSlice, + bias: &CudaSlice, + rows: usize, + ) -> Result<()> { + let total = rows * K_OUT; + let threads = 256_u32; + let blocks = ((total as u32) + threads - 1) / threads; + let rows_i32 = rows as i32; + let cols_i32 = K_OUT as i32; let cfg = LaunchConfig { - grid_dim: (b_size as u32, N_ACTIONS as u32, 1), - block_dim: (Q_N_ATOMS as u32, 1, 1), + grid_dim: (blocks, 1, 1), + block_dim: (threads, 1, 1), shared_mem_bytes: 0, }; - let mut launch = self.stream.launch_builder(&self.fwd_fn); - launch - .arg(&self.w_d) - .arg(&self.b_d) - .arg(h_t) - .arg(&b_i) - .arg(logits_out); unsafe { - launch.launch(cfg).context("dqn_distributional_q_fwd launch")?; + self.stream + .launch_builder(&self.bias_kernels.add_fn) + .arg(y) + .arg(bias) + .arg(&rows_i32) + .arg(&cols_i32) + .launch(cfg) + .context("dqn add_bias_2d")?; } Ok(()) } @@ -361,17 +528,16 @@ impl DqnHead { Ok(()) } - /// Phase E.2 backward chain stage 2 — propagate `grad_logits` into - /// `grad_w_per_batch [B × N_ACTIONS × Q_N_ATOMS × HIDDEN_DIM]`, - /// `grad_b_per_batch [B × N_ACTIONS × Q_N_ATOMS]`, and - /// `grad_h_t [B × HIDDEN_DIM]` (OVERWRITE). The caller reduces - /// `grad_w_per_batch` / `grad_b_per_batch` along axis 0 via the - /// existing `reduce_axis0` kernel to get the final - /// `grad_w [N_ACTIONS × Q_N_ATOMS × HIDDEN_DIM]` / - /// `grad_b [N_ACTIONS × Q_N_ATOMS]`. + /// Phase E.2 backward chain stage 2 (LEGACY — kept for reference; + /// production path is [`backward_gemm`]). /// - /// `grad_h_t` is OVERWRITE: the trainer's `grad_h_accumulate` - /// kernel reads it and folds via scaled +=. + /// Propagates `grad_logits` into per-batch `grad_w_per_batch`, + /// `grad_b_per_batch`, and `grad_h_t` via the hand-written + /// `dqn_grad_w_b_h_t` kernel. Callers must still reduce + /// `grad_w_per_batch` / `grad_b_per_batch` along axis 0 via + /// `reduce_axis0`. Superseded by [`backward_gemm`] which produces + /// the reduced grad_w / grad_b directly via cuBLAS SGEMM + + /// reduce-sum kernel. #[allow(clippy::too_many_arguments)] pub fn backward_to_w_b_h( &self, @@ -383,20 +549,19 @@ impl DqnHead { grad_h_t: &mut CudaSlice, ) -> Result<()> { debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM); - debug_assert_eq!(grad_logits.len(), b_size * N_ACTIONS * Q_N_ATOMS); + debug_assert_eq!(grad_logits.len(), b_size * K_OUT); debug_assert_eq!( grad_w_per_batch.len(), - b_size * N_ACTIONS * Q_N_ATOMS * HIDDEN_DIM + b_size * K_OUT * HIDDEN_DIM ); - debug_assert_eq!(grad_b_per_batch.len(), b_size * N_ACTIONS * Q_N_ATOMS); + debug_assert_eq!(grad_b_per_batch.len(), b_size * K_OUT); debug_assert_eq!(grad_h_t.len(), b_size * HIDDEN_DIM); - let k_out = (N_ACTIONS * Q_N_ATOMS) as usize; let b_i = b_size as i32; let cfg = LaunchConfig { grid_dim: (b_size as u32, 1, 1), block_dim: (HIDDEN_DIM as u32, 1, 1), - shared_mem_bytes: (k_out * std::mem::size_of::()) as u32, + shared_mem_bytes: (K_OUT * std::mem::size_of::()) as u32, }; let mut launch = self.stream.launch_builder(&self.grad_w_b_h_t_fn); launch @@ -413,16 +578,115 @@ impl DqnHead { Ok(()) } - /// Target-network forward: same kernel as the online forward, but - /// reads from `w_target_d` / `b_target_d`. Used by the Bellman - /// target-bootstrap path — `target_logits = Z_target(s_{t+1})` flow - /// into [`project_bellman_target`]. + /// cuBLAS backward: compute grad_w, grad_b, and grad_h_t from + /// `grad_logits` in three SGEMM + one reduce-sum launch. Replaces + /// the hand-written `dqn_grad_w_b_h_t` kernel + `reduce_axis0`. /// - /// Note: at construction the target net starts with the same Xavier - /// draw as the online net (zero-divergence bootstrap), so the first - /// Bellman backup sees `target_logits == online_logits`. Phase E's - /// soft-update controller subsequently blends with τ from - /// `ISV[RL_TARGET_TAU_INDEX]`. + /// Mathematics: + /// grad_h_t [B, HIDDEN_DIM] = grad_logits [B, K_OUT] @ W [K_OUT, HIDDEN_DIM] + /// grad_w [K_OUT, HIDDEN_DIM] = grad_logits^T [K_OUT, B] @ h_t [B, HIDDEN_DIM] + /// grad_b [K_OUT] = sum(grad_logits [B, K_OUT], axis=0) + /// + /// Benefits: + /// - Eliminates the B*K_OUT*HIDDEN_DIM per-batch scratch (7.5 MiB at B=256). + /// - Eliminates two `reduce_axis0` kernel launches. + /// - cuBLAS SGEMM is 5-10x faster than the hand-written kernel for + /// these shapes ([B=256, 128] x [128, 231]). + /// - grad_w / grad_b are produced directly reduced, ready for Adam. + #[allow(clippy::too_many_arguments)] + pub fn backward_gemm( + &self, + h_t: &CudaSlice, + grad_logits: &CudaSlice, + b_size: usize, + grad_w: &mut CudaSlice, + grad_b: &mut CudaSlice, + grad_h_t: &mut CudaSlice, + ) -> Result<()> { + debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM); + debug_assert_eq!(grad_logits.len(), b_size * K_OUT); + debug_assert_eq!(grad_w.len(), K_OUT * HIDDEN_DIM); + debug_assert_eq!(grad_b.len(), K_OUT); + debug_assert_eq!(grad_h_t.len(), b_size * HIDDEN_DIM); + + let gl_ptr = grad_logits.raw_ptr(); + let ht_ptr = h_t.raw_ptr(); + let w_ptr = self.w_d.raw_ptr(); + + // ── grad_h_t [B, HIDDEN_DIM] = grad_logits [B, K_OUT] @ W [K_OUT, HIDDEN_DIM] ── + // Column-major: grad_h_t_col[HIDDEN_DIM, B] = W_col[HIDDEN_DIM, K_OUT] @ gl_col[K_OUT, B] + // transA=N, transB=N, m=HIDDEN_DIM, n=B, k=K_OUT + let gh_ptr = grad_h_t.raw_ptr(); + unsafe { + gemm_f32( + &self.cublas, + cublasOperation_t::CUBLAS_OP_N, // W[K_OUT, HIDDEN_DIM] row-major = col[HIDDEN_DIM, K_OUT] + cublasOperation_t::CUBLAS_OP_N, // gl[B, K_OUT] row-major = col[K_OUT, B] + HIDDEN_DIM as i32, // m + b_size as i32, // n + K_OUT as i32, // k + 1.0, // alpha + w_ptr, // A = W + HIDDEN_DIM as i32, // lda + gl_ptr, // B = grad_logits + K_OUT as i32, // ldb + 0.0, // beta + gh_ptr, // C = grad_h_t + HIDDEN_DIM as i32, // ldc + "dqn_bwd_grad_h_t", + )?; + } + + // ── grad_w [K_OUT, HIDDEN_DIM] = grad_logits^T [K_OUT, B] @ h_t [B, HIDDEN_DIM] ── + // Column-major: grad_w_col[HIDDEN_DIM, K_OUT] = h_t_col[HIDDEN_DIM, B] @ gl_col[K_OUT, B]^T + // transA=N, transB=T, m=HIDDEN_DIM, n=K_OUT, k=B + let gw_ptr = grad_w.raw_ptr(); + unsafe { + gemm_f32( + &self.cublas, + cublasOperation_t::CUBLAS_OP_N, // h_t[B, HIDDEN_DIM] row-major = col[HIDDEN_DIM, B] + cublasOperation_t::CUBLAS_OP_T, // gl[B, K_OUT] row-major = col[K_OUT, B]; transposed + HIDDEN_DIM as i32, // m + K_OUT as i32, // n + b_size as i32, // k + 1.0, // alpha + ht_ptr, // A = h_t + HIDDEN_DIM as i32, // lda + gl_ptr, // B = grad_logits + K_OUT as i32, // ldb + 0.0, // beta + gw_ptr, // C = grad_w + HIDDEN_DIM as i32, // ldc + "dqn_bwd_grad_w", + )?; + } + + // ── grad_b [K_OUT] = sum(grad_logits [B, K_OUT], axis=0) ────── + let threads = 256_u32; + let blocks = ((K_OUT as u32) + threads - 1) / threads; + let rows_i32 = b_size as i32; + let cols_i32 = K_OUT as i32; + let cfg = LaunchConfig { + grid_dim: (blocks, 1, 1), + block_dim: (threads, 1, 1), + shared_mem_bytes: 0, + }; + unsafe { + self.stream + .launch_builder(&self.bias_kernels.reduce_fn) + .arg(grad_logits) + .arg(grad_b) + .arg(&rows_i32) + .arg(&cols_i32) + .launch(cfg) + .context("dqn_bwd_grad_b reduce_sum_axis0")?; + } + Ok(()) + } + + /// Target-network forward via cuBLAS SGEMM (same math as online + /// forward, different weight/bias buffers). Used by the Bellman + /// target-bootstrap path. pub fn forward_target( &self, h_t: &CudaSlice, @@ -430,25 +694,9 @@ impl DqnHead { logits_out: &mut CudaSlice, ) -> Result<()> { debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM); - debug_assert_eq!(logits_out.len(), b_size * N_ACTIONS * Q_N_ATOMS); - - let b_i = b_size as i32; - let cfg = LaunchConfig { - grid_dim: (b_size as u32, N_ACTIONS as u32, 1), - block_dim: (Q_N_ATOMS as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.fwd_fn); - launch - .arg(&self.w_target_d) - .arg(&self.b_target_d) - .arg(h_t) - .arg(&b_i) - .arg(logits_out); - unsafe { - launch.launch(cfg).context("dqn_distributional_q_fwd (target) launch")?; - } - Ok(()) + debug_assert_eq!(logits_out.len(), b_size * K_OUT); + self.forward_gemm(&self.w_target_d, &self.b_target_d, h_t, b_size, logits_out) + .context("dqn forward (target)") } /// Extract per-batch action-row atoms from a full target-net output @@ -539,6 +787,58 @@ impl DqnHead { Ok(()) } + /// Fused `select_action_atoms` + `project_bellman_target` in a single + /// kernel launch. Reads the selected action's atom row from + /// `full_logits_d [B × N_ACTIONS × Q_N_ATOMS]` into shared memory, + /// then runs the Bellman categorical projection in-place — eliminating + /// the intermediate `[B × Q_N_ATOMS]` global memory buffer and one + /// kernel launch overhead. + /// + /// Produces identical output to the sequential + /// `select_action_atoms` → `project_bellman_target` pair. + #[allow(clippy::too_many_arguments)] + pub fn fused_select_and_project_bellman( + &self, + full_logits_d: &CudaSlice, + actions_d: &CudaSlice, + rewards_d: &CudaSlice, + dones_d: &CudaSlice, + n_step_gammas_d: &CudaSlice, + isv_dev_ptr: &u64, + b_size: usize, + target_dist_d: &mut CudaSlice, + ) -> Result<()> { + debug_assert_eq!(full_logits_d.len(), b_size * N_ACTIONS * Q_N_ATOMS); + debug_assert_eq!(actions_d.len(), b_size); + debug_assert_eq!(rewards_d.len(), b_size); + debug_assert_eq!(dones_d.len(), b_size); + debug_assert_eq!(n_step_gammas_d.len(), b_size); + debug_assert_eq!(target_dist_d.len(), b_size * Q_N_ATOMS); + + let b_i = b_size as i32; + let cfg = LaunchConfig { + grid_dim: (b_size as u32, 1, 1), + block_dim: (Q_N_ATOMS as u32, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.bellman_fused_fn); + launch + .arg(full_logits_d) + .arg(actions_d) + .arg(rewards_d) + .arg(dones_d) + .arg(n_step_gammas_d) + .arg(isv_dev_ptr) + .arg(&b_i) + .arg(target_dist_d); + unsafe { + launch + .launch(cfg) + .context("bellman_fused_select_project launch")?; + } + Ok(()) + } + /// Stream used to launch all kernels owned by this head. Phase E's /// training loop reads this when scheduling the soft-update kernel. pub fn stream(&self) -> &Arc { diff --git a/crates/ml-alpha/src/rl/iqn.rs b/crates/ml-alpha/src/rl/iqn.rs index 8ae7c2d1f..1aeac3708 100644 --- a/crates/ml-alpha/src/rl/iqn.rs +++ b/crates/ml-alpha/src/rl/iqn.rs @@ -12,21 +12,39 @@ //! //! where α is read from `ISV[RL_IQN_ENSEMBLE_ALPHA_INDEX=544]`. //! -//! ## Architecture +//! ## Architecture (cuBLAS-accelerated pipeline) +//! +//! The forward pass is split into custom kernels + cuBLAS SGEMMs: //! //! ```text -//! 1. Quantile embedding: -//! cos_feats[i] = cos((i+1) × π × τ) for i=0..EMBED_DIM-1 -//! phi(τ) = ReLU(W_embed × cos_feats + b_embed) → [HIDDEN_DIM] +//! 1. rl_iqn_tau_cos_features (custom): +//! Sample tau ~ U(0,1), compute cos_features[j] = cos((j+1)*pi*tau) +//! → [B*N_TAU, EMBED_DIM] //! -//! 2. Element-wise product: -//! combined = h_t ⊙ phi(τ) → [HIDDEN_DIM] +//! 2. cuBLAS SGEMM: +//! embed_out = cos_features @ W_embed → [B*N_TAU, HIDDEN_DIM] //! -//! 3. Action value projection: -//! Q(s, τ, a) = W_out × combined + b_out → [N_ACTIONS] +//! 3. rl_iqn_relu_hadamard (custom): +//! combined = h_t ⊙ ReLU(embed_out + b_embed) → [B*N_TAU, HIDDEN_DIM] +//! +//! 4. cuBLAS SGEMM: +//! q_raw = combined @ W_out → [B*N_TAU, N_ACTIONS] +//! +//! 5. rl_iqn_bias_add_q (custom): +//! q_values = q_raw + b_out → [B*N_TAU, N_ACTIONS] //! ``` //! -//! Expected Q for action selection: `E_IQN[a] = mean_τ Q(s, τ, a)`. +//! The backward pass uses one cuBLAS SGEMM for the largest matmul: +//! +//! ```text +//! 1. cuBLAS SGEMM: +//! grad_combined = grad_q @ W_out^T → [B*N_TAU, HIDDEN_DIM] +//! +//! 2. rl_iqn_backward (custom): +//! Per-batch backward accumulation. Recomputes phi(tau) and combined +//! inline from (tau, W_embed, b_embed, h_t). Reads grad_combined +//! from cuBLAS. Produces per-batch grad_w_out/b_out/w_embed/b_embed. +//! ``` //! //! ## Loss //! @@ -53,14 +71,18 @@ //! * `feedback_no_nvrtc.md` — pre-compiled cubins via `build.rs`. //! * `feedback_isv_for_adaptive_bounds.md` — N_TAU, ensemble α, and LR //! live in ISV slots 543-545. -//! * `feedback_cpu_is_read_only.md` — forward/backward are GPU kernels. +//! * `feedback_cpu_is_read_only.md` — forward/backward are GPU kernels +//! + cuBLAS SGEMMs. use std::sync::Arc; use anyhow::{Context, Result}; +use cudarc::cublas::CudaBlas; +use cudarc::cublas::sys::{self as cublas_sys, cublasOperation_t}; use cudarc::driver::{ - CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut, LaunchConfig, PushKernelArg, + CudaFunction, CudaModule, CudaSlice, CudaStream, }; +use cudarc::driver::sys::CUstream; use ml_core::cuda_autograd::init::scoped_init_seed; use ml_core::device::MlDevice; use rand::{Rng, SeedableRng}; @@ -69,6 +91,7 @@ use rand_chacha::ChaCha8Rng; use crate::heads::HIDDEN_DIM; use crate::pinned_mem::MappedF32Buffer; use crate::rl::common::N_ACTIONS; +use crate::trainer::raw_launch::{RawArgs, raw_launch}; /// Quantile embedding dimension — number of cosine basis functions. /// Matches the `EMBED_DIM` define in `cuda/rl_iqn_forward.cu`. @@ -78,6 +101,10 @@ pub const EMBED_DIM: usize = 64; /// via `RL_IQN_N_TAU_INDEX`. Matches the `N_TAU` default in the kernel. const DEFAULT_N_TAU: usize = 32; +/// cuBLAS workspace size — 8 MiB, prevents internal cudaMalloc during +/// CUDA Graph stream capture. +const CUBLAS_WORKSPACE_BYTES: usize = 8 * 1024 * 1024; + const IQN_FWD_CUBIN: &[u8] = include_bytes!(concat!( env!("OUT_DIR"), "/rl_iqn_forward.cubin" @@ -103,6 +130,9 @@ pub struct IqnHeadConfig { /// `DEFAULT_N_TAU` (32); overridden at runtime by /// `ISV[RL_IQN_N_TAU_INDEX]` in the training loop. pub n_tau: usize, + /// Maximum batch size for pre-allocated scratch buffers. + /// Must be >= any b_size passed to forward/backward. + pub max_batch_size: usize, } impl Default for IqnHeadConfig { @@ -111,91 +141,178 @@ impl Default for IqnHeadConfig { hidden_dim: HIDDEN_DIM, seed: 0x19A, n_tau: DEFAULT_N_TAU, + max_batch_size: 256, } } } -/// IQN distributional Q-head. Owns device weights for the quantile -/// embedding (`w_embed`, `b_embed`) and output projection (`w_out`, -/// `b_out`), plus a target network copy for Bellman bootstrapping. +// ── raw pointer helpers (same pattern as ml-core linear.rs) ────────── + +/// F32 × F32 → F32 GEMM via `cublasGemmEx` with F32 internal accumulation. /// -/// The head produces `Q(s, τ, a)` for sampled τ ∈ U(0,1). Expected Q -/// for action selection is `E[Q(s,a)] = mean_τ Q(s,τ,a)`, combined -/// with C51's E[Q] via the ensemble α weight from ISV. +/// Uses `CUBLAS_GEMM_DFALT` algo — deterministic, safe for CUDA Graph capture. +/// +/// # Safety +/// All device pointers must be valid and dimensions must be correct. +unsafe fn gemm_ex_f32( + cublas: &CudaBlas, + transa: cublasOperation_t, + transb: cublasOperation_t, + m: i32, + n: i32, + k: i32, + a_ptr: u64, + lda: i32, + b_ptr: u64, + ldb: i32, + c_ptr: u64, + ldc: i32, + label: &str, +) -> Result<()> { + let alpha = 1.0_f32; + let beta = 0.0_f32; + cudarc::cublas::result::gemm_ex( + *cublas.handle(), + transa, + transb, + m, + n, + k, + (&alpha as *const f32).cast(), + a_ptr as *const std::ffi::c_void, + cublas_sys::cudaDataType_t::CUDA_R_32F, + lda, + b_ptr as *const std::ffi::c_void, + cublas_sys::cudaDataType_t::CUDA_R_32F, + ldb, + (&beta as *const f32).cast(), + c_ptr as *mut std::ffi::c_void, + cublas_sys::cudaDataType_t::CUDA_R_32F, + ldc, + cublas_sys::cublasComputeType_t::CUBLAS_COMPUTE_32F, + cublas_sys::cublasGemmAlgo_t::CUBLAS_GEMM_DFALT, + ) + .map_err(|e| anyhow::anyhow!("cublasGemmEx IQN {label}: {e:?}"))?; + Ok(()) +} + +/// IQN distributional Q-head with cuBLAS-accelerated matmuls. +/// +/// Owns device weights for the quantile embedding (`w_embed`, `b_embed`) +/// and output projection (`w_out`, `b_out`), plus a target network copy +/// for Bellman bootstrapping. pub struct IqnHead { cfg: IqnHeadConfig, stream: Arc, + raw_stream: CUstream, + // ── cuBLAS ─────────────────────────────────────────────────────── + cublas: CudaBlas, + _cublas_workspace: CudaSlice, + + // ── Forward kernels (split pipeline) ───────────────────────────── _fwd_module: Arc, - /// Forward kernel: quantile embedding + action-value projection. - pub fwd_fn: CudaFunction, - /// Expected-Q reduction kernel: mean over tau dimension. + /// Stage 1: tau sampling + cosine basis computation. + pub tau_cos_features_fn: CudaFunction, + /// Stage 3: bias-add → ReLU → hadamard product with h_t. + pub relu_hadamard_fn: CudaFunction, + /// Stage 5: bias addition to cuBLAS output projection result. + pub bias_add_q_fn: CudaFunction, + /// Expected-Q reduction: mean over tau dimension. pub expected_q_fn: CudaFunction, _loss_module: Arc, - /// Loss kernel: quantile Huber loss forward + backward. pub loss_fn: CudaFunction, + // ── Backward kernels ───────────────────────────────────────────── _bwd_module: Arc, - /// Backward kernel: grad_output → per-batch grad_w_out/b_out/w_embed/b_embed. + /// Per-batch backward accumulation with cuBLAS-precomputed grad_combined. pub bwd_fn: CudaFunction, + /// Element-wise backward through hadamard + ReLU (for cached forward path). + pub bwd_relu_hadamard_fn: CudaFunction, - /// Quantile embedding weight `[EMBED_DIM, HIDDEN_DIM]`, row-major. + // ── Online weights ─────────────────────────────────────────────── pub w_embed_d: CudaSlice, - /// Quantile embedding bias `[HIDDEN_DIM]`. pub b_embed_d: CudaSlice, - /// Output projection weight `[HIDDEN_DIM, N_ACTIONS]`, row-major. pub w_out_d: CudaSlice, - /// Output projection bias `[N_ACTIONS]`. pub b_out_d: CudaSlice, - /// Target-network weights — same shapes as online. Soft-updated - /// using the same τ from `ISV[RL_TARGET_TAU_INDEX]` as C51. + // ── Target-network weights ─────────────────────────────────────── pub w_embed_target_d: CudaSlice, pub b_embed_target_d: CudaSlice, pub w_out_target_d: CudaSlice, pub b_out_target_d: CudaSlice, + + // ── Pre-allocated scratch for forward/backward (graph-capture safe) ─ + scratch_cos_features: CudaSlice, + scratch_embed_out: CudaSlice, + scratch_combined: CudaSlice, + scratch_embed_pre_relu: CudaSlice, + scratch_grad_combined: CudaSlice, } impl IqnHead { /// Allocate device weights, load cubins, and cache kernel handles. - /// Online and target networks start with IDENTICAL Xavier draws so - /// the first Bellman backup sees zero-divergence (same bootstrap - /// principle as C51). pub fn new(dev: &MlDevice, cfg: IqnHeadConfig) -> Result { let stream: Arc = dev.cuda_stream().context("iqn_head stream")?.clone(); let ctx = dev.cuda_context().context("iqn_head ctx")?; + // ── cuBLAS handle + pre-allocated workspace ────────────────── + let cublas = CudaBlas::new(Arc::clone(&stream)) + .map_err(|e| anyhow::anyhow!("IqnHead: cuBLAS init: {e}"))?; + let cublas_workspace = stream + .alloc_zeros::(CUBLAS_WORKSPACE_BYTES) + .map_err(|e| anyhow::anyhow!("IqnHead: cuBLAS workspace: {e}"))?; + unsafe { + let ws_ptr = cublas_workspace.raw_ptr(); + cudarc::cublas::sys::cublasSetWorkspace_v2( + *cublas.handle(), + ws_ptr as *mut std::ffi::c_void, + CUBLAS_WORKSPACE_BYTES, + ) + .result() + .map_err(|e| anyhow::anyhow!("IqnHead: cublasSetWorkspace_v2: {e:?}"))?; + } + + // ── Forward kernel symbols ─────────────────────────────────── let fwd_module = ctx .load_cubin(IQN_FWD_CUBIN.to_vec()) .context("load rl_iqn_forward cubin")?; - let fwd_fn = fwd_module - .load_function("rl_iqn_forward") - .context("load rl_iqn_forward fn")?; + let tau_cos_features_fn = fwd_module + .load_function("rl_iqn_tau_cos_features") + .context("load rl_iqn_tau_cos_features")?; + let relu_hadamard_fn = fwd_module + .load_function("rl_iqn_relu_hadamard") + .context("load rl_iqn_relu_hadamard")?; + let bias_add_q_fn = fwd_module + .load_function("rl_iqn_bias_add_q") + .context("load rl_iqn_bias_add_q")?; let expected_q_fn = fwd_module .load_function("rl_iqn_expected_q") - .context("load rl_iqn_expected_q fn")?; + .context("load rl_iqn_expected_q")?; let loss_module = ctx .load_cubin(IQN_LOSS_CUBIN.to_vec()) .context("load rl_iqn_loss cubin")?; let loss_fn = loss_module .load_function("rl_iqn_loss_fwd") - .context("load rl_iqn_loss_fwd fn")?; + .context("load rl_iqn_loss_fwd")?; + // ── Backward kernel symbols ────────────────────────────────── let bwd_module = ctx .load_cubin(IQN_BWD_CUBIN.to_vec()) .context("load rl_iqn_backward cubin")?; let bwd_fn = bwd_module .load_function("rl_iqn_backward") - .context("load rl_iqn_backward fn")?; + .context("load rl_iqn_backward")?; + let bwd_relu_hadamard_fn = bwd_module + .load_function("rl_iqn_bwd_relu_hadamard") + .context("load rl_iqn_bwd_relu_hadamard")?; - // Per pearl_scoped_init_seed_for_reproducibility: install the - // scoped seed guard BEFORE drawing any Xavier samples. + // ── Weight initialisation ──────────────────────────────────── let _seed_guard = scoped_init_seed(cfg.seed); let mut rng = ChaCha8Rng::seed_from_u64(cfg.seed); - // Xavier uniform scaled by 0.01 — initial Q outputs near zero. let embed_scale = 0.01_f32 * (6.0_f32 / (EMBED_DIM + cfg.hidden_dim) as f32).sqrt(); let w_embed_host: Vec = (0..EMBED_DIM * cfg.hidden_dim) @@ -215,22 +332,42 @@ impl IqnHead { let w_out_d = upload(&stream, &w_out_host)?; let b_out_d = upload(&stream, &b_out_host)?; - // Target network: identical init for zero-divergence bootstrap. let w_embed_target_d = upload(&stream, &w_embed_host)?; let b_embed_target_d = upload(&stream, &b_embed_host)?; let w_out_target_d = upload(&stream, &w_out_host)?; let b_out_target_d = upload(&stream, &b_out_host)?; + let raw_stream = stream.cu_stream(); + + let max_m = cfg.max_batch_size * cfg.n_tau; + let hd = cfg.hidden_dim; + let scratch_cos_features = stream.alloc_zeros::(max_m * EMBED_DIM) + .context("iqn scratch_cos_features")?; + let scratch_embed_out = stream.alloc_zeros::(max_m * hd) + .context("iqn scratch_embed_out")?; + let scratch_combined = stream.alloc_zeros::(max_m * hd) + .context("iqn scratch_combined")?; + let scratch_embed_pre_relu = stream.alloc_zeros::(max_m * hd) + .context("iqn scratch_embed_pre_relu")?; + let scratch_grad_combined = stream.alloc_zeros::(max_m * hd) + .context("iqn scratch_grad_combined")?; + Ok(Self { cfg, stream, + raw_stream, + cublas, + _cublas_workspace: cublas_workspace, _fwd_module: fwd_module, - fwd_fn, + tau_cos_features_fn, + relu_hadamard_fn, + bias_add_q_fn, expected_q_fn, _loss_module: loss_module, loss_fn, _bwd_module: bwd_module, bwd_fn, + bwd_relu_hadamard_fn, w_embed_d, b_embed_d, w_out_d, @@ -239,16 +376,17 @@ impl IqnHead { b_embed_target_d, w_out_target_d, b_out_target_d, + scratch_cos_features, + scratch_embed_out, + scratch_combined, + scratch_embed_pre_relu, + scratch_grad_combined, }) } - /// Fused tau-sampling + forward pass: samples tau ~ U(0,1) inline - /// via xorshift32 PRNG (eliminating the separate `rl_sample_tau` - /// kernel launch), then computes Q(s, τ, a) for all (batch, tau, - /// action) triples. - /// - /// Output: `tau [B, N_TAU]` (sampled quantile fractions, written by - /// the kernel), `q_values [B, N_TAU, N_ACTIONS]`. + // ── Forward ────────────────────────────────────────────────────── + + /// cuBLAS-accelerated forward pass with online weights. pub fn forward( &self, prng_state: &mut CudaSlice, @@ -258,39 +396,15 @@ impl IqnHead { n_tau: usize, q_values_out: &mut CudaSlice, ) -> Result<()> { - debug_assert_eq!(prng_state.len(), b_size); - debug_assert_eq!(h_t.len(), b_size * self.cfg.hidden_dim); - debug_assert_eq!(tau.len(), b_size * n_tau); - debug_assert_eq!(q_values_out.len(), b_size * n_tau * N_ACTIONS); - - let b_i = b_size as i32; - let n_tau_i = n_tau as i32; - let launch_cfg = LaunchConfig { - grid_dim: (b_size as u32, n_tau as u32, 1), - block_dim: (HIDDEN_DIM as u32, 1, 1), - shared_mem_bytes: (HIDDEN_DIM * std::mem::size_of::()) as u32, - }; - let mut launch = self.stream.launch_builder(&self.fwd_fn); - launch - .arg(prng_state) - .arg(h_t) - .arg(tau) - .arg(&self.w_embed_d) - .arg(&self.b_embed_d) - .arg(&self.w_out_d) - .arg(&self.b_out_d) - .arg(&b_i) - .arg(&n_tau_i) - .arg(q_values_out); - unsafe { - launch.launch(launch_cfg).context("rl_iqn_forward launch")?; - } - Ok(()) + self.forward_inner( + prng_state, h_t, tau, + &self.w_embed_d, &self.b_embed_d, + &self.w_out_d, &self.b_out_d, + b_size, n_tau, q_values_out, + ) } - /// Fused tau-sampling + target-network forward: same kernel as - /// online forward, but reads from the target weights. Used for the - /// Bellman bootstrap. + /// cuBLAS-accelerated forward pass with target weights. pub fn forward_target( &self, prng_state: &mut CudaSlice, @@ -299,43 +413,156 @@ impl IqnHead { b_size: usize, n_tau: usize, q_values_out: &mut CudaSlice, + ) -> Result<()> { + self.forward_inner( + prng_state, h_t, tau, + &self.w_embed_target_d, &self.b_embed_target_d, + &self.w_out_target_d, &self.b_out_target_d, + b_size, n_tau, q_values_out, + ) + } + + /// Internal forward — parameterised over weight slices. + /// + /// Pipeline: + /// 1. `rl_iqn_tau_cos_features` — tau + cos basis + /// 2. cuBLAS SGEMM: `cos_features @ W_embed` → `embed_out` + /// 3. `rl_iqn_relu_hadamard` — bias + ReLU + hadamard + /// 4. cuBLAS SGEMM: `combined @ W_out` → `q_raw` + /// 5. `rl_iqn_bias_add_q` — add b_out + #[allow(clippy::too_many_arguments)] + fn forward_inner( + &self, + prng_state: &mut CudaSlice, + h_t: &CudaSlice, + tau: &mut CudaSlice, + w_embed: &CudaSlice, + b_embed: &CudaSlice, + w_out: &CudaSlice, + b_out: &CudaSlice, + b_size: usize, + n_tau: usize, + q_values_out: &mut CudaSlice, ) -> Result<()> { debug_assert_eq!(prng_state.len(), b_size); debug_assert_eq!(h_t.len(), b_size * self.cfg.hidden_dim); debug_assert_eq!(tau.len(), b_size * n_tau); debug_assert_eq!(q_values_out.len(), b_size * n_tau * N_ACTIONS); - let b_i = b_size as i32; - let n_tau_i = n_tau as i32; - let launch_cfg = LaunchConfig { - grid_dim: (b_size as u32, n_tau as u32, 1), - block_dim: (HIDDEN_DIM as u32, 1, 1), - shared_mem_bytes: (HIDDEN_DIM * std::mem::size_of::()) as u32, - }; - let mut launch = self.stream.launch_builder(&self.fwd_fn); - launch - .arg(prng_state) - .arg(h_t) - .arg(tau) - .arg(&self.w_embed_target_d) - .arg(&self.b_embed_target_d) - .arg(&self.w_out_target_d) - .arg(&self.b_out_target_d) - .arg(&b_i) - .arg(&n_tau_i) - .arg(q_values_out); - unsafe { - launch - .launch(launch_cfg) - .context("rl_iqn_forward (target) launch")?; + let m = b_size * n_tau; + let hd = self.cfg.hidden_dim; + + // ── Stage 1: tau sampling + cos features ───────────────────── + debug_assert!(m * EMBED_DIM <= self.scratch_cos_features.len(), + "IQN scratch overflow: m={m} > max_m={}", self.scratch_cos_features.len() / EMBED_DIM); + let cos_ptr = self.scratch_cos_features.raw_ptr(); + let embed_ptr = self.scratch_embed_out.raw_ptr(); + let comb_ptr = self.scratch_combined.raw_ptr(); + let pre_relu_ptr = self.scratch_embed_pre_relu.raw_ptr(); + { + let mut args = RawArgs::new(); + args.push_ptr(prng_state.raw_ptr()); + args.push_i32(b_size as i32); + args.push_i32(n_tau as i32); + args.push_ptr(tau.raw_ptr()); + args.push_ptr(cos_ptr); + let mut ptrs = args.build_arg_ptrs(); + unsafe { + raw_launch( + self.tau_cos_features_fn.cu_function(), + (b_size as u32, n_tau as u32, 1), + (EMBED_DIM as u32, 1, 1), + std::mem::size_of::() as u32, + self.raw_stream, + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_iqn_tau_cos_features: {:?}", e))?; + } } + + // ── Stage 2: cuBLAS SGEMM — embed_out = cos_features @ W_embed ── + { + let w_ptr = w_embed.raw_ptr(); + unsafe { + gemm_ex_f32( + &self.cublas, + cublasOperation_t::CUBLAS_OP_N, + cublasOperation_t::CUBLAS_OP_N, + hd as i32, m as i32, EMBED_DIM as i32, + w_ptr, hd as i32, + cos_ptr, EMBED_DIM as i32, + embed_ptr, hd as i32, + "embed_fwd", + )?; + } + } + + // ── Stage 3: ReLU + hadamard ───────────────────────────────── + { + let block_x = 128u32.min(hd as u32); + let grid_y = (hd as u32 + block_x - 1) / block_x; + let mut args = RawArgs::new(); + args.push_ptr(embed_ptr); + args.push_ptr(b_embed.raw_ptr()); + args.push_ptr(h_t.raw_ptr()); + args.push_i32(m as i32); + args.push_i32(b_size as i32); + args.push_i32(n_tau as i32); + args.push_ptr(comb_ptr); + args.push_ptr(pre_relu_ptr); + let mut ptrs = args.build_arg_ptrs(); + unsafe { + raw_launch( + self.relu_hadamard_fn.cu_function(), + (m as u32, grid_y, 1), + (block_x, 1, 1), + 0, + self.raw_stream, + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_iqn_relu_hadamard: {:?}", e))?; + } + } + + // ── Stage 4: cuBLAS SGEMM — q_raw = combined @ W_out ──────── + { + let wo_ptr = w_out.raw_ptr(); + let q_ptr = q_values_out.raw_ptr(); + unsafe { + gemm_ex_f32( + &self.cublas, + cublasOperation_t::CUBLAS_OP_N, + cublasOperation_t::CUBLAS_OP_N, + N_ACTIONS as i32, m as i32, hd as i32, + wo_ptr, N_ACTIONS as i32, + comb_ptr, hd as i32, + q_ptr, N_ACTIONS as i32, + "out_proj_fwd", + )?; + } + } + + // ── Stage 5: bias add ──────────────────────────────────────── + { + let mut args = RawArgs::new(); + args.push_ptr(q_values_out.raw_ptr()); + args.push_ptr(b_out.raw_ptr()); + args.push_i32(m as i32); + let mut ptrs = args.build_arg_ptrs(); + unsafe { + raw_launch( + self.bias_add_q_fn.cu_function(), + (m as u32, 1, 1), + (N_ACTIONS as u32, 1, 1), + 0, + self.raw_stream, + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_iqn_bias_add_q: {:?}", e))?; + } + } + Ok(()) } /// Compute expected Q values: `E_IQN[b,a] = mean_τ Q(s,τ,a)`. - /// - /// Input: `q_values [B, N_TAU, N_ACTIONS]` from [`forward`]. - /// Output: `expected_q [B, N_ACTIONS]`. pub fn expected_q( &self, q_values: &CudaSlice, @@ -346,32 +573,26 @@ impl IqnHead { debug_assert_eq!(q_values.len(), b_size * n_tau * N_ACTIONS); debug_assert_eq!(expected_q_out.len(), b_size * N_ACTIONS); - let b_i = b_size as i32; - let n_tau_i = n_tau as i32; - let launch_cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (N_ACTIONS as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.expected_q_fn); - launch - .arg(q_values) - .arg(&b_i) - .arg(&n_tau_i) - .arg(expected_q_out); + let mut args = RawArgs::new(); + args.push_ptr(q_values.raw_ptr()); + args.push_i32(b_size as i32); + args.push_i32(n_tau as i32); + args.push_ptr(expected_q_out.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(launch_cfg) - .context("rl_iqn_expected_q launch")?; + raw_launch( + self.expected_q_fn.cu_function(), + (b_size as u32, 1, 1), + (N_ACTIONS as u32, 1, 1), + 0, + self.raw_stream, + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_iqn_expected_q: {:?}", e))?; } Ok(()) } /// Quantile Huber loss forward + backward for the taken action. - /// - /// Computes per-batch loss and gradient w.r.t. the online Q values. - /// The gradient is only non-zero at the taken action's quantile - /// slice — all other actions receive zero gradient. #[allow(clippy::too_many_arguments)] pub fn compute_loss( &self, @@ -392,36 +613,39 @@ impl IqnHead { debug_assert_eq!(loss_per_batch.len(), b_size); debug_assert_eq!(grad_online_q.len(), b_size * n_tau_online * N_ACTIONS); - let b_i = b_size as i32; - let n_tau_online_i = n_tau_online as i32; - let n_tau_target_i = n_tau_target as i32; - let launch_cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (n_tau_online as u32, 1, 1), - // Shared memory for the tree-reduce over online quantiles. - shared_mem_bytes: (n_tau_online * std::mem::size_of::()) as u32, - }; - let mut launch = self.stream.launch_builder(&self.loss_fn); - launch - .arg(online_q) - .arg(target_q) - .arg(tau_online) - .arg(actions_taken) - .arg(&b_i) - .arg(&n_tau_online_i) - .arg(&n_tau_target_i) - .arg(loss_per_batch) - .arg(grad_online_q); + let mut args = RawArgs::new(); + args.push_ptr(online_q.raw_ptr()); + args.push_ptr(target_q.raw_ptr()); + args.push_ptr(tau_online.raw_ptr()); + args.push_ptr(actions_taken.raw_ptr()); + args.push_i32(b_size as i32); + args.push_i32(n_tau_online as i32); + args.push_i32(n_tau_target as i32); + args.push_ptr(loss_per_batch.raw_ptr()); + args.push_ptr(grad_online_q.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(launch_cfg).context("rl_iqn_loss_fwd launch")?; + raw_launch( + self.loss_fn.cu_function(), + (b_size as u32, 1, 1), + (n_tau_online as u32, 1, 1), + (n_tau_online * std::mem::size_of::()) as u32, + self.raw_stream, + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_iqn_loss_fwd: {:?}", e))?; } Ok(()) } - /// Backward through the forward pass: given `grad_output [B, N_TAU, N_ACTIONS]` - /// (from `compute_loss`), produce per-batch gradients for w_out, b_out, - /// w_embed, b_embed. Caller is responsible for `reduce_axis0` across - /// batches and feeding the reduced gradients to Adam. + // ── Backward ───────────────────────────────────────────────────── + + /// cuBLAS-accelerated backward through the forward pass. + /// + /// Pipeline: + /// 1. cuBLAS SGEMM: `grad_combined = grad_q @ W_out^T` + /// 2. `rl_iqn_backward`: per-batch accumulation using cuBLAS- + /// precomputed `grad_combined` (recomputes phi/combined inline + /// from tau + weights, same as original kernel). #[allow(clippy::too_many_arguments)] pub fn backward( &self, @@ -443,31 +667,61 @@ impl IqnHead { debug_assert_eq!(grad_w_embed_pb.len(), b_size * EMBED_DIM * self.cfg.hidden_dim); debug_assert_eq!(grad_b_embed_pb.len(), b_size * self.cfg.hidden_dim); - let b_i = b_size as i32; - let n_tau_i = n_tau as i32; - let launch_cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (self.cfg.hidden_dim as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.bwd_fn); - launch - .arg(h_t) - .arg(tau) - .arg(&self.w_embed_d) - .arg(&self.b_embed_d) - .arg(&self.w_out_d) - .arg(grad_output) - .arg(&b_i) - .arg(&n_tau_i) - .arg(grad_w_out_pb) - .arg(grad_b_out_pb) - .arg(grad_w_embed_pb) - .arg(grad_b_embed_pb); - unsafe { - launch - .launch(launch_cfg) - .context("rl_iqn_backward launch")?; + let m = b_size * n_tau; + let hd = self.cfg.hidden_dim; + + // ── cuBLAS: grad_combined = grad_q @ W_out^T ───────────────── + // Row-major: C[M, hd] = A[M, N_ACTIONS] @ B^T[N_ACTIONS, hd] + // cuBLAS col-major: C_col[hd, M] = W_out_col^T[hd, N_ACTIONS] @ grad_q_col[N_ACTIONS, M] + // W_out is row-major [hd, N_ACTIONS] = col-major [N_ACTIONS, hd]. + // transA=T on col-major [N_ACTIONS, hd] → [hd, N_ACTIONS]. + // transB=N on grad_q col-major [N_ACTIONS, M]. + // m=hd, n=M, k=N_ACTIONS, lda=N_ACTIONS, ldb=N_ACTIONS, ldc=hd + debug_assert!(m * hd <= self.scratch_grad_combined.len()); + let gc_ptr = self.scratch_grad_combined.raw_ptr(); + { + let wo_ptr = self.w_out_d.raw_ptr(); + let gq_ptr = grad_output.raw_ptr(); + unsafe { + gemm_ex_f32( + &self.cublas, + cublasOperation_t::CUBLAS_OP_T, + cublasOperation_t::CUBLAS_OP_N, + hd as i32, m as i32, N_ACTIONS as i32, + wo_ptr, N_ACTIONS as i32, + gq_ptr, N_ACTIONS as i32, + gc_ptr, hd as i32, + "grad_combined_bwd", + )?; + } + } + + // ── Per-batch backward kernel ──────────────────────────────── + { + let mut args = RawArgs::new(); + args.push_ptr(h_t.raw_ptr()); + args.push_ptr(tau.raw_ptr()); + args.push_ptr(self.w_embed_d.raw_ptr()); + args.push_ptr(self.b_embed_d.raw_ptr()); + args.push_ptr(gc_ptr); + args.push_ptr(grad_output.raw_ptr()); + args.push_i32(b_size as i32); + args.push_i32(n_tau as i32); + args.push_ptr(grad_w_out_pb.raw_ptr()); + args.push_ptr(grad_b_out_pb.raw_ptr()); + args.push_ptr(grad_w_embed_pb.raw_ptr()); + args.push_ptr(grad_b_embed_pb.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); + unsafe { + raw_launch( + self.bwd_fn.cu_function(), + (b_size as u32, 1, 1), + (self.cfg.hidden_dim as u32, 1, 1), + 0, + self.raw_stream, + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_iqn_backward: {:?}", e))?; + } } Ok(()) } @@ -495,20 +749,20 @@ fn upload(stream: &Arc, host: &[f32]) -> Result> { let staging = unsafe { MappedF32Buffer::new(n) } .map_err(|e| anyhow::anyhow!("iqn_head upload staging: {e}"))?; staging.write_from_slice(host); - let mut dst = stream + let dst = stream .alloc_zeros::(n) .context("iqn_head upload alloc")?; if n > 0 { let nbytes = n * std::mem::size_of::(); unsafe { - let (dst_ptr, _g) = dst.device_ptr_mut(stream); - cudarc::driver::result::memcpy_dtod_async( + let dst_ptr = dst.raw_ptr(); + crate::trainer::raw_launch::raw_memcpy_dtod_async( dst_ptr, staging.dev_ptr, nbytes, stream.cu_stream(), ) - .context("iqn_head upload DtoD")?; + .map_err(|e| anyhow::anyhow!("iqn_head upload DtoD: {e:?}"))?; } } Ok(dst) diff --git a/crates/ml-alpha/src/rl/outcome_head.rs b/crates/ml-alpha/src/rl/outcome_head.rs index bb2dddc08..3f2516843 100644 --- a/crates/ml-alpha/src/rl/outcome_head.rs +++ b/crates/ml-alpha/src/rl/outcome_head.rs @@ -38,11 +38,13 @@ pub struct OutcomeHead { _module_ce: Arc, _module_label: Arc, _module_bwd: Arc, + _module_fused: Arc, kernel_fwd: CudaFunction, kernel_ce: CudaFunction, kernel_label: CudaFunction, kernel_fill_sentinel: CudaFunction, kernel_bwd: CudaFunction, + kernel_fused: CudaFunction, } impl OutcomeHead { @@ -100,6 +102,16 @@ impl OutcomeHead { .load_function("rl_outcome_bwd") .map_err(|e| anyhow!("OutcomeHead: bwd kernel resolve: {e}"))?; + static CUBIN_FUSED: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_outcome_fused.cubin")); + let module_fused = stream + .context() + .load_cubin(CUBIN_FUSED.to_vec()) + .map_err(|e| anyhow!("OutcomeHead: fused cubin load: {e}"))?; + let kernel_fused = module_fused + .load_function("rl_outcome_fused") + .map_err(|e| anyhow!("OutcomeHead: fused kernel resolve: {e}"))?; + // Xavier init W at 0.01x scale: near-zero for near-uniform softmax at init. let w_d = ml_core::cuda_autograd::init::near_zero_xavier( HIDDEN_DIM, K_CLASSES, &stream, @@ -156,11 +168,13 @@ impl OutcomeHead { _module_ce: module_ce, _module_label: module_label, _module_bwd: module_bwd, + _module_fused: module_fused, kernel_fwd, kernel_ce, kernel_label, kernel_fill_sentinel, kernel_bwd, + kernel_fused, }) } @@ -278,6 +292,62 @@ impl OutcomeHead { Ok(()) } + /// Fused forward + CE loss + backward in a single kernel launch. + /// + /// Eliminates two global-memory round-trips for logits and grad_logits + /// by keeping both in shared memory. Replaces the sequential triple + /// `forward()` -> `compute_loss()` -> `backward()`. + /// + /// Labels must be set via `assign_labels()` before calling. + /// The caller is responsible for the same post-steps as `backward()`: + /// 1. `reduce_axis0` on `grad_w_per_batch_d` and `grad_b_per_batch_d`. + /// 2. Adam step on `self.w_d` / `self.b_d` with the reduced grads. + /// 3. `grad_h_accumulate` to fold `grad_h_t_d` into the encoder grad. + #[allow(clippy::too_many_arguments)] + pub fn forward_loss_backward( + &mut self, + h_t: &CudaSlice, + grad_w_per_batch_d: &mut CudaSlice, + grad_b_per_batch_d: &mut CudaSlice, + grad_h_t_d: &mut CudaSlice, + actual_b: usize, + ) -> Result<()> { + if actual_b > self.b_size { + return Err(anyhow!( + "OutcomeHead::forward_loss_backward: actual_b ({actual_b}) > allocated b_size ({})", + self.b_size + )); + } + debug_assert_eq!(h_t.len(), actual_b * HIDDEN_DIM); + debug_assert_eq!(grad_w_per_batch_d.len(), actual_b * HIDDEN_DIM * K_CLASSES); + debug_assert_eq!(grad_b_per_batch_d.len(), actual_b * K_CLASSES); + debug_assert_eq!(grad_h_t_d.len(), actual_b * HIDDEN_DIM); + + let b_i32 = actual_b as i32; + // shared memory: s_logits[3] + s_grad_logits[3] = 6 floats = 24 bytes + let cfg = LaunchConfig { + grid_dim: (actual_b as u32, 1, 1), + block_dim: (HIDDEN_DIM as u32, 1, 1), + shared_mem_bytes: 0, // statically declared in kernel + }; + unsafe { + self.stream + .launch_builder(&self.kernel_fused) + .arg(h_t) + .arg(&self.w_d) + .arg(&self.b_d) + .arg(&self.labels_d) + .arg(&b_i32) + .arg(&mut self.loss_pb_d) + .arg(grad_w_per_batch_d) + .arg(grad_b_per_batch_d) + .arg(grad_h_t_d) + .launch(cfg) + .map_err(|e| anyhow!("rl_outcome_fused launch: {e}"))?; + } + Ok(()) + } + /// Assign labels from reward/done signals. /// /// `rewards` and `dones` must be device slices of length `actual_b`. diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index 74371d202..c3c8cd734 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -33,9 +33,9 @@ //! batch log π / entropy diagnostics. (V-loss path moved out — see #4) //! 4. Bellman target build (kernel path, item 2): //! a. `dqn_head.forward_target(...)` → `q_target_logits [B × N_ACTIONS × Q_N_ATOMS]` -//! b. `dqn_head.select_action_atoms(...)` → `[B × Q_N_ATOMS]` slice -//! c. `dqn_head.project_bellman_target(...)` → `target_dist [B × Q_N_ATOMS]`, -//! reading γ from ISV[RL_GAMMA_INDEX=400] +//! b. `dqn_head.fused_select_and_project_bellman(...)` → `target_dist [B × Q_N_ATOMS]` +//! (action-atom extraction + C51 projection in one kernel launch, +//! reading γ from ISV[RL_GAMMA_INDEX=400]) //! 5. `dqn_head.backward_logits(...)` → `q_loss [1]` + `grad_logits`. //! 6. `dqn_head.backward_to_w_b_h(...)` → per-batch grad_w / grad_b //! scratch + grad_h_t (OVERWRITE). @@ -1044,10 +1044,10 @@ pub struct IntegratedTrainer { pub ss_pi_log_prob_d: CudaSlice, pub ss_entropy_d: CudaSlice, - // Q-head per-batch + reduced gradient buffers. + // Q-head gradient buffers. backward_gemm produces grad_w / grad_b + // directly reduced (no per-batch scratch needed — cuBLAS SGEMM + // computes grad_logits^T @ h_t as a single GEMM). pub ss_q_grad_logits_d: CudaSlice, - pub ss_q_grad_w_per_batch_d: CudaSlice, - pub ss_q_grad_b_per_batch_d: CudaSlice, pub ss_q_grad_h_t_d: CudaSlice, pub ss_q_grad_w_d: CudaSlice, pub ss_q_grad_b_d: CudaSlice, @@ -1079,7 +1079,6 @@ pub struct IntegratedTrainer { // Bellman projection scratch. pub ss_q_target_full_d: CudaSlice, - pub ss_q_target_action_d: CudaSlice, pub ss_target_dist_d: CudaSlice, // FRD backward chain scratch (step_synthetic only). @@ -1144,6 +1143,7 @@ impl IntegratedTrainer { hidden_dim: HIDDEN_DIM, seed: cfg.dqn_seed.wrapping_add(100), n_tau: 32, + max_batch_size: cfg.perception.n_batch, }, ) .context("IqnHead::new")?; @@ -2073,14 +2073,10 @@ impl IntegratedTrainer { let ss_entropy_d = stream.alloc_zeros::(b_size) .context("alloc ss_entropy_d")?; - // Q-head per-batch + reduced gradient buffers. + // Q-head gradient buffers. backward_gemm produces grad_w / grad_b + // directly reduced via cuBLAS SGEMM — no per-batch scratch needed. let ss_q_grad_logits_d = stream.alloc_zeros::(b_size * k_dqn_ss) .context("alloc ss_q_grad_logits_d")?; - let ss_q_grad_w_per_batch_d = stream - .alloc_zeros::(b_size * k_dqn_ss * HIDDEN_DIM) - .context("alloc ss_q_grad_w_per_batch_d")?; - let ss_q_grad_b_per_batch_d = stream.alloc_zeros::(b_size * k_dqn_ss) - .context("alloc ss_q_grad_b_per_batch_d")?; let ss_q_grad_h_t_d = stream.alloc_zeros::(b_size * HIDDEN_DIM) .context("alloc ss_q_grad_h_t_d")?; let ss_q_grad_w_d = stream.alloc_zeros::(k_dqn_ss * HIDDEN_DIM) @@ -2140,8 +2136,6 @@ impl IntegratedTrainer { // Bellman projection scratch. let ss_q_target_full_d = stream.alloc_zeros::(b_size * k_dqn_ss) .context("alloc ss_q_target_full_d")?; - let ss_q_target_action_d = stream.alloc_zeros::(b_size * Q_N_ATOMS) - .context("alloc ss_q_target_action_d")?; let ss_target_dist_d = stream.alloc_zeros::(b_size * Q_N_ATOMS) .context("alloc ss_target_dist_d")?; @@ -2492,8 +2486,6 @@ impl IntegratedTrainer { ss_pi_log_prob_d, ss_entropy_d, ss_q_grad_logits_d, - ss_q_grad_w_per_batch_d, - ss_q_grad_b_per_batch_d, ss_q_grad_h_t_d, ss_q_grad_w_d, ss_q_grad_b_d, @@ -2515,7 +2507,6 @@ impl IntegratedTrainer { ss_outcome_grad_w_d, ss_outcome_grad_b_d, ss_q_target_full_d, - ss_q_target_action_d, ss_target_dist_d, ss_frd_grad_logits_d, ss_frd_loss_mapped, @@ -3835,7 +3826,6 @@ impl IntegratedTrainer { // element, and buffers were alloc_zeros'd at construction. The // memsets replayed on every graph launch, burning ~36 × memset // bandwidth per step for zero correctness benefit. - let k_dqn = N_ACTIONS * Q_N_ATOMS; // ── Step 4 (Phase R7a): inputs are trainer-owned device buffers, // populated by step_with_lobsim's GPU pipeline before this call. @@ -3920,16 +3910,9 @@ impl IntegratedTrainer { .forward_target(&self.sampled_h_tp1_d, b_size, &mut self.ss_q_target_full_d) .context("dqn_head.forward_target(sampled_h_tp1) [R7d off-policy]")?; self.dqn_head - .select_action_atoms( + .fused_select_and_project_bellman( &self.ss_q_target_full_d, &self.sampled_next_actions_d, - b_size, - &mut self.ss_q_target_action_d, - ) - .context("dqn_head.select_action_atoms (sampled_next_actions)")?; - self.dqn_head - .project_bellman_target( - &self.ss_q_target_action_d, &self.sampled_rewards_d, &self.sampled_dones_d, &self.sampled_n_step_gammas_d, @@ -3937,7 +3920,7 @@ impl IntegratedTrainer { b_size, &mut self.ss_target_dist_d, ) - .context("dqn_head.project_bellman_target (sampled rewards/dones)")?; + .context("dqn_head.fused_select_and_project_bellman (sampled rewards/dones)")?; // ── Step 6b: DQN backward (logits → grad_w/b/h_t) ──────────── self.dqn_head @@ -3954,29 +3937,26 @@ impl IntegratedTrainer { // Host read deferred to after graph capture region. self.dqn_head - .backward_to_w_b_h( + .backward_gemm( &self.sampled_h_t_d, &self.ss_q_grad_logits_d, b_size, - &mut self.ss_q_grad_w_per_batch_d, - &mut self.ss_q_grad_b_per_batch_d, + &mut self.ss_q_grad_w_d, + &mut self.ss_q_grad_b_d, &mut self.ss_q_grad_h_t_d, ) - .context("dqn_head.backward_to_w_b_h(sampled_h_t) [R7d off-policy]")?; + .context("dqn_head.backward_gemm(sampled_h_t) [R7d off-policy]")?; // R7d stop-grad: ss_q_grad_h_t_d is the gradient wrt SAMPLED h_t // (a past-step encoder output). Accumulating it into the // shared encoder via grad_h_t_combined_d would poison the // encoder with stale-state gradient signal. Standard pattern // for off-policy + shared encoder (SAC / R2D2 / IMPALA do the - // same). The buffer is pre-allocated for backward_to_w_b_h + // same). The buffer is pre-allocated for backward_gemm // to write into; we just don't FOLD it into the encoder grad // combine below. Encoder learns from π + V (on-policy) + // BCE/aux (supervised via step_batched) only. let _ = &self.ss_q_grad_h_t_d; - reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut self.ss_q_grad_w_d)?; - reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_b_per_batch_d, b_size, k_dqn, &mut self.ss_q_grad_b_d)?; - // ── Step 7: PPO backward (logits → grad_w/b/h_t) ───────────── self.policy_head .surrogate_backward_logits( @@ -4122,27 +4102,23 @@ impl IntegratedTrainer { } } - // ── Step 9c: Outcome head forward + CE loss + backward ──────── + // ── Step 9c: Outcome head fused forward + CE + backward ──────── // K=3 outcome classifier trained from reward/done labels assigned - // in step_with_lobsim. Forward and loss run on sampled_h_t_d. - // Backward: grad_logits → grad_W/b (per-batch), grad_h_t. - // reduce_axis0 → batch-summed grad_W/b. Adam on outcome W/b. + // in step_with_lobsim. Single fused kernel: fwd → CE → bwd with + // logits and grad_logits kept in shared memory (eliminates two + // global round-trips). reduce_axis0 → batch-summed grad_W/b. + // Adam on outcome W/b. // Lambda from ISV[RL_OUTCOME_AUX_LAMBDA_INDEX] scales the // encoder-upstream grad_h_t contribution in Step 10. { let h_t_borrow_for_outcome: &CudaSlice = &self.sampled_h_t_d; - self.outcome_head.forward(h_t_borrow_for_outcome, b_size) - .context("step_synthetic: outcome_head.forward")?; - self.outcome_head.compute_loss(b_size) - .context("step_synthetic: outcome_head.compute_loss")?; - // Backward: grad_logits → per-batch grad_W, grad_b, grad_h_t. - self.outcome_head.backward( + self.outcome_head.forward_loss_backward( h_t_borrow_for_outcome, &mut self.ss_outcome_grad_w_pb_d, &mut self.ss_outcome_grad_b_pb_d, &mut self.ss_outcome_grad_h_t_d, b_size, - ).context("step_synthetic: outcome_head.backward")?; + ).context("step_synthetic: outcome_head.forward_loss_backward")?; // Reduce-axis-0: batch-summed weight and bias gradients. let k_classes = crate::rl::outcome_head::K_CLASSES; reduce_axis0_free( @@ -4596,7 +4572,7 @@ impl IntegratedTrainer { /// 2. Double-DQN argmax on online Q at h_tp1 → /// `sampled_next_actions_d` /// 3. Bellman target via `forward_target` (target net) + - /// `select_action_atoms` + `project_bellman_target` + /// `fused_select_and_project_bellman` /// 4. Q backward: logits → grad_w/b/h_t. Q's grad_h_t is /// discarded (R7d stop-grad: SAMPLED h_t is a past-step /// encoder output; folding its gradient into the encoder @@ -4622,8 +4598,6 @@ impl IntegratedTrainer { debug_assert_eq!(self.sampled_h_t_d.len(), b_size * HIDDEN_DIM); debug_assert_eq!(self.sampled_h_tp1_d.len(), b_size * HIDDEN_DIM); - let k_dqn = N_ACTIONS * Q_N_ATOMS; - // Zero persistent per-iter scratch — raw memset bypasses cudarc's // bind_to_thread + event tracking (~200us overhead per call). { @@ -4633,16 +4607,10 @@ impl IntegratedTrainer { .map_err(|e| anyhow::anyhow!("zero ss_q_loss: {:?}", e))?; raw_memset_d8_zero(self.ss_q_target_full_d.raw_ptr(), self.ss_q_target_full_d.num_bytes(), s) .map_err(|e| anyhow::anyhow!("zero ss_q_target_full_d: {:?}", e))?; - raw_memset_d8_zero(self.ss_q_target_action_d.raw_ptr(), self.ss_q_target_action_d.num_bytes(), s) - .map_err(|e| anyhow::anyhow!("zero ss_q_target_action_d: {:?}", e))?; raw_memset_d8_zero(self.ss_target_dist_d.raw_ptr(), self.ss_target_dist_d.num_bytes(), s) .map_err(|e| anyhow::anyhow!("zero ss_target_dist_d: {:?}", e))?; raw_memset_d8_zero(self.ss_q_grad_logits_d.raw_ptr(), self.ss_q_grad_logits_d.num_bytes(), s) .map_err(|e| anyhow::anyhow!("zero ss_q_grad_logits_d: {:?}", e))?; - raw_memset_d8_zero(self.ss_q_grad_w_per_batch_d.raw_ptr(), self.ss_q_grad_w_per_batch_d.num_bytes(), s) - .map_err(|e| anyhow::anyhow!("zero ss_q_grad_w_per_batch_d: {:?}", e))?; - raw_memset_d8_zero(self.ss_q_grad_b_per_batch_d.raw_ptr(), self.ss_q_grad_b_per_batch_d.num_bytes(), s) - .map_err(|e| anyhow::anyhow!("zero ss_q_grad_b_per_batch_d: {:?}", e))?; raw_memset_d8_zero(self.ss_q_grad_h_t_d.raw_ptr(), self.ss_q_grad_h_t_d.num_bytes(), s) .map_err(|e| anyhow::anyhow!("zero ss_q_grad_h_t_d: {:?}", e))?; raw_memset_d8_zero(self.ss_q_grad_w_d.raw_ptr(), self.ss_q_grad_w_d.num_bytes(), s) @@ -4821,16 +4789,9 @@ impl IntegratedTrainer { .forward_target(&self.sampled_h_tp1_d, b_size, &mut self.ss_q_target_full_d) .context("dqn_replay_step: dqn_head.forward_target(sampled_h_tp1)")?; self.dqn_head - .select_action_atoms( + .fused_select_and_project_bellman( &self.ss_q_target_full_d, &self.sampled_next_actions_d, - b_size, - &mut self.ss_q_target_action_d, - ) - .context("dqn_replay_step: dqn_head.select_action_atoms")?; - self.dqn_head - .project_bellman_target( - &self.ss_q_target_action_d, &self.sampled_rewards_d, &self.sampled_dones_d, &self.sampled_n_step_gammas_d, @@ -4838,7 +4799,7 @@ impl IntegratedTrainer { b_size, &mut self.ss_target_dist_d, ) - .context("dqn_replay_step: dqn_head.project_bellman_target")?; + .context("dqn_replay_step: dqn_head.fused_select_and_project_bellman")?; // ── 4. Q backward (logits → grad_w/b/h_t) ─────────────────── self.dqn_head @@ -4862,23 +4823,20 @@ impl IntegratedTrainer { let l_q = l_q_host / (b_size as f32).max(1.0); self.dqn_head - .backward_to_w_b_h( + .backward_gemm( &self.sampled_h_t_d, &self.ss_q_grad_logits_d, b_size, - &mut self.ss_q_grad_w_per_batch_d, - &mut self.ss_q_grad_b_per_batch_d, + &mut self.ss_q_grad_w_d, + &mut self.ss_q_grad_b_d, &mut self.ss_q_grad_h_t_d, ) - .context("dqn_replay_step: dqn_head.backward_to_w_b_h")?; + .context("dqn_replay_step: dqn_head.backward_gemm")?; // R7d stop-grad: discard ss_q_grad_h_t_d (sampled h_t is past-step // encoder output; folding its gradient into the encoder would // poison live training). Same semantics as step_synthetic. let _ = &self.ss_q_grad_h_t_d; - reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut self.ss_q_grad_w_d)?; - reduce_axis0_free(&self.stream, &self.reduce_axis0_fn, &self.ss_q_grad_b_per_batch_d, b_size, k_dqn, &mut self.ss_q_grad_b_d)?; - // ── 5. Q Adam (uses LR set by step_synthetic; we don't re-fire // the LR controller here — that runs once per env step). self.dqn_w_adam @@ -4983,23 +4941,18 @@ impl IntegratedTrainer { } } - // ── 8. Outcome head forward + CE loss + backward ───────────── - // Forward: logits = sampled_h_t × W + b. - self.outcome_head.forward(&self.sampled_h_t_d, b_size) - .context("dqn_replay_step: outcome_head.forward")?; - // CE loss + grad w.r.t. logits (masked by sentinel -1 labels). - self.outcome_head.compute_loss(b_size) - .context("dqn_replay_step: outcome_head.compute_loss")?; - // Backward: grad_logits → per-batch grad_W, grad_b, grad_h_t. + // ── 8. Outcome head fused forward + CE + backward ───────────── + // Single fused kernel: fwd → CE → bwd with logits and grad_logits + // in shared memory (eliminates two global round-trips). // grad_h_t is computed but discarded (no encoder backward in // replay steps — same stop-grad rationale as Q's grad_h_t). - self.outcome_head.backward( + self.outcome_head.forward_loss_backward( &self.sampled_h_t_d, &mut self.ss_outcome_grad_w_pb_d, &mut self.ss_outcome_grad_b_pb_d, &mut self.ss_outcome_grad_h_t_d, b_size, - ).context("dqn_replay_step: outcome_head.backward")?; + ).context("dqn_replay_step: outcome_head.forward_loss_backward")?; { let k_classes = crate::rl::outcome_head::K_CLASSES; reduce_axis0_free(