diff --git a/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs b/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs index 7dc596cb0..986ebb3c2 100644 --- a/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs +++ b/crates/ml/src/cuda_pipeline/gpu_iqn_head.rs @@ -338,6 +338,8 @@ impl GpuIqnHead { dqn_dones_buf: &CudaSlice, ) -> Result { let b = self.config.batch_size; + let shared_h1_i32 = self.config.shared_h1 as i32; + let hidden_dim_i32 = self.config.hidden_dim as i32; // Layer 3a: τ values are fixed midpoints, pre-computed in constructor. // No per-step sampling → deterministic → CUDA Graph compatible. @@ -358,6 +360,8 @@ impl GpuIqnHead { .arg(dqn_actions_buf) .arg(&mut self.branch_actions) .arg(&batch_size_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(decode_config) .map_err(|e| MLError::ModelError(format!("IQN decode_actions kernel: {e}")))?; } @@ -388,6 +392,8 @@ impl GpuIqnHead { target_dueling: &DuelingWeightSet, ) -> Result { let b = self.config.batch_size; + let shared_h1_i32 = self.config.shared_h1 as i32; + let hidden_dim_i32 = self.config.hidden_dim as i32; // 4. Compute target h_s2 via trunk forward kernel let shmem_bytes = self.config.shared_h1 * 4; @@ -412,6 +418,8 @@ impl GpuIqnHead { .arg(&mut self.target_h_s2) .arg(&(b as i32)) .arg(&state_dim_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(trunk_config) .map_err(|e| MLError::ModelError(format!("IQN trunk forward kernel: {e}")))?; } @@ -458,6 +466,8 @@ impl GpuIqnHead { .arg(&mut self.save_q_online) .arg(&mut self.save_q_target) .arg(&batch_size_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(fwd_config) .map_err(|e| MLError::ModelError(format!("IQN forward+loss kernel: {e}")))?; } @@ -484,6 +494,8 @@ impl GpuIqnHead { .arg(&mut self.grad_buf) .arg(&mut self.d_h_s2_buf) // IQN trunk gradient output .arg(&batch_size_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(bwd_config) .map_err(|e| MLError::ModelError(format!("IQN backward kernel: {e}")))?; } @@ -504,6 +516,8 @@ impl GpuIqnHead { .arg(&self.grad_buf) .arg(&mut self.grad_norm_buf) .arg(&total_params_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(norm_config) .map_err(|e| MLError::ModelError(format!("IQN grad_norm kernel: {e}")))?; } @@ -542,6 +556,8 @@ impl GpuIqnHead { .arg(&mgn) .arg(&t) .arg(&total_params_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(adam_config) .map_err(|e| MLError::ModelError(format!("IQN adam kernel: {e}")))?; } @@ -564,6 +580,8 @@ impl GpuIqnHead { }; let n_i32 = n as i32; + let shared_h1_i32 = self.config.shared_h1 as i32; + let hidden_dim_i32 = self.config.hidden_dim as i32; // Safety: target_params and online_params have total_params elements each. unsafe { @@ -573,6 +591,8 @@ impl GpuIqnHead { .arg(&self.online_params) .arg(&tau) .arg(&n_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(config) .map_err(|e| MLError::ModelError(format!("IQN EMA kernel: {e}")))?; } @@ -627,6 +647,8 @@ impl GpuIqnHead { let b = batch_size; let n = self.config.num_quantiles; let tba = self.config.total_branch_actions(); + let shared_h1_i32 = self.config.shared_h1 as i32; + let hidden_dim_i32 = self.config.hidden_dim as i32; // 1. Sample τ values for inference let total_taus = (b * n) as i32; @@ -644,6 +666,8 @@ impl GpuIqnHead { .arg(&mut self.online_taus) .arg(&total_taus) .arg(&rng_step) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(tau_config) .map_err(|e| MLError::ModelError(format!("IQN CVaR sample_taus: {e}")))?; } @@ -663,6 +687,8 @@ impl GpuIqnHead { .arg(&self.online_params) .arg(&mut self.save_q_online) .arg(&batch_i32) + .arg(&shared_h1_i32) + .arg(&hidden_dim_i32) .launch(fwd_config) .map_err(|e| MLError::ModelError(format!("IQN CVaR forward: {e}")))?; } diff --git a/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu b/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu index 59dc59522..afacff86e 100644 --- a/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu +++ b/crates/ml/src/cuda_pipeline/iqn_dual_head_kernel.cu @@ -8,8 +8,8 @@ * Architecture per sample: * For each τ_i ∈ (0,1), i = 1..N: * cos_feat = [cos(π·1·τ), ..., cos(π·D·τ)] [IQN_EMBED_DIM] - * embed = ReLU(W_embed × cos_feat + b_embed) [IQN_HIDDEN] - * combined = h_s2 ⊙ embed [IQN_HIDDEN] + * embed = ReLU(W_embed × cos_feat + b_embed) [hidden_dim] + * combined = h_s2 ⊙ embed [hidden_dim] * per-branch: q_d = W_bd × combined + b_bd [n_d] scalars * * Loss: Quantile Huber loss (Dabney et al., 2018b) @@ -24,8 +24,11 @@ * Grad norm: grid=(ceil(total_params/256), 1, 1), block=(256, 1, 1) * Adam: grid=(ceil(total_params/256), 1, 1), block=(256, 1, 1) * - * Compile-time defines (injected via NVRTC dim_overrides): - * IQN_HIDDEN -- hidden dim = SHARED_H2 from DQN trunk (default 256) + * Runtime parameters (passed to each kernel): + * shared_h1 -- first shared hidden layer width (was compile-time IQN_SHARED_H1) + * hidden_dim -- hidden dim = SHARED_H2 from DQN trunk (was compile-time IQN_HIDDEN) + * + * Compile-time defines (still used for array sizing and non-hidden constants): * IQN_EMBED_DIM -- cosine embedding dimension (default 64) * IQN_NUM_QUANTILES -- τ samples per forward pass (default 32) * IQN_KAPPA -- Huber threshold for quantile loss (default 1.0) @@ -34,7 +37,10 @@ * BRANCH_2_SIZE -- urgency actions (default 3) */ -/* ── Compile-time defaults (overridden by NVRTC injection) ──────────── */ +/* ── Compile-time defaults ──────────────────────────────────────────── */ +/* IQN_HIDDEN and IQN_SHARED_H1 are NOW runtime kernel parameters. + * The #define is kept ONLY for register array sizing (max upper bound). + * Kernels use the runtime hidden_dim / shared_h1 params for all logic. */ #ifndef IQN_HIDDEN #define IQN_HIDDEN 256 #endif @@ -60,10 +66,13 @@ #define TOTAL_BRANCH_ACTIONS (BRANCH_0_SIZE + BRANCH_1_SIZE + BRANCH_2_SIZE) /* ── Weight layout (flat f32 buffer) ─────────────────────────────────── */ -/* W_embed [IQN_HIDDEN, IQN_EMBED_DIM], b_embed [IQN_HIDDEN] - * W_b0 [BRANCH_0_SIZE, IQN_HIDDEN], b_b0 [BRANCH_0_SIZE] - * W_b1 [BRANCH_1_SIZE, IQN_HIDDEN], b_b1 [BRANCH_1_SIZE] - * W_b2 [BRANCH_2_SIZE, IQN_HIDDEN], b_b2 [BRANCH_2_SIZE] */ +/* W_embed [hidden_dim, IQN_EMBED_DIM], b_embed [hidden_dim] + * W_b0 [BRANCH_0_SIZE, hidden_dim], b_b0 [BRANCH_0_SIZE] + * W_b1 [BRANCH_1_SIZE, hidden_dim], b_b1 [BRANCH_1_SIZE] + * W_b2 [BRANCH_2_SIZE, hidden_dim], b_b2 [BRANCH_2_SIZE] + * + * NOTE: IQN_OFF_* macros below use compile-time IQN_HIDDEN. They are kept + * only for documentation. Kernels compute offsets from runtime hidden_dim. */ #define IQN_W_EMBED_SIZE (IQN_HIDDEN * IQN_EMBED_DIM) #define IQN_B_EMBED_SIZE (IQN_HIDDEN) @@ -79,7 +88,7 @@ + IQN_W_B1_SIZE + IQN_B_B1_SIZE \ + IQN_W_B2_SIZE + IQN_B_B2_SIZE) -/* Offsets into the flat parameter buffer */ +/* Offsets into the flat parameter buffer (compile-time, for documentation only) */ #define IQN_OFF_W_EMBED 0 #define IQN_OFF_B_EMBED (IQN_OFF_W_EMBED + IQN_W_EMBED_SIZE) #define IQN_OFF_W_B0 (IQN_OFF_B_EMBED + IQN_B_EMBED_SIZE) @@ -89,9 +98,22 @@ #define IQN_OFF_W_B2 (IQN_OFF_B_B1 + IQN_B_B1_SIZE) #define IQN_OFF_B_B2 (IQN_OFF_W_B2 + IQN_W_B2_SIZE) +/* Max distributed array size: ceil(IQN_HIDDEN / 32) — used for register arrays */ +#define IQN_DIST_MAX (((IQN_HIDDEN) + 31) / 32) + /* Distributed array size: ceil(dim / 32) elements per lane */ #define IQN_DIST(dim) (((dim) + 31) / 32) +#ifndef IQN_STATE_DIM +#define IQN_STATE_DIM 48 +#endif +#ifndef IQN_SHARED_H1 +#define IQN_SHARED_H1 256 +#endif + +/* Max shared_h1 for register array sizing */ +#define IQN_SHARED_H1_MAX (((IQN_SHARED_H1) + 31) / 32) + /* ── Device helpers ──────────────────────────────────────────────────── */ /** Warp-reduce sum across all 32 lanes (butterfly). */ @@ -128,6 +150,19 @@ __device__ __forceinline__ float quantile_huber_grad(float tau, float delta, flo return -weight * huber_grad; } +/** Compute runtime weight offsets from hidden_dim. Returns offsets in an array: + * [0]=W_EMBED, [1]=B_EMBED, [2]=W_B0, [3]=B_B0, [4]=W_B1, [5]=B_B1, [6]=W_B2, [7]=B_B2 */ +__device__ __forceinline__ void iqn_compute_offsets(int hidden_dim, int offsets[8]) { + offsets[0] = 0; /* W_EMBED */ + offsets[1] = offsets[0] + hidden_dim * IQN_EMBED_DIM; /* B_EMBED */ + offsets[2] = offsets[1] + hidden_dim; /* W_B0 */ + offsets[3] = offsets[2] + BRANCH_0_SIZE * hidden_dim; /* B_B0 */ + offsets[4] = offsets[3] + BRANCH_0_SIZE; /* W_B1 */ + offsets[5] = offsets[4] + BRANCH_1_SIZE * hidden_dim; /* B_B1 */ + offsets[6] = offsets[5] + BRANCH_1_SIZE; /* W_B2 */ + offsets[7] = offsets[6] + BRANCH_2_SIZE * hidden_dim; /* B_B2 */ +} + /* ── Forward + Loss Kernel ──────────────────────────────────────────── */ /** * Per-sample IQN forward pass + quantile Huber loss. @@ -146,8 +181,8 @@ __device__ __forceinline__ float quantile_huber_grad(float tau, float delta, flo extern "C" __global__ void iqn_forward_loss_kernel( /* Inputs */ - const float* __restrict__ h_s2, /* [B, IQN_HIDDEN] trunk activations */ - const float* __restrict__ target_h_s2, /* [B, IQN_HIDDEN] target trunk activations */ + const float* __restrict__ h_s2, /* [B, hidden_dim] trunk activations */ + const float* __restrict__ target_h_s2, /* [B, hidden_dim] target trunk activations */ const float* __restrict__ taus, /* [B, N] online τ samples ∈ (0,1) */ const float* __restrict__ target_taus, /* [B, N] target τ samples */ const int* __restrict__ actions, /* [B, 3] branch actions (exposure, order, urgency) */ @@ -155,26 +190,32 @@ void iqn_forward_loss_kernel( const float* __restrict__ dones, /* [B] */ float gamma, /* Weights */ - const float* __restrict__ online_params, /* [IQN_TOTAL_PARAMS] online IQN weights */ - const float* __restrict__ target_params, /* [IQN_TOTAL_PARAMS] target IQN weights */ + const float* __restrict__ online_params, /* online IQN weights */ + const float* __restrict__ target_params, /* target IQN weights */ /* Outputs */ float* __restrict__ per_sample_loss, /* [B] per-sample loss (for PER) */ float* __restrict__ total_loss, /* [1] batch-mean loss (atomicAdd) */ /* Activation saves for backward */ - float* __restrict__ save_embed, /* [B, N, IQN_HIDDEN] post-ReLU embeddings */ - float* __restrict__ save_combined, /* [B, N, IQN_HIDDEN] h_s2 ⊙ embed */ + float* __restrict__ save_embed, /* [B, N, hidden_dim] post-ReLU embeddings */ + float* __restrict__ save_combined, /* [B, N, hidden_dim] h_s2 ⊙ embed */ float* __restrict__ save_q_online, /* [B, N, TOTAL_BRANCH_ACTIONS] quantile Q-values */ float* __restrict__ save_q_target, /* [B, N, TOTAL_BRANCH_ACTIONS] target Q-values */ - int batch_size + int batch_size, + int shared_h1, /* runtime: first hidden layer width (unused here, for consistency) */ + int hidden_dim /* runtime: IQN hidden dim = SHARED_H2 */ ) { int sample = blockIdx.x; if (sample >= batch_size) return; int lane = threadIdx.x; + /* ── Compute runtime weight offsets ── */ + int off[8]; + iqn_compute_offsets(hidden_dim, off); + /* ── Pointers into this sample's data ── */ - const float* my_h_s2 = h_s2 + sample * IQN_HIDDEN; - const float* my_target_h_s2 = target_h_s2 + sample * IQN_HIDDEN; + const float* my_h_s2 = h_s2 + sample * hidden_dim; + const float* my_target_h_s2 = target_h_s2 + sample * hidden_dim; const float* my_taus = taus + sample * IQN_NUM_QUANTILES; const float* my_target_taus = target_taus + sample * IQN_NUM_QUANTILES; @@ -185,31 +226,31 @@ void iqn_forward_loss_kernel( float done = dones[sample]; /* Weight pointers (same for all samples) */ - const float* w_embed = online_params + IQN_OFF_W_EMBED; - const float* b_embed = online_params + IQN_OFF_B_EMBED; - const float* w_b0 = online_params + IQN_OFF_W_B0; - const float* b_b0 = online_params + IQN_OFF_B_B0; - const float* w_b1 = online_params + IQN_OFF_W_B1; - const float* b_b1 = online_params + IQN_OFF_B_B1; - const float* w_b2 = online_params + IQN_OFF_W_B2; - const float* b_b2 = online_params + IQN_OFF_B_B2; + const float* w_embed = online_params + off[0]; + const float* b_embed = online_params + off[1]; + const float* w_b0 = online_params + off[2]; + const float* b_b0 = online_params + off[3]; + const float* w_b1 = online_params + off[4]; + const float* b_b1 = online_params + off[5]; + const float* w_b2 = online_params + off[6]; + const float* b_b2 = online_params + off[7]; - const float* tw_embed = target_params + IQN_OFF_W_EMBED; - const float* tb_embed = target_params + IQN_OFF_B_EMBED; - const float* tw_b0 = target_params + IQN_OFF_W_B0; - const float* tb_b0 = target_params + IQN_OFF_B_B0; - const float* tw_b1 = target_params + IQN_OFF_W_B1; - const float* tb_b1 = target_params + IQN_OFF_B_B1; - const float* tw_b2 = target_params + IQN_OFF_W_B2; - const float* tb_b2 = target_params + IQN_OFF_B_B2; + const float* tw_embed = target_params + off[0]; + const float* tb_embed = target_params + off[1]; + const float* tw_b0 = target_params + off[2]; + const float* tb_b0 = target_params + off[3]; + const float* tw_b1 = target_params + off[4]; + const float* tb_b1 = target_params + off[5]; + const float* tw_b2 = target_params + off[6]; + const float* tb_b2 = target_params + off[7]; /* ── Load h_s2 into distributed registers ── */ - float h_dist[IQN_DIST(IQN_HIDDEN)]; - for (int d = lane; d < IQN_HIDDEN; d += 32) + float h_dist[IQN_DIST_MAX]; + for (int d = lane; d < hidden_dim; d += 32) h_dist[d / 32] = my_h_s2[d]; - float h_target_dist[IQN_DIST(IQN_HIDDEN)]; - for (int d = lane; d < IQN_HIDDEN; d += 32) + float h_target_dist[IQN_DIST_MAX]; + for (int d = lane; d < hidden_dim; d += 32) h_target_dist[d / 32] = my_target_h_s2[d]; /* ── Process each ONLINE quantile ── */ @@ -221,9 +262,9 @@ void iqn_forward_loss_kernel( /* 1. Cosine features: cos(π·(d+1)·τ) for d=0..IQN_EMBED_DIM-1 */ /* 2. Linear embedding: embed[h] = Σ_d W_embed[h,d]·cos_feat[d] + b_embed[h] */ - /* Process via warp-cooperative matvec: each lane handles IQN_HIDDEN/32 outputs */ - float embed_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + /* Process via warp-cooperative matvec: each lane handles hidden_dim/32 outputs */ + float embed_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) { float acc = b_embed[h]; const float* w_row = w_embed + h * IQN_EMBED_DIM; for (int d = 0; d < IQN_EMBED_DIM; d++) { @@ -235,13 +276,13 @@ void iqn_forward_loss_kernel( } /* 3. Element-wise product: combined = h_s2 ⊙ embed */ - float comb_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) + float comb_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) comb_dist[h / 32] = h_dist[h / 32] * embed_dist[h / 32]; /* Save activations for backward pass */ - int save_offset = (sample * IQN_NUM_QUANTILES + t) * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + int save_offset = (sample * IQN_NUM_QUANTILES + t) * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) { save_embed[save_offset + h] = embed_dist[h / 32]; save_combined[save_offset + h] = comb_dist[h / 32]; } @@ -253,9 +294,9 @@ void iqn_forward_loss_kernel( float q0_taken = 0.0f; for (int a = 0; a < BRANCH_0_SIZE; a++) { float acc = b_b0[a]; - const float* w_row = w_b0 + a * IQN_HIDDEN; + const float* w_row = w_b0 + a * hidden_dim; float partial = 0.0f; - for (int h = lane; h < IQN_HIDDEN; h += 32) + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; acc += iqn_warp_sum(partial); if (lane == 0) { @@ -269,9 +310,9 @@ void iqn_forward_loss_kernel( float q1_taken = 0.0f; for (int a = 0; a < BRANCH_1_SIZE; a++) { float acc = b_b1[a]; - const float* w_row = w_b1 + a * IQN_HIDDEN; + const float* w_row = w_b1 + a * hidden_dim; float partial = 0.0f; - for (int h = lane; h < IQN_HIDDEN; h += 32) + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; acc += iqn_warp_sum(partial); if (lane == 0) { @@ -285,9 +326,9 @@ void iqn_forward_loss_kernel( float q2_taken = 0.0f; for (int a = 0; a < BRANCH_2_SIZE; a++) { float acc = b_b2[a]; - const float* w_row = w_b2 + a * IQN_HIDDEN; + const float* w_row = w_b2 + a * hidden_dim; float partial = 0.0f; - for (int h = lane; h < IQN_HIDDEN; h += 32) + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; acc += iqn_warp_sum(partial); if (lane == 0) { @@ -308,8 +349,8 @@ void iqn_forward_loss_kernel( float tau = my_target_taus[t]; /* Cosine embedding + linear (target weights) */ - float embed_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + float embed_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) { float acc = tb_embed[h]; const float* w_row = tw_embed + h * IQN_EMBED_DIM; for (int d = 0; d < IQN_EMBED_DIM; d++) { @@ -320,8 +361,8 @@ void iqn_forward_loss_kernel( } /* Element-wise product with target h_s2 */ - float comb_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) + float comb_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) comb_dist[h / 32] = h_target_dist[h / 32] * embed_dist[h / 32]; /* Per-branch Q-values for taken actions */ @@ -331,8 +372,8 @@ void iqn_forward_loss_kernel( { float acc = tb_b0[a0]; float partial = 0.0f; - const float* w_row = tw_b0 + a0 * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) + const float* w_row = tw_b0 + a0 * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; q0 = acc + iqn_warp_sum(partial); } @@ -340,8 +381,8 @@ void iqn_forward_loss_kernel( { float acc = tb_b1[a1]; float partial = 0.0f; - const float* w_row = tw_b1 + a1 * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) + const float* w_row = tw_b1 + a1 * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; q1 = acc + iqn_warp_sum(partial); } @@ -349,8 +390,8 @@ void iqn_forward_loss_kernel( { float acc = tb_b2[a2]; float partial = 0.0f; - const float* w_row = tw_b2 + a2 * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) + const float* w_row = tw_b2 + a2 * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; q2 = acc + iqn_warp_sum(partial); } @@ -402,9 +443,9 @@ void iqn_forward_loss_kernel( extern "C" __global__ void iqn_backward_kernel( /* Saved activations (from forward pass) */ - const float* __restrict__ h_s2, /* [B, IQN_HIDDEN] */ - const float* __restrict__ save_embed, /* [B, N, IQN_HIDDEN] */ - const float* __restrict__ save_combined, /* [B, N, IQN_HIDDEN] */ + const float* __restrict__ h_s2, /* [B, hidden_dim] */ + const float* __restrict__ save_embed, /* [B, N, hidden_dim] */ + const float* __restrict__ save_combined, /* [B, N, hidden_dim] */ const float* __restrict__ save_q_online, /* [B, N, TOTAL_BRANCH_ACTIONS] */ const float* __restrict__ save_q_target, /* [B, N, TOTAL_BRANCH_ACTIONS] */ const float* __restrict__ taus, /* [B, N] */ @@ -412,29 +453,35 @@ void iqn_backward_kernel( /* Weights (for backward through linear layers) */ const float* __restrict__ online_params, /* Gradient output */ - float* __restrict__ grad_buf, /* [IQN_TOTAL_PARAMS] accumulated gradients */ - float* __restrict__ d_h_s2_out, /* [B, IQN_HIDDEN] dL/d(h_s2) for trunk gradient */ - int batch_size + float* __restrict__ grad_buf, /* accumulated gradients */ + float* __restrict__ d_h_s2_out, /* [B, hidden_dim] dL/d(h_s2) for trunk gradient */ + int batch_size, + int shared_h1, /* runtime: first hidden layer width (unused here) */ + int hidden_dim /* runtime: IQN hidden dim = SHARED_H2 */ ) { int sample = blockIdx.x; if (sample >= batch_size) return; int lane = threadIdx.x; - const float* my_h_s2 = h_s2 + sample * IQN_HIDDEN; + /* ── Compute runtime weight offsets ── */ + int off[8]; + iqn_compute_offsets(hidden_dim, off); + + const float* my_h_s2 = h_s2 + sample * hidden_dim; const float* my_taus = taus + sample * IQN_NUM_QUANTILES; int a0 = actions[sample * 3 + 0]; int a1 = actions[sample * 3 + 1]; int a2 = actions[sample * 3 + 2]; - const float* w_embed = online_params + IQN_OFF_W_EMBED; - const float* w_b0 = online_params + IQN_OFF_W_B0; - const float* w_b1 = online_params + IQN_OFF_W_B1; - const float* w_b2 = online_params + IQN_OFF_W_B2; + const float* w_embed = online_params + off[0]; + const float* w_b0 = online_params + off[2]; + const float* w_b1 = online_params + off[4]; + const float* w_b2 = online_params + off[6]; /* Load h_s2 into distributed registers */ - float h_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) + float h_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) h_dist[h / 32] = my_h_s2[h]; /* Pre-compute online and target Q-values for taken actions across all quantiles */ @@ -458,7 +505,7 @@ void iqn_backward_kernel( for (int ti = 0; ti < IQN_NUM_QUANTILES; ti++) { float tau_i = my_taus[ti]; - int save_offset = (sample * IQN_NUM_QUANTILES + ti) * IQN_HIDDEN; + int save_offset = (sample * IQN_NUM_QUANTILES + ti) * hidden_dim; /* Compute dL/dq_online for this quantile: * dL/dq_i = (1/N²) Σ_j dρ_τi(δ_ij)/dq_i */ @@ -473,63 +520,63 @@ void iqn_backward_kernel( dL_dq = __shfl_sync(0xFFFFFFFF, dL_dq, 0); /* Load saved activations */ - float embed_dist[IQN_DIST(IQN_HIDDEN)]; - float comb_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + float embed_dist[IQN_DIST_MAX]; + float comb_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) { embed_dist[h / 32] = save_embed[save_offset + h]; comb_dist[h / 32] = save_combined[save_offset + h]; } /* ── Gradient through branch output layers (taken action only) ── */ /* dL/d(combined) = Σ_d W_bd[a_d, :] × dL/dq (only for taken actions) */ - float dL_dcomb[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) + float dL_dcomb[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) dL_dcomb[h / 32] = 0.0f; /* Branch 0: accumulate gradient */ - for (int h = lane; h < IQN_HIDDEN; h += 32) { - dL_dcomb[h / 32] += w_b0[a0 * IQN_HIDDEN + h] * dL_dq; + for (int h = lane; h < hidden_dim; h += 32) { + dL_dcomb[h / 32] += w_b0[a0 * hidden_dim + h] * dL_dq; /* dL/dW_b0[a0, h] += dL/dq × combined[h] */ - atomicAdd(&grad_buf[IQN_OFF_W_B0 + a0 * IQN_HIDDEN + h], + atomicAdd(&grad_buf[off[2] + a0 * hidden_dim + h], dL_dq * comb_dist[h / 32]); } if (lane == 0) - atomicAdd(&grad_buf[IQN_OFF_B_B0 + a0], dL_dq); + atomicAdd(&grad_buf[off[3] + a0], dL_dq); /* Branch 1 */ - for (int h = lane; h < IQN_HIDDEN; h += 32) { - dL_dcomb[h / 32] += w_b1[a1 * IQN_HIDDEN + h] * dL_dq; - atomicAdd(&grad_buf[IQN_OFF_W_B1 + a1 * IQN_HIDDEN + h], + for (int h = lane; h < hidden_dim; h += 32) { + dL_dcomb[h / 32] += w_b1[a1 * hidden_dim + h] * dL_dq; + atomicAdd(&grad_buf[off[4] + a1 * hidden_dim + h], dL_dq * comb_dist[h / 32]); } if (lane == 0) - atomicAdd(&grad_buf[IQN_OFF_B_B1 + a1], dL_dq); + atomicAdd(&grad_buf[off[5] + a1], dL_dq); /* Branch 2 */ - for (int h = lane; h < IQN_HIDDEN; h += 32) { - dL_dcomb[h / 32] += w_b2[a2 * IQN_HIDDEN + h] * dL_dq; - atomicAdd(&grad_buf[IQN_OFF_W_B2 + a2 * IQN_HIDDEN + h], + for (int h = lane; h < hidden_dim; h += 32) { + dL_dcomb[h / 32] += w_b2[a2 * hidden_dim + h] * dL_dq; + atomicAdd(&grad_buf[off[6] + a2 * hidden_dim + h], dL_dq * comb_dist[h / 32]); } if (lane == 0) - atomicAdd(&grad_buf[IQN_OFF_B_B2 + a2], dL_dq); + atomicAdd(&grad_buf[off[7] + a2], dL_dq); /* ── Gradient through element-wise product ── */ /* combined = h_s2 ⊙ embed * dL/d(embed) = dL/d(combined) ⊙ h_s2 * dL/d(h_s2) = dL/d(combined) ⊙ embed — NOW COMPUTED for trunk gradient */ - float dL_dembed[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + float dL_dembed[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) { dL_dembed[h / 32] = dL_dcomb[h / 32] * h_dist[h / 32]; /* Accumulate dL/d(h_s2) across all quantiles. * This flows IQN's bounded Huber gradient to the shared trunk. */ - atomicAdd(&d_h_s2_out[sample * IQN_HIDDEN + h], + atomicAdd(&d_h_s2_out[sample * hidden_dim + h], dL_dcomb[h / 32] * embed_dist[h / 32]); } /* ── Gradient through ReLU ── */ /* embed = ReLU(pre_relu) → dL/d(pre_relu) = dL/d(embed) × 𝟙{embed > 0} */ - for (int h = lane; h < IQN_HIDDEN; h += 32) { + for (int h = lane; h < hidden_dim; h += 32) { if (embed_dist[h / 32] <= 0.0f) dL_dembed[h / 32] = 0.0f; } @@ -539,14 +586,14 @@ void iqn_backward_kernel( * dL/dW_embed[h,d] += dL/d(pre_relu)[h] × cos(π(d+1)τ) * dL/db_embed[h] += dL/d(pre_relu)[h] */ float tau = my_taus[ti]; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + for (int h = lane; h < hidden_dim; h += 32) { float dL_dpre = dL_dembed[h / 32]; /* Bias gradient */ - atomicAdd(&grad_buf[IQN_OFF_B_EMBED + h], dL_dpre); + atomicAdd(&grad_buf[off[1] + h], dL_dpre); /* Weight gradient: outer product with cosine features */ for (int d = 0; d < IQN_EMBED_DIM; d++) { float cos_val = cosf(3.14159265f * (float)(d + 1) * tau); - atomicAdd(&grad_buf[IQN_OFF_W_EMBED + h * IQN_EMBED_DIM + d], + atomicAdd(&grad_buf[off[0] + h * IQN_EMBED_DIM + d], dL_dpre * cos_val); } } @@ -558,7 +605,9 @@ extern "C" __global__ void iqn_grad_norm_kernel( const float* __restrict__ grads, float* __restrict__ norm_out, /* [1] */ - int total_params + int total_params, + int shared_h1, /* runtime: unused, for consistent interface */ + int hidden_dim /* runtime: unused, for consistent interface */ ) { int tid = blockIdx.x * blockDim.x + threadIdx.x; @@ -597,7 +646,9 @@ void iqn_adam_kernel( float lr, float beta1, float beta2, float eps, float weight_decay, float max_grad_norm, int adam_t, /* step counter (1-indexed) */ - int total_params + int total_params, + int shared_h1, /* runtime: unused, for consistent interface */ + int hidden_dim /* runtime: unused, for consistent interface */ ) { int tid = blockIdx.x * blockDim.x + threadIdx.x; @@ -639,31 +690,37 @@ void iqn_adam_kernel( */ extern "C" __global__ void iqn_forward_kernel( - const float* __restrict__ h_s2, /* [B, IQN_HIDDEN] */ + const float* __restrict__ h_s2, /* [B, hidden_dim] */ const float* __restrict__ taus, /* [B, N] */ - const float* __restrict__ params, /* [IQN_TOTAL_PARAMS] */ + const float* __restrict__ params, /* IQN weights */ float* __restrict__ expected_q, /* [B, TOTAL_BRANCH_ACTIONS] */ - int batch_size + int batch_size, + int shared_h1, /* runtime: unused here, for consistency */ + int hidden_dim /* runtime: IQN hidden dim = SHARED_H2 */ ) { int sample = blockIdx.x; if (sample >= batch_size) return; int lane = threadIdx.x; - const float* my_h_s2 = h_s2 + sample * IQN_HIDDEN; + /* ── Compute runtime weight offsets ── */ + int off[8]; + iqn_compute_offsets(hidden_dim, off); + + const float* my_h_s2 = h_s2 + sample * hidden_dim; const float* my_taus = taus + sample * IQN_NUM_QUANTILES; - const float* w_embed = params + IQN_OFF_W_EMBED; - const float* b_embed = params + IQN_OFF_B_EMBED; - const float* w_b0 = params + IQN_OFF_W_B0; - const float* b_b0 = params + IQN_OFF_B_B0; - const float* w_b1 = params + IQN_OFF_W_B1; - const float* b_b1 = params + IQN_OFF_B_B1; - const float* w_b2 = params + IQN_OFF_W_B2; - const float* b_b2 = params + IQN_OFF_B_B2; + const float* w_embed = params + off[0]; + const float* b_embed = params + off[1]; + const float* w_b0 = params + off[2]; + const float* b_b0 = params + off[3]; + const float* w_b1 = params + off[4]; + const float* b_b1 = params + off[5]; + const float* w_b2 = params + off[6]; + const float* b_b2 = params + off[7]; /* Load h_s2 */ - float h_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) + float h_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) h_dist[h / 32] = my_h_s2[h]; /* Accumulate Q-values across quantiles (for mean) */ @@ -675,8 +732,8 @@ void iqn_forward_kernel( float tau = my_taus[t]; /* Cosine embedding + linear + ReLU */ - float embed_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) { + float embed_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) { float acc = b_embed[h]; const float* w_row = w_embed + h * IQN_EMBED_DIM; for (int d = 0; d < IQN_EMBED_DIM; d++) { @@ -687,16 +744,16 @@ void iqn_forward_kernel( } /* Element-wise product */ - float comb_dist[IQN_DIST(IQN_HIDDEN)]; - for (int h = lane; h < IQN_HIDDEN; h += 32) + float comb_dist[IQN_DIST_MAX]; + for (int h = lane; h < hidden_dim; h += 32) comb_dist[h / 32] = h_dist[h / 32] * embed_dist[h / 32]; /* Branch outputs */ /* Branch 0 */ for (int a = 0; a < BRANCH_0_SIZE; a++) { float partial = 0.0f; - const float* w_row = w_b0 + a * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) + const float* w_row = w_b0 + a * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; float q_val = iqn_warp_sum(partial) + b_b0[a]; q_acc[a] += q_val; @@ -704,8 +761,8 @@ void iqn_forward_kernel( /* Branch 1 */ for (int a = 0; a < BRANCH_1_SIZE; a++) { float partial = 0.0f; - const float* w_row = w_b1 + a * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) + const float* w_row = w_b1 + a * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; float q_val = iqn_warp_sum(partial) + b_b1[a]; q_acc[BRANCH_0_SIZE + a] += q_val; @@ -713,8 +770,8 @@ void iqn_forward_kernel( /* Branch 2 */ for (int a = 0; a < BRANCH_2_SIZE; a++) { float partial = 0.0f; - const float* w_row = w_b2 + a * IQN_HIDDEN; - for (int h = lane; h < IQN_HIDDEN; h += 32) + const float* w_row = w_b2 + a * hidden_dim; + for (int h = lane; h < hidden_dim; h += 32) partial += w_row[h] * comb_dist[h / 32]; float q_val = iqn_warp_sum(partial) + b_b2[a]; q_acc[BRANCH_0_SIZE + BRANCH_1_SIZE + a] += q_val; @@ -738,25 +795,20 @@ void iqn_forward_kernel( * Used to compute target_h_s2 from next_states + target DQN weights. * One warp per sample. Uses shared memory for h1 intermediate. * - * Compile-time defines: IQN_STATE_DIM, IQN_SHARED_H1, IQN_HIDDEN (=SHARED_H2) + * Runtime params: state_dim, shared_h1, hidden_dim (=SHARED_H2) */ -#ifndef IQN_STATE_DIM -#define IQN_STATE_DIM 48 -#endif -#ifndef IQN_SHARED_H1 -#define IQN_SHARED_H1 256 -#endif - extern "C" __global__ void iqn_trunk_forward_kernel( const float* __restrict__ states, /* [B, state_dim] */ - const float* __restrict__ w_s1, /* [IQN_SHARED_H1, state_dim] */ - const float* __restrict__ b_s1, /* [IQN_SHARED_H1] */ - const float* __restrict__ w_s2, /* [IQN_HIDDEN, IQN_SHARED_H1] */ - const float* __restrict__ b_s2, /* [IQN_HIDDEN] */ - float* __restrict__ h_s2_out, /* [B, IQN_HIDDEN] */ + const float* __restrict__ w_s1, /* [shared_h1, state_dim] */ + const float* __restrict__ b_s1, /* [shared_h1] */ + const float* __restrict__ w_s2, /* [hidden_dim, shared_h1] */ + const float* __restrict__ b_s2, /* [hidden_dim] */ + float* __restrict__ h_s2_out, /* [B, hidden_dim] */ int batch_size, - int state_dim /* runtime state dimension (replaces IQN_STATE_DIM) */ + int state_dim, /* runtime state dimension (replaces IQN_STATE_DIM) */ + int shared_h1, /* runtime: first hidden layer width */ + int hidden_dim /* runtime: IQN hidden dim = SHARED_H2 */ ) { int sample = blockIdx.x; @@ -764,12 +816,12 @@ void iqn_trunk_forward_kernel( int lane = threadIdx.x; extern __shared__ float shmem[]; - float* h1 = shmem; /* [IQN_SHARED_H1] */ + float* h1 = shmem; /* [shared_h1] */ const float* my_state = states + sample * state_dim; /* Layer 1: h1 = leaky_relu(W_s1 @ state + b_s1) */ - for (int h = lane; h < IQN_SHARED_H1; h += 32) { + for (int h = lane; h < shared_h1; h += 32) { float acc = b_s1[h]; const float* w_row = w_s1 + h * state_dim; for (int d = 0; d < state_dim; d++) @@ -779,13 +831,13 @@ void iqn_trunk_forward_kernel( __syncwarp(); /* Layer 2: h2 = leaky_relu(W_s2 @ h1 + b_s2) */ - for (int h = lane; h < IQN_HIDDEN; h += 32) { + for (int h = lane; h < hidden_dim; h += 32) { float acc = b_s2[h]; - const float* w_row = w_s2 + h * IQN_SHARED_H1; - for (int d = 0; d < IQN_SHARED_H1; d++) + const float* w_row = w_s2 + h * shared_h1; + for (int d = 0; d < shared_h1; d++) acc += w_row[d] * h1[d]; float val = (acc > 0.0f) ? acc : 0.01f * acc; - h_s2_out[sample * IQN_HIDDEN + h] = val; + h_s2_out[sample * hidden_dim + h] = val; } } @@ -798,7 +850,9 @@ void iqn_ema_kernel( float* __restrict__ target, const float* __restrict__ online, float tau, - int n + int n, + int shared_h1, /* runtime: unused, for consistent interface */ + int hidden_dim /* runtime: unused, for consistent interface */ ) { int i = blockIdx.x * blockDim.x + threadIdx.x; @@ -818,7 +872,9 @@ extern "C" __global__ void iqn_decode_actions_kernel( const int* __restrict__ flat_actions, int* __restrict__ branch_actions, - int batch_size + int batch_size, + int shared_h1, /* runtime: unused, for consistent interface */ + int hidden_dim /* runtime: unused, for consistent interface */ ) { int i = blockIdx.x * blockDim.x + threadIdx.x; @@ -871,7 +927,9 @@ extern "C" __global__ void iqn_sample_taus_kernel( float* __restrict__ taus, unsigned int seed, - int total + int total, + int shared_h1, /* runtime: unused, for consistent interface */ + int hidden_dim /* runtime: unused, for consistent interface */ ) { int i = blockIdx.x * blockDim.x + threadIdx.x;