From 45a2041db4e06079d3fc8cc251fc9084f6506513 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 24 May 2026 20:36:05 +0200 Subject: [PATCH] =?UTF-8?q?feat(rl):=20SP20=20P6=20position=20heat=20cap?= =?UTF-8?q?=20=E2=80=94=20force-flat=20on=20over-leverage?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Last-defense guard: if |position_lots| exceeds the ISV-driven RL_HEAT_CAP_MAX_LOTS (slot 504, default 8 = MAX_UNITS × max_order_size), the kernel overrides actions[b] to FlatFromLong (a3) or FlatFromShort (a4) — full flatten, no partial. Catches runaway pyramid accumulation before it reaches actions_to_market_targets. Override stack ordering (step_with_lobsim): 1. rl_trail_mutate (a7/a8) 2. rl_trail_stop_check → may override to FlatFromLong/Short 3. rl_position_heat_check (THIS) → may override to FlatFromLong/Short 4. actions_to_market_targets → reads final actions[b] Kernel `cuda/rl_position_heat_check.cu`: * 1 block, b_size threads (grid-stride for b_size > 256) * Reads position_lots from pos_state at offset 0 (PosFlat layout) * Cap read from ISV[504]; if cap ≤ 0 → no-op (guard disabled) * Per feedback_no_atomicadd: fired-count diagnostic uses shared-mem flag array + thread-0 serial count (b_size ≤ 256 in practice) * Writes fired-count to ISV[505] for diag ISV slots: * 504: RL_HEAT_CAP_MAX_LOTS_INDEX (seed 8.0) * 505: RL_HEAT_CAP_FIRED_COUNT_INDEX (diagnostic, written per step) * RL_SLOTS_END bumped 505 → 506 Diag (alpha_rl_train): * "heat_cap": { "fired_count": N, "max_lots": 8 } GPU oracle test (trade_management_kernels.rs): * position_heat_cap_overrides_on_breach — long 5 > cap 4 → a3; short -5 < -cap -4 → a4; long 3 ≤ cap 4 → untouched (Hold) Verification (RTX 3050 Ti): * cargo check -p ml-alpha --examples → clean * integrated_trainer_smoke 1/1 → ok * trade_management_kernels 6/6 (was 5/5, +1 heat cap) → ok * audit-rust-consts → 0 flags --- crates/ml-alpha/build.rs | 1 + .../ml-alpha/cuda/rl_position_heat_check.cu | 74 ++++++++++++++++ crates/ml-alpha/examples/alpha_rl_train.rs | 7 ++ crates/ml-alpha/src/rl/isv_slots.rs | 15 +++- crates/ml-alpha/src/trainer/integrated.rs | 85 ++++++++++++++++++- .../tests/trade_management_kernels.rs | 63 ++++++++++++++ 6 files changed, 243 insertions(+), 2 deletions(-) create mode 100644 crates/ml-alpha/cuda/rl_position_heat_check.cu diff --git a/crates/ml-alpha/build.rs b/crates/ml-alpha/build.rs index 708b40194..e98bff8ff 100644 --- a/crates/ml-alpha/build.rs +++ b/crates/ml-alpha/build.rs @@ -79,6 +79,7 @@ const KERNELS: &[&str] = &[ "rl_frd_softmax_ce_grad", // SP20 P3 F.3a: per-(batch, horizon) softmax + CE loss + dL/dlogits; 1 block per (b, h), 21 threads; label = -1 sentinel masks the row "rl_frd_layer2_bwd", // SP20 P3 F.3b: FRD head layer-2 backward — dW2 (per-batch scratch), db2 (per-batch scratch), dhidden (per-batch overwrite); 1 block per batch, 64 threads "rl_frd_layer1_bwd", // SP20 P3 F.3c: FRD head layer-1 backward — dW1 (per-batch scratch), db1 (per-batch scratch), dh_t (per-batch overwrite); applies ReLU mask via cached post-ReLU hidden; 1 block per batch, 128 threads + "rl_position_heat_check", // SP20 P6: position heat cap — force-flat when |position_lots| exceeds ISV-driven max; last defense before actions_to_market_targets ]; // Cache bust v31 — five new reduce / derive kernels populate the input diff --git a/crates/ml-alpha/cuda/rl_position_heat_check.cu b/crates/ml-alpha/cuda/rl_position_heat_check.cu new file mode 100644 index 000000000..58f4e35dc --- /dev/null +++ b/crates/ml-alpha/cuda/rl_position_heat_check.cu @@ -0,0 +1,74 @@ +// rl_position_heat_check.cu — SP20 P6 position heat cap. +// +// Last-defense guard: if the aggregate position exceeds the ISV-driven +// maximum allowed lots, OVERRIDE actions[b] to FlatFromLong (a3) or +// FlatFromShort (a4) to force a full flatten. Unlike the per-unit +// trail-stop check (P5), the heat cap fires on the AGGREGATE position +// — it catches cases where pyramiding (P7) or rapid re-entry +// accumulates lots beyond the risk budget. +// +// Per SP20 §3 P6: "heat cap is the LAST defense; full flat is +// acceptable here" — no partial close, no per-unit selection. The +// entire position is closed at once via the existing FlatFromLong/Short +// plumbing in actions_to_market_targets. +// +// Override stack ordering (step_with_lobsim): +// 1. rl_trail_mutate (a7/a8) +// 2. rl_trail_stop_check → may override to FlatFromLong/Short +// 3. rl_position_heat_check (THIS) → may override to FlatFromLong/Short +// 4. actions_to_market_targets → reads final actions[b] +// +// Per `feedback_no_atomicadd`: per-batch sole-writer pattern. +// Per `feedback_cpu_is_read_only`: pure device-side. +// Per `feedback_isv_for_adaptive_bounds`: cap lives in ISV (slot 504). + +#include + +#define ACTION_FLAT_FROM_LONG 3 +#define ACTION_FLAT_FROM_SHORT 4 +#define RL_HEAT_CAP_MAX_LOTS_INDEX 504 + +// Diagnostic output slot — writes the count of batches where the heat +// cap fired this step. Diag reads post-step for the JSONL "heat_cap" +// block. Per `feedback_no_atomicadd`: use a single-thread serial +// count after a shared-mem flag pass (b_size ≤ 256 makes serial OK). +#define RL_HEAT_CAP_FIRED_COUNT_INDEX 505 + +extern "C" __global__ void rl_position_heat_check( + int* __restrict__ actions, // [B] IN/OUT + const unsigned char* __restrict__ pos_state, // [B * pos_bytes] + float* __restrict__ isv, // IN/OUT (diag write) + int b_size, + int pos_bytes +) { + const int b = blockIdx.x * blockDim.x + threadIdx.x; + + const int cap = (int) isv[RL_HEAT_CAP_MAX_LOTS_INDEX]; + + // Per-thread flag: did this batch entry fire the cap? + __shared__ int s_flags[256]; // max b_size threads + const int tid = threadIdx.x; + s_flags[tid] = 0; + + if (b < b_size && cap > 0) { + const int position_lots = *(const int*)(pos_state + b * pos_bytes); + if (position_lots > cap) { + actions[b] = ACTION_FLAT_FROM_LONG; + s_flags[tid] = 1; + } else if (position_lots < -cap) { + actions[b] = ACTION_FLAT_FROM_SHORT; + s_flags[tid] = 1; + } + } + __syncthreads(); + + // Thread 0 serial-counts the fired flags (b_size ≤ 256, diagnostic + // path — no perf concern). Per feedback_no_atomicadd: no atomicAdd. + if (tid == 0) { + int count = 0; + for (int i = 0; i < b_size && i < 256; ++i) { + count += s_flags[i]; + } + isv[RL_HEAT_CAP_FIRED_COUNT_INDEX] = (float) count; + } +} diff --git a/crates/ml-alpha/examples/alpha_rl_train.rs b/crates/ml-alpha/examples/alpha_rl_train.rs index 5dbb72ce6..e04d608b4 100644 --- a/crates/ml-alpha/examples/alpha_rl_train.rs +++ b/crates/ml-alpha/examples/alpha_rl_train.rs @@ -930,6 +930,13 @@ fn main() -> Result<()> { "position": { "lots": position_lots_host, }, + // SP20 P6 heat cap diag — how many batch entries fired the + // position-size guard this step. 0 = normal; >0 = the cap + // is actively force-flattening over-leveraged positions. + "heat_cap": { + "fired_count": isv[ml_alpha::rl::isv_slots::RL_HEAT_CAP_FIRED_COUNT_INDEX], + "max_lots": isv[ml_alpha::rl::isv_slots::RL_HEAT_CAP_MAX_LOTS_INDEX], + }, // SP20 P3 FRD head diag — per-horizon softmax entropy + argmax-mode // bucket index (averaged across the batch). At init, Xavier × 0.1 // weights → logits ≈ 0 → entropy ≈ ln(21) = 3.044 and the argmax diff --git a/crates/ml-alpha/src/rl/isv_slots.rs b/crates/ml-alpha/src/rl/isv_slots.rs index 2dbe1bfad..96ffaeaac 100644 --- a/crates/ml-alpha/src/rl/isv_slots.rs +++ b/crates/ml-alpha/src/rl/isv_slots.rs @@ -792,6 +792,19 @@ pub const RL_FRD_HORIZON_3_TICKS_INDEX: usize = 502; /// realised σ drifts; only the 21-atom count is structural compile-time. pub const RL_FRD_BUCKET_RANGE_SIGMA_INDEX: usize = 503; +/// SP20 P6 — maximum absolute position (lots) before the heat-cap +/// kernel force-flats the entire position. Default 8 = +/// `MAX_UNITS(4) × max_single_order_size(2)`. The kernel compares +/// `|position_lots|` from `pos_state` against this slot and overrides +/// `actions[b]` to FlatFromLong/Short when exceeded. Last-defense +/// guard against runaway pyramid accumulation. +pub const RL_HEAT_CAP_MAX_LOTS_INDEX: usize = 504; + +/// SP20 P6 — diagnostic: number of batch entries where the heat cap +/// fired this step. Written by `rl_position_heat_check`; read by +/// diag JSONL. Zero when no batch entry exceeds the cap. +pub const RL_HEAT_CAP_FIRED_COUNT_INDEX: usize = 505; + /// Last RL-allocated slot index (exclusive). The integrated trainer /// extends `ISV_TOTAL_DIM` to at least this value at trainer init time. -pub const RL_SLOTS_END: usize = 504; +pub const RL_SLOTS_END: usize = 506; diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index 52ebffcbd..864cbffe5 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -192,6 +192,8 @@ const RL_UNIT_STATE_UPDATE_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/rl_unit_state_update.cubin")); const RL_TRAIL_MUTATE_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/rl_trail_mutate.cubin")); +const RL_POSITION_HEAT_CHECK_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_position_heat_check.cubin")); const RL_TRAIL_STOP_CHECK_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/rl_trail_stop_check.cubin")); const ACTIONS_TO_MARKET_TARGETS_CUBIN: &[u8] = @@ -439,6 +441,8 @@ pub struct IntegratedTrainer { rl_trail_mutate_fn: CudaFunction, _rl_trail_stop_check_module: Arc, rl_trail_stop_check_fn: CudaFunction, + _rl_position_heat_check_module: Arc, + rl_position_heat_check_fn: CudaFunction, // Per-batch per-unit trade state (SP20 P1). // MAX_UNITS=4 reserved; only slot 0 used until SP20 P7 ships @@ -912,6 +916,12 @@ impl IntegratedTrainer { let rl_trail_stop_check_fn = rl_trail_stop_check_module .load_function("rl_trail_stop_check") .context("load rl_trail_stop_check")?; + let rl_position_heat_check_module = ctx + .load_cubin(RL_POSITION_HEAT_CHECK_CUBIN.to_vec()) + .context("load rl_position_heat_check cubin")?; + let rl_position_heat_check_fn = rl_position_heat_check_module + .load_function("rl_position_heat_check") + .context("load rl_position_heat_check")?; let actions_to_market_targets_module = ctx .load_cubin(ACTIONS_TO_MARKET_TARGETS_CUBIN.to_vec()) .context("load actions_to_market_targets cubin")?; @@ -1293,6 +1303,8 @@ impl IntegratedTrainer { rl_trail_mutate_fn, _rl_trail_stop_check_module: rl_trail_stop_check_module, rl_trail_stop_check_fn, + _rl_position_heat_check_module: rl_position_heat_check_module, + rl_position_heat_check_fn, unit_entry_price_d, unit_entry_step_d, unit_lots_d, @@ -1426,7 +1438,7 @@ impl IntegratedTrainer { // (slot, value) pair — pure device write, no HtoD per // `feedback_no_htod_htoh_only_mapped_pinned`. { - let isv_constants: [(usize, f32); 46] = [ + let isv_constants: [(usize, f32); 47] = [ // Static seeds for the adaptive reward-clamp controller — // these are the initial values that // `rl_reward_clamp_controller` will replace once it @@ -1482,6 +1494,9 @@ impl IntegratedTrainer { crate::rl::common::FRD_HORIZON_TICKS[2] as f32), (crate::rl::isv_slots::RL_FRD_BUCKET_RANGE_SIGMA_INDEX, crate::rl::common::FRD_BUCKET_RANGE_SIGMA), + // SP20 P6 heat-cap: max absolute position (lots). + // Default 8 = MAX_UNITS(4) × max_single_order_size(2). + (crate::rl::isv_slots::RL_HEAT_CAP_MAX_LOTS_INDEX, 8.0), (crate::rl::isv_slots::RL_KL_TARGET_INDEX, 0.01), (crate::rl::isv_slots::RL_IMPROVEMENT_THRESHOLD_INDEX, 0.99), (crate::rl::isv_slots::RL_PLATEAU_PATIENCE_INDEX, 1000.0), @@ -2042,6 +2057,45 @@ impl IntegratedTrainer { Ok(()) } + /// SP20 P6: launch `rl_position_heat_check` — if `|position_lots|` + /// exceeds the ISV-driven cap (`RL_HEAT_CAP_MAX_LOTS_INDEX`, slot + /// 504, default 8), OVERRIDE `actions[b]` to FlatFromLong (a3) or + /// FlatFromShort (a4). Last-defense guard — full flat, no partial. + /// Also writes `RL_HEAT_CAP_FIRED_COUNT_INDEX` (slot 505) with the + /// per-step fired count for diag. + pub fn launch_rl_position_heat_check( + &self, + actions_d: &mut CudaSlice, + pos_state_d: &CudaSlice, + b_size: usize, + pos_bytes: usize, + ) -> Result<()> { + debug_assert_eq!(actions_d.len(), b_size); + debug_assert_eq!(pos_state_d.len(), b_size * pos_bytes); + let block = (b_size as u32).min(256); + let grid = ((b_size as u32) + block - 1) / block; + let cfg = LaunchConfig { + grid_dim: (grid.max(1), 1, 1), + block_dim: (block.max(1), 1, 1), + shared_mem_bytes: 256 * std::mem::size_of::() as u32, + }; + let b_size_i = b_size as i32; + let pos_bytes_i = pos_bytes as i32; + let mut launch = self.stream.launch_builder(&self.rl_position_heat_check_fn); + launch + .arg(actions_d) + .arg(pos_state_d) + .arg(&self.isv_d) + .arg(&b_size_i) + .arg(&pos_bytes_i); + unsafe { + launch + .launch(cfg) + .context("rl_position_heat_check launch")?; + } + Ok(()) + } + /// SP20 P1: launch `rl_unit_state_update` — detects per-batch /// position transitions vs `prev_pos_lots_d` and updates per-unit /// state (entry_price, entry_step, lots, initial_r, trail_distance, @@ -3543,6 +3597,35 @@ impl IntegratedTrainer { } } + // ── SP20 P6 position heat cap — AFTER trail_stop_check, BEFORE + // actions_to_market_targets. If |position_lots| exceeds the ISV- + // driven max (slot 504, default 8), OVERRIDE to FlatFromLong/Short. + // 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 cfg = LaunchConfig { + grid_dim: (grid.max(1), 1, 1), + block_dim: (block.max(1), 1, 1), + shared_mem_bytes: 256 * std::mem::size_of::() as u32, + }; + let mut launch = + self.stream.launch_builder(&self.rl_position_heat_check_fn); + launch + .arg(&mut self.actions_d) + .arg(pos_d_ref_heat) + .arg(&self.isv_d) + .arg(&b_size_i) + .arg(&pos_bytes_i); + unsafe { + launch + .launch(cfg) + .context("rl_position_heat_check launch")?; + } + } + { let (pos_d_ref, market_targets_d) = lobsim.pos_and_market_targets_mut(); let cfg = LaunchConfig { diff --git a/crates/ml-alpha/tests/trade_management_kernels.rs b/crates/ml-alpha/tests/trade_management_kernels.rs index 5c1a740f0..4726eb354 100644 --- a/crates/ml-alpha/tests/trade_management_kernels.rs +++ b/crates/ml-alpha/tests/trade_management_kernels.rs @@ -482,3 +482,66 @@ fn trail_stop_check_overrides_action_on_breach() -> Result<()> { eprintln!("PASS — trail-stop breach overrides + no-breach leaves alone + short-side symmetry"); Ok(()) } + +#[test] +#[ignore = "requires CUDA (MlDevice::cuda(0))"] +fn position_heat_cap_overrides_on_breach() -> Result<()> { + let Some((dev, mut trainer)) = build_trainer() else { return Ok(()) }; + let stream = dev.cuda_stream()?.clone(); + let pos_bytes = POS_BYTES; + let b_size = 1; + + // Seed the heat cap to 4 lots. + set_isv_slot(&mut trainer, &stream, ml_alpha::rl::isv_slots::RL_HEAT_CAP_MAX_LOTS_INDEX, 4.0)?; + + // Long 5 lots (exceeds cap 4) — should override to FlatFromLong (a3). + let mut actions_d = upload_i32(&stream, &[2])?; // Hold + let pos_over_d = upload_u8(&stream, &pos_buf(5, 100.0))?; + trainer.launch_rl_position_heat_check( + &mut actions_d, + &pos_over_d, + b_size, + pos_bytes, + )?; + let actions = read_slice_i32_d_pub(&stream, &actions_d, b_size)?; + assert_eq!( + actions[0], 3, + "long 5 > cap 4 should override to FlatFromLong (a3); got {}", + actions[0] + ); + + // Short -5 lots (exceeds cap 4 in abs) — should override to FlatFromShort (a4). + let mut actions_d2 = upload_i32(&stream, &[2])?; + let pos_short_d = upload_u8(&stream, &pos_buf(-5, 100.0))?; + trainer.launch_rl_position_heat_check( + &mut actions_d2, + &pos_short_d, + b_size, + pos_bytes, + )?; + let actions = read_slice_i32_d_pub(&stream, &actions_d2, b_size)?; + assert_eq!( + actions[0], 4, + "short -5 < -cap -4 should override to FlatFromShort (a4); got {}", + actions[0] + ); + + // Within cap (long 3 ≤ cap 4) — action should NOT change. + let mut actions_d3 = upload_i32(&stream, &[2])?; + let pos_within_d = upload_u8(&stream, &pos_buf(3, 100.0))?; + trainer.launch_rl_position_heat_check( + &mut actions_d3, + &pos_within_d, + b_size, + pos_bytes, + )?; + let actions = read_slice_i32_d_pub(&stream, &actions_d3, b_size)?; + assert_eq!( + actions[0], 2, + "long 3 ≤ cap 4 should leave action untouched (a2 Hold); got {}", + actions[0] + ); + + eprintln!("PASS — heat cap overrides on breach + leaves alone within cap"); + Ok(()) +}