perf: kan_grad_reduce — batch-parallel 2-phase (8ms×208 → <0.1ms×208)

Replace single-phase one-thread-per-param serial loop (batch_size iterations
per thread) with kan_grad_reduce_p1 (block sums, grid=(ceil(B/256),total_params),
shared mem) + kan_grad_reduce_p2 (warp-shuffle final reduce, grid=(total_params)).
Allocate partials scratch [ceil(B/256)*total_params] for trunk + 4 branches.
Update CublasBackwardSet constructor signature and both call sites.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-19 12:49:03 +02:00
parent b6f88a5525
commit bf87a8d0bb
3 changed files with 236 additions and 78 deletions

View File

@@ -162,11 +162,17 @@ pub struct CublasBackwardSet {
/// basis contributions without atomicAdd.
kan_gate_backward_kernel: CudaFunction,
/// `kan_grad_reduce(d_coeff_per_elem, d_coeff_span, d_residual_per_elem,
/// d_spline_coeff, d_residual_w, n, adv_h)` —
/// Deterministic reduction of KAN spline coefficient + residual gradients.
/// One thread per output parameter, loops over batch dimension.
kan_grad_reduce_kernel: CudaFunction,
/// `kan_grad_reduce_p1(d_coeff_per_elem, d_coeff_span, d_residual_per_elem,
/// partials, n, adv_h, total_params)` —
/// Phase 1: partial block sums for KAN spline/residual gradients.
/// Grid=(ceil(B/256), total_params), Block=256, Shared=256*sizeof(float).
kan_grad_reduce_p1_kernel: CudaFunction,
/// `kan_grad_reduce_p2(partials, d_spline_coeff, d_residual_w,
/// num_blocks, adv_h, total_params)` —
/// Phase 2: warp-shuffle final reduction for KAN spline/residual gradients.
/// Grid=(total_params), Block=256.
kan_grad_reduce_p2_kernel: CudaFunction,
// ── KAN backward intermediate buffers (allocated once, reused per-call) ──
/// Per-element B-spline basis contributions [B*AH, 4] — f32.
@@ -175,6 +181,8 @@ pub struct CublasBackwardSet {
kan_d_coeff_span: CudaSlice<i32>,
/// Per-element residual_w gradient contribution [B*AH] — f32.
kan_d_residual_per_elem: CudaSlice<f32>,
/// Partials scratch [ceil(B/256), total_params] — f32 (trunk kan reduction).
kan_reduce_partials: CudaSlice<f32>,
// ── Network dimensions (baked at construction) ──
batch_size: usize,
@@ -226,6 +234,8 @@ pub struct CublasBackwardSet {
branch_kan_d_coeff_span: [CudaSlice<i32>; 4],
/// Per-element residual_w gradient contribution [B*AH] per branch — f32.
branch_kan_d_residual_per_elem: [CudaSlice<f32>; 4],
/// Partials scratch [ceil(B/256), total_params] per branch — f32 (2-phase KAN reduce).
branch_kan_reduce_partials: [CudaSlice<f32>; 4],
// ── Cached GEMM descriptors (created once at init, reused per-call) ──
/// Map from (transa, transb, m, n, k, lda, ldb, ldc) → pre-created descriptors + algo.
@@ -248,7 +258,8 @@ impl CublasBackwardSet {
shared: Arc<super::shared_cublas_handle::SharedCublasHandle>,
config: &GpuDqnTrainConfig,
kan_gate_backward_kernel: CudaFunction,
kan_grad_reduce_kernel: CudaFunction,
kan_grad_reduce_p1_kernel: CudaFunction,
kan_grad_reduce_p2_kernel: CudaFunction,
fwd_branch_workspace_ptrs: [u64; 4],
) -> Result<Self, MLError> {
let stream = &shared.stream;
@@ -406,15 +417,31 @@ impl CublasBackwardSet {
alloc_f32(kan_n, "branch_kan_d_resid_pe_3")?,
];
// ── KAN 2-phase reduce partials buffers ────────────────────────
// total_params = adv_h * 9 (8 spline coeffs + 1 residual per neuron)
// num_blocks = ceil(batch_size / 256)
let kan_total_params = config.adv_h * 9;
let kan_num_blocks = (config.batch_size + 255) / 256;
let kan_partials_size = kan_num_blocks * kan_total_params;
let kan_reduce_partials = alloc_f32(kan_partials_size, "kan_reduce_partials")?;
let branch_kan_reduce_partials = [
alloc_f32(kan_partials_size, "branch_kan_reduce_partials_0")?,
alloc_f32(kan_partials_size, "branch_kan_reduce_partials_1")?,
alloc_f32(kan_partials_size, "branch_kan_reduce_partials_2")?,
alloc_f32(kan_partials_size, "branch_kan_reduce_partials_3")?,
];
Ok(Self {
handle: shared,
relu_mask_kernel,
bias_grad_kernel,
kan_gate_backward_kernel,
kan_grad_reduce_kernel,
kan_grad_reduce_p1_kernel,
kan_grad_reduce_p2_kernel,
kan_d_coeff_per_elem,
kan_d_coeff_span,
kan_d_residual_per_elem,
kan_reduce_partials,
batch_size: config.batch_size,
state_dim: config.state_dim,
state_dim_padded: sd_pad,
@@ -437,6 +464,7 @@ impl CublasBackwardSet {
branch_kan_d_coeff_per_elem,
branch_kan_d_coeff_span,
branch_kan_d_residual_per_elem,
branch_kan_reduce_partials,
param_sizes,
gemm_cache,
})
@@ -939,27 +967,48 @@ impl CublasBackwardSet {
.map_err(|e| MLError::ModelError(format!("kan_gate_backward_kernel: {e}")))?;
}
// ── Stage 2: deterministic reduction ──
// ── Stage 2: 2-phase warp-parallel reduction ──
// total_params = adv_h * 8 (coefficients) + adv_h (residual) = adv_h * 9
let batch_size = n / adv_h;
let total_params = adv_h * 9;
let reduce_blocks = ((total_params + 255) / 256) as u32;
let p1_blocks_x = ((batch_size + 255) / 256) as u32;
let total_params_i32 = total_params as i32;
let p1_blocks_x_i32 = p1_blocks_x as i32;
let partials_ptr = raw_f32_ptr(&self.kan_reduce_partials, stream);
unsafe {
// Phase 1: partial block sums — grid=(ceil(B/256), total_params)
stream
.launch_builder(&self.kan_grad_reduce_kernel)
.launch_builder(&self.kan_grad_reduce_p1_kernel)
.arg(&d_coeff_pe_ptr)
.arg(&d_span_ptr)
.arg(&d_resid_pe_ptr)
.arg(&d_spline_coeff)
.arg(&d_residual_w)
.arg(&partials_ptr)
.arg(&n_i32)
.arg(&adv_h_i32)
.arg(&total_params_i32)
.launch(LaunchConfig {
grid_dim: (reduce_blocks, 1, 1),
grid_dim: (p1_blocks_x, total_params as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 256 * 4,
})
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_p1_kernel: {e}")))?;
// Phase 2: final reduction — grid=(total_params)
stream
.launch_builder(&self.kan_grad_reduce_p2_kernel)
.arg(&partials_ptr)
.arg(&d_spline_coeff)
.arg(&d_residual_w)
.arg(&p1_blocks_x_i32)
.arg(&adv_h_i32)
.arg(&total_params_i32)
.launch(LaunchConfig {
grid_dim: (total_params as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
})
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_kernel: {e}")))?;
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_p2_kernel: {e}")))?;
}
Ok(())
@@ -1028,26 +1077,47 @@ impl CublasBackwardSet {
.map_err(|e| MLError::ModelError(format!("kan_gate_backward_kernel branch {d}: {e}")))?;
}
// ── Stage 2: deterministic reduction ──
// ── Stage 2: 2-phase warp-parallel reduction ──
let batch_size = n / adv_h;
let total_params = adv_h * 9;
let reduce_blocks = ((total_params + 255) / 256) as u32;
let p1_blocks_x = ((batch_size + 255) / 256) as u32;
let total_params_i32 = total_params as i32;
let p1_blocks_x_i32 = p1_blocks_x as i32;
let partials_ptr = raw_f32_ptr(&self.branch_kan_reduce_partials[d], stream);
unsafe {
// Phase 1: partial block sums — grid=(ceil(B/256), total_params)
stream
.launch_builder(&self.kan_grad_reduce_kernel)
.launch_builder(&self.kan_grad_reduce_p1_kernel)
.arg(&d_coeff_pe_ptr)
.arg(&d_span_ptr)
.arg(&d_resid_pe_ptr)
.arg(&d_spline_coeff)
.arg(&d_residual_w)
.arg(&partials_ptr)
.arg(&n_i32)
.arg(&adv_h_i32)
.arg(&total_params_i32)
.launch(LaunchConfig {
grid_dim: (reduce_blocks, 1, 1),
grid_dim: (p1_blocks_x, total_params as u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 256 * 4,
})
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_p1_kernel branch {d}: {e}")))?;
// Phase 2: final reduction — grid=(total_params)
stream
.launch_builder(&self.kan_grad_reduce_p2_kernel)
.arg(&partials_ptr)
.arg(&d_spline_coeff)
.arg(&d_residual_w)
.arg(&p1_blocks_x_i32)
.arg(&adv_h_i32)
.arg(&total_params_i32)
.launch(LaunchConfig {
grid_dim: (total_params as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
})
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_kernel branch {d}: {e}")))?;
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_p2_kernel branch {d}: {e}")))?;
}
Ok(())

View File

@@ -3571,74 +3571,122 @@ extern "C" __global__ void kan_gate_backward(
}
/* ================================================================== */
/* Kernel: kan_grad_reduce */
/* Kernel: kan_grad_reduce (2-phase warp-parallel reduction) */
/* ================================================================== */
/**
* Deterministic reduction of KAN spline coefficient + residual gradients.
* Phase 1: Partial block sums for KAN spline coefficient + residual gradients.
*
* One thread per output parameter:
* - param_idx in [0, adv_h * NUM_BASES) → spline coefficient gradient
* - param_idx in [adv_h * NUM_BASES, adv_h * NUM_BASES + adv_h) → residual_w gradient
* Grid: (ceil(batch_size/256), total_params, 1)
* Block: (256, 1, 1), Shared: 256*sizeof(float)
*
* Each thread loops over all batch samples for its neuron and accumulates
* contributions from the per-element buffers written by kan_gate_backward.
* No atomicAdd — fully deterministic sequential reduction.
* blockIdx.y = param_idx in [0, total_params)
* blockIdx.x = batch tile (256 samples per block)
*
* Each block reduces 256 batch samples for one param_idx and writes a
* single partial sum to partials[blockIdx.x * total_params + param_idx].
*
* Grid: ceil(total_params / 256), Block: 256.
* total_params = adv_h * NUM_BASES + adv_h = adv_h * 9.
* n = batch_size * adv_h.
*/
extern "C" __global__ void kan_grad_reduce(
extern "C" __global__ void kan_grad_reduce_p1(
const float* __restrict__ d_coeff_per_elem,
const int* __restrict__ d_coeff_span,
const float* __restrict__ d_residual_per_elem,
float* __restrict__ d_spline_coeff,
float* __restrict__ d_residual_w,
float* __restrict__ partials, /* [ceil(B/256), total_params] */
int n,
int adv_h)
int adv_h,
int total_params)
{
const int NUM_BASES = 8;
int total_params = adv_h * NUM_BASES + adv_h;
int param_idx = blockIdx.x * blockDim.x + threadIdx.x;
if (param_idx >= total_params) return;
extern __shared__ float sdata[];
int param_idx = blockIdx.y;
int tid = threadIdx.x;
int batch_size = n / adv_h;
bool is_residual = (param_idx >= adv_h * NUM_BASES);
int s = blockIdx.x * blockDim.x + tid; /* sample index */
bool is_residual = (param_idx >= adv_h * NUM_BASES);
int neuron, target_k;
if (!is_residual) {
neuron = param_idx / NUM_BASES;
neuron = param_idx / NUM_BASES;
target_k = param_idx % NUM_BASES;
} else {
neuron = param_idx - adv_h * NUM_BASES;
target_k = 0; /* unused */
neuron = param_idx - adv_h * NUM_BASES;
target_k = 0;
}
float sum = 0.0f;
for (int s = 0; s < batch_size; s++) {
float val = 0.0f;
if (s < batch_size) {
int elem_idx = s * adv_h + neuron;
if (is_residual) {
sum += d_residual_per_elem[elem_idx];
val = d_residual_per_elem[elem_idx];
} else {
int span = d_coeff_span[elem_idx];
/* The 4 basis functions map to coefficient indices:
* b0 → span-1, b1 → span, b2 → span+1, b3 → span+2 */
int k0 = span - 1;
int k0 = span - 1;
if (target_k == k0 && k0 >= 0 && k0 < NUM_BASES)
sum += d_coeff_per_elem[elem_idx * 4 + 0];
val = d_coeff_per_elem[elem_idx * 4 + 0];
else if (target_k == span)
sum += d_coeff_per_elem[elem_idx * 4 + 1];
val = d_coeff_per_elem[elem_idx * 4 + 1];
else if (target_k == span + 1 && span + 1 < NUM_BASES)
sum += d_coeff_per_elem[elem_idx * 4 + 2];
val = d_coeff_per_elem[elem_idx * 4 + 2];
else if (target_k == span + 2 && span + 2 < NUM_BASES)
sum += d_coeff_per_elem[elem_idx * 4 + 3];
val = d_coeff_per_elem[elem_idx * 4 + 3];
}
}
if (is_residual)
d_residual_w[neuron] = sum;
else
d_spline_coeff[neuron * NUM_BASES + target_k] = sum;
sdata[tid] = val;
__syncthreads();
for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
if (tid < stride) sdata[tid] += sdata[tid + stride];
__syncthreads();
}
if (tid == 0)
partials[blockIdx.x * total_params + param_idx] = sdata[0];
}
/**
* Phase 2: Final reduction + write for KAN spline coefficient + residual gradients.
*
* Grid: (total_params, 1, 1)
* Block: (256, 1, 1)
*
* blockIdx.x = param_idx. Threads stride over num_blocks partials, then
* warp-shuffle reduce to lane 0, which writes the final gradient.
*/
extern "C" __global__ void kan_grad_reduce_p2(
const float* __restrict__ partials, /* [num_blocks, total_params] */
float* __restrict__ d_spline_coeff,
float* __restrict__ d_residual_w,
int num_blocks,
int adv_h,
int total_params)
{
const int NUM_BASES = 8;
int param_idx = blockIdx.x;
int tid = threadIdx.x;
float sum = 0.0f;
for (int i = tid; i < num_blocks; i += blockDim.x)
sum += partials[i * total_params + param_idx];
/* Warp-level reduction */
for (int offset = 16; offset > 0; offset >>= 1)
sum += __shfl_down_sync(0xFFFFFFFF, sum, offset);
if (tid == 0) {
bool is_residual = (param_idx >= adv_h * NUM_BASES);
if (is_residual) {
int neuron = param_idx - adv_h * NUM_BASES;
d_residual_w[neuron] = sum;
} else {
int neuron = param_idx / NUM_BASES;
int target_k = param_idx % NUM_BASES;
d_spline_coeff[neuron * NUM_BASES + target_k] = sum;
}
}
}
/* ================================================================== */
@@ -4937,6 +4985,18 @@ extern "C" __global__ void homeostatic_regularizer(
/* ───────── Risk-budget branch (5th DQN head) ───────── */
/**
* Risk-budget forward — parallel over hidden neurons.
*
* Grid: (B, 1, 1) — one block per sample
* Block: (AH, 1, 1) — one thread per hidden neuron (AH=128)
* Shared: (SH2 + 13) * sizeof(float) — input vector cached once per block
*
* Each block:
* 1. Cooperatively loads h_s2[i, :] + ISV[12] + predicted_error[i] into shmem.
* 2. Each thread j computes dot(w_risk_fc[j, :], input) + bias → ReLU → h_risk[i,j].
* 3. Thread 0 reads risk_hidden[i,:] and computes sigmoid output.
*/
extern "C" __global__ void risk_budget_forward(
const float* __restrict__ h_s2,
const float* __restrict__ isv_signals_dev_ptr, /* [12] pinned — raw ISV */
@@ -4949,27 +5009,44 @@ extern "C" __global__ void risk_budget_forward(
float* __restrict__ risk_budget_out,
int B, int SH2, int AH
) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
extern __shared__ float sinput[]; /* [SH2 + 13] */
int i = blockIdx.x;
int j = threadIdx.x; /* hidden neuron index */
int tid = threadIdx.x;
if (i >= B) return;
int input_dim = SH2 + 13; /* SH2 + 12 ISV + 1 predicted_error */
const float* h = h_s2 + (long long)i * SH2;
float* h_risk = risk_hidden + (long long)i * AH;
int input_dim = SH2 + 13;
for (int j = 0; j < AH; j++) {
float val = b_risk_fc[j];
for (int k = 0; k < SH2; k++)
val += w_risk_fc[(long long)j * input_dim + k] * h[k];
for (int k = 0; k < 12; k++)
val += w_risk_fc[(long long)j * input_dim + SH2 + k] * isv_signals_dev_ptr[k];
val += w_risk_fc[(long long)j * input_dim + SH2 + 12] * predicted_error[i];
h_risk[j] = fmaxf(val, 0.0f);
/* Cooperatively load input vector into shared memory */
/* h_s2 portion */
for (int k = tid; k < SH2; k += blockDim.x)
sinput[k] = h_s2[(long long)i * SH2 + k];
/* ISV portion */
if (tid < 12)
sinput[SH2 + tid] = isv_signals_dev_ptr[tid];
/* predicted_error scalar */
if (tid == 0)
sinput[SH2 + 12] = predicted_error[i];
__syncthreads();
/* Each thread computes one hidden neuron */
float val = b_risk_fc[j];
const float* w_row = w_risk_fc + (long long)j * input_dim;
for (int k = 0; k < input_dim; k++)
val += w_row[k] * sinput[k];
val = fmaxf(val, 0.0f);
risk_hidden[(long long)i * AH + j] = val;
__syncthreads();
/* Thread 0 computes the scalar output */
if (j == 0) {
const float* h_risk = risk_hidden + (long long)i * AH;
float raw = b_risk_out[0];
for (int k = 0; k < AH; k++)
raw += w_risk_out[k] * h_risk[k];
risk_budget_out[i] = 1.0f / (1.0f + expf(-raw));
}
float raw = b_risk_out[0];
for (int j = 0; j < AH; j++)
raw += w_risk_out[j] * h_risk[j];
risk_budget_out[i] = 1.0f / (1.0f + expf(-raw));
}
extern "C" __global__ void apply_risk_budget(

View File

@@ -883,7 +883,8 @@ pub struct GpuDqnTrainer {
glu_backward_kernel: CudaFunction,
kan_gate_combine_kernel: CudaFunction,
kan_gate_backward_kernel: CudaFunction,
kan_grad_reduce_kernel: CudaFunction,
kan_grad_reduce_p1_kernel: CudaFunction,
kan_grad_reduce_p2_kernel: CudaFunction,
// ── Branch confidence routing (ISV gate × Q-value confidence) ──
branch_confidence_routing_kernel: CudaFunction,
@@ -2621,7 +2622,13 @@ impl GpuDqnTrainer {
.arg(&(batch_size as i32))
.arg(&(self.config.shared_h2 as i32))
.arg(&(self.config.adv_h as i32))
.launch(LaunchConfig::for_num_elems(batch_size as u32))
.launch(LaunchConfig {
// One block per sample, one thread per hidden neuron (AH)
grid_dim: (batch_size as u32, 1, 1),
block_dim: (self.config.adv_h as u32, 1, 1),
// Shared memory: input vector = (SH2 + 13) floats
shared_mem_bytes: ((self.config.shared_h2 + 13) * 4) as u32,
})
.map_err(|e| MLError::ModelError(format!("risk_budget_forward: {e}")))?;
}
Ok(())
@@ -5031,8 +5038,10 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("kan_gate_combine load: {e}")))?;
let kan_gate_backward_kernel = cpbi_module.load_function("kan_gate_backward")
.map_err(|e| MLError::ModelError(format!("kan_gate_backward load: {e}")))?;
let kan_grad_reduce_kernel = cpbi_module.load_function("kan_grad_reduce")
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce load: {e}")))?;
let kan_grad_reduce_p1_kernel = cpbi_module.load_function("kan_grad_reduce_p1")
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_p1 load: {e}")))?;
let kan_grad_reduce_p2_kernel = cpbi_module.load_function("kan_grad_reduce_p2")
.map_err(|e| MLError::ModelError(format!("kan_grad_reduce_p2 load: {e}")))?;
info!("GpuDqnTrainer: Q-attn + selectivity + VSN + GLU + KAN kernels loaded");
// ── Compile CQL penalty kernel (if enabled) ──────────────────────
@@ -5301,7 +5310,8 @@ impl GpuDqnTrainer {
// ── Initialize cuBLAS backward context (required) ──────────
let cublas_backward = CublasBackwardSet::new(
Arc::clone(&shared_cublas), &config, kan_gate_backward_kernel.clone(), kan_grad_reduce_kernel.clone(),
Arc::clone(&shared_cublas), &config, kan_gate_backward_kernel.clone(),
kan_grad_reduce_p1_kernel.clone(), kan_grad_reduce_p2_kernel.clone(),
cublas_forward.branch_workspace_ptrs(),
)?;
info!("GpuDqnTrainer: cuBLAS batched backward initialized (KAN backward wired)");
@@ -6196,7 +6206,8 @@ impl GpuDqnTrainer {
glu_backward_kernel,
kan_gate_combine_kernel,
kan_gate_backward_kernel,
kan_grad_reduce_kernel,
kan_grad_reduce_p1_kernel,
kan_grad_reduce_p2_kernel,
branch_confidence_routing_kernel,
regime_q_gap_buf,
regime_util_pinned,