malkhut(T2): Code World Model — deterministic exchange simulator
CWM core (core.py): price-time priority, sequential level consumption, partial fills, queue position, latency injection, maker/taker fees. Numba acceleration (numba_core.py): JIT hot loops, 1.8x fill speedup. Replay verification (replay_verify.py): binary search, trajectory recording. Supporting: adverse_selection, correlation, latency_model, multi_level, queue_model, spread_dynamics, volatility, hftbacktest_validator.
This commit is contained in:
255
MALKHUT/malkhut/cwm/numba_core.py
Normal file
255
MALKHUT/malkhut/cwm/numba_core.py
Normal file
@@ -0,0 +1,255 @@
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user