diff --git a/crates/ml-dqn/src/gpu_replay_buffer.rs b/crates/ml-dqn/src/gpu_replay_buffer.rs index 852e7ce8c..607caf604 100644 --- a/crates/ml-dqn/src/gpu_replay_buffer.rs +++ b/crates/ml-dqn/src/gpu_replay_buffer.rs @@ -753,9 +753,11 @@ impl GpuReplayBuffer { 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"); - let kernels = get_cast_kernels(stream)?; - eprintln!("PER_DIAG: launching bf16_to_f32 cast"); + 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 {