diff --git a/crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu b/crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu index 1fb05cdd6..777528831 100644 --- a/crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu +++ b/crates/ml-alpha/cuda/rl_fused_reward_pipeline.cu @@ -220,14 +220,8 @@ extern "C" __global__ void rl_fused_reward_pipeline( r -= entry_cost; } - // 2. Short-hold penalty: ONLY on winning quick exits. - // Quick exit on a LOSER is neutral (no penalty, no bonus). - // Quick exit on a WINNER is BAD (bailed on a good wave) — penalize. - // This alone creates the right asymmetry: model learns to hold - // winners (penalty for bailing) and is FREE to exit losers (no - // penalty). No explicit bonus needed — the natural PnL signal - // teaches loss-cutting. - if (done > 0.5f && r > 0.0f && (float)hold_time < min_hold) { + // 2. Short-hold penalty: trade close with hold time below minimum. + if (done > 0.5f && (float)hold_time < min_hold) { r *= penalty; } diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index 7b1f77c02..16708341f 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -360,6 +360,32 @@ struct HotPathPtrs { multires_output: u64, } +/// Cached raw device pointers for lobsim accessors. Eliminates per-call +/// vtable dispatch + CudaSlice borrow overhead on `pos_d()`, `bid_px_d()`, +/// `ask_px_d()` which are called 5-10× per step but always return the same +/// pre-allocated device buffer. Extracted once at step entry, then threaded +/// through all raw_launch sites as plain `u64`. +struct LobPtrs { + pos: u64, + bid_px: u64, + ask_px: u64, + pos_bytes: i32, +} + +impl LobPtrs { + /// Snapshot all stable lobsim device pointers. Safe because the + /// lobsim pre-allocates pos/bid/ask buffers at construction; the + /// returned u64 values are valid for the lifetime of the lobsim. + fn new(lobsim: &dyn RlLobBackend) -> Self { + Self { + pos: lobsim.pos_d().raw_ptr(), + bid_px: lobsim.bid_px_d().raw_ptr(), + ask_px: lobsim.ask_px_d().raw_ptr(), + pos_bytes: lobsim.pos_bytes() as i32, + } + } +} + /// Configuration for [`IntegratedTrainer`]. Wraps a `PerceptionTrainerConfig` /// (the encoder + BCE + aux side) plus RL-specific overrides for the new /// heads. @@ -5207,6 +5233,12 @@ impl IntegratedTrainer { .map_err(|e| anyhow::anyhow!("stream wait train_done: {:?}", e))?; } + // Cache lobsim raw device pointers once per step to eliminate + // per-call vtable dispatch + CudaSlice borrow overhead. The + // underlying buffers are pre-allocated and stable for the + // lifetime of the lobsim. + let lob = LobPtrs::new(lobsim); + // ── Step 0: bump device-resident step counter (ISV[548]). // Must run BEFORE any kernel that reads current_step from ISV. // Single thread, single block — graph-safe (no scalar args change). @@ -5516,17 +5548,15 @@ impl IntegratedTrainer { // ── Gates + log_pi: OUTSIDE graph capture so they fire every // step (not captured as no-ops during warmup). Device-side kernels. { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); args.push_ptr(self.q_logits_d.raw_ptr()); args.push_ptr(self.atom_supports_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); - args.push_i32(pos_bytes_i); + args.push_i32(lob.pos_bytes); let mut ptrs = args.build_arg_ptrs(); unsafe { raw_launch( @@ -5538,16 +5568,14 @@ impl IntegratedTrainer { } } { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); args.push_ptr(self.frd_logits_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); - args.push_i32(pos_bytes_i); + args.push_i32(lob.pos_bytes); let mut ptrs = args.build_arg_ptrs(); unsafe { raw_launch( @@ -5625,7 +5653,7 @@ impl IntegratedTrainer { // Reads current position_lots from lobsim.pos_d for Flat-from-* // actions; reads actions from self.actions_d (just written by // rl_action_kernel above). - let pos_bytes_i = lobsim.pos_bytes() as i32; + let pos_bytes_i = lob.pos_bytes; let b_size_i = b_size as i32; // Stage ts_ns to device buffer for graph-captured multires kernel. @@ -5690,12 +5718,11 @@ impl IntegratedTrainer { // hold time < ISV-driven minimum. Trail stops still fire // (they run AFTER this, safety overrides patience). { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); args.push_ptr(self.steps_since_done_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); args.push_i32(pos_bytes_i); @@ -5714,16 +5741,14 @@ impl IntegratedTrainer { // Runs BEFORE trail_mutate so the structural decay is applied // first, then agent's a7/a8 fine-tuning on top. { - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); let mut args = RawArgs::new(); args.push_ptr(self.unit_trail_distance_d.raw_ptr()); args.push_ptr(self.unit_active_d.raw_ptr()); args.push_ptr(self.unit_entry_price_d.raw_ptr()); args.push_ptr(self.unit_initial_r_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); let mut ptrs = args.build_arg_ptrs(); @@ -5763,12 +5788,10 @@ impl IntegratedTrainer { // FlatFromLong/Short on per-unit breach. Reads shared lobsim // best book for current mid. { - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.unit_active_d.raw_ptr()); args.push_ptr(self.unit_entry_price_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr()); @@ -5793,13 +5816,12 @@ impl IntegratedTrainer { // Reads pos_state for aggregate position; writes diag fired-count // to ISV slot 505. { - let pos_d_ref_heat: &CudaSlice = lobsim.pos_d(); let block = (b_size as u32).min(256); let grid = ((b_size as u32) + block - 1) / block; let smem = 256 * std::mem::size_of::() as u32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); - args.push_ptr(pos_d_ref_heat.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); args.push_i32(pos_bytes_i); @@ -5815,16 +5837,17 @@ impl IntegratedTrainer { } { - let (pos_d_ref, bid_px_d, ask_px_d, market_targets_d) = - lobsim.pos_book_and_market_targets_mut(); + let (_pos_d_ref, market_targets_d) = + lobsim.pos_and_market_targets_mut(); + let market_targets_ptr = market_targets_d.raw_ptr(); let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); - args.push_ptr(market_targets_d.raw_ptr()); + args.push_ptr(lob.pos); + args.push_ptr(market_targets_ptr); args.push_ptr(self.isv_dev_ptr); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.unit_entry_price_d.raw_ptr()); args.push_ptr(self.unit_active_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr()); @@ -5879,8 +5902,9 @@ impl IntegratedTrainer { snapshots: &[Mbp10RawInput], b_size: usize, ) -> Result { + let lob = LobPtrs::new(lobsim); let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; + let pos_bytes_i = lob.pos_bytes; // ── Graph B: post-fill reward/EMA/controller pipeline ───────── // ~20 kernels from extract_realized_pnl_delta through @@ -5906,10 +5930,9 @@ impl IntegratedTrainer { // reward_shaping + raw_rewards snapshot + recent_outcome_update // in a single per-batch kernel (7→1 launch). { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.prev_realized_pnl_d.raw_ptr()); args.push_ptr(self.prev_position_lots_d.raw_ptr()); args.push_ptr(self.rewards_d.raw_ptr()); @@ -5999,11 +6022,6 @@ impl IntegratedTrainer { // Trade context features — derived from unit state + current mid. { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); - let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); args.push_ptr(self.trade_context_d.raw_ptr()); @@ -6012,9 +6030,9 @@ impl IntegratedTrainer { args.push_ptr(self.unit_entry_step_d.raw_ptr()); args.push_ptr(self.unit_initial_r_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.pos); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); args.push_i32(pos_bytes_i); @@ -6054,19 +6072,16 @@ impl IntegratedTrainer { // Multi-resolution streaming features — time-weighted EMA at 3 horizons. { - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); - let b_size_i = b_size as i32; let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); args.push_ptr(self.multires_state_d.raw_ptr()); args.push_ptr(self.multires_output_d.raw_ptr()); args.push_ptr(self.multires_prev_mid_d.raw_ptr()); args.push_ptr(self.multires_prev_ts_ns_d.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); // bid_sz uses bid_px - args.push_ptr(ask_px_d.raw_ptr()); // ask_sz uses ask_px + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); + args.push_ptr(lob.bid_px); // bid_sz uses bid_px + args.push_ptr(lob.ask_px); // ask_sz uses ask_px args.push_ptr(self.isv_dev_ptr); args.push_ptr(self.ts_ns_d.raw_ptr()); args.push_i32(b_size_i); @@ -6371,16 +6386,11 @@ impl IntegratedTrainer { // mid-price into ring buffer and update peak (max for long, // min for short) while a position is open. Flat resets ring. { - let pos_d_ref = lobsim.pos_d(); - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); let grid_x = ((b_size as u32) + 31) / 32; - let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; let mut args = RawArgs::new(); - args.push_ptr(pos_d_ref.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.pos); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.hindsight.mid_ring_d.raw_ptr()); args.push_ptr(self.hindsight.ring_write_idx_d.raw_ptr()); args.push_ptr(self.hindsight.peak_mid_d.raw_ptr()); @@ -6546,13 +6556,11 @@ impl IntegratedTrainer { // (ISV[552] steps). If holding would have been better, inject // a synthetic "should-have-held-longer" transition. { - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); let smem = (2 * 256 * std::mem::size_of::()) as u32; let cap_i = self.gpu_replay.capacity as i32; let mut args = RawArgs::new(); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.isv_dev_ptr); args.push_ptr(self.hindsight.closed_ring_d.raw_ptr()); args.push_ptr(self.hindsight.closed_h_t_d.raw_ptr()); @@ -7014,6 +7022,12 @@ impl IntegratedTrainer { b_size: usize, seq_len: usize, ) -> Result { + // Cache lobsim raw device pointers once per step to eliminate + // per-call vtable dispatch + CudaSlice borrow overhead. The + // underlying buffers are pre-allocated and stable for the + // lifetime of the lobsim. + let lob = LobPtrs::new(lobsim); + // ── Graph A: pre-snapshot kernel pipeline ────────────────────── if self.prefill_graph.is_some() { unsafe { @@ -7212,17 +7226,15 @@ impl IntegratedTrainer { // ── Gates + log_pi: OUTSIDE graph capture (GPU data path). { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); args.push_ptr(self.q_logits_d.raw_ptr()); args.push_ptr(self.atom_supports_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); - args.push_i32(pos_bytes_i); + args.push_i32(lob.pos_bytes); let mut ptrs = args.build_arg_ptrs(); unsafe { raw_launch( @@ -7234,16 +7246,14 @@ impl IntegratedTrainer { } } { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let b_size_i = b_size as i32; - let pos_bytes_i = lobsim.pos_bytes() as i32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); args.push_ptr(self.frd_logits_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); - args.push_i32(pos_bytes_i); + args.push_i32(lob.pos_bytes); let mut ptrs = args.build_arg_ptrs(); unsafe { raw_launch( @@ -7333,7 +7343,7 @@ impl IntegratedTrainer { } } - let pos_bytes_i = lobsim.pos_bytes() as i32; + let pos_bytes_i = lob.pos_bytes; let b_size_i = b_size as i32; // ── Graph A2: post-snapshot / pre-fill kernel pipeline ───────── @@ -7374,12 +7384,11 @@ impl IntegratedTrainer { // Min hold check { - let pos_d_ref: &CudaSlice = lobsim.pos_d(); let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); args.push_ptr(self.steps_since_done_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); args.push_i32(pos_bytes_i); @@ -7396,16 +7405,14 @@ impl IntegratedTrainer { // Asymmetric trail decay { - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); let mut args = RawArgs::new(); args.push_ptr(self.unit_trail_distance_d.raw_ptr()); args.push_ptr(self.unit_active_d.raw_ptr()); args.push_ptr(self.unit_entry_price_d.raw_ptr()); args.push_ptr(self.unit_initial_r_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); let mut ptrs = args.build_arg_ptrs(); @@ -7440,12 +7447,10 @@ impl IntegratedTrainer { // Trail stop check { - let bid_px_d = lobsim.bid_px_d(); - let ask_px_d = lobsim.ask_px_d(); let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.unit_active_d.raw_ptr()); args.push_ptr(self.unit_entry_price_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr()); @@ -7466,13 +7471,12 @@ impl IntegratedTrainer { // Position heat cap { - let pos_d_ref_heat: &CudaSlice = lobsim.pos_d(); let block = (b_size as u32).min(256); let grid = ((b_size as u32) + block - 1) / block; let smem = 256 * std::mem::size_of::() as u32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); - args.push_ptr(pos_d_ref_heat.raw_ptr()); + args.push_ptr(lob.pos); args.push_ptr(self.isv_dev_ptr); args.push_i32(b_size_i); args.push_i32(pos_bytes_i); @@ -7489,16 +7493,17 @@ impl IntegratedTrainer { // actions_to_market_targets { - let (pos_d_ref, bid_px_d, ask_px_d, market_targets_d) = - lobsim.pos_book_and_market_targets_mut(); + let (_pos_d_ref, market_targets_d) = + lobsim.pos_and_market_targets_mut(); + let market_targets_ptr = market_targets_d.raw_ptr(); let grid_x = ((b_size as u32) + 31) / 32; let mut args = RawArgs::new(); args.push_ptr(self.actions_d.raw_ptr()); - args.push_ptr(pos_d_ref.raw_ptr()); - args.push_ptr(market_targets_d.raw_ptr()); + args.push_ptr(lob.pos); + args.push_ptr(market_targets_ptr); args.push_ptr(self.isv_dev_ptr); - args.push_ptr(bid_px_d.raw_ptr()); - args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(lob.bid_px); + args.push_ptr(lob.ask_px); args.push_ptr(self.unit_entry_price_d.raw_ptr()); args.push_ptr(self.unit_active_d.raw_ptr()); args.push_ptr(self.unit_lots_d.raw_ptr());