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:
@@ -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(())
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user