""" Numba-accelerated core functions for MALKHUT CWM. Targets the hottest loops: - fill_from_levels: sequential level consumption (called every transition) - round_tick / round_lot / clip_lots: rounding operations - feature extraction: vectorized operations - replay comparison: deep state comparison Design: numba-friendly inner functions operate on flat arrays, not dataclasses. The CWM calls these from its hot path. """ from __future__ import annotations import numpy as np from numba import njit, prange # ============================================================================== # Fill from levels — sequential level consumption # ============================================================================== @njit(cache=True) def fill_from_levels( bid_prices: np.ndarray, bid_qtys: np.ndarray, ask_prices: np.ndarray, ask_qtys: np.ndarray, qty_desired: float, lot: float, min_qty: float, side_is_buy: bool, ) -> tuple: """ Consume qty from price levels (price-time priority). Returns: (filled_qty, avg_fill_price, remaining_bid_qtys, remaining_ask_qtys) Numba-optimized: operates on flat arrays, no object creation. """ filled = 0.0 total_cost = 0.0 qty_remaining = qty_desired if side_is_buy: # Consume from asks (lowest first — already sorted ascending) new_ask_qtys = ask_qtys.copy() for i in range(len(ask_prices)): if qty_remaining <= 1e-12: break level_qty = new_ask_qtys[i] if level_qty <= 0: continue take = min(qty_remaining, level_qty) # Round to lot take_rounded = round(take / lot) * lot if take_rounded < min_qty: break filled += take_rounded total_cost += take_rounded * ask_prices[i] qty_remaining -= take_rounded new_ask_qtys[i] = level_qty - take_rounded if new_ask_qtys[i] < min_qty: new_ask_qtys[i] = 0.0 return filled, total_cost / filled if filled > 0 else 0.0, bid_qtys, new_ask_qtys else: # Consume from bids (highest first — already sorted descending) new_bid_qtys = bid_qtys.copy() for i in range(len(bid_prices)): if qty_remaining <= 1e-12: break level_qty = new_bid_qtys[i] if level_qty <= 0: continue take = min(qty_remaining, level_qty) take_rounded = round(take / lot) * lot if take_rounded < min_qty: break filled += take_rounded total_cost += take_rounded * bid_prices[i] qty_remaining -= take_rounded new_bid_qtys[i] = level_qty - take_rounded if new_bid_qtys[i] < min_qty: new_bid_qtys[i] = 0.0 return filled, total_cost / filled if filled > 0 else 0.0, new_bid_qtys, ask_qtys # ============================================================================== # Rounding operations # ============================================================================== @njit(cache=True) def round_tick(price: float, tick: float) -> float: return round(price / tick) * tick @njit(cache=True) def round_lot(qty: float, lot: float) -> float: return round(qty / lot) * lot @njit(cache=True) def clip_lots(qty: float, lot: float, min_qty: float) -> float: q = round(qty / lot) * lot return q if q >= min_qty else 0.0 # ============================================================================== # Feature extraction — vectorized # ============================================================================== @njit(cache=True) def extract_features_vectorized( bid_prices: np.ndarray, bid_qtys: np.ndarray, ask_prices: np.ndarray, ask_qtys: np.ndarray, last_trade_price: float, last_trade_qty: float, funding_bps: float, volatility_state: float, pnl_bps: float, mae_bps: float, mfe_bps: float, distance_from_mfe_bps: float, seconds_held: float, time_in_loss_s: float, time_since_deep_mae_s: float, recovery_velocity_bps_per_s: float, adverse_velocity_bps_per_s: float, orderflow_toxicity: float, queue_churn_score: float, cross_venue_lead_score: float, ) -> np.ndarray: """ Extract features as flat array (numba-optimized). Returns 17-element feature vector. """ mid = 0.0 spread_bps = 0.0 if len(bid_prices) > 0 and len(ask_prices) > 0: mid = 0.5 * (bid_prices[0] + ask_prices[0]) spread = ask_prices[0] - bid_prices[0] spread_bps = 10_000.0 * spread / max(mid, 1e-12) bid_qty_sum = 0.0 for i in range(min(5, len(bid_qtys))): bid_qty_sum += bid_qtys[i] ask_qty_sum = 0.0 for i in range(min(5, len(ask_qtys))): ask_qty_sum += ask_qtys[i] imbalance = (bid_qty_sum - ask_qty_sum) / max(bid_qty_sum + ask_qty_sum, 1e-12) features = np.zeros(17, dtype=np.float64) features[0] = mid features[1] = spread_bps features[2] = imbalance features[3] = funding_bps features[4] = volatility_state features[5] = pnl_bps features[6] = mae_bps features[7] = mfe_bps features[8] = distance_from_mfe_bps features[9] = seconds_held features[10] = time_in_loss_s features[11] = time_since_deep_mae_s features[12] = recovery_velocity_bps_per_s features[13] = adverse_velocity_bps_per_s features[14] = orderflow_toxicity features[15] = queue_churn_score features[16] = cross_venue_lead_score return features # ============================================================================== # Replay comparison — vectorized # ============================================================================== @njit(cache=True) def compare_states_vectorized( expected_equity: float, actual_equity: float, expected_bid: float, actual_bid: float, expected_ask: float, actual_ask: float, tolerance_price: float, tolerance_equity: float, ) -> tuple: """ Compare two states as flat values. Returns: (match, field_index, expected_val, actual_val) field_index: -1 if match, 0=equity, 1=bid, 2=ask """ if abs(expected_equity - actual_equity) > tolerance_equity: return (False, 0, expected_equity, actual_equity) if abs(expected_bid - actual_bid) > tolerance_price: return (False, 1, expected_bid, actual_bid) if abs(expected_ask - actual_ask) > tolerance_price: return (False, 2, expected_ask, actual_ask) return (True, -1, 0.0, 0.0) # ============================================================================== # Reward computation — vectorized # ============================================================================== @njit(cache=True) def compute_reward_vectorized( pnl_bps: float, toxicity: float, churn: float, time_in_loss: float, spread_bps: float, inventory_risk: float, tail_risk: float, w_pnl: float, w_toxicity: float, w_inventory: float, w_tail: float, w_time: float, is_maker: bool, maker_fee_bps: float, is_cross: bool, taker_fee_bps: float, is_cancel: bool, adverse_threshold: float, churn_threshold: float, w_queue: float, w_adverse: float, ) -> float: """Compute reward as flat function (numba-optimized).""" reward = 0.0 reward += w_pnl * pnl_bps reward -= w_toxicity * toxicity reward -= w_inventory * inventory_risk reward -= w_tail * tail_risk reward -= w_time * math.log1p(max(time_in_loss, 0.0)) if is_maker: reward += 0.5 * max(0.0, -maker_fee_bps) if is_cross: reward -= spread_bps + max(taker_fee_bps, 0.0) if is_cancel: if toxicity > adverse_threshold: reward += w_adverse * toxicity if churn > churn_threshold: reward += w_queue * churn return reward