diff --git a/crates/ml-alpha/cuda/snap_feature_assemble.cu b/crates/ml-alpha/cuda/snap_feature_assemble.cu index 87ce6c06c..346ae9974 100644 --- a/crates/ml-alpha/cuda/snap_feature_assemble.cu +++ b/crates/ml-alpha/cuda/snap_feature_assemble.cu @@ -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; +} diff --git a/crates/ml-alpha/src/pinned_mem.rs b/crates/ml-alpha/src/pinned_mem.rs index 9e03a4432..c91c70719 100644 --- a/crates/ml-alpha/src/pinned_mem.rs +++ b/crates/ml-alpha/src/pinned_mem.rs @@ -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 { + let num_bytes = len * std::mem::size_of::(); + 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 { + let num_bytes = len * std::mem::size_of::(); + 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); } + } +} diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index 11596cc81..b6b30a280 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -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, _heads_module: Arc, _bce_module: Arc, - 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, - bid_sz_d: CudaSlice, - ask_px_d: CudaSlice, - ask_sz_d: CudaSlice, - prev_bid_sz_d: CudaSlice, - prev_ask_sz_d: CudaSlice, - regime_d: CudaSlice, - snap_feat_d: CudaSlice, + // 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, + bid_sz_all_d: CudaSlice, + ask_px_all_d: CudaSlice, + ask_sz_all_d: CudaSlice, + /// `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, + prev_ask_sz_all_d: CudaSlice, + regime_all_d: CudaSlice, + prev_mid_all_d: CudaSlice, + trade_signed_vol_all_d: CudaSlice, + trade_count_all_d: CudaSlice, + ts_ns_all_d: CudaSlice, + prev_ts_ns_all_d: CudaSlice, // 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, grad_heads_b_d: CudaSlice, - // 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::(10)?, - bid_sz_d: stream.alloc_zeros::(10)?, - ask_px_d: stream.alloc_zeros::(10)?, - ask_sz_d: stream.alloc_zeros::(10)?, - prev_bid_sz_d: stream.alloc_zeros::(10)?, - prev_ask_sz_d: stream.alloc_zeros::(10)?, - regime_d: stream.alloc_zeros::(REGIME_DIM)?, - snap_feat_d: stream.alloc_zeros::(FEATURE_DIM)?, h_new_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * n_hid)?, probs_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * N_HORIZONS)?, labels_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * N_HORIZONS)?, @@ -314,12 +329,43 @@ impl PerceptionTrainer { grad_tau_d: stream.alloc_zeros::(n_hid)?, grad_heads_w_d: stream.alloc_zeros::(N_HORIZONS * n_hid)?, grad_heads_b_d: stream.alloc_zeros::(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::(cfg.n_batch * k * 10)?, + bid_sz_all_d: stream.alloc_zeros::(cfg.n_batch * k * 10)?, + ask_px_all_d: stream.alloc_zeros::(cfg.n_batch * k * 10)?, + ask_sz_all_d: stream.alloc_zeros::(cfg.n_batch * k * 10)?, + prev_bid_sz_all_d: stream.alloc_zeros::(cfg.n_batch * k * 10)?, + prev_ask_sz_all_d: stream.alloc_zeros::(cfg.n_batch * k * 10)?, + regime_all_d: stream.alloc_zeros::(cfg.n_batch * k * REGIME_DIM)?, + prev_mid_all_d: stream.alloc_zeros::(cfg.n_batch * k)?, + trade_signed_vol_all_d: stream.alloc_zeros::(cfg.n_batch * k)?, + trade_count_all_d: stream.alloc_zeros::(cfg.n_batch * k)?, + ts_ns_all_d: stream.alloc_zeros::(cfg.n_batch * k)?, + prev_ts_ns_all_d: stream.alloc_zeros::(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::(); - 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::(); - 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, host: &[f32]) -> Result> { Ok(dst) } -fn upload_into( - stream: &Arc, - host: &[f32], - staging: &MappedF32Buffer, - dst: &mut CudaSlice, -) -> 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::(); - 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, src: &CudaSlice) -> Result> { let n = src.len();