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:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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(¶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::<f32>();
|
||||
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(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user