diff --git a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs index d11b95fc9..9b318b933 100644 --- a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs +++ b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs @@ -2299,21 +2299,6 @@ fn fill_episode_ids_gpu( // Pure cudarc DtoD helpers // --------------------------------------------------------------------------- -/// Allocate a fresh `CudaSlice` and DtoD-copy `n_elems` elements from `src`. -fn dtod_clone_bf16( - stream: &Arc, - src: &CudaSlice, - n_elems: usize, - label: &str, -) -> Result, MLError> { - let mut dst = stream.alloc_zeros::(n_elems) - .map_err(|e| MLError::ModelError(format!("alloc {label} bf16[{n_elems}]: {e}")))?; - let src_view = src.slice(..n_elems); - stream.memcpy_dtod(&src_view, &mut dst) - .map_err(|e| MLError::ModelError(format!("DtoD {label}: {e}")))?; - Ok(dst) -} - /// Allocate a fresh `CudaSlice` and DtoD-copy `n_elems` f32 elements from `src`. fn dtod_clone_f32_native( stream: &Arc, @@ -2344,55 +2329,7 @@ fn dtod_clone_i32( Ok(dst) } -/// Build next_states from states via episode-aware DtoD shift. -/// -/// For n-step returns, `shift` is n (not 1): `next_states[t] = states[t+n]`. -/// When `t + shift >= timesteps`, duplicates the last available state. -fn build_next_states_dtod( - stream: &Arc, - states: &CudaSlice, - n_episodes: usize, - timesteps: usize, - sd: usize, - shift: usize, -) -> Result, MLError> { - let total = n_episodes * timesteps; - let total_elems = total * sd; - let shift = shift.max(1).min(timesteps); // clamp shift to valid range - - if timesteps <= shift { - return dtod_clone_bf16(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 shifted_count = timesteps - shift; - - for ep in 0..n_episodes { - let ep_offset = ep * timesteps * sd; - - // Copy states[ep][shift..timesteps] -> next_states[ep][0..timesteps-shift] - let src_shifted = states.slice((ep_offset + shift * 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}")))?; - - // Fill remaining positions with last available state - let last_src_offset = ep_offset + (timesteps - 1) * sd; - let src_last = states.slice(last_src_offset..(last_src_offset + sd)); - for t in shifted_count..timesteps { - let dst_offset = ep_offset + t * sd; - let mut dst_t = dst.slice_mut(dst_offset..(dst_offset + sd)); - stream.memcpy_dtod(&src_last, &mut dst_t) - .map_err(|e| MLError::ModelError(format!("DtoD next_states fill ep{ep} t{t}: {e}")))?; - } - } - - Ok(dst) -} - -/// #30 Build next-states via DtoD shift (f32 version). +/// #30 Build next-states via GPU kernel (f32 version). /// Build next_states via GPU kernel — zero CPU loops, single kernel launch. /// /// For each element (ep, t, d):