diff --git a/crates/ml-dqn/src/gpu_replay_buffer.rs b/crates/ml-dqn/src/gpu_replay_buffer.rs index dc8368cd3..f3a0aebb8 100644 --- a/crates/ml-dqn/src/gpu_replay_buffer.rs +++ b/crates/ml-dqn/src/gpu_replay_buffer.rs @@ -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>, current_step: usize, pa_buf: CudaSlice, cs_buf: CudaSlice, + // Pre-allocated PER sampling buffers (zero cuMemAlloc after warmup) + sample_thresholds: CudaSlice, + sample_indices_i64: CudaSlice, + sample_indices_u32: CudaSlice, + sample_states: CudaSlice, + sample_next_states: CudaSlice, + sample_actions: CudaSlice, + sample_rewards: CudaSlice, + sample_dones: CudaSlice, + sample_priorities: CudaSlice, + sample_weights: CudaSlice, + sample_max_weight: CudaSlice, + total_sum_buf: CudaSlice, + rng_step: u32, } impl GpuReplayBuffer { pub fn new(config: GpuReplayBufferConfig, stream: &Arc) -> Result { 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::(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 = Vec::with_capacity(batch_size); - { use rand::Rng; let mut r = rand::thread_rng(); - for _ in 0..batch_size { th.push(r.gen::() * 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::(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::(); + 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 { - 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, n: usize, nm: &str) -> Result, MLErro s.alloc_zeros::(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, src: &CudaSlice, n: usize, nm: &str) -> Result, 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, src: &CudaSlice, n: usize, nm: &str) -> Result, 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, src: &CudaSlice, n: usize, nm: &str) -> Result, 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); diff --git a/crates/ml-dqn/src/replay_buffer_type.rs b/crates/ml-dqn/src/replay_buffer_type.rs index 4c3a50c22..0d74136bb 100644 --- a/crates/ml-dqn/src/replay_buffer_type.rs +++ b/crates/ml-dqn/src/replay_buffer_type.rs @@ -188,6 +188,7 @@ impl ReplayBufferType { beta_annealing_steps, epsilon: 1e-6, max_memory_bytes, + max_batch_size: 1024, }; let buffer = GpuReplayBuffer::new(config, stream)?; diff --git a/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs b/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs index 7ae1d6c9f..d0cd7eb95 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/gpu_residency.rs @@ -21,6 +21,7 @@ fn test_buffer_config(capacity: usize, state_dim: usize) -> GpuReplayBufferConfi beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4 * 1024 * 1024 * 1024, + max_batch_size: 1024, } } diff --git a/crates/ml/src/trainers/dqn/smoke_tests/performance.rs b/crates/ml/src/trainers/dqn/smoke_tests/performance.rs index 8af3651f1..1a6f727f7 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/performance.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/performance.rs @@ -66,6 +66,7 @@ async fn test_per_sample_latency() -> anyhow::Result<()> { beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4 * 1024 * 1024 * 1024, + max_batch_size: 1024, }; let mut buf = GpuReplayBuffer::new(config, &stream)?; diff --git a/crates/ml/src/trainers/dqn/smoke_tests/training_stability.rs b/crates/ml/src/trainers/dqn/smoke_tests/training_stability.rs index 761417025..9f139ae03 100644 --- a/crates/ml/src/trainers/dqn/smoke_tests/training_stability.rs +++ b/crates/ml/src/trainers/dqn/smoke_tests/training_stability.rs @@ -57,6 +57,7 @@ async fn test_per_weights_valid() -> anyhow::Result<()> { beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4 * 1024 * 1024 * 1024, + max_batch_size: 1024, }; let mut buf = GpuReplayBuffer::new(config, &stream)?; @@ -101,6 +102,7 @@ async fn test_per_indices_valid() -> anyhow::Result<()> { beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4 * 1024 * 1024 * 1024, + max_batch_size: 1024, }; let mut buf = GpuReplayBuffer::new(config, &stream)?; diff --git a/crates/ml/tests/gpu_per_integration_test.rs b/crates/ml/tests/gpu_per_integration_test.rs index 85aba4277..cca6eaa68 100644 --- a/crates/ml/tests/gpu_per_integration_test.rs +++ b/crates/ml/tests/gpu_per_integration_test.rs @@ -101,6 +101,7 @@ fn test_config(capacity: usize, state_dim: usize) -> GpuReplayBufferConfig { beta_annealing_steps: 1000, epsilon: 1e-6, max_memory_bytes: 4 * 1024 * 1024 * 1024, + max_batch_size: 1024, } }