perf(ml-alpha): fused batched snap_feature_assemble kernel (#4)

Previously the per-step snap_feature path did B*K = 768 single-snapshot
kernel launches (at B=8, K=96) + 768 DtoD copies into the window
tensor. New `snap_feature_assemble_batched` processes all B*K
snapshots in a SINGLE launch and writes outputs directly into the
window tensor's storage.

Per-step CPU work: pack 12 mapped-pinned staging buffers (~150 KB
total host writes), then 10 DtoD copies of the staging → device
buffers. Per-step GPU work: 1 batched kernel launch with B*K
threads (each writes 32 floats to its output row).

Mapped-pinned staging buffers cover the full B*K capacity at trainer
init — no per-step allocation. New `MappedI32Buffer` and
`MappedI64Buffer` types parallel `MappedF32Buffer` to stage
`trade_count` (i32) and `ts_ns` / `prev_ts_ns` (i64) without
violating the no-htod rule (`feedback_no_htod_htoh_only_mapped_pinned.md`).

Dead per-snapshot scratch + helpers (`bid_px_d`, `snap_feat_d`,
`stg_bid_px`, `snap_fn`, `upload_into`, etc.) removed per
`feedback_no_legacy_aliases.md` — the only callers were the
per-snapshot path, gone.

Expected per-step savings: ~5-10 ms launch + DtoD overhead at
B=8, K=96. Over 2000 steps/epoch = 10-20 sec/epoch.

77 ml-alpha tests pass. Synthetic overfit unchanged.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 13:12:32 +02:00
parent ebae67cb6b
commit c70c5cdf21
3 changed files with 445 additions and 130 deletions

View File

@@ -72,3 +72,67 @@ extern "C" __global__ void snap_feature_assemble(
for (int i = 26; i < 32; ++i) out[i] = 0.0f;
}
// ─── Batched variant ──────────────────────────────────────────────────
//
// Process N snapshots in a single kernel launch — one thread per
// snapshot. Replaces N separate single-snapshot launches with one
// fused launch. Used by PerceptionTrainer to compute snap_features
// for all B*K snapshots per training step in a single kernel call
// (eliminates ~768 launches per step at B=8, K=96).
//
// Layout: all inputs are [N, ...] row-major. Output is [N, 32].
// Per-snapshot semantics are IDENTICAL to the single-snapshot kernel
// above — bit-equivalent for any one snapshot in isolation.
extern "C" __global__ void snap_feature_assemble_batched(
const float* __restrict__ bid_px, // [N, 10]
const float* __restrict__ bid_sz, // [N, 10]
const float* __restrict__ ask_px, // [N, 10]
const float* __restrict__ ask_sz, // [N, 10]
const float* __restrict__ prev_bid_sz, // [N, 10]
const float* __restrict__ prev_ask_sz, // [N, 10]
const float* __restrict__ regime, // [N, 6]
const float* __restrict__ prev_mid, // [N]
const float* __restrict__ trade_signed_vol, // [N]
const int* __restrict__ trade_count, // [N]
const long* __restrict__ ts_ns, // [N]
const long* __restrict__ prev_ts_ns, // [N]
float tick_size,
int N,
float* __restrict__ out // [N, 32]
) {
int n = blockIdx.x * blockDim.x + threadIdx.x;
if (n >= N) return;
const float* bx = bid_px + n * 10;
const float* bs = bid_sz + n * 10;
const float* ax = ask_px + n * 10;
const float* as_ = ask_sz + n * 10;
const float* pbs = prev_bid_sz + n * 10;
const float* pas = prev_ask_sz + n * 10;
const float* rg = regime + n * 6;
float* o = out + n * 32;
const float mid = 0.5f * (bx[0] + ax[0]);
const float pm = prev_mid[n];
o[0] = (mid > 0.0f && pm > 0.0f) ? (mid - pm) / tick_size : 0.0f;
o[1] = (ax[0] - bx[0]) / tick_size;
for (int i = 0; i < 5; ++i) {
o[2 + i] = log1pf(bs[i]);
o[7 + i] = log1pf(as_[i]);
const float bid_delta = bs[i] - pbs[i];
const float ask_delta = as_[i] - pas[i];
o[12 + i] = signed_log1p(bid_delta - ask_delta);
}
o[17] = log1pf((float) trade_count[n]);
o[18] = signed_log1p(trade_signed_vol[n]);
const float dt_ms = (float)(ts_ns[n] - prev_ts_ns[n]) * 1e-6f;
o[19] = log1pf(fmaxf(dt_ms, 0.0f));
for (int i = 0; i < 6; ++i) o[20 + i] = rg[i];
for (int i = 26; i < 32; ++i) o[i] = 0.0f;
}

View File

@@ -108,3 +108,88 @@ impl Drop for MappedF32Buffer {
}
}
}
/// Mapped-pinned `i32` buffer. Same allocation pattern as
/// [`MappedF32Buffer`] — the only permitted CPU↔GPU path per
/// `feedback_no_htod_htoh_only_mapped_pinned.md`. Needed by the
/// batched snap_feature path to stage per-snapshot `trade_count`
/// values for kernel consumption.
#[allow(missing_debug_implementations)]
pub struct MappedI32Buffer {
pub host_ptr: *mut i32,
pub dev_ptr: cudarc::driver::sys::CUdeviceptr,
pub len: usize,
}
unsafe impl Send for MappedI32Buffer {}
unsafe impl Sync for MappedI32Buffer {}
impl MappedI32Buffer {
pub unsafe fn new(len: usize) -> Result<Self, String> {
let num_bytes = len * std::mem::size_of::<i32>();
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP
| cudarc::driver::sys::CU_MEMHOSTALLOC_PORTABLE;
let host_ptr = cudarc::driver::result::malloc_host(num_bytes, flags)
.map_err(|e| format!("MappedI32Buffer alloc ({len} i32): {e}"))? as *mut i32;
std::ptr::write_bytes(host_ptr, 0, len);
let mut dev_ptr_raw = MaybeUninit::uninit();
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
dev_ptr_raw.as_mut_ptr(),
host_ptr as *mut c_void,
0,
)
.result()
.map_err(|e| format!("cuMemHostGetDevicePointer (i32 buf): {e}"))?;
Ok(Self { host_ptr, dev_ptr: dev_ptr_raw.assume_init(), len })
}
pub fn host_slice_mut(&mut self) -> &mut [i32] {
unsafe { std::slice::from_raw_parts_mut(self.host_ptr, self.len) }
}
}
impl Drop for MappedI32Buffer {
fn drop(&mut self) {
unsafe { let _ = cudarc::driver::result::free_host(self.host_ptr as *mut c_void); }
}
}
/// Mapped-pinned `i64` buffer. Needed by the batched snap_feature
/// path to stage per-snapshot `ts_ns` / `prev_ts_ns` timestamps.
#[allow(missing_debug_implementations)]
pub struct MappedI64Buffer {
pub host_ptr: *mut i64,
pub dev_ptr: cudarc::driver::sys::CUdeviceptr,
pub len: usize,
}
unsafe impl Send for MappedI64Buffer {}
unsafe impl Sync for MappedI64Buffer {}
impl MappedI64Buffer {
pub unsafe fn new(len: usize) -> Result<Self, String> {
let num_bytes = len * std::mem::size_of::<i64>();
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP
| cudarc::driver::sys::CU_MEMHOSTALLOC_PORTABLE;
let host_ptr = cudarc::driver::result::malloc_host(num_bytes, flags)
.map_err(|e| format!("MappedI64Buffer alloc ({len} i64): {e}"))? as *mut i64;
std::ptr::write_bytes(host_ptr, 0, len);
let mut dev_ptr_raw = MaybeUninit::uninit();
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
dev_ptr_raw.as_mut_ptr(),
host_ptr as *mut c_void,
0,
)
.result()
.map_err(|e| format!("cuMemHostGetDevicePointer (i64 buf): {e}"))?;
Ok(Self { host_ptr, dev_ptr: dev_ptr_raw.assume_init(), len })
}
pub fn host_slice_mut(&mut self) -> &mut [i64] {
unsafe { std::slice::from_raw_parts_mut(self.host_ptr, self.len) }
}
}
impl Drop for MappedI64Buffer {
fn drop(&mut self) {
unsafe { let _ = cudarc::driver::result::free_host(self.host_ptr as *mut c_void); }
}
}

View File

@@ -49,7 +49,7 @@ use crate::heads::{HIDDEN_DIM, N_HORIZONS};
use crate::mamba2_block::{
Mamba2AdamW, Mamba2AdamWConfig, Mamba2BackwardScratch, Mamba2Block, Mamba2BlockConfig,
};
use crate::pinned_mem::MappedF32Buffer;
use crate::pinned_mem::{MappedF32Buffer, MappedI32Buffer, MappedI64Buffer};
use crate::trainer::optim::AdamW;
const SNAP_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/snap_feature_assemble.cubin"));
@@ -113,7 +113,9 @@ pub struct PerceptionTrainer {
_step_module: Arc<CudaModule>,
_heads_module: Arc<CudaModule>,
_bce_module: Arc<CudaModule>,
snap_fn: CudaFunction,
/// Batched snap_feature kernel — processes all B*K snapshots per
/// training step in one launch.
snap_batched_fn: CudaFunction,
bce_fn: CudaFunction,
/// Batched cfc/heads forward + backward kernels — used by
/// `step_batched()` for all training updates regardless of B.
@@ -146,15 +148,39 @@ pub struct PerceptionTrainer {
pub opt_heads_w: AdamW,
pub opt_heads_b: AdamW,
// Per-step scratch for snap_feature_assemble (single-snapshot uploads).
bid_px_d: CudaSlice<f32>,
bid_sz_d: CudaSlice<f32>,
ask_px_d: CudaSlice<f32>,
ask_sz_d: CudaSlice<f32>,
prev_bid_sz_d: CudaSlice<f32>,
prev_ask_sz_d: CudaSlice<f32>,
regime_d: CudaSlice<f32>,
snap_feat_d: CudaSlice<f32>,
// Pre-allocated BATCHED snap_feature staging — sized for max B*K
// snapshots. All staging buffers are MAPPED-PINNED (the only
// permitted CPU↔GPU path per `feedback_no_htod_htoh_only_mapped_pinned.md`).
// Per training step: host writes through staging.host_slice_mut(),
// then one DtoD copy per array → device buffer, then one batched
// snap_feature_assemble launch instead of B*K single-snapshot launches.
bk_capacity: usize,
stg_bid_px_all: MappedF32Buffer, // [B*K, 10]
stg_bid_sz_all: MappedF32Buffer, // [B*K, 10]
stg_ask_px_all: MappedF32Buffer, // [B*K, 10]
stg_ask_sz_all: MappedF32Buffer, // [B*K, 10]
stg_regime_all: MappedF32Buffer, // [B*K, 6]
stg_prev_mid: MappedF32Buffer, // [B*K]
stg_trade_signed_vol: MappedF32Buffer, // [B*K]
stg_trade_count: MappedI32Buffer, // [B*K]
stg_ts_ns: MappedI64Buffer, // [B*K]
stg_prev_ts_ns: MappedI64Buffer, // [B*K]
bid_px_all_d: CudaSlice<f32>,
bid_sz_all_d: CudaSlice<f32>,
ask_px_all_d: CudaSlice<f32>,
ask_sz_all_d: CudaSlice<f32>,
/// `prev_bid_sz_all_d` / `prev_ask_sz_all_d` — zero-init device
/// buffers passed to the snap kernel. The loader doesn't compute
/// per-snapshot prev-OFI yet, so these stay zero (same as the
/// single-snapshot path).
prev_bid_sz_all_d: CudaSlice<f32>,
prev_ask_sz_all_d: CudaSlice<f32>,
regime_all_d: CudaSlice<f32>,
prev_mid_all_d: CudaSlice<f32>,
trade_signed_vol_all_d: CudaSlice<f32>,
trade_count_all_d: CudaSlice<i32>,
ts_ns_all_d: CudaSlice<i64>,
prev_ts_ns_all_d: CudaSlice<i64>,
// Per-K device-resident scratch — pre-allocated once at construction
// so the K loop in step() does ZERO device allocations.
@@ -198,12 +224,9 @@ pub struct PerceptionTrainer {
grad_heads_w_d: CudaSlice<f32>,
grad_heads_b_d: CudaSlice<f32>,
// Mapped-pinned staging for raw input uploads.
stg_bid_px: MappedF32Buffer,
stg_bid_sz: MappedF32Buffer,
stg_ask_px: MappedF32Buffer,
stg_ask_sz: MappedF32Buffer,
stg_regime: MappedF32Buffer,
// Mapped-pinned staging for label uploads (the only remaining
// single-call host→device path; staging covers all K*B labels per
// training step in one DtoD copy).
stg_labels: MappedF32Buffer,
}
@@ -217,7 +240,7 @@ impl PerceptionTrainer {
let step_module = ctx.load_cubin(STEP_CUBIN.to_vec()).context("step cubin")?;
let heads_module = ctx.load_cubin(HEADS_CUBIN.to_vec()).context("heads cubin")?;
let bce_module = ctx.load_cubin(BCE_CUBIN.to_vec()).context("bce cubin")?;
let snap_fn = snap_module.load_function("snap_feature_assemble")?;
let snap_batched_fn = snap_module.load_function("snap_feature_assemble_batched")?;
let bce_fn = bce_module.load_function("bce_multi_horizon_forward_backward")?;
let step_batched_fn = step_module.load_function("cfc_step_batched")?;
let step_bwd_batched_fn = step_module.load_function("cfc_step_backward_batched")?;
@@ -290,14 +313,6 @@ impl PerceptionTrainer {
let k = cfg.seq_len;
Ok(Self {
cfg: cfg.clone(),
bid_px_d: stream.alloc_zeros::<f32>(10)?,
bid_sz_d: stream.alloc_zeros::<f32>(10)?,
ask_px_d: stream.alloc_zeros::<f32>(10)?,
ask_sz_d: stream.alloc_zeros::<f32>(10)?,
prev_bid_sz_d: stream.alloc_zeros::<f32>(10)?,
prev_ask_sz_d: stream.alloc_zeros::<f32>(10)?,
regime_d: stream.alloc_zeros::<f32>(REGIME_DIM)?,
snap_feat_d: stream.alloc_zeros::<f32>(FEATURE_DIM)?,
h_new_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * n_hid)?,
probs_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS)?,
labels_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS)?,
@@ -314,12 +329,43 @@ impl PerceptionTrainer {
grad_tau_d: stream.alloc_zeros::<f32>(n_hid)?,
grad_heads_w_d: stream.alloc_zeros::<f32>(N_HORIZONS * n_hid)?,
grad_heads_b_d: stream.alloc_zeros::<f32>(N_HORIZONS)?,
stg_bid_px: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg: {e}"))?,
stg_bid_sz: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg: {e}"))?,
stg_ask_px: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg: {e}"))?,
stg_ask_sz: unsafe { MappedF32Buffer::new(10) }.map_err(|e| anyhow::anyhow!("stg: {e}"))?,
stg_regime: unsafe { MappedF32Buffer::new(REGIME_DIM) }.map_err(|e| anyhow::anyhow!("stg_regime: {e}"))?,
stg_labels: unsafe { MappedF32Buffer::new(k * cfg.n_batch * N_HORIZONS) }.map_err(|e| anyhow::anyhow!("stg_labels: {e}"))?,
// Batched snap_feature staging — sized for max B*K snapshots.
bk_capacity: cfg.n_batch * k,
stg_bid_px_all: unsafe { MappedF32Buffer::new(cfg.n_batch * k * 10) }
.map_err(|e| anyhow::anyhow!("stg_bid_px_all: {e}"))?,
stg_bid_sz_all: unsafe { MappedF32Buffer::new(cfg.n_batch * k * 10) }
.map_err(|e| anyhow::anyhow!("stg_bid_sz_all: {e}"))?,
stg_ask_px_all: unsafe { MappedF32Buffer::new(cfg.n_batch * k * 10) }
.map_err(|e| anyhow::anyhow!("stg_ask_px_all: {e}"))?,
stg_ask_sz_all: unsafe { MappedF32Buffer::new(cfg.n_batch * k * 10) }
.map_err(|e| anyhow::anyhow!("stg_ask_sz_all: {e}"))?,
stg_regime_all: unsafe { MappedF32Buffer::new(cfg.n_batch * k * REGIME_DIM) }
.map_err(|e| anyhow::anyhow!("stg_regime_all: {e}"))?,
stg_prev_mid: unsafe { MappedF32Buffer::new(cfg.n_batch * k) }
.map_err(|e| anyhow::anyhow!("stg_prev_mid: {e}"))?,
stg_trade_signed_vol: unsafe { MappedF32Buffer::new(cfg.n_batch * k) }
.map_err(|e| anyhow::anyhow!("stg_trade_signed_vol: {e}"))?,
stg_trade_count: unsafe { MappedI32Buffer::new(cfg.n_batch * k) }
.map_err(|e| anyhow::anyhow!("stg_trade_count: {e}"))?,
stg_ts_ns: unsafe { MappedI64Buffer::new(cfg.n_batch * k) }
.map_err(|e| anyhow::anyhow!("stg_ts_ns: {e}"))?,
stg_prev_ts_ns: unsafe { MappedI64Buffer::new(cfg.n_batch * k) }
.map_err(|e| anyhow::anyhow!("stg_prev_ts_ns: {e}"))?,
bid_px_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * 10)?,
bid_sz_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * 10)?,
ask_px_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * 10)?,
ask_sz_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * 10)?,
prev_bid_sz_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * 10)?,
prev_ask_sz_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * 10)?,
regime_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k * REGIME_DIM)?,
prev_mid_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k)?,
trade_signed_vol_all_d: stream.alloc_zeros::<f32>(cfg.n_batch * k)?,
trade_count_all_d: stream.alloc_zeros::<i32>(cfg.n_batch * k)?,
ts_ns_all_d: stream.alloc_zeros::<i64>(cfg.n_batch * k)?,
prev_ts_ns_all_d: stream.alloc_zeros::<i64>(cfg.n_batch * k)?,
w_in_d,
w_rec_d,
b_d,
@@ -331,7 +377,7 @@ impl PerceptionTrainer {
_step_module: step_module,
_heads_module: heads_module,
_bce_module: bce_module,
snap_fn,
snap_batched_fn,
bce_fn,
step_batched_fn,
step_bwd_batched_fn,
@@ -422,57 +468,131 @@ impl PerceptionTrainer {
);
}
// ── 1. Build the batched window tensor [B, K, FEATURE_DIM] by
// running snap_feature_assemble per (b, k) snapshot and
// packing into the right offset. K * B snap launches are
// stream-ordered; single sync at the end since Mamba2's
// forward depends on the full window.
// ── 1. Build the batched window tensor [B, K, FEATURE_DIM] in
// ONE fused snap_feature_assemble_batched launch (was
// B*K separate launches). Per-step host work: pack the
// 12 mapped-pinned staging buffers (~150 KB total); then
// one DtoD per array → device; then one kernel launch
// with B*K threads. Output written directly into the
// window tensor's storage.
let mut window_tensor = GpuTensor::zeros(
&[b_sz, k_seq, FEATURE_DIM],
&self.stream,
).map_err(|e| anyhow::anyhow!("window alloc: {e}"))?;
let snap_cfg = LaunchConfig {
grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0,
};
let snap_nbytes = FEATURE_DIM * std::mem::size_of::<f32>();
for (b_idx, snapshots) in snapshots_batch.iter().enumerate() {
for (k, snap) in snapshots.iter().enumerate() {
upload_into(&self.stream, &snap.bid_px, &self.stg_bid_px, &mut self.bid_px_d)?;
upload_into(&self.stream, &snap.bid_sz, &self.stg_bid_sz, &mut self.bid_sz_d)?;
upload_into(&self.stream, &snap.ask_px, &self.stg_ask_px, &mut self.ask_px_d)?;
upload_into(&self.stream, &snap.ask_sz, &self.stg_ask_sz, &mut self.ask_sz_d)?;
upload_into(&self.stream, &snap.regime, &self.stg_regime, &mut self.regime_d)?;
let prev_mid = snap.prev_mid;
let trade_signed_vol = snap.trade_signed_vol;
let trade_count = snap.trade_count as i32;
let ts_ns = snap.ts_ns as i64;
let prev_ts_ns = snap.prev_ts_ns as i64;
let tick_size = ES_TICK_SIZE;
{
let mut launch = self.stream.launch_builder(&self.snap_fn);
launch
.arg(&self.bid_px_d).arg(&self.bid_sz_d).arg(&self.ask_px_d).arg(&self.ask_sz_d)
.arg(&self.prev_bid_sz_d).arg(&self.prev_ask_sz_d)
.arg(&self.regime_d)
.arg(&prev_mid).arg(&trade_signed_vol).arg(&trade_count)
.arg(&ts_ns).arg(&prev_ts_ns).arg(&tick_size)
.arg(&mut self.snap_feat_d);
unsafe { launch.launch(snap_cfg).context("snap fwd")?; }
}
unsafe {
let (src_ptr, _g1) = self.snap_feat_d.device_ptr(&self.stream);
let (dst_base, _g2) = window_tensor.data_mut().device_ptr_mut(&self.stream);
// window_tensor layout [B, K, F] row-major.
let dst_offset_ptr = dst_base
+ ((b_idx * k_seq + k) * snap_nbytes) as u64;
cudarc::driver::result::memcpy_dtod_async(
dst_offset_ptr, src_ptr, snap_nbytes, self.stream.cu_stream(),
).context("window pack DtoD")?;
let total_snaps = b_sz * k_seq;
debug_assert!(total_snaps <= self.bk_capacity);
// Pack mapped-pinned staging from the input batch (host writes).
{
let bid_px_h = self.stg_bid_px_all.host_slice_mut();
let bid_sz_h = self.stg_bid_sz_all.host_slice_mut();
let ask_px_h = self.stg_ask_px_all.host_slice_mut();
let ask_sz_h = self.stg_ask_sz_all.host_slice_mut();
let regime_h = self.stg_regime_all.host_slice_mut();
for (b_idx, snapshots) in snapshots_batch.iter().enumerate() {
for (k, snap) in snapshots.iter().enumerate() {
let n = b_idx * k_seq + k;
let base10 = n * 10;
let base6 = n * REGIME_DIM;
for i in 0..10 {
bid_px_h[base10 + i] = snap.bid_px[i];
bid_sz_h[base10 + i] = snap.bid_sz[i];
ask_px_h[base10 + i] = snap.ask_px[i];
ask_sz_h[base10 + i] = snap.ask_sz[i];
}
for i in 0..REGIME_DIM {
regime_h[base6 + i] = snap.regime[i];
}
}
}
}
self.stream.synchronize().context("window pack sync")?;
{
let prev_mid_h = self.stg_prev_mid.host_slice_mut();
let tsv_h = self.stg_trade_signed_vol.host_slice_mut();
let tc_h = self.stg_trade_count.host_slice_mut();
let ts_ns_h = self.stg_ts_ns.host_slice_mut();
let prev_ts_ns_h = self.stg_prev_ts_ns.host_slice_mut();
for (b_idx, snapshots) in snapshots_batch.iter().enumerate() {
for (k, snap) in snapshots.iter().enumerate() {
let n = b_idx * k_seq + k;
prev_mid_h[n] = snap.prev_mid;
tsv_h[n] = snap.trade_signed_vol;
tc_h[n] = snap.trade_count as i32;
ts_ns_h[n] = snap.ts_ns as i64;
prev_ts_ns_h[n] = snap.prev_ts_ns as i64;
}
}
}
// DtoD: staging (mapped-pinned, device-visible) → device buffers.
// The prev_bid_sz_all / prev_ask_sz_all buffers stay zero-init
// (loader doesn't currently compute prev-snapshot sizes — same
// as the single-snapshot path).
let n10 = total_snaps * 10;
let n6 = total_snaps * REGIME_DIM;
let n1 = total_snaps;
unsafe {
let s = self.stream.cu_stream();
// f32 arrays
let (d, _g) = self.bid_px_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_bid_px_all.dev_ptr, n10 * 4, s)
.context("bid_px_all dtod")?;
let (d, _g) = self.bid_sz_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_bid_sz_all.dev_ptr, n10 * 4, s)
.context("bid_sz_all dtod")?;
let (d, _g) = self.ask_px_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_ask_px_all.dev_ptr, n10 * 4, s)
.context("ask_px_all dtod")?;
let (d, _g) = self.ask_sz_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_ask_sz_all.dev_ptr, n10 * 4, s)
.context("ask_sz_all dtod")?;
let (d, _g) = self.regime_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_regime_all.dev_ptr, n6 * 4, s)
.context("regime_all dtod")?;
let (d, _g) = self.prev_mid_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_prev_mid.dev_ptr, n1 * 4, s)
.context("prev_mid_all dtod")?;
let (d, _g) = self.trade_signed_vol_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_trade_signed_vol.dev_ptr, n1 * 4, s)
.context("tsv_all dtod")?;
// i32 / i64 arrays
let (d, _g) = self.trade_count_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_trade_count.dev_ptr, n1 * 4, s)
.context("trade_count_all dtod")?;
let (d, _g) = self.ts_ns_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_ts_ns.dev_ptr, n1 * 8, s)
.context("ts_ns_all dtod")?;
let (d, _g) = self.prev_ts_ns_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_prev_ts_ns.dev_ptr, n1 * 8, s)
.context("prev_ts_ns_all dtod")?;
}
// Single batched snap_feature_assemble_batched launch — writes
// directly into window_tensor's storage at [b_idx * K + k] * 32.
let tick_size = ES_TICK_SIZE;
let n_total_i32 = total_snaps as i32;
let snap_block: u32 = 128;
let snap_grid: u32 = (total_snaps as u32).div_ceil(snap_block);
let snap_cfg = LaunchConfig {
grid_dim: (snap_grid, 1, 1),
block_dim: (snap_block, 1, 1),
shared_mem_bytes: 0,
};
unsafe {
let mut launch = self.stream.launch_builder(&self.snap_batched_fn);
launch
.arg(&self.bid_px_all_d).arg(&self.bid_sz_all_d)
.arg(&self.ask_px_all_d).arg(&self.ask_sz_all_d)
.arg(&self.prev_bid_sz_all_d).arg(&self.prev_ask_sz_all_d)
.arg(&self.regime_all_d)
.arg(&self.prev_mid_all_d).arg(&self.trade_signed_vol_all_d)
.arg(&self.trade_count_all_d)
.arg(&self.ts_ns_all_d).arg(&self.prev_ts_ns_all_d)
.arg(&tick_size).arg(&n_total_i32)
.arg(window_tensor.data_mut());
launch.launch(snap_cfg).context("snap_batched fwd")?;
}
self.stream.synchronize().context("snap_batched sync")?;
// ── 2. Mamba2 per-step forward → h_enriched_seq [B, K, HIDDEN_DIM].
let (h_enriched_seq, cache) = self
@@ -817,47 +937,112 @@ impl PerceptionTrainer {
b_sz, snapshots_batch.len(), labels_batch.len()
);
// Build snap_feature window [B, K, FEATURE_DIM].
// Build snap_feature window [B, K, FEATURE_DIM] — fused batched
// pack+upload+launch path (same as step_batched).
let mut window_tensor = GpuTensor::zeros(&[b_sz, k_seq, FEATURE_DIM], &self.stream)
.map_err(|e| anyhow::anyhow!("eval window alloc: {e}"))?;
let snap_cfg = LaunchConfig {
grid_dim: (1, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0,
};
let snap_nbytes = FEATURE_DIM * std::mem::size_of::<f32>();
for (b_idx, snapshots) in snapshots_batch.iter().enumerate() {
for (k, snap) in snapshots.iter().enumerate() {
upload_into(&self.stream, &snap.bid_px, &self.stg_bid_px, &mut self.bid_px_d)?;
upload_into(&self.stream, &snap.bid_sz, &self.stg_bid_sz, &mut self.bid_sz_d)?;
upload_into(&self.stream, &snap.ask_px, &self.stg_ask_px, &mut self.ask_px_d)?;
upload_into(&self.stream, &snap.ask_sz, &self.stg_ask_sz, &mut self.ask_sz_d)?;
upload_into(&self.stream, &snap.regime, &self.stg_regime, &mut self.regime_d)?;
let prev_mid = snap.prev_mid;
let trade_signed_vol = snap.trade_signed_vol;
let trade_count = snap.trade_count as i32;
let ts_ns = snap.ts_ns as i64;
let prev_ts_ns = snap.prev_ts_ns as i64;
let tick_size = ES_TICK_SIZE;
{
let mut launch = self.stream.launch_builder(&self.snap_fn);
launch
.arg(&self.bid_px_d).arg(&self.bid_sz_d).arg(&self.ask_px_d).arg(&self.ask_sz_d)
.arg(&self.prev_bid_sz_d).arg(&self.prev_ask_sz_d).arg(&self.regime_d)
.arg(&prev_mid).arg(&trade_signed_vol).arg(&trade_count)
.arg(&ts_ns).arg(&prev_ts_ns).arg(&tick_size)
.arg(&mut self.snap_feat_d);
unsafe { launch.launch(snap_cfg).context("eval snap fwd")?; }
}
unsafe {
let (src_ptr, _g1) = self.snap_feat_d.device_ptr(&self.stream);
let (dst_base, _g2) = window_tensor.data_mut().device_ptr_mut(&self.stream);
let dst_offset_ptr = dst_base + ((b_idx * k_seq + k) * snap_nbytes) as u64;
cudarc::driver::result::memcpy_dtod_async(
dst_offset_ptr, src_ptr, snap_nbytes, self.stream.cu_stream(),
).context("eval window pack")?;
let total_snaps = b_sz * k_seq;
debug_assert!(total_snaps <= self.bk_capacity);
{
let bid_px_h = self.stg_bid_px_all.host_slice_mut();
let bid_sz_h = self.stg_bid_sz_all.host_slice_mut();
let ask_px_h = self.stg_ask_px_all.host_slice_mut();
let ask_sz_h = self.stg_ask_sz_all.host_slice_mut();
let regime_h = self.stg_regime_all.host_slice_mut();
for (b_idx, snapshots) in snapshots_batch.iter().enumerate() {
for (k, snap) in snapshots.iter().enumerate() {
let n = b_idx * k_seq + k;
let base10 = n * 10;
let base6 = n * REGIME_DIM;
for i in 0..10 {
bid_px_h[base10 + i] = snap.bid_px[i];
bid_sz_h[base10 + i] = snap.bid_sz[i];
ask_px_h[base10 + i] = snap.ask_px[i];
ask_sz_h[base10 + i] = snap.ask_sz[i];
}
for i in 0..REGIME_DIM {
regime_h[base6 + i] = snap.regime[i];
}
}
}
}
self.stream.synchronize().context("eval window sync")?;
{
let prev_mid_h = self.stg_prev_mid.host_slice_mut();
let tsv_h = self.stg_trade_signed_vol.host_slice_mut();
let tc_h = self.stg_trade_count.host_slice_mut();
let ts_ns_h = self.stg_ts_ns.host_slice_mut();
let prev_ts_ns_h = self.stg_prev_ts_ns.host_slice_mut();
for (b_idx, snapshots) in snapshots_batch.iter().enumerate() {
for (k, snap) in snapshots.iter().enumerate() {
let n = b_idx * k_seq + k;
prev_mid_h[n] = snap.prev_mid;
tsv_h[n] = snap.trade_signed_vol;
tc_h[n] = snap.trade_count as i32;
ts_ns_h[n] = snap.ts_ns as i64;
prev_ts_ns_h[n] = snap.prev_ts_ns as i64;
}
}
}
let n10 = total_snaps * 10;
let n6 = total_snaps * REGIME_DIM;
let n1 = total_snaps;
unsafe {
let s = self.stream.cu_stream();
let (d, _g) = self.bid_px_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_bid_px_all.dev_ptr, n10 * 4, s)
.context("eval bid_px_all dtod")?;
let (d, _g) = self.bid_sz_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_bid_sz_all.dev_ptr, n10 * 4, s)
.context("eval bid_sz_all dtod")?;
let (d, _g) = self.ask_px_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_ask_px_all.dev_ptr, n10 * 4, s)
.context("eval ask_px_all dtod")?;
let (d, _g) = self.ask_sz_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_ask_sz_all.dev_ptr, n10 * 4, s)
.context("eval ask_sz_all dtod")?;
let (d, _g) = self.regime_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_regime_all.dev_ptr, n6 * 4, s)
.context("eval regime_all dtod")?;
let (d, _g) = self.prev_mid_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_prev_mid.dev_ptr, n1 * 4, s)
.context("eval prev_mid_all dtod")?;
let (d, _g) = self.trade_signed_vol_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_trade_signed_vol.dev_ptr, n1 * 4, s)
.context("eval tsv_all dtod")?;
let (d, _g) = self.trade_count_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_trade_count.dev_ptr, n1 * 4, s)
.context("eval trade_count_all dtod")?;
let (d, _g) = self.ts_ns_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_ts_ns.dev_ptr, n1 * 8, s)
.context("eval ts_ns_all dtod")?;
let (d, _g) = self.prev_ts_ns_all_d.device_ptr_mut(&self.stream);
cudarc::driver::result::memcpy_dtod_async(d, self.stg_prev_ts_ns.dev_ptr, n1 * 8, s)
.context("eval prev_ts_ns_all dtod")?;
}
let tick_size = ES_TICK_SIZE;
let n_total_i32 = total_snaps as i32;
let snap_block: u32 = 128;
let snap_grid: u32 = (total_snaps as u32).div_ceil(snap_block);
let snap_cfg = LaunchConfig {
grid_dim: (snap_grid, 1, 1),
block_dim: (snap_block, 1, 1),
shared_mem_bytes: 0,
};
unsafe {
let mut launch = self.stream.launch_builder(&self.snap_batched_fn);
launch
.arg(&self.bid_px_all_d).arg(&self.bid_sz_all_d)
.arg(&self.ask_px_all_d).arg(&self.ask_sz_all_d)
.arg(&self.prev_bid_sz_all_d).arg(&self.prev_ask_sz_all_d)
.arg(&self.regime_all_d)
.arg(&self.prev_mid_all_d).arg(&self.trade_signed_vol_all_d)
.arg(&self.trade_count_all_d)
.arg(&self.ts_ns_all_d).arg(&self.prev_ts_ns_all_d)
.arg(&tick_size).arg(&n_total_i32)
.arg(window_tensor.data_mut());
launch.launch(snap_cfg).context("eval snap_batched fwd")?;
}
self.stream.synchronize().context("eval snap_batched sync")?;
let (h_enriched_seq, _cache) = self
.mamba2
@@ -1015,25 +1200,6 @@ fn upload(stream: &Arc<CudaStream>, host: &[f32]) -> Result<CudaSlice<f32>> {
Ok(dst)
}
fn upload_into(
stream: &Arc<CudaStream>,
host: &[f32],
staging: &MappedF32Buffer,
dst: &mut CudaSlice<f32>,
) -> Result<()> {
let n = host.len();
assert_eq!(n, dst.len());
staging.write_from_slice(host);
if n > 0 {
let nbytes = n * std::mem::size_of::<f32>();
unsafe {
let (dst_ptr, _g) = dst.device_ptr_mut(stream);
cudarc::driver::result::memcpy_dtod_async(dst_ptr, staging.dev_ptr, nbytes, stream.cu_stream())
.context("upload_into DtoD")?;
}
}
Ok(())
}
fn download(stream: &Arc<CudaStream>, src: &CudaSlice<f32>) -> Result<Vec<f32>> {
let n = src.len();