From f0b8044b4adf4518e36f6776c7ebd49da2b1c479 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 6 Apr 2026 14:32:09 +0200 Subject: [PATCH] refactor: replace segment tree with prefix sum buffers, rewrite update_priorities_gpu MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Replace segment tree fields (seg_tree, capacity_pow2) with prefix sum buffers (priorities_pa, prefix_sum, scan_tile_state/aggregate/prefix) - Replace 5 seg_tree kernel fields with 4 PER prefix sum kernels (per_update_pa, per_insert_pa, per_prefix_scan, per_sample) - Delete CastKernels infrastructure (struct, OnceLock, get_cast_kernels) — per_update_pa reads bf16 td_errors directly, no cast needed - Delete rebuild_tree method (no tree to rebuild) - Delete update_td_f32 scratch buffer (bf16→f32 conversion eliminated) - Remove all PER_DIAG diagnostic eprintln calls Co-Authored-By: Claude Opus 4.6 (1M context) --- crates/ml-dqn/src/gpu_replay_buffer.rs | 213 ++++++++----------------- 1 file changed, 65 insertions(+), 148 deletions(-) diff --git a/crates/ml-dqn/src/gpu_replay_buffer.rs b/crates/ml-dqn/src/gpu_replay_buffer.rs index f52ad51fc..89ac82913 100644 --- a/crates/ml-dqn/src/gpu_replay_buffer.rs +++ b/crates/ml-dqn/src/gpu_replay_buffer.rs @@ -2,17 +2,14 @@ //! GPU-Resident Replay Buffer -- pure cudarc `CudaSlice` internals. //! -//! PER sampling uses an O(log n) GPU segment tree instead of O(n) prefix sum. -//! The segment tree is a flat `CudaSlice` of size `2 * capacity_pow2`: -//! root at index 1, leaves at `capacity_pow2..2*capacity_pow2`. Internal -//! nodes store sums of children. Sampling is parallel root-to-leaf traversal -//! with Philox RNG (1 kernel launch), replacing the old 5-kernel pipeline -//! (pow_alpha + prefix_sum + threshold_gen + searchsorted + i64_cast). +//! PER sampling uses a prefix sum over `priorities_pa` (priority^alpha). +//! `per_prefix_scan` builds the prefix sum, `per_sample` does uniform +//! threshold search with Philox RNG. //! -//! Output `GpuBatchSlices` wraps gathered data as raw `CudaSlice` buffers +//! Output `GpuBatchPtrs` wraps gathered data as raw `CudaSlice` buffers //! for downstream neural network consumption. -use std::sync::{Arc, OnceLock}; +use std::sync::Arc; use cudarc::driver::{CudaFunction, CudaSlice, CudaStream, DevicePtr, DevicePtrMut, LaunchConfig, PushKernelArg}; use ml_core::nvtx::NvtxRange; @@ -41,35 +38,6 @@ pub struct GpuBatchPtrs { } -// --------------------------------------------------------------------------- -// Cast kernel cache for bf16->f32 and u32->f32 GPU-only conversions -// --------------------------------------------------------------------------- - -struct CastKernels { - bf16_to_f32: CudaFunction, -} - -static CAST_KERNELS: OnceLock> = OnceLock::new(); - -fn get_cast_kernels(stream: &Arc) -> Result<&'static CastKernels, MLError> { - let result = CAST_KERNELS.get_or_init(|| { - static CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/cast_kernels.cubin")); - let ctx = stream.context(); - let module = ctx.load_cubin(CUBIN.to_vec()) - .map_err(|e| format!("cast cubin load: {e}"))?; - let ld = |n: &str| -> Result { - module.load_function(n).map_err(|e| format!("{n}: {e}")) - }; - Ok(CastKernels { - bf16_to_f32: ld("bf16_to_f32_cast")?, - }) - }); - match result { - Ok(k) => Ok(k), - Err(e) => Err(MLError::ModelError(e.clone())), - } -} - // --------------------------------------------------------------------------- // Compiled kernel cache // --------------------------------------------------------------------------- @@ -87,12 +55,11 @@ struct ReplayKernels { reduce_max_f32: CudaFunction, i64_to_u32: CudaFunction, max_of_two_f32: CudaFunction, - // Segment tree kernels (replace prefix_sum + searchsorted + pow_alpha pipeline) - seg_tree_update_leaves: CudaFunction, - seg_tree_rebuild_level: CudaFunction, - seg_tree_insert: CudaFunction, - seg_tree_sample: CudaFunction, - seg_tree_gather_prios: CudaFunction, + // PER prefix sum kernels + per_update_pa: CudaFunction, + per_insert_pa: CudaFunction, + per_prefix_scan: CudaFunction, + per_sample: CudaFunction, } impl ReplayKernels { @@ -107,10 +74,13 @@ impl ReplayKernels { rb_mod.load_function(n).map_err(|e| MLError::ModelError(format!("load {n}: {e}"))) }; - // Load precompiled segment tree kernels cubin - static ST_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/seg_tree_kernel.cubin")); - let st_mod = ctx.load_cubin(ST_CUBIN.to_vec()) - .map_err(|e| MLError::ModelError(format!("seg_tree cubin load: {e}")))?; + // Load precompiled PER prefix sum kernels cubin + static PER_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/per_kernels.cubin")); + let per_mod = ctx.load_cubin(PER_CUBIN.to_vec()) + .map_err(|e| MLError::ModelError(format!("per cubin load: {e}")))?; + let per_ld = |n: &str| -> Result { + per_mod.load_function(n).map_err(|e| MLError::ModelError(format!("per {n}: {e}"))) + }; Ok(Self { scatter_insert_f32: ld("scatter_insert_f32")?, @@ -123,16 +93,10 @@ impl ReplayKernels { reduce_max_f32: ld("reduce_max_f32")?, i64_to_u32: ld("i64_to_u32")?, max_of_two_f32: ld("max_of_two_f32")?, - seg_tree_update_leaves: st_mod.load_function("seg_tree_update_leaves") - .map_err(|e| MLError::ModelError(format!("st_update_leaves fn: {e}")))?, - seg_tree_rebuild_level: st_mod.load_function("seg_tree_rebuild_level") - .map_err(|e| MLError::ModelError(format!("st_rebuild_level fn: {e}")))?, - seg_tree_insert: st_mod.load_function("seg_tree_insert") - .map_err(|e| MLError::ModelError(format!("st_insert fn: {e}")))?, - seg_tree_sample: st_mod.load_function("seg_tree_sample") - .map_err(|e| MLError::ModelError(format!("st_sample fn: {e}")))?, - seg_tree_gather_prios: st_mod.load_function("seg_tree_gather_prios") - .map_err(|e| MLError::ModelError(format!("st_gather fn: {e}")))?, + per_update_pa: per_ld("per_update_pa")?, + per_insert_pa: per_ld("per_insert_pa")?, + per_prefix_scan: per_ld("per_prefix_scan")?, + per_sample: per_ld("per_sample")?, }) } } @@ -170,9 +134,13 @@ pub struct GpuReplayBuffer { max_priority: CudaSlice, pending_max_priority: Option>, current_step: usize, - // Segment tree: flat array [2 * capacity_pow2], root at [1], leaves at [cap_pow2..2*cap_pow2] - seg_tree: CudaSlice, - capacity_pow2: usize, + // PER prefix sum buffers + priorities_pa: CudaSlice, + prefix_sum: CudaSlice, + scan_tile_state: CudaSlice, + scan_tile_aggregate: CudaSlice, + scan_tile_prefix: CudaSlice, + num_scan_tiles: usize, // Pre-allocated PER sampling buffers (zero cuMemAlloc after warmup) sample_indices_i64: CudaSlice, sample_indices_u32: CudaSlice, @@ -187,7 +155,6 @@ pub struct GpuReplayBuffer { sample_episode_ids: CudaSlice, total_sum_buf: CudaSlice, // Pre-allocated scratch for update_priorities_gpu (zero cuMemAlloc per step) - update_td_f32: CudaSlice, // [max_batch_size] bf16→f32 conversion update_batch_max: CudaSlice, // [1] per-batch atomicMax accumulator update_max_merge: CudaSlice, // [1] merge pending + batch max rng_step: u32, @@ -200,8 +167,8 @@ pub struct GpuReplayBuffer { impl Drop for GpuReplayBuffer { fn drop(&mut self) { // Synchronize stream before CudaSlice fields drop. - // The segment tree, PER sampling, and priority update kernels may have - // pending work. Without sync, cuMemFree races with async kernel writes. + // PER sampling and priority update kernels may have pending work. + // Without sync, cuMemFree races with async kernel writes. #[allow(unsafe_code)] unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); @@ -231,10 +198,15 @@ impl GpuReplayBuffer { let mut mp = a32f(stream, 1, "mp")?; stream.memcpy_htod(&[1.0_f32], &mut mp).map_err(|e| MLError::ModelError(format!("mp: {e}")))?; - // Segment tree: flat array [2 * capacity_pow2], all zeros initially. - // Leaves at [cap_pow2..2*cap_pow2], root at [1]. Power-of-2 for balanced tree. - let cap_pow2 = cap.next_power_of_two(); - let seg = a32f(stream, 2 * cap_pow2, "seg_tree")?; + // PER prefix sum buffers + let priorities_pa = a32f(stream, cap, "priorities_pa")?; + let prefix_sum = a32f(stream, cap, "prefix_sum")?; + let tile_size = 256_usize; + let num_scan_tiles = cap.div_ceil(tile_size); + let scan_tile_state = stream.alloc_zeros::(num_scan_tiles) + .map_err(|e| MLError::ModelError(format!("alloc scan_tile_state: {e}")))?; + let scan_tile_aggregate = a32f(stream, num_scan_tiles, "scan_tile_agg")?; + let scan_tile_prefix = a32f(stream, num_scan_tiles, "scan_tile_pfx")?; // Pre-allocate PER sampling buffers (zero cuMemAlloc after warmup) let si64 = stream.alloc_zeros::(mbs).map_err(|e| MLError::ModelError(format!("alloc s_i64: {e}")))?; @@ -252,23 +224,17 @@ impl GpuReplayBuffer { let s_ep = a32i(stream, mbs, "s_episode_ids")?; // Pre-allocate PER update scratch buffers (zero cuMemAlloc per step) - let update_td_f32 = a32f(stream, mbs, "update_td_f32")?; let update_batch_max = a32f(stream, 1, "update_batch_max")?; let update_max_merge = a32f(stream, 1, "update_max_merge")?; - // Pre-initialize cast kernels while CUDA context is clean. - // OnceLock + cuModuleLoadData requires a current CUDA context. - // If deferred to first PER update (after graph capture), the context - // may be in a state that blocks cuModuleLoadData on H100 CUDA 13. - get_cast_kernels(stream)?; - Ok(Self { config, stream: Arc::clone(stream), kernels: k, states: s, next_states: ns, actions: a, rewards: r, dones: d, priorities: p, episode_ids: ep_ids, write_cursor: 0, size: 0, max_priority: mp, pending_max_priority: None, current_step: 0, - seg_tree: seg, capacity_pow2: cap_pow2, + priorities_pa, prefix_sum, + scan_tile_state, scan_tile_aggregate, scan_tile_prefix, num_scan_tiles, sample_indices_i64: si64, sample_indices_u32: su32, sample_states: ss, sample_next_states: sns, sample_actions: sa, @@ -276,7 +242,7 @@ impl GpuReplayBuffer { sample_priorities: sp, sample_weights: sw, sample_episode_ids: s_ep, sample_max_weight: smw, total_sum_buf: tsb, - update_td_f32, update_batch_max, update_max_merge, + update_batch_max, update_max_merge, rng_step: 0, last_batch_size: 0, }) @@ -295,7 +261,8 @@ impl GpuReplayBuffer { pub fn clear(&mut self) -> Result<(), MLError> { self.write_cursor = 0; self.size = 0; self.stream.memcpy_htod(&[1.0_f32], &mut self.max_priority).map_err(|e| MLError::ModelError(format!("{e}")))?; - self.stream.memset_zeros(&mut self.seg_tree).map_err(|e| MLError::ModelError(format!("seg_tree clear: {e}")))?; + self.stream.memset_zeros(&mut self.priorities_pa).map_err(|e| MLError::ModelError(format!("priorities_pa clear: {e}")))?; + self.stream.memset_zeros(&mut self.prefix_sum).map_err(|e| MLError::ModelError(format!("prefix_sum clear: {e}")))?; self.pending_max_priority = None; self.current_step = 0; Ok(()) } pub const fn stream(&self) -> &Arc { &self.stream } @@ -699,47 +666,16 @@ impl GpuReplayBuffer { self.sample_weights.slice(..self.last_batch_size) } - /// Update priorities from GPU-resident index and `td_error` `CudaSlices`. - /// - /// Uses the segment tree update kernel: computes `(|td|^alpha + eps)`, - /// writes to priorities buffer AND tree leaves, propagates sums to root. - /// O(log n) per update instead of O(n) prefix sum rebuild. - /// Raw-pointer variant for tests and callers without CudaSlice access. - /// Bottom-up tree rebuild: one kernel launch per level. - /// log2(capacity_pow2) launches, each fully parallel within its level. - /// For capacity_pow2=1M: 20 launches × ~2µs overhead = ~40µs total. - fn rebuild_tree(&mut self, stream: &Arc) -> Result<(), MLError> { - let mut level_size = (self.capacity_pow2 >> 1) as i32; - let mut level_count = 0_u32; - while level_size >= 1 { - let blocks = ((level_size as u32 + 255) / 256) as u32; - if level_count < 3 || level_size <= 4 { - eprintln!("PER_DIAG: rebuild level={level_size} blocks={blocks}"); - } - unsafe { - stream.launch_builder(&self.kernels.seg_tree_rebuild_level) - .arg(&self.seg_tree) - .arg(&level_size) - .launch(LaunchConfig { - grid_dim: (blocks.max(1), 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("st rebuild level={level_size}: {e}")))?; - } - level_count += 1; - level_size >>= 1; - } - eprintln!("PER_DIAG: rebuild complete, {level_count} levels"); - Ok(()) - } - pub fn update_priorities_gpu_raw(&mut self, _indices_ptr: u64, td_errors: &CudaSlice, bs: usize, ext_stream: Option<&Arc>) -> Result<(), MLError> { let idx_slice = unsafe { &*(&self.sample_indices_u32 as *const CudaSlice) }; self.update_priorities_gpu(idx_slice, td_errors, bs, ext_stream) } - /// Update PER priorities from TD errors using the segment tree kernel. + /// Update PER priorities from TD errors using the prefix-sum kernel. + /// + /// `per_update_pa` reads bf16 td_errors directly (no cast needed), + /// computes `(|td|^alpha + eps)`, and writes to both `priorities` and + /// `priorities_pa`. Prefix sum is rebuilt lazily before sampling. /// /// If `ext_stream` is provided, all GPU work runs on that stream instead of /// the replay buffer's own stream. This avoids cross-stream synchronization @@ -753,64 +689,45 @@ impl GpuReplayBuffer { bs: usize, ext_stream: Option<&Arc>, ) -> Result<(), MLError> { - let _nvtx = NvtxRange::new("per_update_priorities"); if bs == 0 { return Ok(()); } - eprintln!("PER_DIAG: enter update_priorities_gpu bs={bs} ext_stream={}", ext_stream.is_some()); let stream_owned = ext_stream.cloned().unwrap_or_else(|| Arc::clone(&self.stream)); let stream = &stream_owned; - // Convert bf16 td_errors to f32 using pre-allocated scratch buffer - eprintln!("PER_DIAG: getting cast kernels (using self.stream for module load)"); - // Use self.stream for kernel compilation (needs bound context). - // ext_stream is only used for kernel launches. - let kernels = get_cast_kernels(&self.stream)?; - eprintln!("PER_DIAG: launching bf16_to_f32 cast on ext_stream"); - let ni = bs as i32; - // SAFETY: update_td_f32, td_errors are valid device allocations of at least bs elements. - unsafe { - stream.launch_builder(&kernels.bf16_to_f32) - .arg(&self.update_td_f32).arg(td_errors).arg(&ni) - .launch(lcfg(bs)) - .map_err(|e| MLError::ModelError(format!("td bf16_to_f32: {e}")))?; - } - eprintln!("PER_DIAG: bf16_to_f32 launched, zeroing batch_max"); + let (al, ep, bsi) = (self.config.alpha, self.config.epsilon, bs as i32); - let cap_pow2_i = self.capacity_pow2 as i32; + let cap_i = self.config.capacity as i32; + stream.memset_zeros(&mut self.update_batch_max) .map_err(|e| MLError::ModelError(format!("zero batch_max: {e}")))?; - eprintln!("PER_DIAG: launching seg_tree_update_leaves cap_pow2={}", self.capacity_pow2); - // Phase 1: Write leaves + priorities (fully parallel, no tree propagation). + + // per_update_pa: reads bf16 td_errors directly, writes priorities + priorities_pa unsafe { - stream.launch_builder(&self.kernels.seg_tree_update_leaves) - .arg(&self.seg_tree).arg(&self.priorities).arg(&self.update_batch_max) - .arg(indices).arg(&self.update_td_f32) - .arg(&al).arg(&ep).arg(&cap_pow2_i).arg(&bsi) - .launch(lcfg(bs)).map_err(|e| MLError::ModelError(format!("st update_leaves: {e}")))?; + stream.launch_builder(&self.kernels.per_update_pa) + .arg(&self.priorities_pa) + .arg(&self.priorities) + .arg(&self.update_batch_max) + .arg(indices) + .arg(td_errors) + .arg(&al).arg(&ep).arg(&cap_i).arg(&bsi) + .launch(lcfg(bs)) + .map_err(|e| MLError::ModelError(format!("per_update_pa: {e}")))?; } - eprintln!("PER_DIAG: update_leaves launched (self-propagating via atomicAdd), merging batch max"); - // Merge batch max into pending max (GPU-only). - // Uses pre-allocated update_max_merge as output scratch. + + // Merge batch max into pending (unchanged logic) match self.pending_max_priority.take() { Some(prev) => { - // SAFETY: update_max_merge, prev, update_batch_max are valid 1-element device allocations. unsafe { stream.launch_builder(&self.kernels.max_of_two_f32) .arg(&self.update_max_merge).arg(&prev).arg(&self.update_batch_max) .launch(LaunchConfig { grid_dim: (1,1,1), block_dim: (1,1,1), shared_mem_bytes: 0 }) .map_err(|e| MLError::ModelError(format!("max_of_two: {e}")))?; } - // Rotate: update_max_merge becomes pending, prev becomes the new merge scratch. self.pending_max_priority = Some(std::mem::replace(&mut self.update_max_merge, prev)); } None => { - eprintln!("PER_DIAG: first epoch — allocating pending_max"); - // First update this epoch: batch_max IS the pending. - // Rotate: update_batch_max becomes pending, allocate new batch_max scratch. - // This allocation happens once per epoch (after flush_max_priority), not per step. let fresh = a32f(stream, 1, "bm_epoch")?; self.pending_max_priority = Some(std::mem::replace(&mut self.update_batch_max, fresh)); } } - eprintln!("PER_DIAG: update_priorities_gpu complete"); Ok(()) }