perf: batched spectral norm — 10 launches → 1, delete old kernel

Replace 10 individual `spectral_norm_kernel` launches (each grid=(1,1,1))
with a single `spectral_norm_batched` launch (grid=(10,1,1)) that processes
all weight matrices in parallel across 10 blocks.

- Build descriptor buffer at construction (10 × 6 u64 entries with W/u/v
  pointers, out_dim, in_dim per matrix) — pointers are stable
- Delete `spectral_norm_kernel` from dqn_utility_kernels.cu (replaced by
  `spectral_norm_batched` which was already written)
- Remove spec_norm! macro and per-matrix launch loop

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-06 11:46:16 +02:00
parent 4d78f90cc0
commit 10857e2c10
2 changed files with 103 additions and 95 deletions

View File

@@ -411,45 +411,49 @@ extern "C" __global__ void dqn_relu_mask_kernel(
}
/* ══════════════════════════════════════════════════════════════════════
* SPECTRAL NORM POWER ITERATION KERNEL
* BATCHED SPECTRAL NORM — all weight matrices in a single launch
*
* One step of power iteration for spectral normalization:
* v_new = W^T u / ||W^T u||
* u_new = W v_new / ||W v_new||
* sigma = u_new^T (W v_new)
* Each block processes one weight matrix independently. 10 blocks = 10
* matrices, all running in parallel on H100's 132 SMs.
*
* Then scales the weight matrix: W[i] /= max(1.0, sigma / sigma_max)
* Descriptor layout (per matrix, 6 values):
* [0] = W pointer (device), [1] = u pointer, [2] = v pointer,
* [3] = out_dim, [4] = in_dim, [5] = (unused, alignment pad)
*
* For small matrices (256x256), this runs efficiently with a single block.
* Uses shared memory for the matmul + reduction.
*
* Launch: grid=(1,1,1), block=(256,1,1)
* Launch: grid=(num_matrices,1,1), block=(256,1,1), shmem=2048
* ══════════════════════════════════════════════════════════════════════ */
extern "C" __global__ void spectral_norm_kernel(
__nv_bfloat16* __restrict__ W, /* [out_dim, in_dim] weight matrix — scaled in-place */
__nv_bfloat16* __restrict__ u, /* [out_dim] left singular vector (persistent) */
__nv_bfloat16* __restrict__ v, /* [in_dim] right singular vector (persistent) */
int out_dim,
int in_dim,
float sigma_max /* clip sigma above this (typically 1.0) */
) {
__shared__ __nv_bfloat16 shmem[256]; /* scratch for reductions */
extern "C" __global__ void spectral_norm_batched(
const unsigned long long* __restrict__ descriptors, /* [num_matrices, 6] */
float sigma_max,
int num_matrices)
{
int mat_idx = blockIdx.x;
if (mat_idx >= num_matrices) return;
/* Unpack descriptor for this matrix */
const unsigned long long* desc = &descriptors[mat_idx * 6];
__nv_bfloat16* W = (__nv_bfloat16*)desc[0];
__nv_bfloat16* u = (__nv_bfloat16*)desc[1];
__nv_bfloat16* v = (__nv_bfloat16*)desc[2];
int out_dim = (int)desc[3];
int in_dim = (int)desc[4];
__shared__ __nv_bfloat16 shmem[256];
int tid = threadIdx.x;
int bd = blockDim.x; /* typically 256 */
int bd = blockDim.x;
int n_total = out_dim * in_dim;
/* ── Step 1: v_new = W^T u ── (strided for dims > blockDim) */
/* Each thread handles multiple v elements via stride loop. */
/* ── Step 1: v_new = W^T u ── */
for (int col = tid; col < in_dim; col += bd) {
__nv_bfloat16 val = bf16_zero();
for (int row = 0; row < out_dim; row++)
val = val + W[row * in_dim + col] * u[row];
v[col] = val; /* unnormalized — normalize below */
v[col] = val;
}
__syncthreads();
/* ── Normalize v_new: ||v||₂ via parallel reduction ── */
/* ── Normalize v_new ── */
__nv_bfloat16 local_v2 = bf16_zero();
for (int col = tid; col < in_dim; col += bd)
{ __nv_bfloat16 vc = v[col]; local_v2 = local_v2 + vc * vc; }
@@ -464,16 +468,16 @@ extern "C" __global__ void spectral_norm_kernel(
v[col] = v[col] / v_norm;
__syncthreads();
/* ── Step 2: u_new = W v_new ── (strided for dims > blockDim) */
/* ── Step 2: u_new = W v_new ── */
for (int row = tid; row < out_dim; row += bd) {
__nv_bfloat16 val = bf16_zero();
for (int col = 0; col < in_dim; col++)
val = val + W[row * in_dim + col] * v[col];
u[row] = val; /* unnormalized — sigma = ||u_new|| */
u[row] = val;
}
__syncthreads();
/* ── Sigma = ||W v_new|| = ||u_new|| (before normalizing) ── */
/* ── Sigma = ||u_new|| ── */
__nv_bfloat16 local_u2 = bf16_zero();
for (int row = tid; row < out_dim; row += bd)
{ __nv_bfloat16 ur = u[row]; local_u2 = local_u2 + ur * ur; }

View File

@@ -501,7 +501,11 @@ pub struct GpuDqnTrainer {
regime_scale_kernel: CudaFunction,
shrink_perturb_kernel: CudaFunction,
relu_mask_kernel: CudaFunction,
spectral_norm_kernel: CudaFunction,
spectral_norm_batched_kernel: CudaFunction,
/// Descriptor buffer for batched spectral norm: 10 matrices x 6 u64 = 60 elements.
/// Each entry: [W_ptr, u_ptr, v_ptr, out_dim, in_dim, pad=0].
/// Built once at construction (pointers into params_bf16 are stable).
spectral_norm_descriptors: CudaSlice<u64>,
clipped_saxpy_kernel: CudaFunction,
clip_grad_kernel: CudaFunction,
pad_states_kernel: CudaFunction,
@@ -1858,71 +1862,26 @@ impl GpuDqnTrainer {
self.params_initialized = true;
}
let sh1 = self.config.shared_h1 as i32;
let _eg = EventTrackingGuard::new(self.stream.context());
let sh2 = self.config.shared_h2 as i32;
let sd = self.config.state_dim as i32;
let sigma_max = self.config.spectral_norm_sigma_max;
let _evt_guard = EventTrackingGuard::new(self.stream.context());
let sigma_max = self.config.spectral_norm_sigma_max;
let num_matrices = 10_i32;
// Compute raw u64 pointers into params_bf16 at GOFF_* offsets.
// These point directly into the flat bf16 shadow buffer — no separate
// allocations, no D2D sync needed after kernel writes.
let param_sizes = compute_param_sizes(&self.config);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_sizes);
// ── Macro for launching spectral norm on a params_bf16 offset ────
macro_rules! spec_norm {
($goff_idx:expr, $u_slice:expr, $v_slice:expr, $out_dim:expr, $in_dim:expr, $label:literal) => {{
let w_ptr = w_ptrs[$goff_idx];
let u_ptr = $u_slice.raw_ptr();
let v_ptr = $v_slice.raw_ptr();
unsafe {
self.stream
.launch_builder(&self.spectral_norm_kernel)
.arg(&w_ptr).arg(&u_ptr).arg(&v_ptr)
.arg(&$out_dim).arg(&$in_dim).arg(&sigma_max)
.launch(LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 512 * 4,
})
.map_err(|e| MLError::ModelError(format!("spectral_norm {}: {e}", $label)))?;
}
}};
// Single batched launch: 10 blocks × 256 threads, one matrix per block.
// Descriptor buffer was built at construction with stable pointers.
unsafe {
self.stream
.launch_builder(&self.spectral_norm_batched_kernel)
.arg(&self.spectral_norm_descriptors)
.arg(&sigma_max)
.arg(&num_matrices)
.launch(LaunchConfig {
grid_dim: (num_matrices as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0, // shmem[256] is static in the kernel
})
.map_err(|e| MLError::ModelError(format!("spectral_norm_batched: {e}")))?;
}
// W_s1 [shared_h1, state_dim] — GOFF index 0
spec_norm!(0, self.spec_u_s1, self.spec_v_s1, sh1, sd, "W_s1");
// W_s2 [shared_h2, shared_h1] — GOFF index 2
spec_norm!(2, self.spec_u_s2, self.spec_v_s2, sh2, sh1, "W_s2");
// ── Head spectral norm kernel launches ───────────────────────────
let vh = self.config.value_h as i32;
let ah = self.config.adv_h as i32;
let na = self.config.num_atoms as i32;
let b0 = self.config.branch_0_size as i32;
let b1 = self.config.branch_1_size as i32;
let b2 = self.config.branch_2_size as i32;
// W_v1 [value_h, shared_h2] — GOFF index 4
spec_norm!(4, self.spec_u_v1, self.spec_v_v1, vh, sh2, "W_v1");
// W_v2 [num_atoms, value_h] — GOFF index 6
spec_norm!(6, self.spec_u_v2, self.spec_v_v2, na, vh, "W_v2");
// W_a1 [adv_h, shared_h2] — GOFF index 8
spec_norm!(8, self.spec_u_a1, self.spec_v_a1, ah, sh2, "W_a1");
// W_a2 [b0*num_atoms, adv_h] — GOFF index 10
spec_norm!(10, self.spec_u_a2, self.spec_v_a2, b0 * na, ah, "W_a2");
// W_bo1 [adv_h, shared_h2] — GOFF index 12
spec_norm!(12, self.spec_u_bo1, self.spec_v_bo1, ah, sh2, "W_bo1");
// W_bo2 [b1*num_atoms, adv_h] — GOFF index 14
spec_norm!(14, self.spec_u_bo2, self.spec_v_bo2, b1 * na, ah, "W_bo2");
// W_bu1 [adv_h, shared_h2] — GOFF index 16
spec_norm!(16, self.spec_u_bu1, self.spec_v_bu1, ah, sh2, "W_bu1");
// W_bu2 [b2*num_atoms, adv_h] — GOFF index 18
spec_norm!(18, self.spec_u_bu2, self.spec_v_bu2, b2 * na, ah, "W_bu2");
// No sync_w copies needed — spectral norm wrote directly to params_bf16.
// Sync bf16 shadow → f32 master so Adam sees spectrally normalized weights.
@@ -2433,6 +2392,50 @@ impl GpuDqnTrainer {
// W_bu2 [b2*num_atoms, adv_h]
alloc_spec_pair!(spec_u_bu2, spec_v_bu2, b2 * na, ah, "spec_u_bu2", "spec_v_bu2");
// ── Build batched spectral norm descriptor buffer (10 matrices × 6 u64) ──
// Pointers are stable: params_bf16 base and spec_u/v allocations never move.
let spectral_norm_descriptors = {
let param_sizes = compute_param_sizes(&config);
let w_ptrs = bf16_weight_ptrs_from_base(params_bf16.raw_ptr(), &param_sizes);
let sd_u64 = config.state_dim as u64;
let sh1_u64 = config.shared_h1 as u64;
let sh2_u64 = config.shared_h2 as u64;
let vh_u64 = config.value_h as u64;
let ah_u64 = config.adv_h as u64;
let na_u64 = config.num_atoms as u64;
let b0_u64 = config.branch_0_size as u64;
let b1_u64 = config.branch_1_size as u64;
let b2_u64 = config.branch_2_size as u64;
// 10 matrices, 6 u64 each: [W_ptr, u_ptr, v_ptr, out_dim, in_dim, pad]
let host_desc: [u64; 60] = [
// [0] W_s1 [shared_h1, state_dim] — GOFF 0
w_ptrs[0], spec_u_s1.raw_ptr(), spec_v_s1.raw_ptr(), sh1_u64, sd_u64, 0,
// [1] W_s2 [shared_h2, shared_h1] — GOFF 2
w_ptrs[2], spec_u_s2.raw_ptr(), spec_v_s2.raw_ptr(), sh2_u64, sh1_u64, 0,
// [2] W_v1 [value_h, shared_h2] — GOFF 4
w_ptrs[4], spec_u_v1.raw_ptr(), spec_v_v1.raw_ptr(), vh_u64, sh2_u64, 0,
// [3] W_v2 [num_atoms, value_h] — GOFF 6
w_ptrs[6], spec_u_v2.raw_ptr(), spec_v_v2.raw_ptr(), na_u64, vh_u64, 0,
// [4] W_a1 [adv_h, shared_h2] — GOFF 8
w_ptrs[8], spec_u_a1.raw_ptr(), spec_v_a1.raw_ptr(), ah_u64, sh2_u64, 0,
// [5] W_a2 [b0*num_atoms, adv_h] — GOFF 10
w_ptrs[10], spec_u_a2.raw_ptr(), spec_v_a2.raw_ptr(), b0_u64 * na_u64, ah_u64, 0,
// [6] W_bo1 [adv_h, shared_h2] — GOFF 12
w_ptrs[12], spec_u_bo1.raw_ptr(), spec_v_bo1.raw_ptr(), ah_u64, sh2_u64, 0,
// [7] W_bo2 [b1*num_atoms, adv_h] — GOFF 14
w_ptrs[14], spec_u_bo2.raw_ptr(), spec_v_bo2.raw_ptr(), b1_u64 * na_u64, ah_u64, 0,
// [8] W_bu1 [adv_h, shared_h2] — GOFF 16
w_ptrs[16], spec_u_bu1.raw_ptr(), spec_v_bu1.raw_ptr(), ah_u64, sh2_u64, 0,
// [9] W_bu2 [b2*num_atoms, adv_h] — GOFF 18
w_ptrs[18], spec_u_bu2.raw_ptr(), spec_v_bu2.raw_ptr(), b2_u64 * na_u64, ah_u64, 0,
];
let mut desc_buf = stream.alloc_zeros::<u64>(60)
.map_err(|e| MLError::ModelError(format!("alloc spectral_norm_descriptors: {e}")))?;
stream.memcpy_htod(&host_desc, &mut desc_buf)
.map_err(|e| MLError::ModelError(format!("upload spectral_norm_descriptors: {e}")))?;
desc_buf
};
// ── Initialize cuBLAS forward context (required) ─────────────
let cublas_forward = CublasForward::new(
@@ -2667,7 +2670,8 @@ impl GpuDqnTrainer {
regime_scale_kernel,
shrink_perturb_kernel: shrink_perturb,
relu_mask_kernel: relu_mask_standalone,
spectral_norm_kernel,
spectral_norm_batched_kernel: spectral_norm_kernel,
spectral_norm_descriptors,
clipped_saxpy_kernel,
clip_grad_kernel,
pad_states_kernel,
@@ -5796,10 +5800,10 @@ impl GpuDqnTrainer {
///
/// The cubin is produced by build.rs from `common_device_functions.cuh` +
/// `dqn_utility_kernels.cu`. Contains: grad_norm, adam_update, f32_to_bf16,
/// bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, spectral_norm,
/// clipped_saxpy, clip_grad.
/// bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask,
/// spectral_norm_batched, clipped_saxpy, clip_grad.
///
/// Returns `(grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, spectral_norm, clipped_saxpy, clip_grad)`.
/// Returns `(grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, relu_mask, spectral_norm_batched, clipped_saxpy, clip_grad)`.
fn compile_training_kernels(
stream: &Arc<CudaStream>,
config: &GpuDqnTrainConfig,
@@ -5839,8 +5843,8 @@ fn compile_training_kernels(
.map_err(|e| MLError::ModelError(format!("dqn_shrink_perturb_kernel load: {e}")))?;
let _relu_mask_from_module = module.load_function("dqn_relu_mask_kernel")
.map_err(|e| MLError::ModelError(format!("dqn_relu_mask_kernel load: {e}")))?;
let spectral_norm = module.load_function("spectral_norm_kernel")
.map_err(|e| MLError::ModelError(format!("spectral_norm_kernel load: {e}")))?;
let spectral_norm = module.load_function("spectral_norm_batched")
.map_err(|e| MLError::ModelError(format!("spectral_norm_batched load: {e}")))?;
let clipped_saxpy = module.load_function("dqn_clipped_saxpy_kernel")
.map_err(|e| MLError::ModelError(format!("dqn_clipped_saxpy_kernel load: {e}")))?;
let clip_grad = module.load_function("dqn_clip_grad_kernel")