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:
jgrusewski
2026-04-05 23:10:33 +02:00
parent e00203b9a0
commit 5f16a86a4a

View File

@@ -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):