|
|
|
|
@@ -206,6 +206,7 @@ struct ReplayKernels {
|
|
|
|
|
max_of_two_f32: CudaFunction,
|
|
|
|
|
searchsorted: CudaFunction,
|
|
|
|
|
prefix_sum: CudaFunction,
|
|
|
|
|
per_gen_thresholds: CudaFunction,
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
impl ReplayKernels {
|
|
|
|
|
@@ -227,6 +228,10 @@ impl ReplayKernels {
|
|
|
|
|
include_str!("prefix_sum_kernel.cu"), ctx,
|
|
|
|
|
).map_err(|e| MLError::ModelError(format!("ps compile: {e}")))?;
|
|
|
|
|
let ps_mod = ctx.load_module(ps_ptx).map_err(|e| MLError::ModelError(format!("ps mod: {e}")))?;
|
|
|
|
|
let pt_ptx = ml_core::cuda_compile::compile_ptx_for_device(
|
|
|
|
|
include_str!("per_threshold_kernel.cu"), ctx,
|
|
|
|
|
).map_err(|e| MLError::ModelError(format!("pt compile: {e}")))?;
|
|
|
|
|
let pt_mod = ctx.load_module(pt_ptx).map_err(|e| MLError::ModelError(format!("pt mod: {e}")))?;
|
|
|
|
|
|
|
|
|
|
Ok(Self {
|
|
|
|
|
scatter_insert_f32: ld("scatter_insert_f32")?,
|
|
|
|
|
@@ -247,6 +252,8 @@ impl ReplayKernels {
|
|
|
|
|
.map_err(|e| MLError::ModelError(format!("ss fn: {e}")))?,
|
|
|
|
|
prefix_sum: ps_mod.load_function("prefix_sum_kernel")
|
|
|
|
|
.map_err(|e| MLError::ModelError(format!("ps fn: {e}")))?,
|
|
|
|
|
per_gen_thresholds: pt_mod.load_function("per_generate_thresholds")
|
|
|
|
|
.map_err(|e| MLError::ModelError(format!("pt fn: {e}")))?,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
@@ -266,6 +273,8 @@ pub struct GpuReplayBufferConfig {
|
|
|
|
|
pub capacity: usize, pub state_dim: 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 {
|
|
|
|
|
@@ -280,11 +289,26 @@ pub struct GpuReplayBuffer {
|
|
|
|
|
pending_max_priority: Option<CudaSlice<f32>>,
|
|
|
|
|
current_step: usize,
|
|
|
|
|
pa_buf: CudaSlice<f32>, cs_buf: CudaSlice<f32>,
|
|
|
|
|
// Pre-allocated PER sampling buffers (zero cuMemAlloc after warmup)
|
|
|
|
|
sample_thresholds: CudaSlice<f32>,
|
|
|
|
|
sample_indices_i64: CudaSlice<i64>,
|
|
|
|
|
sample_indices_u32: CudaSlice<u32>,
|
|
|
|
|
sample_states: CudaSlice<u16>,
|
|
|
|
|
sample_next_states: CudaSlice<u16>,
|
|
|
|
|
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>,
|
|
|
|
|
total_sum_buf: CudaSlice<f32>,
|
|
|
|
|
rng_step: u32,
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
impl GpuReplayBuffer {
|
|
|
|
|
pub fn new(config: GpuReplayBufferConfig, stream: &Arc<CudaStream>) -> Result<Self, MLError> {
|
|
|
|
|
let (cap, sd) = (config.capacity, config.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)]
|
|
|
|
|
@@ -305,12 +329,34 @@ impl GpuReplayBuffer {
|
|
|
|
|
|
|
|
|
|
let pa = a32f(stream, cap, "pa")?;
|
|
|
|
|
let cs = a32f(stream, cap, "cs")?;
|
|
|
|
|
|
|
|
|
|
// Pre-allocate PER sampling buffers (zero cuMemAlloc after warmup)
|
|
|
|
|
let st = a32f(stream, mbs, "s_thresh")?;
|
|
|
|
|
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 = a16(stream, mbs * sd, "s_states")?;
|
|
|
|
|
let sns = a16(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")?;
|
|
|
|
|
|
|
|
|
|
Ok(Self {
|
|
|
|
|
config, stream: Arc::clone(stream), kernels: k,
|
|
|
|
|
states: s, next_states: ns, actions: a, rewards: r, dones: d, priorities: p,
|
|
|
|
|
write_cursor: 0, size: 0, max_priority: mp,
|
|
|
|
|
pending_max_priority: None, current_step: 0,
|
|
|
|
|
pa_buf: pa, cs_buf: cs,
|
|
|
|
|
sample_thresholds: st, 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_max_weight: smw, total_sum_buf: tsb,
|
|
|
|
|
rng_step: 0,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
@@ -416,100 +462,159 @@ impl GpuReplayBuffer {
|
|
|
|
|
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, self.config.state_dim);
|
|
|
|
|
let (al, nb) = (self.config.alpha, -self.current_beta());
|
|
|
|
|
let (ni, bsi) = (n as i32, batch_size as i32);
|
|
|
|
|
|
|
|
|
|
// Step 1: pow_alpha on priorities -> pa_buf (GPU-only)
|
|
|
|
|
// SAFETY: pa_buf and priorities are valid device allocations on same context.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.pow_alpha_f32)
|
|
|
|
|
.arg(&self.pa_buf).arg(&self.priorities).arg(&al).arg(&ni)
|
|
|
|
|
.launch(lcfg(n)).map_err(|e| MLError::ModelError(format!("pow: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Step 2: prefix sum (GPU-only)
|
|
|
|
|
self.pfx_sum(n)?;
|
|
|
|
|
let ts = self.cs_total(n)?;
|
|
|
|
|
if ts <= 0.0 || !ts.is_finite() { return Err(MLError::ModelError(format!("bad total_sum: {ts}"))); }
|
|
|
|
|
let mut th: Vec<f32> = Vec::with_capacity(batch_size);
|
|
|
|
|
{ use rand::Rng; let mut r = rand::thread_rng();
|
|
|
|
|
for _ in 0..batch_size { th.push(r.gen::<f32>() * ts); } }
|
|
|
|
|
let mut tb = a32f(&self.stream, batch_size, "t")?;
|
|
|
|
|
self.stream.memcpy_htod(&th, &mut tb).map_err(|e| MLError::ModelError(format!("t: {e}")))?;
|
|
|
|
|
let mut i64 = self.stream.alloc_zeros::<i64>(batch_size).map_err(|e| MLError::ModelError(format!("i64: {e}")))?;
|
|
|
|
|
// SAFETY: cs_buf, tb, i64 are valid device allocations. Sizes checked above.
|
|
|
|
|
|
|
|
|
|
// Step 3: DtoD copy of cs_buf[n-1] -> total_sum_buf (replaces memcpy_dtoh)
|
|
|
|
|
{
|
|
|
|
|
let cs_last = self.cs_buf.slice((n - 1)..n);
|
|
|
|
|
let num_bytes = std::mem::size_of::<f32>();
|
|
|
|
|
let src_ptr = {
|
|
|
|
|
let (ptr, guard) = cs_last.device_ptr(&self.stream);
|
|
|
|
|
let _no_drop = std::mem::ManuallyDrop::new(guard);
|
|
|
|
|
ptr
|
|
|
|
|
};
|
|
|
|
|
let dst_ptr = {
|
|
|
|
|
let (ptr, guard) = self.total_sum_buf.device_ptr_mut(&self.stream);
|
|
|
|
|
let _no_drop = std::mem::ManuallyDrop::new(guard);
|
|
|
|
|
ptr
|
|
|
|
|
};
|
|
|
|
|
// SAFETY: src_ptr and dst_ptr are valid device pointers. num_bytes = sizeof(f32).
|
|
|
|
|
unsafe {
|
|
|
|
|
cudarc::driver::result::memcpy_dtod_async(
|
|
|
|
|
dst_ptr, src_ptr, num_bytes, self.stream.cu_stream(),
|
|
|
|
|
).map_err(|e| MLError::ModelError(format!("cs dtod: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Step 4: GPU-resident Philox RNG for thresholds (replaces CPU rand + memcpy_htod)
|
|
|
|
|
self.rng_step = self.rng_step.wrapping_add(1);
|
|
|
|
|
let seed = self.rng_step;
|
|
|
|
|
// SAFETY: sample_thresholds, total_sum_buf are valid device allocations of sufficient size.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.per_gen_thresholds)
|
|
|
|
|
.arg(&mut self.sample_thresholds).arg(&self.total_sum_buf)
|
|
|
|
|
.arg(&seed).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("rng: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Step 5: searchsorted into pre-allocated i64 buffer
|
|
|
|
|
// SAFETY: cs_buf, sample_thresholds, sample_indices_i64 are valid device allocations.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.searchsorted)
|
|
|
|
|
.arg(&self.cs_buf).arg(&tb).arg(&mut i64).arg(&ni).arg(&bsi)
|
|
|
|
|
.arg(&self.cs_buf).arg(&self.sample_thresholds).arg(&mut self.sample_indices_i64).arg(&ni).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("ss: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
let mut i32b = a32u(&self.stream, batch_size, "i32")?;
|
|
|
|
|
// SAFETY: i32b and i64 are valid device allocations of at least batch_size elements.
|
|
|
|
|
|
|
|
|
|
// Step 6: i64 -> u32 indices into pre-allocated buffer
|
|
|
|
|
// SAFETY: sample_indices_u32 and sample_indices_i64 are valid device allocations.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.i64_to_u32)
|
|
|
|
|
.arg(&mut i32b).arg(&i64).arg(&bsi)
|
|
|
|
|
.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}")))?;
|
|
|
|
|
}
|
|
|
|
|
let (mut gs, mut gn) = (a16(&self.stream, batch_size * sd, "gs")?, a16(&self.stream, batch_size * sd, "gn")?);
|
|
|
|
|
let (mut ga, mut gr, mut gd) = (a32u(&self.stream, batch_size, "ga")?, a32f(&self.stream, batch_size, "gr")?, a32f(&self.stream, batch_size, "gd")?);
|
|
|
|
|
|
|
|
|
|
// Step 7: gather into pre-allocated buffers
|
|
|
|
|
let sdi = sd as i32;
|
|
|
|
|
// SAFETY: gs, states, i64, sdi, bsi are valid device allocations. Gather indices within buffer size.
|
|
|
|
|
// SAFETY: sample_states, states, sample_indices_i64 are valid device allocations.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_bf16_rows).arg(&mut gs).arg(&self.states).arg(&i64).arg(&sdi).arg(&bsi)
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_bf16_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: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
// SAFETY: same as above for next_states gather.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_bf16_rows).arg(&mut gn).arg(&self.next_states).arg(&i64).arg(&sdi).arg(&bsi)
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_bf16_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: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
// SAFETY: same context, actions buffer valid.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_u32).arg(&mut ga).arg(&self.actions).arg(&i64).arg(&bsi)
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_u32)
|
|
|
|
|
.arg(&mut self.sample_actions).arg(&self.actions).arg(&self.sample_indices_i64).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g a: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
// SAFETY: same context, rewards buffer valid.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_f32).arg(&mut gr).arg(&self.rewards).arg(&i64).arg(&bsi)
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_f32)
|
|
|
|
|
.arg(&mut self.sample_rewards).arg(&self.rewards).arg(&self.sample_indices_i64).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g r: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
// SAFETY: same context, dones buffer valid.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_f32).arg(&mut gd).arg(&self.dones).arg(&i64).arg(&bsi)
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_f32)
|
|
|
|
|
.arg(&mut self.sample_dones).arg(&self.dones).arg(&self.sample_indices_i64).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g d: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
let mut sp = a32f(&self.stream, batch_size, "sp")?;
|
|
|
|
|
let mut wt = a32f(&self.stream, batch_size, "wt")?;
|
|
|
|
|
// SAFETY: sp, pa_buf, i64 are valid device allocations. Gather indices within buffer size.
|
|
|
|
|
|
|
|
|
|
// Step 8: gather sampled priorities for IS weight computation
|
|
|
|
|
// SAFETY: sample_priorities, pa_buf, sample_indices_i64 are valid device allocations.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_f32).arg(&mut sp).arg(&self.pa_buf).arg(&i64).arg(&bsi)
|
|
|
|
|
self.stream.launch_builder(&self.kernels.gather_f32)
|
|
|
|
|
.arg(&mut self.sample_priorities).arg(&self.pa_buf).arg(&self.sample_indices_i64).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("g sp: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
// SAFETY: wt, sp are valid device allocations. IS weight computation is bounds-checked.
|
|
|
|
|
|
|
|
|
|
// Step 9: IS weights via GPU-resident total_sum (zero CPU readback)
|
|
|
|
|
// is_weights_f32 now reads total_sum from GPU pointer, not scalar arg.
|
|
|
|
|
// SAFETY: sample_weights, sample_priorities, total_sum_buf are valid device allocations.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.is_weights_f32).arg(&mut wt).arg(&sp).arg(&ts).arg(&nb).arg(&(n as i32)).arg(&bsi)
|
|
|
|
|
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(&(n as i32)).arg(&bsi)
|
|
|
|
|
.launch(lcfg(batch_size)).map_err(|e| MLError::ModelError(format!("isw: {e}")))?;
|
|
|
|
|
}
|
|
|
|
|
let mut mw = a32f(&self.stream, 1, "mw")?;
|
|
|
|
|
self.stream.memcpy_htod(&[0.0_f32], &mut mw).map_err(|e| MLError::ModelError(format!("mw: {e}")))?;
|
|
|
|
|
|
|
|
|
|
// Step 10: reduce max weight (memset_zeros replaces memcpy_htod)
|
|
|
|
|
self.stream.memset_zeros(&mut self.sample_max_weight)
|
|
|
|
|
.map_err(|e| MLError::ModelError(format!("mw zero: {e}")))?;
|
|
|
|
|
let rt = 256_u32.min(batch_size as u32).max(1);
|
|
|
|
|
let rb = (batch_size as u32).div_ceil(rt);
|
|
|
|
|
// SAFETY: wt, mw are valid device allocations. Reduction kernel uses shared memory within limits.
|
|
|
|
|
// SAFETY: sample_weights, sample_max_weight are valid device allocations.
|
|
|
|
|
unsafe {
|
|
|
|
|
self.stream.launch_builder(&self.kernels.reduce_max_f32).arg(&wt).arg(&mut mw).arg(&bsi)
|
|
|
|
|
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}")))?;
|
|
|
|
|
}
|
|
|
|
|
// SAFETY: wt, mw are valid. Normalization kernel divides element-wise.
|
|
|
|
|
|
|
|
|
|
// Step 11: 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 wt).arg(&mw).arg(&bsi)
|
|
|
|
|
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}")))?;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Return DtoD clones of pre-allocated slices sized to actual batch_size.
|
|
|
|
|
// The caller owns the returned GpuBatchSlices (consumed by into_gpu_batch),
|
|
|
|
|
// so we DtoD-clone the relevant portions. All copies are async on the stream.
|
|
|
|
|
Ok(GpuBatchSlices {
|
|
|
|
|
states: gs,
|
|
|
|
|
next_states: gn,
|
|
|
|
|
actions: ga,
|
|
|
|
|
rewards: gr,
|
|
|
|
|
dones: gd,
|
|
|
|
|
weights: wt,
|
|
|
|
|
indices: i32b,
|
|
|
|
|
states: dtod_clone_u16(&self.stream, &self.sample_states, batch_size * sd, "o_s")?,
|
|
|
|
|
next_states: dtod_clone_u16(&self.stream, &self.sample_next_states, batch_size * sd, "o_n")?,
|
|
|
|
|
actions: dtod_clone_u32(&self.stream, &self.sample_actions, batch_size, "o_act")?,
|
|
|
|
|
rewards: dtod_clone_f32(&self.stream, &self.sample_rewards, batch_size, "o_r")?,
|
|
|
|
|
dones: dtod_clone_f32(&self.stream, &self.sample_dones, batch_size, "o_d")?,
|
|
|
|
|
weights: dtod_clone_f32(&self.stream, &self.sample_weights, batch_size, "o_w")?,
|
|
|
|
|
indices: dtod_clone_u32(&self.stream, &self.sample_indices_u32, batch_size, "o_i")?,
|
|
|
|
|
batch_size,
|
|
|
|
|
state_dim: sd,
|
|
|
|
|
})
|
|
|
|
|
@@ -666,16 +771,6 @@ impl GpuReplayBuffer {
|
|
|
|
|
Ok(())
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// Read prefix-sum total (single scalar readback -- exits system for RNG threshold generation).
|
|
|
|
|
fn cs_total(&self, n: usize) -> Result<f32, MLError> {
|
|
|
|
|
if n == 0 { return Ok(0.0); }
|
|
|
|
|
let v = self.cs_buf.slice((n-1)..n);
|
|
|
|
|
let mut h = [0.0_f32];
|
|
|
|
|
// ALLOWED: single scalar readback (1 float) used to generate random thresholds on CPU.
|
|
|
|
|
// The total_sum value exits the GPU pipeline entirely to seed uniform sampling.
|
|
|
|
|
self.stream.memcpy_dtoh(&v, &mut h).map_err(|e| MLError::ModelError(format!("{e}")))?; // gpu-exit: 1 scalar for RNG seed
|
|
|
|
|
Ok(h[0])
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
impl std::fmt::Debug for GpuReplayBuffer {
|
|
|
|
|
@@ -697,6 +792,30 @@ fn a16(s: &Arc<CudaStream>, n: usize, nm: &str) -> Result<CudaSlice<u16>, MLErro
|
|
|
|
|
s.alloc_zeros::<u16>(n).map_err(|e| MLError::ModelError(format!("alloc {nm}: {e}")))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// DtoD clone of first `n` elements from `src` into a new allocation (async on stream).
|
|
|
|
|
fn dtod_clone_f32(s: &Arc<CudaStream>, src: &CudaSlice<f32>, n: usize, nm: &str) -> Result<CudaSlice<f32>, MLError> {
|
|
|
|
|
let mut dst = a32f(s, n, nm)?;
|
|
|
|
|
let sv = src.slice(..n);
|
|
|
|
|
s.memcpy_dtod(&sv, &mut dst).map_err(|e| MLError::ModelError(format!("dtod {nm}: {e}")))?;
|
|
|
|
|
Ok(dst)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// DtoD clone of first `n` elements from `src` into a new allocation (async on stream).
|
|
|
|
|
fn dtod_clone_u32(s: &Arc<CudaStream>, src: &CudaSlice<u32>, n: usize, nm: &str) -> Result<CudaSlice<u32>, MLError> {
|
|
|
|
|
let mut dst = a32u(s, n, nm)?;
|
|
|
|
|
let sv = src.slice(..n);
|
|
|
|
|
s.memcpy_dtod(&sv, &mut dst).map_err(|e| MLError::ModelError(format!("dtod {nm}: {e}")))?;
|
|
|
|
|
Ok(dst)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
/// DtoD clone of first `n` elements from `src` into a new allocation (async on stream).
|
|
|
|
|
fn dtod_clone_u16(s: &Arc<CudaStream>, src: &CudaSlice<u16>, n: usize, nm: &str) -> Result<CudaSlice<u16>, MLError> {
|
|
|
|
|
let mut dst = a16(s, n, nm)?;
|
|
|
|
|
let sv = src.slice(..n);
|
|
|
|
|
s.memcpy_dtod(&sv, &mut dst).map_err(|e| MLError::ModelError(format!("dtod {nm}: {e}")))?;
|
|
|
|
|
Ok(dst)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
#[cfg(test)]
|
|
|
|
|
mod tests {
|
|
|
|
|
use super::*;
|
|
|
|
|
@@ -710,13 +829,13 @@ mod tests {
|
|
|
|
|
|
|
|
|
|
#[test]
|
|
|
|
|
fn test_creation() {
|
|
|
|
|
let c = GpuReplayBufferConfig { capacity: 1000, state_dim: 48, alpha: 0.6, beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 100_000, epsilon: 1e-6, max_memory_bytes: 4<<30 };
|
|
|
|
|
let c = GpuReplayBufferConfig { capacity: 1000, state_dim: 48, 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, state_dim: 4, alpha: 0.6, beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4<<30 };
|
|
|
|
|
let c = GpuReplayBufferConfig { capacity: 100, state_dim: 4, 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(); }
|
|
|
|
|
@@ -726,7 +845,7 @@ mod tests {
|
|
|
|
|
}
|
|
|
|
|
#[test]
|
|
|
|
|
fn test_clear() {
|
|
|
|
|
let c = GpuReplayBufferConfig { capacity: 100, state_dim: 4, alpha: 0.6, beta_start: 0.4, beta_max: 1.0, beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4<<30 };
|
|
|
|
|
let c = GpuReplayBufferConfig { capacity: 100, state_dim: 4, 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);
|
|
|
|
|
|