fix: replace raw memcpy_dtod_async with safe stream.memcpy_dtod

Fixes cudarc 0.19 event tracking corruption caused by mixing safe
device_ptr/device_ptr_mut wrappers with raw memcpy_dtod_async.

The anti-pattern: manually calling device_ptr() to get raw pointers,
then using low-level memcpy_dtod_async, then dropping SyncOnDrop
guards — this bypasses cudarc's event management and corrupts the
synchronization state of CudaSlice objects.

Fixed in 7 call sites across 3 files:
- gpu_experience_collector.rs: dtod_clone_f32, dtod_clone_i32,
  build_next_states_dtod (replaced pointer arithmetic with slice views)
- gpu_weights.rs: extract_one, sync_one
- training_loop.rs: cuda_slice_to_tensor_f32

Investigation ongoing: deadlock persists in PER insert path
(cuda_slice_to_tensor_f32 -> stream.synchronize()). The memcpy fix
is correct but there's an additional issue in the CudaSlice->GpuTensor
conversion that needs further debugging.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-19 23:24:42 +01:00
parent 9f5b88c81d
commit 2856daa8e9
2 changed files with 21 additions and 63 deletions

View File

@@ -1524,17 +1524,8 @@ fn dtod_clone_f32(
let mut dst = stream.alloc_zeros::<f32>(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::<f32>();
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::<i32>(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::<i32>();
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::<f32>(total_elems)
.map_err(|e| MLError::ModelError(format!("alloc next_states f32[{total_elems}]: {e}")))?;
let bytes_per_f32 = std::mem::size_of::<f32>();
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)
}

View File

@@ -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::<f32>(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::<f32>();
dtod_copy_checked(dst_ptr, src_ptr, num_bytes, stream, name, "copy")?;
}
stream.memcpy_dtod(&param.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::<f32>();
dtod_copy_checked(dst_ptr, src_ptr, num_bytes, stream, name, "sync")?;
stream.memcpy_dtod(&param.data, buf)
.map_err(|e| MLError::ModelError(format!("DtoD sync {name}: {e}")))?;
Ok(())
}