chore: remove dead build_next_states_dtod + dtod_clone_bf16
Both were unused (bf16 path replaced by f32 path in #30). build_next_states_dtod had the same CPU loop anti-pattern (795K memcpy calls at 4096 episodes) that caused the H100 hang. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -2299,21 +2299,6 @@ fn fill_episode_ids_gpu(
|
||||
// Pure cudarc DtoD helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Allocate a fresh `CudaSlice<half::bf16>` and DtoD-copy `n_elems` elements from `src`.
|
||||
fn dtod_clone_bf16(
|
||||
stream: &Arc<CudaStream>,
|
||||
src: &CudaSlice<half::bf16>,
|
||||
n_elems: usize,
|
||||
label: &str,
|
||||
) -> Result<CudaSlice<half::bf16>, MLError> {
|
||||
let mut dst = stream.alloc_zeros::<half::bf16>(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<f32>` and DtoD-copy `n_elems` f32 elements from `src`.
|
||||
fn dtod_clone_f32_native(
|
||||
stream: &Arc<CudaStream>,
|
||||
@@ -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<CudaStream>,
|
||||
states: &CudaSlice<half::bf16>,
|
||||
n_episodes: usize,
|
||||
timesteps: usize,
|
||||
sd: usize,
|
||||
shift: usize,
|
||||
) -> Result<CudaSlice<half::bf16>, 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::<half::bf16>(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):
|
||||
|
||||
Reference in New Issue
Block a user