refactor: replace segment tree with prefix sum buffers, rewrite update_priorities_gpu

- 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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-06 14:32:09 +02:00
parent 867ff1aa7c
commit f0b8044b4a

View File

@@ -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<f32>` 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<Result<CastKernels, String>> = OnceLock::new();
fn get_cast_kernels(stream: &Arc<CudaStream>) -> 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<CudaFunction, String> {
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<CudaFunction, MLError> {
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<f32>,
pending_max_priority: Option<CudaSlice<f32>>,
current_step: usize,
// Segment tree: flat array [2 * capacity_pow2], root at [1], leaves at [cap_pow2..2*cap_pow2]
seg_tree: CudaSlice<f32>,
capacity_pow2: usize,
// PER prefix sum buffers
priorities_pa: CudaSlice<f32>,
prefix_sum: CudaSlice<f32>,
scan_tile_state: CudaSlice<u32>,
scan_tile_aggregate: CudaSlice<f32>,
scan_tile_prefix: CudaSlice<f32>,
num_scan_tiles: usize,
// Pre-allocated PER sampling buffers (zero cuMemAlloc after warmup)
sample_indices_i64: CudaSlice<i64>,
sample_indices_u32: CudaSlice<u32>,
@@ -187,7 +155,6 @@ pub struct GpuReplayBuffer {
sample_episode_ids: CudaSlice<i32>,
total_sum_buf: CudaSlice<f32>,
// Pre-allocated scratch for update_priorities_gpu (zero cuMemAlloc per step)
update_td_f32: CudaSlice<f32>, // [max_batch_size] bf16→f32 conversion
update_batch_max: CudaSlice<f32>, // [1] per-batch atomicMax accumulator
update_max_merge: CudaSlice<f32>, // [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::<u32>(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::<i64>(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<CudaStream> { &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<CudaStream>) -> 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<half::bf16>, bs: usize, ext_stream: Option<&Arc<CudaStream>>) -> Result<(), MLError> {
let idx_slice = unsafe { &*(&self.sample_indices_u32 as *const CudaSlice<u32>) };
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<CudaStream>>,
) -> 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(())
}