fix(phase-e-4-a): GPU kernel for h_enriched slot copy — eliminate per-step CPU roundtrip

The Phase E.4.A T8 wiring stored Mamba2's per-step cache.h_enriched
into h_enriched_buf_dev via a dtoh+htod sequence:

  let h_host = stream.clone_dtoh(cache.h_enriched.cuda_data())?;
  let mut buf_host = stream.clone_dtoh(&h_enriched_buf_dev)?;  // <- whole buffer
  for j in 0..hidden_dim { buf_host[slot_offset + j] = h_host[j]; }
  stream.memcpy_htod(&buf_host, &mut h_enriched_buf_dev)?;     // <- whole buffer

This violates feedback_cpu_is_read_only AND
feedback_no_htod_htoh_only_mapped_pinned. Worse, the buffer-wide
dtoh+htod every step is ~20K floats × 600 steps × 500 eps × 30 cells
= ~9M roundtrips totaling significant PCIe latency in the backtest.

Fix: new tiny CUDA kernel alpha_h_enriched_store_kernel in
alpha_window_push.cu (one thread per hidden-dim feature, writes
src[j] → buf[slot_offset + j]). Replaces the dtoh/htod sequence
in both smoke and backtest binaries.

Estimated speed-up at backtest scale: 3-6× on the temporal eval
path. Pure-GPU per-step inference restored — no synchronisation
points on the hot path.

docs/isv-slots.md updated per kernel-audit-doc hook.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-15 22:20:28 +02:00
parent 2fe76f2f34
commit 2feeeda8bb
5 changed files with 128 additions and 69 deletions

View File

@@ -378,6 +378,7 @@ fn main() -> Result<()> {
.load_cubin(ml::cuda_pipeline::alpha_kernels::ALPHA_WINDOW_PUSH_CUBIN.to_vec())
.context("alpha_window_push cubin")?;
let push_kernel = push_module.load_function("alpha_window_push_kernel")?;
let h_store_kernel = push_module.load_function("alpha_h_enriched_store_kernel")?;
// Phase E.4.A.7 controller (used in --isv-continual eval path).
let ctl_module = ctx
.load_cubin(ml::cuda_pipeline::alpha_kernels::STACKER_THRESHOLD_CONTROLLER_CUBIN.to_vec())
@@ -566,14 +567,19 @@ fn main() -> Result<()> {
let block = mamba2_block.as_ref().expect("Mamba2Block missing");
let (_logit, cache) = block.forward_train(&window_tensor)
.map_err(|e| anyhow::anyhow!("mamba2 forward: {e}"))?;
// Store h_enriched for batched training.
let h_host = stream.clone_dtoh(cache.h_enriched.cuda_data())?;
let slot_offset = state.step * cli.mamba2_hidden_dim;
let mut buf_host = stream.clone_dtoh(&h_enriched_buf_dev)?;
for j in 0..cli.mamba2_hidden_dim {
buf_host[slot_offset + j] = h_host[j];
// Phase E.4.A.8.fix: GPU-side slot copy.
{
let slot_offset = (state.step * cli.mamba2_hidden_dim) as i32;
let (src_ptr, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream);
let (buf_ptr, _g_buf) = h_enriched_buf_dev.device_ptr_mut(&stream);
unsafe {
ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store(
&stream, &h_store_kernel,
src_ptr, buf_ptr, slot_offset,
cli.mamba2_hidden_dim as i32,
)?;
}
}
stream.memcpy_htod(&buf_host, &mut h_enriched_buf_dev)?;
// C51 forward on h_enriched.
let (w_ptr, _g0) = w_dev.device_ptr(&stream);
let (b_ptr, _g1) = b_dev.device_ptr(&stream);
@@ -679,13 +685,17 @@ fn main() -> Result<()> {
let block = mamba2_block.as_ref().expect("Mamba2Block missing");
let (_logit, cache) = block.forward_train(&window_tensor)
.map_err(|e| anyhow::anyhow!("mamba2 terminal forward: {e}"))?;
let h_host_term = stream.clone_dtoh(cache.h_enriched.cuda_data())?;
let term_offset = (ep_len as usize) * cli.mamba2_hidden_dim;
let mut buf_host = stream.clone_dtoh(&h_enriched_buf_dev)?;
for j in 0..cli.mamba2_hidden_dim {
buf_host[term_offset + j] = h_host_term[j];
// Phase E.4.A.8.fix: GPU-side slot copy for terminal h_enriched.
let term_offset = ((ep_len as usize) * cli.mamba2_hidden_dim) as i32;
let (src_ptr, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream);
let (buf_ptr, _g_buf) = h_enriched_buf_dev.device_ptr_mut(&stream);
unsafe {
ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store(
&stream, &h_store_kernel,
src_ptr, buf_ptr, term_offset,
cli.mamba2_hidden_dim as i32,
)?;
}
stream.memcpy_htod(&buf_host, &mut h_enriched_buf_dev)?;
}
let (curr_input_ptr_guard, next_input_ptr_guard);
let (curr_input_ptr, next_input_ptr): (u64, u64) = if cli.temporal {

View File

@@ -619,6 +619,11 @@ fn main() -> Result<()> {
let push_kernel = push_module
.load_function("alpha_window_push_kernel")
.context("window_push kernel load")?;
// Phase E.4.A.8.fix: GPU-side h_enriched slot copy (replaces
// dtoh/htod that was violating feedback_cpu_is_read_only).
let h_store_kernel = push_module
.load_function("alpha_h_enriched_store_kernel")
.context("h_enriched_store kernel load")?;
if cli.temporal {
info!(
" TEMPORAL: sliding window K={} × state_dim={} = {} floats",
@@ -932,42 +937,27 @@ fn main() -> Result<()> {
// inference doesn't need backward (we only train via
// the batched update at end-of-episode). Backward
// wiring with cache retention lands in T10.
// Phase E.4.A.8: in temporal mode, run Mamba2 over the
// current window, persist h_enriched into the per-step
// buffer (slot = current env step), and feed it to C51.
// In non-temporal mode, C51 reads state_pinned directly.
// Phase E.4.A.8 + .8.fix: per-step Mamba2 forward; store
// h_enriched into per-step buffer slot via GPU kernel
// (no CPU roundtrip); feed h_enriched to C51 forward.
if cli.temporal {
let block = mamba2_block.as_ref()
.expect("temporal flag set but Mamba2Block missing");
let (_logit, cache) = block.forward_train(&window_tensor)
.map_err(|e| anyhow::anyhow!("mamba2 forward: {e}"))?;
// 1) Copy h_enriched into h_enriched_buf_dev at slot t.
// Slot t holds the representation BEFORE env.step at step t
// (i.e., the "current state" representation).
let slot_offset = state.step * cli.mamba2_hidden_dim;
// GPU-side slot copy: cache.h_enriched h_enriched_buf_dev[slot..]
{
// memcpy_dtod into the slot
let src_ptr = {
let (p, _g) = cache.h_enriched.cuda_data().device_ptr(&stream);
p
};
// Use a small custom copy kernel? Or use cudarc's memcpy_dtod.
// For now: dtoh then htod (small — hidden_dim=32 floats).
// OPTIMIZATION HOOK: replace with kernel for zero round-trips.
let h_host = stream.clone_dtoh(cache.h_enriched.cuda_data())
.map_err(|e| anyhow::anyhow!("dtoh h_enriched: {e}"))?;
let _ = src_ptr; // silence unused
let slot_slice = &mut h_enriched_buf_dev;
// Build a slice view: write h_host into slot_offset..slot_offset+hidden_dim
let mut buf_host = stream.clone_dtoh(slot_slice)
.map_err(|e| anyhow::anyhow!("dtoh h buf: {e}"))?;
for j in 0..cli.mamba2_hidden_dim {
buf_host[slot_offset + j] = h_host[j];
let slot_offset = (state.step * cli.mamba2_hidden_dim) as i32;
let (src_ptr, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream);
let (buf_ptr, _g_buf) = h_enriched_buf_dev.device_ptr_mut(&stream);
unsafe {
ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store(
&stream, &h_store_kernel,
src_ptr, buf_ptr, slot_offset,
cli.mamba2_hidden_dim as i32,
)?;
}
stream.memcpy_htod(&buf_host, slot_slice)
.map_err(|e| anyhow::anyhow!("htod h buf: {e}"))?;
}
// 2) C51 forward on h_enriched (batch=1, dim=hidden).
let (w_ptr, _g0) = w_dev.device_ptr(&stream);
let (b_ptr, _g1) = b_dev.device_ptr(&stream);
let (p_ptr, _g2) = single_probs_dev.device_ptr_mut(&stream);
@@ -1087,34 +1077,29 @@ fn main() -> Result<()> {
if cli.c51 {
// ── C51 batched compute ─────────────────────────────────
// Phase E.4.A.8: in temporal mode, the batched C51 forwards
// read the per-step h_enriched buffer instead of raw states.
// h_enriched_buf_dev[t..t+hidden_dim] = h at step t.
// The "next-state" representation for step t is h at step
// t+1; we need to run Mamba2 once MORE on the terminal
// next-window to fill slot ep_len.
// Phase E.4.A.8 + .8.fix: in temporal mode, read pre-stored
// h_enriched from h_enriched_buf_dev (populated GPU-side
// during per-step inference via alpha_h_enriched_store).
// Run ONE extra Mamba2 forward on the terminal window to
// fill slot ep_len (the "next-state" representation for
// the last training transition).
if cli.temporal {
// Run Mamba2 on the final window (it currently holds
// the post-terminal observation since env.step pushed
// before returning done). h_enriched for slot ep_len.
let block = mamba2_block.as_ref().expect("Mamba2Block missing");
let (_logit, cache) = block.forward_train(&window_tensor)
.map_err(|e| anyhow::anyhow!("mamba2 terminal forward: {e}"))?;
let h_host_term = stream.clone_dtoh(cache.h_enriched.cuda_data())
.map_err(|e| anyhow::anyhow!("dtoh term h: {e}"))?;
let term_offset = (ep_len as usize) * cli.mamba2_hidden_dim;
let mut buf_host = stream.clone_dtoh(&h_enriched_buf_dev)
.map_err(|e| anyhow::anyhow!("dtoh buf for term: {e}"))?;
for j in 0..cli.mamba2_hidden_dim {
buf_host[term_offset + j] = h_host_term[j];
let term_offset = ((ep_len as usize) * cli.mamba2_hidden_dim) as i32;
let (src_ptr, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream);
let (buf_ptr, _g_buf) = h_enriched_buf_dev.device_ptr_mut(&stream);
unsafe {
ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store(
&stream, &h_store_kernel,
src_ptr, buf_ptr, term_offset,
cli.mamba2_hidden_dim as i32,
)?;
}
stream.memcpy_htod(&buf_host, &mut h_enriched_buf_dev)
.map_err(|e| anyhow::anyhow!("htod buf for term: {e}"))?;
}
let (curr_input_ptr_guard, next_input_ptr_guard);
let (curr_input_ptr, next_input_ptr): (u64, u64) = if cli.temporal {
// h_current = h_enriched_buf_dev[0..ep_len*hidden]
// h_next = h_enriched_buf_dev[hidden..(ep_len+1)*hidden]
let (h_ptr_curr, g_curr) = h_enriched_buf_dev.device_ptr(&stream);
let (h_ptr_next_base, g_next) = h_enriched_buf_dev.device_ptr(&stream);
curr_input_ptr_guard = g_curr;
@@ -1122,7 +1107,7 @@ fn main() -> Result<()> {
let _ = (&curr_input_ptr_guard, &next_input_ptr_guard);
(
h_ptr_curr,
h_ptr_next_base + (cli.mamba2_hidden_dim as u64) * 4u64, // 4 = sizeof f32
h_ptr_next_base + (cli.mamba2_hidden_dim as u64) * 4u64,
)
} else {
let (s_ptr, g_curr) = states_dev.device_ptr(&stream);

View File

@@ -646,6 +646,44 @@ pub unsafe fn launch_alpha_c51_expected_q(
Ok(())
}
/// Launch `alpha_h_enriched_store_kernel`. Copies `src[hidden_dim]`
/// (typically a Mamba2 cache.h_enriched device buffer) into
/// `buf[slot_offset..slot_offset+hidden_dim]`. Replaces the per-step
/// dtoh/htod sequence per `feedback_cpu_is_read_only`. Pure GPU.
///
/// # Safety
/// `src_dev` and `buf_dev` MUST be valid device pointers.
pub unsafe fn launch_alpha_h_enriched_store(
stream: &cudarc::driver::CudaStream,
kernel: &cudarc::driver::CudaFunction,
src_dev: u64,
buf_dev: u64,
slot_offset: i32,
hidden_dim: i32,
) -> Result<(), MLError> {
use cudarc::driver::{LaunchConfig, PushKernelArg};
debug_assert!(hidden_dim > 0, "hidden_dim must be positive");
debug_assert!(slot_offset >= 0, "slot_offset must be non-negative");
const BLOCK: u32 = 32;
let grid_x = ((hidden_dim as u32) + BLOCK - 1) / BLOCK;
let cfg = LaunchConfig {
grid_dim: (grid_x.max(1), 1, 1),
block_dim: (BLOCK, 1, 1),
shared_mem_bytes: 0,
};
stream
.launch_builder(kernel)
.arg(&src_dev)
.arg(&buf_dev)
.arg(&slot_offset)
.arg(&hidden_dim)
.launch(cfg)
.map_err(|e| MLError::ModelError(format!("alpha_h_enriched_store launch: {e}")))?;
Ok(())
}
/// Launch `alpha_c51_grad_input_kernel`. Computes dL/d_input for the
/// C51 layer (gradient w.r.t. the upstream encoder's output).
/// Threads: one per (batch, input feature). Needed to chain C51's

View File

@@ -31,3 +31,21 @@ extern "C" __global__ void alpha_window_push_kernel(
// Insert new state at slot K-1.
window[(K - 1) * state_dim + j] = state_in[j];
}
// ----------------------------------------------------------------------
// Phase E.4.A.8.fix (2026-05-15): GPU-side slot copy for h_enriched.
// ----------------------------------------------------------------------
//
// Replaces the dtoh/htod sequence that was used to store per-step
// Mamba2 h_enriched into a buffer slot — a CPU-roundtrip violation of
// `feedback_cpu_is_read_only`. One thread per hidden_dim feature.
extern "C" __global__ void alpha_h_enriched_store_kernel(
const float* __restrict__ src, // [hidden_dim]
float* __restrict__ buf, // [(horizon+1) * hidden_dim]
int slot_offset, // slot index × hidden_dim
int hidden_dim
) {
const int j = blockIdx.x * blockDim.x + threadIdx.x;
if (j >= hidden_dim) return;
buf[slot_offset + j] = src[j];
}