feat(phase-e-4-a): wire sliding-window buffer in smoke (buffer-only)
Phase E.4.A Task 7: maintain a GPU-resident circular window buffer in the smoke binary's --temporal path. Per-step: 1. mapped-pinned state_pinned write (existing) 2. alpha_window_push_kernel writes state into window[head_idx] 3. head_idx = (head_idx + 1) % window_k 4. C51 forward proceeds against state_pinned (consumer of window wires in T8 — Mamba2 over the window) On episode reset: zero the buffer and reset head_idx so Mamba2 sees clean zero-context for the first window_k-1 steps. CLI: --temporal flag + --window-k (default 16, kernel max 32 per mamba2_alpha_kernel constraint). Validation: 100-episode smoke with --temporal produced bit-identical R_mean / rvr / kill-criteria values to the C51-flat baseline run — confirms buffer maintenance has zero side effect on the existing C51 path. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -337,6 +337,17 @@ struct Cli {
|
||||
/// half-tick gap to the Phase 1d.4 baseline.
|
||||
#[arg(long, default_value_t = false)]
|
||||
real_spread: bool,
|
||||
/// Phase E.4.A: enable the temporal-encoder stack (sliding window
|
||||
/// buffer + Mamba2 SSM + C51 head). When false, runs the existing
|
||||
/// C51-flat path on the current state vector only. The buffer is
|
||||
/// maintained but unused until Task 8 wires Mamba2 over it.
|
||||
#[arg(long, default_value_t = false)]
|
||||
temporal: bool,
|
||||
/// Sliding-window length K (in snapshots). 16 starting point; the
|
||||
/// kernel max is `MAMBA2_KERNEL_SEQ_MAX = 32`. Sweep 16 / 32 in
|
||||
/// later iterations once Mamba2 forward is wired (T8).
|
||||
#[arg(long, default_value_t = 16)]
|
||||
window_k: usize,
|
||||
/// Output JSON path for final verdict + ISV readings.
|
||||
#[arg(long, default_value = "config/ml/alpha_dqn_h600_smoke.json")]
|
||||
out_path: PathBuf,
|
||||
@@ -588,6 +599,19 @@ fn main() -> Result<()> {
|
||||
let c51_module = ctx
|
||||
.load_cubin(ml::cuda_pipeline::alpha_kernels::ALPHA_C51_CUBIN.to_vec())
|
||||
.context("alpha_c51 cubin load")?;
|
||||
// Phase E.4.A.6: circular-buffer push kernel for sliding window.
|
||||
let push_module = ctx
|
||||
.load_cubin(ml::cuda_pipeline::alpha_kernels::ALPHA_WINDOW_PUSH_CUBIN.to_vec())
|
||||
.context("alpha_window_push cubin load")?;
|
||||
let push_kernel = push_module
|
||||
.load_function("alpha_window_push_kernel")
|
||||
.context("window_push kernel load")?;
|
||||
if cli.temporal {
|
||||
info!(
|
||||
" TEMPORAL: sliding window K={} × state_dim={} = {} floats",
|
||||
cli.window_k, STATE_DIM, cli.window_k * STATE_DIM,
|
||||
);
|
||||
}
|
||||
let c51_fwd_kernel = c51_module
|
||||
.load_function("alpha_c51_forward_kernel")
|
||||
.context("c51 forward load")?;
|
||||
@@ -777,6 +801,14 @@ fn main() -> Result<()> {
|
||||
let state_pinned = unsafe { MappedF32::new(STATE_DIM)? };
|
||||
let action_pinned = unsafe { MappedI32::new()? };
|
||||
|
||||
// Phase E.4.A.7: GPU-resident sliding-window state buffer for the
|
||||
// temporal encoder. Allocated unconditionally; only populated when
|
||||
// cli.temporal is set. Capacity: window_k × STATE_DIM floats.
|
||||
let mut window_dev = stream
|
||||
.alloc_zeros::<f32>(cli.window_k * STATE_DIM)
|
||||
.context("alloc window buffer")?;
|
||||
let mut head_idx: i32 = 0;
|
||||
|
||||
// Kill-criteria inputs: action_counts (i32 × n_actions),
|
||||
// scalar_inputs (f32 × 3: rollout_R_mean, q_init_norm, q_early_norm).
|
||||
let mut kc_action_counts_dev = stream.alloc_zeros::<i32>(N_ACTIONS).context("alloc kc actions")?;
|
||||
@@ -804,6 +836,15 @@ fn main() -> Result<()> {
|
||||
env.reset_at(env_seed, start_cursor);
|
||||
let mut state = EpisodeState::new();
|
||||
|
||||
// Phase E.4.A.7: clear the temporal window on episode start so
|
||||
// Mamba2 sees zero-history for the first window_k-1 steps.
|
||||
if cli.temporal {
|
||||
stream
|
||||
.memset_zeros(&mut window_dev)
|
||||
.context("zero window on episode reset")?;
|
||||
head_idx = 0;
|
||||
}
|
||||
|
||||
let mut states_host: Vec<f32> = Vec::with_capacity(cli.horizon * STATE_DIM);
|
||||
let mut next_states_host: Vec<f32> = Vec::with_capacity(cli.horizon * STATE_DIM);
|
||||
let mut actions_host: Vec<i32> = Vec::with_capacity(cli.horizon);
|
||||
@@ -818,6 +859,20 @@ fn main() -> Result<()> {
|
||||
let s_vec = state_as_vec(&env, &state);
|
||||
let action: u8 = if cli.c51 {
|
||||
state_pinned.write(&s_vec);
|
||||
// Phase E.4.A.7: push current state into the circular
|
||||
// window buffer. The buffer is populated for ALL
|
||||
// --temporal runs; consumer (Mamba2) wires in T8.
|
||||
if cli.temporal {
|
||||
let (w_ptr, _g1) = window_dev.device_ptr_mut(&stream);
|
||||
unsafe {
|
||||
ml::cuda_pipeline::alpha_kernels::launch_alpha_window_push(
|
||||
&stream, &push_kernel,
|
||||
state_pinned.dev_u64(), w_ptr,
|
||||
head_idx, state_dim_i,
|
||||
)?;
|
||||
}
|
||||
head_idx = (head_idx + 1) % (cli.window_k as i32);
|
||||
}
|
||||
{
|
||||
let (w_ptr, _g0) = w_dev.device_ptr(&stream);
|
||||
let (b_ptr, _g1) = b_dev.device_ptr(&stream);
|
||||
|
||||
Reference in New Issue
Block a user