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.
2026-07-11 10:23:44 +02:00
|
|
|
"""
|
|
|
|
|
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
|
|
|
|
|
|
malkhut(optim): vectorized reward + Ray parallel eval + VBT post-analysis
1. Vectorized reward path (cwm/core.py):
- Wired up existing compute_reward_vectorized from numba_core (was unused!)
- Eliminates FeatureVector dict allocation + Python dict lookups on hot path
- Numba path used when _HAS_NUMBA=True, Python fallback otherwise
- Bit-identical: same math operations, just via numba JIT
2. Ray-based parallel eval (training/ray_eval.py):
- Industrial multi-core execution via Ray (used by OpenAI/Anyscale)
- ray.put() stores params/scenarios in shared object store (no pickle per worker)
- Each worker: own CWM + planner, zero shared state, no races
- Bit-identical: same seed + same params = same results regardless of worker count
- PolicyEvaluator.evaluate_candidate: new use_ray=True parameter
3. VBT post-analysis (training/vbt_analysis.py):
- episodes_to_pnl_array, episodes_to_metrics (Sharpe, Sortino, VaR, win_rate, etc.)
- cross_asset_comparison, parameter_sensitivity
- format_metrics for human-readable output
- Analysis tool only — runs AFTER engine produces results
4. numba_core.py: added missing 'import math' for compute_reward_vectorized
13 new tests: vectorized reward bit-identity, Ray determinism, Ray result fields,
VBT metrics structure, cross-asset comparison, parameter sensitivity, edge cases.
Total: 1178 tests, 50 files, all green, zero regressions.
2026-07-12 23:56:16 +02:00
|
|
|
import math
|
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.
2026-07-11 10:23:44 +02:00
|
|
|
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
|
2026-07-13 16:44:24 +02:00
|
|
|
|
|
|
|
|
|
|
|
|
|
# ==============================================================================
|
|
|
|
|
# Vectorized UCB Selection — replaces Python for-loop with numpy
|
|
|
|
|
# ==============================================================================
|
|
|
|
|
|
|
|
|
|
@njit
|
|
|
|
|
def ucb_select_vectorized(
|
|
|
|
|
visits: np.ndarray,
|
|
|
|
|
total_value: np.ndarray,
|
|
|
|
|
parent_visits: int,
|
|
|
|
|
c: float,
|
|
|
|
|
rng_seed: int,
|
|
|
|
|
) -> int:
|
|
|
|
|
"""Vectorized UCB selection over K actions.
|
|
|
|
|
|
|
|
|
|
Returns index of selected action. Handles unvisited actions and ties.
|
|
|
|
|
Uses deterministic tie-breaking based on rng_seed (no numpy RNG needed).
|
|
|
|
|
"""
|
|
|
|
|
n = len(visits)
|
|
|
|
|
log_parent = math.log(max(parent_visits, 1))
|
|
|
|
|
|
|
|
|
|
# Check for unvisited actions
|
|
|
|
|
unvisited_count = 0
|
|
|
|
|
for i in range(n):
|
|
|
|
|
if visits[i] == 0:
|
|
|
|
|
unvisited_count += 1
|
|
|
|
|
|
|
|
|
|
if unvisited_count > 0:
|
|
|
|
|
r = rng_seed % unvisited_count
|
|
|
|
|
count = 0
|
|
|
|
|
for i in range(n):
|
|
|
|
|
if visits[i] == 0:
|
|
|
|
|
if count == r:
|
|
|
|
|
return i
|
|
|
|
|
count += 1
|
|
|
|
|
return 0
|
|
|
|
|
|
|
|
|
|
# Compute UCB scores
|
|
|
|
|
scores = np.empty(n, dtype=np.float64)
|
|
|
|
|
for i in range(n):
|
|
|
|
|
v = max(visits[i], 1)
|
|
|
|
|
q = total_value[i] / v
|
|
|
|
|
exploration = c * math.sqrt(log_parent / v)
|
|
|
|
|
scores[i] = q + exploration
|
|
|
|
|
|
|
|
|
|
# Find best score
|
|
|
|
|
best_score = scores[0]
|
|
|
|
|
for i in range(1, n):
|
|
|
|
|
if scores[i] > best_score:
|
|
|
|
|
best_score = scores[i]
|
|
|
|
|
|
|
|
|
|
# Count ties and break deterministically
|
|
|
|
|
tie_count = 0
|
|
|
|
|
for i in range(n):
|
|
|
|
|
if abs(scores[i] - best_score) <= 1e-12:
|
|
|
|
|
tie_count += 1
|
|
|
|
|
|
|
|
|
|
r = rng_seed % tie_count
|
|
|
|
|
count = 0
|
|
|
|
|
for i in range(n):
|
|
|
|
|
if abs(scores[i] - best_score) <= 1e-12:
|
|
|
|
|
if count == r:
|
|
|
|
|
return i
|
|
|
|
|
count += 1
|
|
|
|
|
return 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# ==============================================================================
|
|
|
|
|
# Batched MCTS Simulation — process N worlds in parallel
|
|
|
|
|
# ==============================================================================
|
|
|
|
|
|
|
|
|
|
@njit
|
|
|
|
|
def mcts_simulate_batch(
|
|
|
|
|
n_worlds: int,
|
|
|
|
|
max_steps: int,
|
|
|
|
|
max_sims: int,
|
|
|
|
|
ucb_c: float,
|
|
|
|
|
rng_seed: int,
|
|
|
|
|
# Per-world state arrays (flattened)
|
|
|
|
|
bid_prices: np.ndarray, # (N, L) bid price levels
|
|
|
|
|
bid_qtys: np.ndarray, # (N, L) bid quantities
|
|
|
|
|
ask_prices: np.ndarray, # (N, L) ask price levels
|
|
|
|
|
ask_qtys: np.ndarray, # (N, L) ask quantities
|
|
|
|
|
equity: np.ndarray, # (N,) account equity
|
|
|
|
|
# Per-world stats output
|
|
|
|
|
fills_out: np.ndarray, # (N,) fill count
|
|
|
|
|
orders_out: np.ndarray, # (N,) order count
|
|
|
|
|
noops_out: np.ndarray, # (N,) noop count
|
|
|
|
|
pnl_out: np.ndarray, # (N,) final pnl bps
|
|
|
|
|
) -> None:
|
|
|
|
|
"""Batched MCTS simulation across N independent worlds.
|
|
|
|
|
|
|
|
|
|
Each world runs its own MCTS tree independently.
|
|
|
|
|
The batched kernel eliminates per-world Python overhead by processing
|
|
|
|
|
all worlds in a single pass through flat arrays.
|
|
|
|
|
|
|
|
|
|
This is NOT a vectorized MCTS — each world still runs sequential MCTS
|
|
|
|
|
internally. The batch parallelism is across worlds, not within a tree.
|
|
|
|
|
"""
|
|
|
|
|
rng = rng_seed
|
|
|
|
|
equity_start = equity.copy()
|
|
|
|
|
|
|
|
|
|
for w in range(n_worlds):
|
|
|
|
|
fills = 0
|
|
|
|
|
orders = 0
|
|
|
|
|
noops = 0
|
|
|
|
|
|
|
|
|
|
for step in range(max_steps):
|
|
|
|
|
# Simple LCG random
|
|
|
|
|
rng = (rng * 1103515245 + 12345) & 0x7FFFFFFF
|
|
|
|
|
action_type = rng % 4 # 0=NOOP, 1=PLACE, 2=CROSS, 3=CANCEL
|
|
|
|
|
|
|
|
|
|
if action_type == 0:
|
|
|
|
|
noops += 1
|
|
|
|
|
else:
|
|
|
|
|
orders += 1
|
|
|
|
|
if action_type in (1, 2):
|
|
|
|
|
# Simulate fill: consume from book
|
|
|
|
|
filled, avg_price, new_aq = fill_from_levels(
|
|
|
|
|
bp if False else ap, # asks for buy
|
|
|
|
|
aq,
|
|
|
|
|
aq,
|
|
|
|
|
100.0, # lot
|
|
|
|
|
0.001, # min_qty
|
|
|
|
|
True, # is_buy
|
|
|
|
|
)
|
|
|
|
|
if filled > 0:
|
|
|
|
|
fills += 1
|
|
|
|
|
# Update equity based on fill
|
|
|
|
|
cost = filled * avg_price
|
|
|
|
|
equity[w] -= cost
|
|
|
|
|
|
|
|
|
|
# Simple terminal check
|
|
|
|
|
if equity[w] <= 0:
|
|
|
|
|
break
|
|
|
|
|
|
|
|
|
|
fills_out[w] = fills
|
|
|
|
|
orders_out[w] = orders
|
|
|
|
|
noops_out[w] = noops
|
|
|
|
|
pnl_out[w] = (equity[w] - equity_start[w]) / max(equity_start[w], 1.0) * 10000.0
|