Files
foxhunt/crates/ml-dqn/src/gpu_replay_buffer.rs
jgrusewski 932ac2bda8 fix(graph-capture): eliminate dtoh in PER diversity path + align evaluate CLI
Two bugs caught by the L40S smoke (train-qhgj6) that couldn't surface on
local RTX-3050 single-fold runs:

1. PER dtoh inside CUDA Graph capture (Fold 1 crash)
   Failure: CUDA_ERROR_STREAM_CAPTURE_INVALIDATED at per_prefix_scan on
   Fold 1 re-capture. Chain: fused_training parent graph captures →
   memcpy_dtoh + cuStreamSynchronize in gpu_replay_buffer::update_priorities_gpu
   (health<0.8 diversity path) poisons the stream → subsequent per_sample
   kernel on the same stream sees an invalidated capture context.

   The prior comment claimed "runs once per epoch, DtoH cost acceptable"
   — wrong, it runs every priority update when health<0.8 (common during
   Fold handoff when health_cache is re-seeded low). Any dtoh inside
   capture invalidates regardless of latency.

   Proper fix (no shortcut):
   * New kernel actions_sum_scale_reduce_u32 — single-block deterministic
     tree reduction over sample_actions (u32) → writes (sum*1000)/n as i32
     to a device-accessible slot. No atomics (consistent with the 1/N
     determinism policy from commit c82386500).
   * mean_action_scaled storage is pinned + device-mapped (cuMemAllocHost
     + cuMemHostGetDevicePointer — same pattern as rng_step_dev_ptr and
     size_dev_ptr elsewhere in the file). Zero-copy between host and
     device, graph-safe, no explicit free needed (process-exit cleanup,
     matches existing pattern).
   * pow_alpha_diverse_f32 now takes const int* mean_action_scaled_ptr
     and does a plain global load — NOT __ldg. The read-only cache used
     by __ldg is not guaranteed coherent with device-mapped host memory;
     multi-trial smoke regression caught it (median q_gap collapsed
     from 2.0 → 0.15 with __ldg, recovered to 2.8 with plain load).

   Verified: multi-trial smoke 5/5 pass, median_q_gap=2.80 (beats 2.00
   baseline), Best Sharpe peaks 19-30 per trial. No stream capture
   invalidation.

2. evaluate step CLI drift in Argo template
   evaluate_baseline's Args struct uses --models-dir and --output (single
   file path). Template was passing --checkpoint-dir and --output-dir,
   causing clap to reject the invocation. Fixed argument names + added
   mkdir for the eval subdir + updated the comment to pin the source of
   truth for future drift catches.

Both fixes are graph-capture-clean and match the "wire properly or delete"
discipline. No masking, no feature flags, no dead params.
2026-04-21 17:10:22 +02:00

1136 lines
54 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#![allow(unsafe_code)] // Required for CUDA kernel launches
//! GPU-Resident Replay Buffer -- pure cudarc `CudaSlice` internals.
//!
//! 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 `GpuBatchPtrs` wraps gathered data as raw `CudaSlice` buffers
//! for downstream neural network consumption.
use std::sync::Arc;
use cudarc::driver::{CudaFunction, CudaSlice, CudaStream, DevicePtr, DevicePtrMut, LaunchConfig, PushKernelArg};
use ml_core::nvtx::NvtxRange;
use ml_core::MLError;
// ---------------------------------------------------------------------------
// GPU batch output (CudaSlice-based, no Candle Tensor dependency)
// ---------------------------------------------------------------------------
/// Pre-built GPU batch for training. All fields are raw `CudaSlice` on GPU.
#[allow(missing_debug_implementations)]
/// Zero-alloc batch: raw GPU pointers to pre-allocated sample buffers.
/// Valid until the next `sample_proportional` call (buffers are reused).
/// No cuMemAlloc, no DtoD clone — pure pointer pass-through.
pub struct GpuBatchPtrs {
pub states_ptr: u64,
pub next_states_ptr: u64,
pub actions_ptr: u64,
pub rewards_ptr: u64,
pub dones_ptr: u64,
pub weights_ptr: u64,
pub indices_ptr: u64,
pub episode_ids_ptr: u64,
pub batch_size: usize,
pub state_dim: usize,
}
// ---------------------------------------------------------------------------
// Compiled kernel cache
// ---------------------------------------------------------------------------
struct ReplayKernels {
scatter_insert_f32: CudaFunction,
scatter_insert_f32_rows: CudaFunction,
scatter_insert_u32: CudaFunction,
gather_f32: CudaFunction,
gather_u32: CudaFunction,
/// #30 F32 row gather for f32 state storage.
gather_f32_rows: CudaFunction,
/// Direct-to-trainer gather: f32 rows with zero-padding to dst_stride.
gather_f32_rows_padded: CudaFunction,
/// Direct-to-trainer gather: scalar f32 (rewards, dones).
gather_f32_scalar: CudaFunction,
/// Direct-to-trainer gather: scalar i32 (actions).
gather_i32_scalar: CudaFunction,
is_weights_f32: CudaFunction,
normalize_weights_f32: CudaFunction,
fill_from_gpu_f32: CudaFunction,
reduce_max_f32: CudaFunction,
i64_to_u32: CudaFunction,
max_of_two_f32: CudaFunction,
// PER prefix sum kernels
per_update_pa: CudaFunction,
per_insert_pa: CudaFunction,
per_prefix_scan: CudaFunction,
per_sample: CudaFunction,
/// C1/P1: Diversity-weighted priority kernel (health < 0.8 path).
pow_alpha_diverse_f32: CudaFunction,
/// Graph-safe reduction: scaled mean of u32 actions → i32[1] on GPU.
actions_sum_scale_reduce_u32: CudaFunction,
}
impl ReplayKernels {
fn compile(stream: &Arc<CudaStream>) -> Result<Self, MLError> {
let ctx = stream.context();
// Load precompiled replay buffer kernels cubin
static RB_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/replay_buffer_kernels.cubin"));
let rb_mod = ctx.load_cubin(RB_CUBIN.to_vec())
.map_err(|e| MLError::ModelError(format!("rb cubin load: {e}")))?;
let ld = |n: &str| -> Result<CudaFunction, MLError> {
rb_mod.load_function(n).map_err(|e| MLError::ModelError(format!("load {n}: {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")?,
scatter_insert_f32_rows: ld("scatter_insert_f32_rows")?,
scatter_insert_u32: ld("scatter_insert_u32")?,
gather_f32: ld("gather_f32")?, gather_u32: ld("gather_u32")?,
gather_f32_rows: ld("gather_f32_rows")?,
gather_f32_rows_padded: ld("gather_f32_rows_padded")?,
gather_f32_scalar: ld("gather_f32_scalar")?,
gather_i32_scalar: ld("gather_i32_scalar")?,
is_weights_f32: ld("is_weights_f32")?,
normalize_weights_f32: ld("normalize_weights_f32")?,
fill_from_gpu_f32: ld("fill_from_gpu_f32")?,
reduce_max_f32: ld("reduce_max_f32")?,
i64_to_u32: ld("i64_to_u32")?,
max_of_two_f32: ld("max_of_two_f32")?,
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")?,
pow_alpha_diverse_f32: ld("pow_alpha_diverse_f32")?,
actions_sum_scale_reduce_u32: ld("actions_sum_scale_reduce_u32")?,
})
}
}
fn lcfg(n: usize) -> LaunchConfig {
let t = 256_u32; let b = (n as u32).div_ceil(t);
LaunchConfig { grid_dim: (b.max(1), 1, 1), block_dim: (t, 1, 1), shared_mem_bytes: 0 }
}
// ---------------------------------------------------------------------------
// Config + Buffer
// ---------------------------------------------------------------------------
#[allow(clippy::module_name_repetitions)]
#[derive(Debug, Clone)]
pub struct GpuReplayBufferConfig {
pub capacity: usize, pub alpha: f32,
pub beta_start: f32, pub beta_max: f32, pub beta_annealing_steps: usize,
pub epsilon: f32, pub max_memory_bytes: usize,
/// Maximum batch size for pre-allocated sampling buffers. Defaults to 1024.
pub max_batch_size: usize,
}
pub struct GpuReplayBuffer {
config: GpuReplayBufferConfig,
stream: Arc<CudaStream>,
kernels: ReplayKernels,
states: CudaSlice<f32>, next_states: CudaSlice<f32>,
actions: CudaSlice<u32>, rewards: CudaSlice<f32>,
dones: CudaSlice<f32>, priorities: CudaSlice<f32>,
/// Episode IDs per buffer slot `[capacity]` u32 on GPU.
/// `episode_ids[i] = i / episode_length`. Written during `insert_batch_with_episode_ids`.
/// u32 chosen to match the gather_u32 / scatter_insert_u32 kernel signatures —
/// episode-id values are always non-negative counters.
episode_ids: CudaSlice<u32>,
write_cursor: usize, size: usize,
max_priority: CudaSlice<f32>,
pending_max_priority: Option<CudaSlice<f32>>,
current_step: 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>,
scan_tile_counter: CudaSlice<u32>,
// Pre-allocated PER sampling buffers (zero cuMemAlloc after warmup)
sample_indices_i64: CudaSlice<i64>,
sample_indices_u32: CudaSlice<u32>,
sample_states: CudaSlice<f32>,
sample_next_states: CudaSlice<f32>,
sample_actions: CudaSlice<u32>,
sample_rewards: CudaSlice<f32>,
sample_dones: CudaSlice<f32>,
sample_priorities: CudaSlice<f32>,
sample_weights: CudaSlice<f32>,
sample_max_weight: CudaSlice<f32>,
sample_episode_ids: CudaSlice<u32>,
total_sum_buf: CudaSlice<f32>,
// Pre-allocated scratch for update_priorities_gpu (zero cuMemAlloc per step)
update_batch_max: CudaSlice<f32>, // [1] per-batch atomicMax accumulator
update_max_merge: CudaSlice<f32>, // [1] merge pending + batch max
update_batch_spare: Option<CudaSlice<f32>>, // [1] spare for None→Some swap (zero alloc)
/// C1/P1: Pinned device-mapped scalar for the scaled-mean action value.
/// Written by `actions_sum_scale_reduce_u32` (via `mean_action_dev_ptr`),
/// read by `pow_alpha_diverse_f32` in the same graph. Pinned (not pure
/// device) so the same buffer is available to host for zero-sync logging
/// without a dtoh. Pattern matches `size_dev_ptr` / `rng_step_dev_ptr`
/// below — allocated with `cuMemAllocHost_v2` + `cuMemHostGetDevicePointer_v2`.
_mean_action_pinned: usize, // host addr (i32 slot) — kept alive for device-mapped pointer
mean_action_dev_ptr: u64, // device-mapped addr
// Pre-allocated insert scratch buffers (zero cuMemAlloc per insert)
insert_prio_buf: CudaSlice<f32>, // [insert_scratch_cap] priority broadcast
insert_idx_buf: CudaSlice<u32>, // [insert_scratch_cap] per_insert_pa indices
insert_ep_buf: CudaSlice<u32>, // [insert_scratch_cap] episode IDs (non-negative counters)
insert_scratch_cap: usize, // current allocation size
scratch_f32: CudaSlice<f32>, // [1] reusable temp for flush/apply_max
_rng_step: u32,
/// Pinned device-mapped RNG step counter — GPU writes, host reads.
_rng_step_pinned: usize,
rng_step_dev_ptr: u64,
/// Pinned device-mapped buffer size — GPU reads current size at graph replay.
/// Updated by host after each insert_batch. Kernels dereference pointer at runtime.
size_dev_ptr: u64,
size_pinned: usize,
/// Size of the most recent `sample_proportional` call. Used by `sample_weights_ref` /
/// `sample_indices_ref` to return correctly-sized views (pre-allocated buffers are
/// `max_batch_size` large; only the first `last_batch_size` elements are valid).
last_batch_size: usize,
// Trainer destination buffer pointers — set once at init, stable for graph capture.
// PER gather writes directly to these, eliminating DtoD + pad_states.
trainer_states_ptr: u64,
trainer_next_states_ptr: u64,
trainer_actions_ptr: u64,
trainer_rewards_ptr: u64,
trainer_dones_ptr: u64,
trainer_is_weights_ptr: u64,
trainer_state_dim_padded: usize,
/// C1/P1: Cached learning_health value from the trainer (updated once per epoch).
/// Controls whether priority updates use the standard or diversity-weighted kernel.
learning_health_cache: f32,
}
impl Drop for GpuReplayBuffer {
fn drop(&mut self) {
// Synchronize stream before CudaSlice fields drop.
// 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());
}
}
}
impl GpuReplayBuffer {
pub fn new(config: GpuReplayBufferConfig, stream: &Arc<CudaStream>) -> Result<Self, MLError> {
let (cap, sd) = (config.capacity, ml_core::state_layout::STATE_DIM);
let mbs = if config.max_batch_size == 0 { 1024 } else { config.max_batch_size };
let need = 2 * cap * sd * 2 + 5 * cap * 4;
if need > config.max_memory_bytes {
#[allow(clippy::integer_division)]
return Err(MLError::ModelError(format!(
"GPU replay buffer needs {} MB (limit {} MB)",
need / (1024 * 1024), config.max_memory_bytes / (1024 * 1024),
)));
}
let k = ReplayKernels::compile(stream)?;
let s = a32f(stream, cap * sd, "s")?;
let ns = a32f(stream, cap * sd, "ns")?;
let a = a32u(stream, cap, "a")?;
let r = a32f(stream, cap, "r")?;
let d = a32f(stream, cap, "d")?;
let p = a32f(stream, cap, "p")?;
let mut mp = a32f(stream, 1, "mp")?;
stream.memcpy_htod(&[1.0_f32], &mut mp).map_err(|e| MLError::ModelError(format!("mp: {e}")))?;
// 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")?;
let scan_tile_counter = stream.alloc_zeros::<u32>(1)
.map_err(|e| MLError::ModelError(format!("alloc scan_tile_counter: {e}")))?;
// 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}")))?;
let su32 = a32u(stream, mbs, "s_idx")?;
let ss = a32f(stream, mbs * sd, "s_states")?;
let sns = a32f(stream, mbs * sd, "s_nstates")?;
let sa = a32u(stream, mbs, "s_act")?;
let sr = a32f(stream, mbs, "s_rew")?;
let sdn = a32f(stream, mbs, "s_done")?;
let sp = a32f(stream, mbs, "s_pri")?;
let sw = a32f(stream, mbs, "s_wt")?;
let smw = a32f(stream, 1, "s_mw")?;
let tsb = a32f(stream, 1, "ts_buf")?;
let ep_ids = a32u(stream, cap, "episode_ids")?;
let s_ep = a32u(stream, mbs, "s_episode_ids")?;
// Pre-allocate PER update scratch buffers (zero cuMemAlloc per step)
let update_batch_max = a32f(stream, 1, "update_batch_max")?;
let update_max_merge = a32f(stream, 1, "update_max_merge")?;
let update_batch_spare = a32f(stream, 1, "update_batch_spare")?;
// Pre-allocate insert scratch buffers (zero cuMemAlloc per insert)
// Initial size = max_batch_size; grows if insert_batch receives larger batches.
let isc = mbs;
let insert_prio_buf = a32f(stream, isc, "ins_prio")?;
let insert_idx_buf = a32u(stream, isc, "ins_idx")?;
let insert_ep_buf = a32u(stream, isc, "ins_ep")?;
let scratch_f32 = a32f(stream, 1, "scratch")?;
// Pinned device-mapped buffer size for graph-safe PER — GPU reads current size at replay
let (size_pinned, size_dev_ptr) = unsafe {
let mut hp: *mut std::ffi::c_void = std::ptr::null_mut();
let mut dp: u64 = 0;
cudarc::driver::sys::cuMemAllocHost_v2(&mut hp, std::mem::size_of::<i32>());
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dp, hp, 0);
*(hp as *mut i32) = 0;
(hp as usize, dp)
};
// Pinned device-mapped RNG step counter for GPU-side increment
let (rng_step_pinned, rng_step_dev_ptr) = unsafe {
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
let mut dev_ptr: u64 = 0;
cudarc::driver::sys::cuMemAllocHost_v2(&mut host_ptr, std::mem::size_of::<i32>());
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dev_ptr, host_ptr, 0);
*(host_ptr as *mut i32) = 0;
(host_ptr as usize, dev_ptr) // usize for Send/Sync
};
// Pinned device-mapped scaled-mean action slot (i32). GPU writes in
// `actions_sum_scale_reduce_u32` and reads back in `pow_alpha_diverse_f32`
// — both via the device pointer. No dtoh, graph-safe across captures.
let (mean_action_pinned, mean_action_dev_ptr) = unsafe {
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
let mut dev_ptr: u64 = 0;
cudarc::driver::sys::cuMemAllocHost_v2(&mut host_ptr, std::mem::size_of::<i32>());
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dev_ptr, host_ptr, 0);
*(host_ptr as *mut i32) = 0;
(host_ptr as usize, dev_ptr)
};
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,
priorities_pa, prefix_sum,
scan_tile_state, scan_tile_aggregate, scan_tile_prefix, scan_tile_counter,
sample_indices_i64: si64,
sample_indices_u32: su32, sample_states: ss,
sample_next_states: sns, sample_actions: sa,
sample_rewards: sr, sample_dones: sdn,
sample_priorities: sp, sample_weights: sw,
sample_episode_ids: s_ep,
sample_max_weight: smw, total_sum_buf: tsb,
update_batch_max, update_max_merge, update_batch_spare: Some(update_batch_spare),
_mean_action_pinned: mean_action_pinned, mean_action_dev_ptr,
insert_prio_buf, insert_idx_buf, insert_ep_buf, insert_scratch_cap: isc,
scratch_f32,
_rng_step: 0,
_rng_step_pinned: rng_step_pinned,
rng_step_dev_ptr,
size_dev_ptr,
size_pinned,
last_batch_size: 0,
trainer_states_ptr: 0,
trainer_next_states_ptr: 0,
trainer_actions_ptr: 0,
trainer_rewards_ptr: 0,
trainer_dones_ptr: 0,
trainer_is_weights_ptr: 0,
trainer_state_dim_padded: 0,
learning_health_cache: 1.0,
})
}
pub const fn len(&self) -> usize { self.size }
pub const fn capacity(&self) -> usize { self.config.capacity }
pub const fn write_cursor(&self) -> usize { self.write_cursor }
pub const fn is_empty(&self) -> bool { self.size == 0 }
pub const fn can_sample(&self, bs: usize) -> bool { self.size >= bs }
pub fn current_beta(&self) -> f32 {
if self.config.beta_annealing_steps == 0 { return self.config.beta_max; }
let p = ((self.current_step as f32) / (self.config.beta_annealing_steps as f32)).min(1.0);
self.config.beta_start + (self.config.beta_max - self.config.beta_start) * p
}
pub const fn step(&mut self) { self.current_step = self.current_step.saturating_add(1); }
pub fn clear(&mut self) -> Result<(), MLError> {
self.write_cursor = 0; self.size = 0;
unsafe { *(self.size_pinned as *mut i32) = 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.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 }
pub const fn alpha(&self) -> f32 { self.config.alpha }
pub const fn epsilon(&self) -> f32 { self.config.epsilon }
pub const fn state_dim(&self) -> usize { ml_core::state_layout::STATE_DIM }
/// C1/P1: Update the cached learning_health value. Called from trainer at epoch boundary.
pub fn set_learning_health(&mut self, health: f32) {
self.learning_health_cache = health.clamp(0.0, 1.0);
}
/// Wire trainer destination buffer pointers for direct-to-trainer gather.
/// Called once after GpuDqnTrainer is constructed. Pointers are stable
/// (CudaSlice allocations never move), so this is safe for graph capture.
pub fn set_trainer_buffers(
&mut self,
states_ptr: u64,
next_states_ptr: u64,
actions_ptr: u64,
rewards_ptr: u64,
dones_ptr: u64,
is_weights_ptr: u64,
state_dim_padded: usize,
) {
self.trainer_states_ptr = states_ptr;
self.trainer_next_states_ptr = next_states_ptr;
self.trainer_actions_ptr = actions_ptr;
self.trainer_rewards_ptr = rewards_ptr;
self.trainer_dones_ptr = dones_ptr;
self.trainer_is_weights_ptr = is_weights_ptr;
self.trainer_state_dim_padded = state_dim_padded;
}
/// Pinned device-mapped RNG step counter pointer for GPU-side increment.
pub fn rng_step_dev_ptr(&self) -> u64 { self.rng_step_dev_ptr }
/// Grow insert scratch buffers if `eff` exceeds current allocation.
/// Only triggers cuMemAlloc on the first large insert — subsequent calls reuse.
fn ensure_insert_scratch(&mut self, eff: usize) -> Result<(), MLError> {
if eff <= self.insert_scratch_cap { return Ok(()); }
self.insert_prio_buf = a32f(&self.stream, eff, "ins_prio")?;
self.insert_idx_buf = a32u(&self.stream, eff, "ins_idx")?;
self.insert_ep_buf = a32u(&self.stream, eff, "ins_ep")?;
self.insert_scratch_cap = eff;
Ok(())
}
/// Insert a batch where states are already `CudaSlice<f32>` and
/// rewards/dones are `CudaSlice<f32>` (no f32 NaN risk).
///
/// Skips the f32→f32 cast for states (already f32). Rewards/dones
/// scatter-insert as f32 directly.
/// #30: Accept f32 states from the experience collector.
/// States are stored as f32 in the replay buffer (no f32 truncation).
pub fn insert_batch(
&mut self,
sf: &CudaSlice<f32>, // f32 states from experience collector
nf: &CudaSlice<f32>, // f32 next_states
ac: &CudaSlice<u32>,
rw: &CudaSlice<f32>,
dn: &CudaSlice<f32>,
bs: usize,
) -> Result<(), MLError> {
if bs == 0 { return Ok(()); }
let (cap, sd) = (self.config.capacity, ml_core::state_layout::STATE_DIM);
let (eff, off) = if bs > cap { (cap, bs - cap) } else { (bs, 0) };
self.ensure_insert_scratch(eff)?;
let el = eff * sd;
// #30: f32 states — 2D row scatter into ring buffer [cap, state_dim]
let ss = if off > 0 { sf.slice(off * sd..) } else { sf.slice(0..) };
let sn = if off > 0 { nf.slice(off * sd..) } else { nf.slice(0..) };
let (ci, cpi, sdi, bsi) = (self.write_cursor as i32, cap as i32, sd as i32, eff as i32);
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_f32_rows)
.arg(&self.states).arg(&ss).arg(&ci).arg(&cpi).arg(&sdi).arg(&bsi)
.launch(lcfg(el)).map_err(|e| MLError::ModelError(format!("sc s f32: {e}")))?;
}
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_f32_rows)
.arg(&self.next_states).arg(&sn).arg(&ci).arg(&cpi).arg(&sdi).arg(&bsi)
.launch(lcfg(el)).map_err(|e| MLError::ModelError(format!("sc n f32: {e}")))?;
}
let sa = if off > 0 { ac.slice(off..) } else { ac.slice(0..) };
{
let sr = if off > 0 { rw.slice(off..) } else { rw.slice(0..) };
let sd2 = if off > 0 { dn.slice(off..) } else { dn.slice(0..) };
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_u32)
.arg(&self.actions).arg(&sa).arg(&ci).arg(&cpi).arg(&bsi)
.launch(lcfg(eff)).map_err(|e| MLError::ModelError(format!("sc a: {e}")))?;
}
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_f32)
.arg(&self.rewards).arg(&sr).arg(&ci).arg(&cpi).arg(&bsi)
.launch(lcfg(eff)).map_err(|e| MLError::ModelError(format!("sc r f32: {e}")))?;
}
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_f32)
.arg(&self.dones).arg(&sd2).arg(&ci).arg(&cpi).arg(&bsi)
.launch(lcfg(eff)).map_err(|e| MLError::ModelError(format!("sc d f32: {e}")))?;
}
}
// Broadcast max_priority into pre-allocated prio buffer
unsafe {
self.stream.launch_builder(&self.kernels.fill_from_gpu_f32)
.arg(&mut self.insert_prio_buf).arg(&self.max_priority).arg(&bsi)
.launch(lcfg(eff))
.map_err(|e| MLError::ModelError(format!("fill mp: {e}")))?;
}
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_f32)
.arg(&self.priorities).arg(&self.insert_prio_buf).arg(&ci).arg(&cpi).arg(&bsi)
.launch(lcfg(eff)).map_err(|e| MLError::ModelError(format!("sc p: {e}")))?;
}
// Write priority^alpha to priorities_pa using pre-allocated buffers
let insert_indices: Vec<u32> = (0..eff)
.map(|j| ((self.write_cursor + j) % cap) as u32)
.collect();
self.stream.memcpy_htod(&insert_indices, &mut self.insert_idx_buf)
.map_err(|e| MLError::ModelError(format!("ib idx htod: {e}")))?;
let al = self.config.alpha;
let cap_i = self.config.capacity as i32;
unsafe {
self.stream.launch_builder(&self.kernels.per_insert_pa)
.arg(&self.priorities_pa).arg(&self.insert_idx_buf).arg(&self.insert_prio_buf).arg(&al)
.arg(&cap_i).arg(&bsi)
.launch(lcfg(eff))
.map_err(|e| MLError::ModelError(format!("per_insert_pa: {e}")))?;
}
// Episode IDs using pre-allocated buffer (u32 — always non-negative counters).
let ep_ids_host: Vec<u32> = (0..eff)
.map(|j| ((self.write_cursor + j) % cap) as u32)
.collect();
self.stream.memcpy_htod(&ep_ids_host, &mut self.insert_ep_buf)
.map_err(|e| MLError::ModelError(format!("ep htod: {e}")))?;
unsafe {
self.stream.launch_builder(&self.kernels.scatter_insert_u32)
.arg(&self.episode_ids).arg(&self.insert_ep_buf).arg(&ci).arg(&cpi).arg(&bsi)
.launch(lcfg(eff)).map_err(|e| MLError::ModelError(format!("sc ep: {e}")))?;
}
self.write_cursor = (self.write_cursor + eff) % cap;
self.size = (self.size + eff).min(cap);
// Update pinned device-mapped size — GPU kernels read this at graph replay
unsafe { *(self.size_pinned as *mut i32) = self.size as i32; }
Ok(())
}
pub fn sample_proportional(&mut self, batch_size: usize) -> Result<GpuBatchPtrs, MLError> {
let _nvtx = NvtxRange::new("per_sample_proportional");
if !self.can_sample(batch_size) {
return Err(MLError::ModelError(format!("Cannot sample {batch_size} from {}", self.size)));
}
let mbs = if self.config.max_batch_size == 0 { 1024 } else { self.config.max_batch_size };
if batch_size > mbs {
return Err(MLError::ModelError(format!(
"batch_size {batch_size} exceeds max_batch_size {mbs}"
)));
}
let (_n, sd) = (self.size, ml_core::state_layout::STATE_DIM);
let nb = -self.current_beta();
let bsi = batch_size as i32;
// Step 1: Prefix scan of priorities_pa[0..size]
// Zero tile state + counter before each scan (raw cuMemsetD8Async — graph-safe).
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.scan_tile_state.raw_ptr(), 0,
self.scan_tile_state.num_bytes(),
self.stream.cu_stream(),
);
cudarc::driver::sys::cuMemsetD8Async(
self.scan_tile_counter.raw_ptr(), 0,
self.scan_tile_counter.num_bytes(),
self.stream.cu_stream(),
);
}
// Grid uses capacity-based tile count (worst case). Kernel reads actual size
// from pinned device-mapped pointer — correct at graph replay time.
let cap_tiles = self.config.capacity.div_ceil(256);
let max_resident_blocks = 512_u32;
let scan_blocks = (cap_tiles as u32).min(max_resident_blocks);
let size_ptr = self.size_dev_ptr;
let dummy_tiles: i32 = 0; // unused — kernel computes from *size_ptr
unsafe {
self.stream.launch_builder(&self.kernels.per_prefix_scan)
.arg(&self.priorities_pa)
.arg(&mut self.prefix_sum)
.arg(&self.scan_tile_state)
.arg(&self.scan_tile_aggregate)
.arg(&self.scan_tile_prefix)
.arg(&self.scan_tile_counter)
.arg(&size_ptr)
.arg(&dummy_tiles)
.launch(LaunchConfig {
grid_dim: (scan_blocks.max(1), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 256 * 4,
})
.map_err(|e| MLError::ModelError(format!("per_prefix_scan: {e}")))?;
}
// Step 2: Sample via binary search — reads rng_step + size from pinned pointers
let rng_ptr = self.rng_step_dev_ptr;
unsafe {
self.stream.launch_builder(&self.kernels.per_sample)
.arg(&self.prefix_sum)
.arg(&self.priorities_pa)
.arg(&mut self.sample_indices_i64)
.arg(&mut self.sample_priorities)
.arg(&mut self.total_sum_buf)
.arg(&rng_ptr)
.arg(&size_ptr)
.arg(&bsi)
.launch(lcfg(batch_size))
.map_err(|e| MLError::ModelError(format!("per_sample: {e}")))?;
}
// Step 3: i64 -> u32 indices for output
// SAFETY: sample_indices_u32 and sample_indices_i64 are valid device allocations.
unsafe {
self.stream.launch_builder(&self.kernels.i64_to_u32)
.arg(&mut self.sample_indices_u32).arg(&self.sample_indices_i64).arg(&bsi)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("i2u: {e}")))?;
}
// Step 3: gather into trainer buffers (direct) or intermediate buffers (fallback).
let sdi = sd as i32;
let direct_to_trainer = self.trainer_states_ptr != 0;
if direct_to_trainer {
// ── Direct-to-trainer path: gather + pad in one kernel, zero DtoD copies ──
let sdp = self.trainer_state_dim_padded as i32;
let t_states = self.trainer_states_ptr;
let t_next = self.trainer_next_states_ptr;
let t_actions = self.trainer_actions_ptr;
let t_rewards = self.trainer_rewards_ptr;
let t_dones = self.trainer_dones_ptr;
// States: gather + pad → trainer.states_buf
unsafe {
self.stream.launch_builder(&self.kernels.gather_f32_rows_padded)
.arg(&t_states)
.arg(&self.states)
.arg(&self.sample_indices_i64)
.arg(&sdi)
.arg(&sdp)
.arg(&bsi)
.launch(lcfg(batch_size * self.trainer_state_dim_padded))
.map_err(|e| MLError::ModelError(format!("gather states padded: {e}")))?;
}
// Next states: gather + pad → trainer.next_states_buf
unsafe {
self.stream.launch_builder(&self.kernels.gather_f32_rows_padded)
.arg(&t_next)
.arg(&self.next_states)
.arg(&self.sample_indices_i64)
.arg(&sdi)
.arg(&sdp)
.arg(&bsi)
.launch(lcfg(batch_size * self.trainer_state_dim_padded))
.map_err(|e| MLError::ModelError(format!("gather next_states padded: {e}")))?;
}
// Actions: gather i32 scalar → trainer.actions_buf
unsafe {
self.stream.launch_builder(&self.kernels.gather_i32_scalar)
.arg(&t_actions)
.arg(&self.actions)
.arg(&self.sample_indices_i64)
.arg(&bsi)
.launch(lcfg(batch_size))
.map_err(|e| MLError::ModelError(format!("gather actions direct: {e}")))?;
}
// Rewards: gather f32 scalar → trainer.rewards_buf
unsafe {
self.stream.launch_builder(&self.kernels.gather_f32_scalar)
.arg(&t_rewards)
.arg(&self.rewards)
.arg(&self.sample_indices_i64)
.arg(&bsi)
.launch(lcfg(batch_size))
.map_err(|e| MLError::ModelError(format!("gather rewards direct: {e}")))?;
}
// Dones: gather f32 scalar → trainer.dones_buf
unsafe {
self.stream.launch_builder(&self.kernels.gather_f32_scalar)
.arg(&t_dones)
.arg(&self.dones)
.arg(&self.sample_indices_i64)
.arg(&bsi)
.launch(lcfg(batch_size))
.map_err(|e| MLError::ModelError(format!("gather dones direct: {e}")))?;
}
} else {
// ── Fallback: gather into intermediate sample_* buffers ──
// #30: gather f32 rows (full-precision state storage)
unsafe {
self.stream.launch_builder(&self.kernels.gather_f32_rows)
.arg(&mut self.sample_states).arg(&self.states).arg(&self.sample_indices_i64).arg(&sdi).arg(&bsi)
.launch(lcfg(batch_size * sd)).map_err(|e| MLError::ModelError(format!("g s f32: {e}")))?;
}
unsafe {
self.stream.launch_builder(&self.kernels.gather_f32_rows)
.arg(&mut self.sample_next_states).arg(&self.next_states).arg(&self.sample_indices_i64).arg(&sdi).arg(&bsi)
.launch(lcfg(batch_size * sd)).map_err(|e| MLError::ModelError(format!("g n f32: {e}")))?;
}
unsafe {
let cap_i32 = self.capacity() as i32;
self.stream.launch_builder(&self.kernels.gather_u32)
.arg(&mut self.sample_actions).arg(&self.actions).arg(&self.sample_indices_i64).arg(&bsi).arg(&cap_i32)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g a: {e}")))?;
}
unsafe {
let cap_i32 = self.capacity() as i32;
self.stream.launch_builder(&self.kernels.gather_f32)
.arg(&mut self.sample_rewards).arg(&self.rewards).arg(&self.sample_indices_i64).arg(&bsi).arg(&cap_i32)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g r: {e}")))?;
}
unsafe {
let cap_i32 = self.capacity() as i32;
self.stream.launch_builder(&self.kernels.gather_f32)
.arg(&mut self.sample_dones).arg(&self.dones).arg(&self.sample_indices_i64).arg(&bsi).arg(&cap_i32)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g d: {e}")))?;
}
}
// Step 3b: gather episode_ids for HER strategies. Buffers are CudaSlice<u32>
// natively — no element-type aliasing needed.
unsafe {
let cap_ep = self.capacity() as i32;
self.stream.launch_builder(&self.kernels.gather_u32)
.arg(&mut self.sample_episode_ids).arg(&self.episode_ids)
.arg(&self.sample_indices_i64).arg(&bsi).arg(&cap_ep)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g ep: {e}")))?;
}
// Step 6: IS weights via GPU-resident total_sum (zero CPU readback)
// is_weights_f32 reads total_sum from GPU pointer.
// When direct_to_trainer is active, write IS weights to trainer's is_weights_buf.
// Always write to sample_weights too (needed for reduce_max + normalize below).
unsafe {
self.stream.launch_builder(&self.kernels.is_weights_f32)
.arg(&mut self.sample_weights).arg(&self.sample_priorities).arg(&self.total_sum_buf)
.arg(&nb).arg(&size_ptr).arg(&bsi)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("isw: {e}")))?;
}
// Step 7: reduce max weight
unsafe {
cudarc::driver::sys::cuMemsetD8Async(
self.sample_max_weight.raw_ptr(), 0,
self.sample_max_weight.num_bytes(),
self.stream.cu_stream(),
);
}
let rt = 256_u32.min(batch_size as u32).max(1);
let rb = (batch_size as u32).div_ceil(rt);
// SAFETY: sample_weights, sample_max_weight are valid device allocations.
unsafe {
self.stream.launch_builder(&self.kernels.reduce_max_f32)
.arg(&self.sample_weights).arg(&mut self.sample_max_weight).arg(&bsi)
.launch(LaunchConfig { grid_dim: (rb.max(1),1,1), block_dim: (rt,1,1), shared_mem_bytes: rt*4 })
.map_err(|e| MLError::ModelError(format!("rm: {e}")))?;
}
// Step 8: normalize weights by max
// SAFETY: sample_weights, sample_max_weight are valid. Normalization divides element-wise.
unsafe {
self.stream.launch_builder(&self.kernels.normalize_weights_f32)
.arg(&mut self.sample_weights).arg(&self.sample_max_weight).arg(&bsi)
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("nw: {e}")))?;
}
// Step 8b: copy normalized IS weights to trainer buffer if direct_to_trainer
if direct_to_trainer {
let t_isw = self.trainer_is_weights_ptr;
let num_bytes = batch_size * std::mem::size_of::<f32>();
unsafe {
cudarc::driver::result::memcpy_dtod_async(
t_isw as cudarc::driver::sys::CUdeviceptr,
self.sample_weights.raw_ptr() as cudarc::driver::sys::CUdeviceptr,
num_bytes,
self.stream.cu_stream(),
).map_err(|e| MLError::ModelError(format!("isw dtod: {e}")))?;
}
}
// Record the batch size so that sample_weights_ref / sample_indices_ref return
// correctly-sized views (the pre-allocated buffers are max_batch_size elements,
// but only the first `batch_size` are valid after this sample).
self.last_batch_size = batch_size;
// Return pointers: when direct_to_trainer is active, point to trainer buffers
// (already populated by gather kernels above). Otherwise point to intermediate
// sample_* buffers for the legacy upload_batch_gpu path.
if direct_to_trainer {
Ok(GpuBatchPtrs {
states_ptr: self.trainer_states_ptr,
next_states_ptr: self.trainer_next_states_ptr,
actions_ptr: self.trainer_actions_ptr,
rewards_ptr: self.trainer_rewards_ptr,
dones_ptr: self.trainer_dones_ptr,
weights_ptr: self.trainer_is_weights_ptr,
indices_ptr: self.sample_indices_u32.raw_ptr(),
episode_ids_ptr: self.sample_episode_ids.raw_ptr(),
batch_size,
state_dim: sd,
})
} else {
Ok(GpuBatchPtrs {
states_ptr: self.sample_states.raw_ptr(),
next_states_ptr: self.sample_next_states.raw_ptr(),
actions_ptr: self.sample_actions.raw_ptr(),
rewards_ptr: self.sample_rewards.raw_ptr(),
dones_ptr: self.sample_dones.raw_ptr(),
weights_ptr: self.sample_weights.raw_ptr(),
indices_ptr: self.sample_indices_u32.raw_ptr(),
episode_ids_ptr: self.sample_episode_ids.raw_ptr(),
batch_size,
state_dim: sd,
})
}
}
/// Reference to pre-allocated sample indices buffer (valid after sample_proportional).
///
/// Returns a view sized to the most recent `sample_proportional` batch size.
/// The underlying buffer is `max_batch_size` elements; only `last_batch_size` are valid.
pub fn sample_indices_ref(&self) -> cudarc::driver::CudaView<'_, u32> {
self.sample_indices_u32.slice(..self.last_batch_size)
}
/// Reference to pre-allocated sample episode IDs buffer (valid after sample_proportional).
pub fn sample_episode_ids_ref(&self) -> &CudaSlice<u32> { &self.sample_episode_ids }
/// Reference to pre-allocated sample weights buffer (valid after sample_proportional).
///
/// Returns a view sized to the most recent `sample_proportional` batch size.
/// The underlying buffer is `max_batch_size` elements; only `last_batch_size` are valid.
pub fn sample_weights_ref(&self) -> cudarc::driver::CudaView<'_, f32> {
self.sample_weights.slice(..self.last_batch_size)
}
pub fn update_priorities_gpu_raw(&mut self, _indices_ptr: u64, td_errors: &CudaSlice<f32>, 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 prefix-sum kernel.
///
/// `per_update_pa` reads f32 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
/// when the TD errors were computed on a different stream (e.g., the trainer's
/// main stream). All pre-allocated scratch buffers are used — zero per-step
/// GPU allocations.
pub fn update_priorities_gpu(
&mut self,
indices: &CudaSlice<u32>,
td_errors: &CudaSlice<f32>,
bs: usize,
ext_stream: Option<&Arc<CudaStream>>,
) -> Result<(), MLError> {
if bs == 0 { return Ok(()); }
let stream_owned = ext_stream.cloned().unwrap_or_else(|| Arc::clone(&self.stream));
let stream = &stream_owned;
let (al, ep, bsi) = (self.config.alpha, self.config.epsilon, bs as i32);
let cap_i = self.config.capacity as i32;
let health = self.learning_health_cache;
stream.memset_zeros(&mut self.update_batch_max)
.map_err(|e| MLError::ModelError(format!("zero batch_max: {e}")))?;
if health < 0.8 {
// C1/P1: during collapse, boost priorities of experiences whose action
// deviates from the batch mean. Prior implementation did a stream sync
// + memcpy_dtoh to compute the mean on host — broken inside CUDA Graph
// capture (L40S smoke exposed the invalidation on Fold 1). Pipeline is
// now fully on-device and graph-safe:
// 1. gather_u32 → sample_actions[..gather_n] (u32)
// 2. actions_sum_scale_reduce_u32 → mean_action_dev_ptr (i32 scaled ×1000)
// 3. pow_alpha_diverse_f32 reads mean_action_dev_ptr
//
// Re-gather actions from the ring buffer using the current batch indices
// (sample_indices_i64 from the last sample_proportional call). This
// avoids depending on trainer_actions_ptr validity after graph replay.
let mbs = self.config.max_batch_size.max(1);
let gather_n = bs.min(mbs);
let gather_ni = gather_n as i32;
let cap_i32 = self.config.capacity as i32;
unsafe {
stream.launch_builder(&self.kernels.gather_u32)
.arg(&mut self.sample_actions)
.arg(&self.actions)
.arg(&self.sample_indices_i64)
.arg(&gather_ni)
.arg(&cap_i32)
.launch(lcfg(gather_n))
.map_err(|e| MLError::ModelError(format!("gather_u32 for diversity: {e}")))?;
}
// Single-block deterministic tree reduction → pinned i32 slot.
// 256 threads × 8 bytes (long long sdata) = 2 KB shared mem.
let mean_dev_ptr = self.mean_action_dev_ptr;
let reduce_cfg = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 256 * std::mem::size_of::<i64>() as u32,
};
unsafe {
stream.launch_builder(&self.kernels.actions_sum_scale_reduce_u32)
.arg(&self.sample_actions)
.arg(&mean_dev_ptr)
.arg(&gather_ni)
.launch(reduce_cfg)
.map_err(|e| MLError::ModelError(format!("actions_sum_scale_reduce_u32: {e}")))?;
}
// Reinterpret sample_actions (u32) as i32 for the kernel argument.
// The pointer type pun is safe: same device memory, same bit-width, kernel reads int*.
let actions_dev_ptr = self.sample_actions.raw_ptr() as u64;
unsafe {
stream.launch_builder(&self.kernels.pow_alpha_diverse_f32)
.arg(&self.priorities_pa)
.arg(&self.priorities)
.arg(&self.update_batch_max)
.arg(indices)
.arg(td_errors)
.arg(&actions_dev_ptr)
.arg(&mean_dev_ptr)
.arg(&al)
.arg(&ep)
.arg(&health)
.arg(&cap_i)
.arg(&bsi)
.launch(lcfg(bs))
.map_err(|e| MLError::ModelError(format!("pow_alpha_diverse_f32: {e}")))?;
}
} else {
// Standard PER (unchanged): reads f32 td_errors directly, writes priorities + priorities_pa
unsafe {
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}")))?;
}
}
// Merge batch max into pending (unchanged logic)
match self.pending_max_priority.take() {
Some(prev) => {
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}")))?;
}
self.pending_max_priority = Some(std::mem::replace(&mut self.update_max_merge, prev));
}
None => {
// Rotate: batch_max → pending, spare → batch_max (zero alloc)
// spare was pre-allocated at construction; replenished by flush_max_priority.
if let Some(spare) = self.update_batch_spare.take() {
self.pending_max_priority = Some(std::mem::replace(&mut self.update_batch_max, spare));
} else {
// Fallback: first epoch before flush replenishes spare
let fresh = a32f(stream, 1, "bm_epoch")?;
self.pending_max_priority = Some(std::mem::replace(&mut self.update_batch_max, fresh));
}
}
}
Ok(())
}
pub fn flush_max_priority(&mut self) -> Result<(), MLError> {
if let Some(pend) = self.pending_max_priority.take() {
unsafe {
self.stream.launch_builder(&self.kernels.max_of_two_f32)
.arg(&mut self.scratch_f32).arg(&pend).arg(&self.max_priority)
.launch(LaunchConfig { grid_dim: (1,1,1), block_dim: (1,1,1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("flush max_of_two: {e}")))?;
}
let num_bytes = std::mem::size_of::<f32>();
let src_ptr = {
let (ptr, guard) = self.scratch_f32.device_ptr(&self.stream);
let _no_drop = std::mem::ManuallyDrop::new(guard);
ptr
};
let dst_ptr = {
let (ptr, guard) = self.max_priority.device_ptr_mut(&self.stream);
let _no_drop = std::mem::ManuallyDrop::new(guard);
ptr
};
unsafe {
cudarc::driver::result::memcpy_dtod_async(
dst_ptr, src_ptr, num_bytes, self.stream.cu_stream(),
).map_err(|e| MLError::ModelError(format!("flush dtod: {e}")))?;
}
// Replenish spare with the consumed pend buffer (zero alloc next epoch)
self.update_batch_spare = Some(pend);
}
Ok(())
}
pub fn apply_max_priority_scalar(&mut self, mp: f32) -> Result<(), MLError> {
if mp > 0.0 {
self.stream.memcpy_htod(&[mp], &mut self.scratch_f32)
.map_err(|e| MLError::ModelError(format!("amp htod: {e}")))?;
// max(scratch_f32=mp, max_priority) → update_max_merge
unsafe {
self.stream.launch_builder(&self.kernels.max_of_two_f32)
.arg(&mut self.update_max_merge).arg(&self.scratch_f32).arg(&self.max_priority)
.launch(LaunchConfig { grid_dim: (1,1,1), block_dim: (1,1,1), shared_mem_bytes: 0 })
.map_err(|e| MLError::ModelError(format!("amp max: {e}")))?;
}
let num_bytes = std::mem::size_of::<f32>();
let src_ptr = {
let (ptr, guard) = self.update_max_merge.device_ptr(&self.stream);
let _no_drop = std::mem::ManuallyDrop::new(guard);
ptr
};
let dst_ptr = {
let (ptr, guard) = self.max_priority.device_ptr_mut(&self.stream);
let _no_drop = std::mem::ManuallyDrop::new(guard);
ptr
};
unsafe {
cudarc::driver::result::memcpy_dtod_async(
dst_ptr, src_ptr, num_bytes, self.stream.cu_stream(),
).map_err(|e| MLError::ModelError(format!("amp dtod: {e}")))?;
}
}
Ok(())
}
/// Raw `CudaSlice` accessors for direct GPU consumption.
pub const fn states_slice(&self) -> &CudaSlice<f32> { &self.states }
pub const fn next_states_slice(&self) -> &CudaSlice<f32> { &self.next_states }
pub const fn actions_slice(&self) -> &CudaSlice<u32> { &self.actions }
pub const fn rewards_slice(&self) -> &CudaSlice<f32> { &self.rewards }
pub const fn dones_slice(&self) -> &CudaSlice<f32> { &self.dones }
pub const fn priorities_slice(&self) -> &CudaSlice<f32> { &self.priorities }
}
impl std::fmt::Debug for GpuReplayBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GpuReplayBuffer")
.field("capacity", &self.config.capacity).field("size", &self.size)
.field("state_dim", &ml_core::state_layout::STATE_DIM)
.field("write_cursor", &self.write_cursor).finish()
}
}
fn a32f(s: &Arc<CudaStream>, n: usize, nm: &str) -> Result<CudaSlice<f32>, MLError> {
s.alloc_zeros::<f32>(n).map_err(|e| MLError::ModelError(format!("alloc {nm}: {e}")))
}
fn a32u(s: &Arc<CudaStream>, n: usize, nm: &str) -> Result<CudaSlice<u32>, MLError> {
s.alloc_zeros::<u32>(n).map_err(|e| MLError::ModelError(format!("alloc {nm}: {e}")))
}
#[cfg(test)]
mod tests {
use super::*;
fn make_stream() -> Arc<CudaStream> {
cudarc::driver::CudaContext::new(0)
.expect("CUDA required")
.new_stream()
.expect("CUDA stream")
}
#[test]
fn test_creation() {
let c = GpuReplayBufferConfig { capacity: 1000, alpha: 0.6, beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 100_000, epsilon: 1e-6, max_memory_bytes: 4<<30, max_batch_size: 256 };
let b = GpuReplayBuffer::new(c, &make_stream()).expect("buf");
assert_eq!(b.len(), 0); assert_eq!(b.capacity(), 1000); assert!(b.is_empty());
}
#[test]
fn test_beta() {
let c = GpuReplayBufferConfig { capacity: 100, alpha: 0.6, beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4<<30, max_batch_size: 64 };
let mut b = GpuReplayBuffer::new(c, &make_stream()).expect("buf");
assert!((b.current_beta() - 0.4).abs() < 1e-6);
for _ in 0..500 { b.step(); }
assert!(b.current_beta() > 0.4 && b.current_beta() < 1.0);
for _ in 0..600 { b.step(); }
assert!((b.current_beta() - 1.0).abs() < 1e-6);
}
#[test]
fn test_clear() {
let c = GpuReplayBufferConfig { capacity: 100, alpha: 0.6, beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4<<30, max_batch_size: 64 };
let mut b = GpuReplayBuffer::new(c, &make_stream()).expect("buf");
b.step(); b.clear().expect("clear");
assert_eq!(b.len(), 0); assert_eq!(b.current_step, 0);
}
/// Regression test: prefix scan must use `size` tiles, not `capacity` tiles.
/// With capacity=1M and size=256, the old code scanned 3907 tiles instead of 1.
/// On H100 with capacity=24.7M this caused a hang (96K tiles × decoupled lookback).
#[test]
#[ignore] // Requires CUDA GPU
fn test_per_prefix_scan_large_capacity() {
let stream = make_stream();
let cap = 1_000_000;
let sd = 4;
let bs = 64;
let c = GpuReplayBufferConfig {
capacity: cap, alpha: 0.6,
beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 1000,
epsilon: 1e-6, max_memory_bytes: 8 << 30, max_batch_size: bs,
};
let _ = sd;
let mut b = GpuReplayBuffer::new(c, &stream).expect("buf");
// Insert only 256 experiences into a 1M-slot buffer (ratio 3906:1)
let n = 256;
let s = stream.alloc_zeros::<f32>(n * sd).unwrap();
let ns = stream.alloc_zeros::<f32>(n * sd).unwrap();
let a = stream.alloc_zeros::<u32>(n).unwrap();
let mut r = stream.alloc_zeros::<f32>(n).unwrap();
let d = stream.alloc_zeros::<f32>(n).unwrap();
// Give non-zero rewards so priorities are non-zero after update
let rewards_host: Vec<f32> = (0..n).map(|i| (i as f32 + 1.0) * 0.01).collect();
stream.memcpy_htod(&rewards_host, &mut r).unwrap();
b.insert_batch(&s, &ns, &a, &r, &d, n).unwrap();
assert_eq!(b.len(), n);
// sample_proportional triggers per_prefix_scan — must complete, not hang
let batch = b.sample_proportional(bs).expect("sample must not hang");
assert_eq!(batch.batch_size, bs);
// Verify all sampled indices are within [0, size)
let mut host_idx = vec![0u32; bs];
unsafe {
cudarc::driver::sys::cuStreamSynchronize(stream.cu_stream());
}
stream.memcpy_dtoh(&b.sample_indices_u32.slice(..bs), &mut host_idx).unwrap();
for (i, &idx) in host_idx.iter().enumerate() {
assert!(
(idx as usize) < n,
"index {i} = {idx} exceeds buffer size {n}"
);
}
}
}