diff --git a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs index fae17054a..7ef94913d 100644 --- a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs +++ b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs @@ -1524,17 +1524,8 @@ fn dtod_clone_f32( let mut dst = stream.alloc_zeros::(n_elems) .map_err(|e| MLError::ModelError(format!("alloc {label} f32[{n_elems}]: {e}")))?; let src_view = src.slice(..n_elems); - let (src_ptr, _src_sync) = src_view.device_ptr(stream); - let (dst_ptr, _dst_sync) = dst.device_ptr_mut(stream); - let num_bytes = n_elems * std::mem::size_of::(); - unsafe { - cudarc::driver::result::memcpy_dtod_async( - dst_ptr, src_ptr, num_bytes, stream.cu_stream(), - ).map_err(|e| MLError::ModelError(format!("DtoD {label}: {e}")))?; - } - // Drop sync guards before moving dst out (SyncOnDrop borrows dst) - drop(_dst_sync); - drop(_src_sync); + stream.memcpy_dtod(&src_view, &mut dst) + .map_err(|e| MLError::ModelError(format!("DtoD {label}: {e}")))?; Ok(dst) } @@ -1548,16 +1539,8 @@ fn dtod_clone_i32( let mut dst = stream.alloc_zeros::(n_elems) .map_err(|e| MLError::ModelError(format!("alloc {label} i32[{n_elems}]: {e}")))?; let src_view = src.slice(..n_elems); - let (src_ptr, _src_sync) = src_view.device_ptr(stream); - let (dst_ptr, _dst_sync) = dst.device_ptr_mut(stream); - let num_bytes = n_elems * std::mem::size_of::(); - unsafe { - cudarc::driver::result::memcpy_dtod_async( - dst_ptr, src_ptr, num_bytes, stream.cu_stream(), - ).map_err(|e| MLError::ModelError(format!("DtoD {label}: {e}")))?; - } - drop(_dst_sync); - drop(_src_sync); + stream.memcpy_dtod(&src_view, &mut dst) + .map_err(|e| MLError::ModelError(format!("DtoD {label}: {e}")))?; Ok(dst) } @@ -1578,48 +1561,30 @@ fn build_next_states_dtod( let total_elems = total * sd; if timesteps <= 1 { - // Single timestep per episode -- next_states == states return dtod_clone_f32(stream, states, total_elems, "next_states_trivial"); } let mut dst = stream.alloc_zeros::(total_elems) .map_err(|e| MLError::ModelError(format!("alloc next_states f32[{total_elems}]: {e}")))?; - let bytes_per_f32 = std::mem::size_of::(); - let bytes_per_state = sd * bytes_per_f32; - let shifted_count = timesteps - 1; // states[t+1] for t in 0..timesteps-1 - - // Get base pointers once -- then compute offsets manually in the loop. - // This avoids repeated borrows of `dst` and `states` inside the loop. - let (src_base, _src_sync) = states.device_ptr(stream); - let (dst_base, _dst_sync) = dst.device_ptr_mut(stream); + let shifted_count = timesteps - 1; for ep in 0..n_episodes { - let ep_byte_offset = (ep * timesteps * sd * bytes_per_f32) as u64; + let ep_offset = ep * timesteps * sd; // Copy states[ep][1..timesteps] -> next_states[ep][0..timesteps-1] - let src_shifted_ptr = src_base + ep_byte_offset + (sd * bytes_per_f32) as u64; - let dst_shifted_ptr = dst_base + ep_byte_offset; - unsafe { - cudarc::driver::result::memcpy_dtod_async( - dst_shifted_ptr, src_shifted_ptr, - shifted_count * bytes_per_state, stream.cu_stream(), - ).map_err(|e| MLError::ModelError(format!("DtoD next_states shift ep{ep}: {e}")))?; - } + let src_shifted = states.slice((ep_offset + sd)..(ep_offset + timesteps * sd)); + let mut dst_shifted = dst.slice_mut(ep_offset..(ep_offset + shifted_count * sd)); + stream.memcpy_dtod(&src_shifted, &mut dst_shifted) + .map_err(|e| MLError::ModelError(format!("DtoD next_states shift ep{ep}: {e}")))?; - // Duplicate last timestep: states[ep][timesteps-1] -> next_states[ep][timesteps-1] - let last_byte_offset = ep_byte_offset + ((timesteps - 1) * sd * bytes_per_f32) as u64; - let src_last_ptr = src_base + last_byte_offset; - let dst_last_ptr = dst_base + last_byte_offset; - unsafe { - cudarc::driver::result::memcpy_dtod_async( - dst_last_ptr, src_last_ptr, bytes_per_state, stream.cu_stream(), - ).map_err(|e| MLError::ModelError(format!("DtoD next_states last ep{ep}: {e}")))?; - } + // Duplicate last timestep: states[ep][T-1] -> next_states[ep][T-1] + let last_offset = ep_offset + (timesteps - 1) * sd; + let src_last = states.slice(last_offset..(last_offset + sd)); + let mut dst_last = dst.slice_mut(last_offset..(last_offset + sd)); + stream.memcpy_dtod(&src_last, &mut dst_last) + .map_err(|e| MLError::ModelError(format!("DtoD next_states last ep{ep}: {e}")))?; } - // Drop sync guards before moving dst out - drop(_dst_sync); - drop(_src_sync); Ok(dst) } diff --git a/crates/ml/src/cuda_pipeline/gpu_weights.rs b/crates/ml/src/cuda_pipeline/gpu_weights.rs index efbf422c7..0f951479b 100644 --- a/crates/ml/src/cuda_pipeline/gpu_weights.rs +++ b/crates/ml/src/cuda_pipeline/gpu_weights.rs @@ -795,16 +795,12 @@ fn extract_one( .ok_or_else(|| MLError::ModelError(format!("Missing weight: {name}")))?; let n_elems = param.data.len(); - let buf = stream + let mut buf = stream .alloc_zeros::(n_elems) .map_err(|e| MLError::ModelError(format!("Alloc {name}: {e}")))?; - { - let (src_ptr, _src_sync) = param.data.device_ptr(stream); - let (dst_ptr, _dst_sync) = buf.device_ptr(stream); - let num_bytes = n_elems * std::mem::size_of::(); - dtod_copy_checked(dst_ptr, src_ptr, num_bytes, stream, name, "copy")?; - } + stream.memcpy_dtod(¶m.data, &mut buf) + .map_err(|e| MLError::ModelError(format!("DtoD extract {name}: {e}")))?; Ok(buf) } @@ -821,11 +817,8 @@ fn sync_one( .get(name) .ok_or_else(|| MLError::ModelError(format!("Missing weight: {name}")))?; - let n_elems = param.data.len(); - let (src_ptr, _src_sync) = param.data.device_ptr(stream); - let (dst_ptr, _dst_sync) = buf.device_ptr(stream); - let num_bytes = n_elems * std::mem::size_of::(); - dtod_copy_checked(dst_ptr, src_ptr, num_bytes, stream, name, "sync")?; + stream.memcpy_dtod(¶m.data, buf) + .map_err(|e| MLError::ModelError(format!("DtoD sync {name}: {e}")))?; Ok(()) }