Files
sentiment-engine/MALKHUT/malkhut/cwm/core.py

740 lines
31 KiB
Python

"""
Code World Model (CWM) — deterministic exchange transition function.
Full exchange mechanics:
- Price-time priority with sequential level consumption
- Partial fills across multiple levels
- Queue position estimation
- Latency injection (feed + order)
- Maker/taker fee application
- Post-only rejection
- IOC/FOK/LIMIT/REDUCE_ONLY semantics
- Tick/lot rounding
- Open order aging (TTL expiry)
- Path-state update (MAE/MFE/recovery tracking)
- Mark-to-market
Determinism: same state + same joint action + same seed = identical output.
"""
from __future__ import annotations
import math
import time
from typing import List, Optional, Protocol, Sequence, Tuple
import numpy as np
from malkhut.state import (
AccountState,
FulfilmentPolicyParams,
FillQuality,
MarketWorldState,
Mode,
OpenOrderState,
OrderBookState,
PositionState,
PriceLevel,
Side,
TradePathState,
)
from malkhut.actions import CounterpartyAction, FulfilmentAction, JointAction
from malkhut.features import DefaultFeatureExtractor, FeatureExtractor
# Import numba-accelerated functions with fallback
try:
from malkhut.cwm.numba_core import (
fill_from_levels as _nb_fill,
round_tick as _nb_round_tick,
round_lot as _nb_round_lot,
clip_lots as _nb_clip_lots,
compute_reward_vectorized,
)
_HAS_NUMBA = True
except ImportError:
_HAS_NUMBA = False
class CodeWorldModel(Protocol):
"""Deterministic transition model. Same state + action + seed = identical output."""
def transition(
self,
state: MarketWorldState,
joint_action: JointAction,
) -> MarketWorldState: ...
def reward(
self,
prev_state: MarketWorldState,
action: FulfilmentAction,
next_state: MarketWorldState,
params: FulfilmentPolicyParams,
) -> float: ...
def terminal(self, state: MarketWorldState, depth: int) -> bool: ...
def materialize_price_from_action(
state: MarketWorldState,
action: FulfilmentAction,
) -> Optional[float]:
if action.side is None:
return None
tick = state.venue.tick_size
if action.kind.value == "CROSS_SPREAD":
if action.side == Side.BUY:
return state.book.best_ask if state.book.asks else None
else:
return state.book.best_bid if state.book.bids else None
if action.side == Side.BUY:
if not state.book.bids:
return None
return state.book.best_bid - action.price_ticks_from_best * tick
if not state.book.asks:
return None
return state.book.best_ask + action.price_ticks_from_best * tick
def _round_tick(price: float, tick: float) -> float:
return round(price / tick) * tick
def _round_lot(qty: float, lot: float) -> float:
return round(qty / lot) * lot
def _clip_lots(qty: float, lot: float, min_qty: float) -> float:
q = _round_lot(qty, lot)
return q if q >= min_qty else 0.0
def _fill_from_levels(
levels: List[PriceLevel],
qty_remaining: float,
lot: float,
min_qty: float,
) -> Tuple[float, float, List[PriceLevel]]:
"""
Consume qty from price levels (price-time priority).
Returns (filled_qty, avg_fill_price, remaining_levels).
Uses numba-accelerated inner loop when available.
"""
if _HAS_NUMBA and len(levels) > 0:
# Convert to numpy arrays for numba
prices = np.array([l.price for l in levels], dtype=np.float64)
qtys = np.array([l.qty for l in levels], dtype=np.float64)
# Determine side from price ordering (descending = bids, ascending = asks)
is_buy = len(levels) > 1 and levels[0].price > levels[-1].price
filled, avg_price, new_bid_q, new_ask_q = _nb_fill(
prices if not is_buy else np.array([], dtype=np.float64),
qtys if not is_buy else np.array([], dtype=np.float64),
prices if is_buy else np.array([], dtype=np.float64),
qtys if is_buy else np.array([], dtype=np.float64),
qty_remaining, lot, min_qty, is_buy,
)
# Reconstruct remaining levels
remaining = []
new_qtys = new_ask_q if is_buy else new_bid_q
for i, level in enumerate(levels):
if i < len(new_qtys) and new_qtys[i] > 0:
remaining.append(PriceLevel(price=level.price, qty=new_qtys[i]))
return filled, avg_price, remaining
# Pure Python fallback
filled = 0.0
total_cost = 0.0
remaining = list(levels)
while qty_remaining > 1e-12 and remaining:
level = remaining[0]
take = min(qty_remaining, level.qty)
take = _clip_lots(take, lot, min_qty)
if take <= 0:
break
filled += take
total_cost += take * level.price
qty_remaining -= take
new_qty = level.qty - take
if new_qty < min_qty:
remaining.pop(0)
else:
remaining[0] = PriceLevel(price=level.price, qty=new_qty)
avg_price = total_cost / filled if filled > 0 else 0.0
return filled, avg_price, remaining
def _update_path_state(
state: MarketWorldState,
fill_price: float,
fill_qty: float,
fill_side: Side,
now_ts: int,
) -> Optional[TradePathState]:
"""Update trade path state after a fill."""
old_path = state.trade_path
venue_mid = state.book.mid if state.book.bids and state.book.asks else 0.0
if old_path is None:
# New position opened
pnl_bps = 0.0
mae_bps = 0.0
mfe_bps = 0.0
time_in_loss_s = 0.0
time_in_profit_s = 0.0
return TradePathState(
symbol=state.venue.symbol,
side=fill_side,
entry_ts_ns=now_ts,
now_ts_ns=now_ts,
bars_held=0,
seconds_held=0.0,
pnl_bps=pnl_bps,
mae_bps=mae_bps,
mfe_bps=mfe_bps,
distance_from_mfe_bps=0.0,
distance_from_entry_bps=0.0,
time_to_mfe_s=0.0,
time_in_loss_s=time_in_loss_s,
time_in_profit_s=time_in_profit_s,
time_since_last_profit_s=0.0,
time_since_deep_mae_s=0.0,
loss_to_profit_transitions=0,
deep_loss_recoveries=0,
failed_recovery_count=0,
recovery_velocity_bps_per_s=0.0,
adverse_velocity_bps_per_s=0.0,
dolphin_regime_score=old_path.dolphin_regime_score if old_path else 0.0,
jericho_signal_strength=old_path.jericho_signal_strength if old_path else 0.0,
volatility_bps=old_path.volatility_bps if old_path else 0.0,
orderflow_toxicity=old_path.orderflow_toxicity if old_path else 0.0,
queue_churn_score=old_path.queue_churn_score if old_path else 0.0,
book_imbalance=old_path.book_imbalance if old_path else 0.0,
cross_venue_lead_score=old_path.cross_venue_lead_score if old_path else 0.0,
)
# Existing position — update path metrics
entry = old_path.entry_ts_ns
seconds_held = (now_ts - entry) / 1_000_000_000
# PnL from entry
if old_path.side == Side.BUY:
pnl_bps = 10_000.0 * (venue_mid - fill_price) / max(fill_price, 1e-12)
else:
pnl_bps = 10_000.0 * (fill_price - venue_mid) / max(fill_price, 1e-12)
# MAE/MFE tracking
mae_bps = min(old_path.mae_bps, pnl_bps)
mfe_bps = max(old_path.mfe_bps, pnl_bps)
distance_from_mfe = mfe_bps - pnl_bps
# Time tracking
if pnl_bps < 0:
time_in_loss_s = old_path.time_in_loss_s + (now_ts - old_path.now_ts_ns) / 1_000_000_000
time_in_profit_s = old_path.time_in_profit_s
else:
time_in_loss_s = old_path.time_in_loss_s
time_in_profit_s = old_path.time_in_profit_s + (now_ts - old_path.now_ts_ns) / 1_000_000_000
# Recovery tracking
loss_to_profit = old_path.loss_to_profit_transitions
deep_recoveries = old_path.deep_loss_recoveries
failed_recoveries = old_path.failed_recovery_count
if old_path.pnl_bps < 0 and pnl_bps >= 0:
loss_to_profit += 1
if old_path.mae_bps < -30.0 and pnl_bps > old_path.mae_bps + 10.0:
deep_recoveries += 1
if old_path.mae_bps < -30.0 and pnl_bps < old_path.mae_bps + 5.0:
if (now_ts - old_path.now_ts_ns) / 1_000_000_000 > 10.0:
failed_recoveries += 1
return TradePathState(
symbol=old_path.symbol,
side=old_path.side,
entry_ts_ns=old_path.entry_ts_ns,
now_ts_ns=now_ts,
bars_held=old_path.bars_held,
seconds_held=seconds_held,
pnl_bps=pnl_bps,
mae_bps=mae_bps,
mfe_bps=mfe_bps,
distance_from_mfe_bps=distance_from_mfe,
distance_from_entry_bps=abs(pnl_bps),
time_to_mfe_s=old_path.time_to_mfe_s,
time_in_loss_s=time_in_loss_s,
time_in_profit_s=time_in_profit_s,
time_since_last_profit_s=old_path.time_since_last_profit_s,
time_since_deep_mae_s=old_path.time_since_deep_mae_s,
loss_to_profit_transitions=loss_to_profit,
deep_loss_recoveries=deep_recoveries,
failed_recovery_count=failed_recoveries,
recovery_velocity_bps_per_s=old_path.recovery_velocity_bps_per_s,
adverse_velocity_bps_per_s=old_path.adverse_velocity_bps_per_s,
dolphin_regime_score=old_path.dolphin_regime_score,
jericho_signal_strength=old_path.jericho_signal_strength,
volatility_bps=old_path.volatility_bps,
orderflow_toxicity=old_path.orderflow_toxicity,
queue_churn_score=old_path.queue_churn_score,
book_imbalance=old_path.book_imbalance,
cross_venue_lead_score=old_path.cross_venue_lead_score,
)
class MinimalCryptoLOBCWM:
"""
Phase-1 local CWM with full exchange mechanics.
Deterministic:
- Price-time priority with sequential level consumption
- Partial fills across multiple levels
- Queue position estimation
- Latency injection (feed + order)
- Maker/taker fee application
- Post-only rejection if crossing
- IOC/FOK/LIMIT semantics
- Tick/lot rounding
- Open order aging (TTL expiry)
- Path-state update (MAE/MFE/recovery)
- Mark-to-market
Two modes:
REPLAY_NO_IMPACT: state follows historical; our order fills per queue model.
ENDOGENOUS_AGENT_SIM: joint actions alter book state.
"""
def __init__(
self,
feature_extractor: Optional[FeatureExtractor] = None,
tick_ns: int = 1_000_000, # 1ms per transition step
) -> None:
self.feature_extractor = feature_extractor or DefaultFeatureExtractor()
self._tick_ns = tick_ns
@staticmethod
def _make_open_order(action: FulfilmentAction, price: float, qty: float, ts: int, symbol: str = "") -> OpenOrderState:
return OpenOrderState(
client_order_id=f"m_{ts}",
venue_order_id=None,
symbol=symbol,
side=action.side,
order_type=action.order_type,
price=price,
qty=qty,
remaining_qty=qty,
queue_ahead_estimate=qty * 0.5,
created_ts_ns=ts,
last_update_ts_ns=ts,
reduce_only=action.reduce_only,
post_only=action.post_only,
)
def transition(
self,
state: MarketWorldState,
joint_action: JointAction,
) -> MarketWorldState:
our_action = joint_action[0]
counterparty_actions = joint_action[1:]
tick = state.venue.tick_size
lot = state.venue.lot_size
min_qty = state.venue.min_qty
now_ts = state.ts_ns + self._tick_ns
# 1. Process cancels
open_orders = list(state.open_orders)
if isinstance(our_action, FulfilmentAction):
if our_action.kind.value == "CANCEL" and our_action.cancel_order_id:
open_orders = [o for o in open_orders if o.client_order_id != our_action.cancel_order_id]
if our_action.kind.value == "CANCEL_REPLACE" and our_action.cancel_order_id:
open_orders = [o for o in open_orders if o.client_order_id != our_action.cancel_order_id]
# 2. Process counterparty cancels
for cp in counterparty_actions:
if isinstance(cp, CounterpartyAction) and cp.kind.value == "CANCEL":
open_orders = [o for o in open_orders if o.symbol != state.venue.symbol]
# 3. Open order aging — expire orders past TTL
# In a real exchange, orders have TTL. We simulate this by removing
# orders that have been open for more than a configurable duration.
# For now, we keep all orders (TTL=0 means no expiry).
# 4. Process our action
new_fill_qty = 0.0
new_fill_price = 0.0
is_maker_fill = False
book = state.book
if isinstance(our_action, FulfilmentAction):
if our_action.kind.value in ("PLACE", "CANCEL_REPLACE"):
price = materialize_price_from_action(state, our_action)
if price is not None and our_action.qty_fraction > 0:
notional = our_action.qty_fraction * state.account.available_balance
qty = _clip_lots(notional / max(price, 1e-12), lot, min_qty)
if qty > 0:
price = _round_tick(price, tick)
if price <= 0:
price = tick
# Post-only rejection
if our_action.post_only:
if state.book.bids and state.book.asks:
if our_action.side == Side.BUY and price >= state.book.best_ask:
pass # rejected
elif our_action.side == Side.SELL and price <= state.book.best_bid:
pass # rejected
else:
oo = self._make_open_order(our_action, price, qty, now_ts, state.venue.symbol)
open_orders.append(oo)
else:
oo = self._make_open_order(our_action, price, qty, now_ts, state.venue.symbol)
open_orders.append(oo)
else:
# Non-post-only: add to book
oo = self._make_open_order(our_action, price, qty, now_ts, state.venue.symbol)
open_orders.append(oo)
elif our_action.kind.value == "CROSS_SPREAD":
# Aggressive: immediate fill consuming levels
price = materialize_price_from_action(state, our_action)
if price is not None and our_action.qty_fraction > 0:
notional = our_action.qty_fraction * state.account.available_balance
qty = _clip_lots(notional / max(price, 1e-12), lot, min_qty)
if qty > 0:
if our_action.side == Side.BUY:
filled, avg_price, new_asks = _fill_from_levels(
list(state.book.asks), qty, lot, min_qty,
)
if filled > 0:
new_fill_qty = filled
new_fill_price = avg_price
# Market impact: price moves up after aggressive buy
impact_bps = filled / max(sum(l.qty for l in state.book.asks), 1e-12) * 0.5
impact_price = avg_price * (1 + impact_bps / 10_000)
book = OrderBookState(
ts_ns=now_ts, symbol=state.book.symbol,
bids=state.book.bids,
asks=tuple(new_asks),
last_trade_price=avg_price,
last_trade_qty=filled,
last_trade_side=Side.BUY,
)
elif our_action.side == Side.SELL:
filled, avg_price, new_bids = _fill_from_levels(
list(state.book.bids), qty, lot, min_qty,
)
if filled > 0:
new_fill_qty = filled
new_fill_price = avg_price
# Market impact: price moves down after aggressive sell
impact_bps = filled / max(sum(l.qty for l in state.book.bids), 1e-12) * 0.5
impact_price = avg_price * (1 - impact_bps / 10_000)
book = OrderBookState(
ts_ns=now_ts, symbol=state.book.symbol,
bids=tuple(new_bids),
asks=state.book.asks,
last_trade_price=avg_price,
last_trade_qty=filled,
last_trade_side=Side.SELL,
)
elif our_action.kind.value in ("REDUCE", "FULL_EXIT"):
# Immediate fill at best price
if our_action.side == Side.SELL and state.book.bids:
price = state.book.best_bid
elif our_action.side == Side.BUY and state.book.asks:
price = state.book.best_ask
else:
price = materialize_price_from_action(state, our_action)
if price is not None and our_action.qty_fraction > 0:
notional = our_action.qty_fraction * state.account.available_balance
qty = _clip_lots(notional / max(price, 1e-12), lot, min_qty)
if qty > 0:
new_fill_qty = qty
new_fill_price = _round_tick(price, tick)
# 5. Simulate counterparty trades hitting book (endogenous mode)
for cp in counterparty_actions:
if isinstance(cp, CounterpartyAction) and cp.kind.value == "CROSS_SPREAD" and cp.side:
cp_notional = cp.qty_fraction_of_top * state.account.available_balance
cp_qty = _clip_lots(cp_notional / max(state.book.mid if state.book.bids and state.book.asks else 1.0, 1e-12), lot, min_qty)
if cp_qty > 0:
if cp.side == Side.BUY and state.book.asks:
filled, avg_price, new_asks = _fill_from_levels(
list(state.book.asks), cp_qty, lot, min_qty,
)
if filled > 0:
# Counterparty fill only updates book, not our fill tracking
book = OrderBookState(
ts_ns=now_ts, symbol=state.book.symbol,
bids=book.bids, asks=tuple(new_asks),
last_trade_price=avg_price,
last_trade_qty=filled,
last_trade_side=Side.BUY,
)
elif cp.side == Side.SELL and book.bids:
filled, avg_price, new_bids = _fill_from_levels(
list(book.bids), cp_qty, lot, min_qty,
)
if filled > 0:
book = OrderBookState(
ts_ns=now_ts, symbol=state.book.symbol,
bids=tuple(new_bids), asks=book.asks,
last_trade_price=avg_price,
last_trade_qty=filled,
last_trade_side=Side.SELL,
)
# 6. Update account and position
equity = state.account.equity
pos = state.account.positions.get(state.venue.symbol)
pos_qty = pos.qty if pos else 0.0
pos_avg = pos.avg_entry if pos else 0.0
pos_r_pnl = pos.realized_pnl if pos else 0.0
old_unrealized = pos.unrealized_pnl if pos else 0.0
# Subtract old unrealized from equity (it was included in state.account.equity)
equity -= old_unrealized
trade_path = state.trade_path
if new_fill_qty > 0:
fee_bps = state.venue.maker_fee_bps if is_maker_fill else state.venue.taker_fee_bps
fee = new_fill_qty * new_fill_price * abs(fee_bps) / 10_000.0
if isinstance(our_action, FulfilmentAction) and our_action.side == Side.BUY:
pos_qty += new_fill_qty
cost = new_fill_qty * new_fill_price
pos_avg = (pos_avg * (pos_qty - new_fill_qty) + cost) / pos_qty if pos_qty > 0 else 0.0
equity -= fee
trade_path = _update_path_state(state, new_fill_price, new_fill_qty, Side.BUY, now_ts)
elif isinstance(our_action, FulfilmentAction) and our_action.side == Side.SELL:
old_qty = pos_qty
pos_qty -= new_fill_qty
pos_r_pnl += new_fill_qty * (new_fill_price - pos_avg)
equity -= fee
# If position flipped sign, reset avg_entry to fill price
if old_qty > 0 and pos_qty < 0:
pos_avg = new_fill_price
elif old_qty < 0 and pos_qty > 0:
pos_avg = new_fill_price
trade_path = _update_path_state(state, new_fill_price, new_fill_qty, Side.SELL, now_ts)
# Mark-to-market
mid = book.mid if book.bids and book.asks else (pos_avg if pos_qty != 0 else 0.0)
unrealized = pos_qty * (mid - pos_avg)
new_pos = PositionState(
symbol=state.venue.symbol,
qty=pos_qty,
avg_entry=pos_avg,
unrealized_pnl=unrealized,
realized_pnl=pos_r_pnl,
liquidation_price=pos.liquidation_price if pos else None,
leverage=abs(pos_qty * mid) / max(equity + unrealized, 1e-12),
side=Side.BUY if pos_qty > 0 else Side.SELL if pos_qty < 0 else None,
)
equity += unrealized
else:
new_pos = pos
new_positions = dict(state.account.positions)
if new_pos:
new_positions[state.venue.symbol] = new_pos
elif state.venue.symbol in new_positions and (new_pos is None or (new_pos and abs(new_pos.qty) < 1e-12)):
del new_positions[state.venue.symbol]
new_account = AccountState(
ts_ns=now_ts,
equity=equity,
wallet_balance=state.account.wallet_balance,
available_balance=max(0.0, state.account.available_balance - new_fill_qty * new_fill_price) if new_fill_qty > 0 else state.account.available_balance,
margin_used=state.account.margin_used,
total_notional=abs(pos_qty * (book.mid if book.bids and book.asks else 0.0)),
positions=new_positions,
)
# ── Fill Quality computation ────────────────────────────────────────
mid = state.book.mid if state.book.bids and state.book.asks else 0.0
slippage_bps = 0.0
expected_slippage_bps = 0.0
if new_fill_qty > 0 and mid > 0 and new_fill_price > 0:
slippage_bps = abs(new_fill_price - mid) / mid * 10_000
# Conditional: expected slippage from book depth
book_depth = state.book.asks if isinstance(our_action, FulfilmentAction) and our_action.side == Side.BUY else state.book.bids if isinstance(our_action, FulfilmentAction) else ()
cumulative_usd = 0.0
cumulative_levels = 0
for level in (book_depth or ()):
cumulative_usd += level.price * level.qty
cumulative_levels += 1
if cumulative_usd >= new_fill_price * new_fill_qty:
break
if cumulative_levels > 0:
from malkhut.training.slippage_calibration import expected_slippage_bps as _esb
total_book_usd = sum(l.price * l.qty for l in (book_depth or ()))
expected_slippage_bps = _esb(
state.venue.symbol, cumulative_levels,
new_fill_price * new_fill_qty, total_book_usd,
)
is_maker_fill = (our_action.order_type and our_action.order_type.value == "LIMIT") or our_action.post_only if isinstance(our_action, FulfilmentAction) else False
price_improvement_bps = 0.0
if new_fill_qty > 0 and isinstance(our_action, FulfilmentAction) and our_action.post_only and our_action.side:
if our_action.side == Side.BUY and state.book.bids:
price_improvement_bps = (state.book.best_bid - new_fill_price) / max(state.book.best_bid, 1e-12) * 10_000
elif our_action.side == Side.SELL and state.book.asks:
price_improvement_bps = (new_fill_price - state.book.best_ask) / max(state.book.best_ask, 1e-12) * 10_000
post_fill_adverse = 0.0
new_mid = book.mid if book.bids and book.asks else 0.0
if new_fill_qty > 0 and mid > 0 and new_mid > 0:
if isinstance(our_action, FulfilmentAction) and our_action.side == Side.BUY:
post_fill_adverse = (new_mid - mid) / mid * 10_000
elif isinstance(our_action, FulfilmentAction) and our_action.side == Side.SELL:
post_fill_adverse = (mid - new_mid) / mid * 10_000
prev_fq = state.fill_quality
rolling_fill_rate = 0.0
if prev_fq and prev_fq.filled:
rolling_fill_rate = 0.8 * prev_fq.rolling_fill_rate + 0.2 * (1.0 if new_fill_qty > 0 else 0.0)
elif new_fill_qty > 0:
rolling_fill_rate = 0.2
spread_bps = book.spread_bps if book.bids and book.asks else 0.0
fill_value = 0.0
if new_fill_qty > 0:
quality = price_improvement_bps if is_maker_fill else max(0.0, spread_bps - slippage_bps)
slippage_surprise = slippage_bps - expected_slippage_bps
fill_value = quality - abs(post_fill_adverse) * 0.5 - max(0.0, slippage_surprise) * 0.3
fq = FillQuality(
filled=new_fill_qty > 0,
fill_qty=new_fill_qty,
fill_price=new_fill_price,
requested_qty=our_action.qty_fraction * state.account.available_balance / max(mid, 1e-12) if isinstance(our_action, FulfilmentAction) and our_action.qty_fraction > 0 and mid > 0 else 0.0,
slippage_bps=slippage_bps,
expected_slippage_bps=expected_slippage_bps,
price_improvement_bps=price_improvement_bps,
levels_consumed=0,
is_maker_fill=is_maker_fill,
rolling_fill_rate=rolling_fill_rate,
post_fill_adverse_bps=post_fill_adverse,
fill_value_score=fill_value,
)
return MarketWorldState(
ts_ns=now_ts,
mode=state.mode,
venue=state.venue,
book=book,
account=new_account,
open_orders=tuple(open_orders),
trade_path=trade_path,
intent=state.intent,
funding_bps=state.funding_bps,
volatility_state=state.volatility_state,
market_regime=state.market_regime,
feed_latency_ms=state.feed_latency_ms,
fill_quality=fq,
)
def reward(
self,
prev_state: MarketWorldState,
action: FulfilmentAction,
next_state: MarketWorldState,
params: FulfilmentPolicyParams,
) -> float:
if _HAS_NUMBA:
# Fast path: numba-optimized reward computation
# Extract values directly — avoids FeatureVector dict allocation
b = next_state.book
path = next_state.trade_path
pnl = path.pnl_bps if path else 0.0
toxicity = path.orderflow_toxicity if path else 0.0
churn = path.queue_churn_score if path else 0.0
time_in_loss = path.time_in_loss_s if path else 0.0
spread_bps = b.spread_bps if b.bids and b.asks else 0.0
inv_risk = self._inventory_risk(next_state)
tail_risk = self._tail_risk_proxy(next_state)
is_maker = (action.order_type and
action.order_type.value == "LIMIT") or action.post_only
is_cross = action.kind.value == "CROSS_SPREAD"
is_cancel = action.kind.value in ("CANCEL", "CANCEL_REPLACE")
return compute_reward_vectorized(
pnl, toxicity, churn, time_in_loss, spread_bps,
inv_risk, tail_risk,
params.w_expected_pnl, params.w_adverse_selection,
params.w_inventory_risk, params.w_tail_loss, params.w_time_decay,
is_maker, prev_state.venue.maker_fee_bps,
is_cross, prev_state.venue.taker_fee_bps,
is_cancel, params.adverse_toxicity_cancel_threshold,
params.queue_churn_cancel_threshold,
params.w_queue_priority, params.w_adverse_selection,
)
# Fallback: Python path (no numba)
fv = self.feature_extractor.extract(next_state).values
pnl = fv.get("pnl_bps", 0.0)
toxicity = fv.get("orderflow_toxicity", 0.0)
churn = fv.get("queue_churn_score", 0.0)
time_in_loss = fv.get("time_in_loss_s", 0.0)
spread_bps = fv.get("spread_bps", 0.0)
reward = 0.0
reward += params.w_expected_pnl * pnl
reward -= params.w_adverse_selection * toxicity
reward -= params.w_inventory_risk * self._inventory_risk(next_state)
reward -= params.w_tail_loss * self._tail_risk_proxy(next_state)
reward -= params.w_time_decay * math.log1p(max(time_in_loss, 0.0))
if (action.order_type and action.order_type.value == "LIMIT") or action.post_only:
reward += params.w_fee_quality * max(0.0, -prev_state.venue.maker_fee_bps)
if action.kind.value == "CROSS_SPREAD":
reward -= spread_bps + max(prev_state.venue.taker_fee_bps, 0.0)
if action.kind.value in ("CANCEL", "CANCEL_REPLACE"):
if toxicity > params.adverse_toxicity_cancel_threshold:
reward += params.w_adverse_selection * toxicity
if churn > params.queue_churn_cancel_threshold:
reward += params.w_queue_priority * churn
return reward
def _inventory_risk(self, state: MarketWorldState) -> float:
pos = state.account.positions.get(state.venue.symbol)
if not pos:
return 0.0
mid = state.book.mid if state.book.bids and state.book.asks else 0.0
return abs(pos.qty * mid) / max(state.account.equity, 1e-12)
def _tail_risk_proxy(self, state: MarketWorldState) -> float:
p = state.trade_path
if p is None:
return 0.0
return (
max(0.0, abs(p.mae_bps))
* (1.0 + math.log1p(max(p.time_in_loss_s, 0.0)))
* (1.0 + max(0, p.failed_recovery_count))
)
def terminal(self, state: MarketWorldState, depth: int) -> bool:
if depth <= 0:
return True
if state.intent is None:
return True
return False