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:
461
MALKHUT/malkhut/cwm/replay_verify.py
Normal file
461
MALKHUT/malkhut/cwm/replay_verify.py
Normal file
@@ -0,0 +1,461 @@
|
||||
"""
|
||||
Replay verification — mandatory before trusting CWM.
|
||||
|
||||
Three verification modes:
|
||||
1. Historical replay: venue data → ReplayStep → CWM transition → compare
|
||||
2. Self-play replay: persist trajectory → deterministic re-run → exact match
|
||||
3. hftbacktest comparison: CWM vs known replay engine for queue/fill validation
|
||||
|
||||
Design rule from spec:
|
||||
"Replay correctness before search depth.
|
||||
A wrong CWM plus deep search creates confident nonsense."
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable, List, Optional, Protocol, Sequence, Tuple
|
||||
|
||||
from malkhut.state import (
|
||||
AccountState, MarketWorldState, Mode, OpenOrderState, OrderBookState,
|
||||
PositionState, PriceLevel, Side, VenueRules,
|
||||
)
|
||||
from malkhut.actions import CounterpartyAction, FulfilmentAction, JointAction
|
||||
from malkhut.cwm.core import CodeWorldModel
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Data types
|
||||
# ==============================================================================
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReplayStep:
|
||||
"""One step in a replay trajectory."""
|
||||
before: MarketWorldState
|
||||
joint_action: JointAction
|
||||
after_ground_truth: MarketWorldState
|
||||
step_index: int = 0
|
||||
metadata: Mapping[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReplayMismatch:
|
||||
"""One field mismatch between predicted and ground truth."""
|
||||
index: int
|
||||
field: str
|
||||
expected: Any
|
||||
actual: Any
|
||||
severity: str # "critical", "warning", "info"
|
||||
tolerance: float = 0.0
|
||||
|
||||
@property
|
||||
def is_critical(self) -> bool:
|
||||
return self.severity == "critical"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReplayResult:
|
||||
"""Complete result of a replay verification run."""
|
||||
passed: bool
|
||||
mismatches: List[ReplayMismatch]
|
||||
steps_verified: int
|
||||
total_steps: int
|
||||
first_mismatch_index: Optional[int]
|
||||
duration_ns: int
|
||||
trajectory_hash: str
|
||||
|
||||
@property
|
||||
def match_rate(self) -> float:
|
||||
return self.steps_verified / max(self.total_steps, 1)
|
||||
|
||||
@property
|
||||
def critical_count(self) -> int:
|
||||
return sum(1 for m in self.mismatches if m.is_critical)
|
||||
|
||||
@property
|
||||
def warning_count(self) -> int:
|
||||
return sum(1 for m in self.mismatches if m.severity == "warning")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TrajectoryRecord:
|
||||
"""One step in a persisted trajectory for self-play verification."""
|
||||
step_index: int
|
||||
before_hash: str
|
||||
action_hash: str
|
||||
after_hash: str
|
||||
ts_ns: int
|
||||
symbol: str
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Deep state comparison
|
||||
# ==============================================================================
|
||||
|
||||
def _hash_state(state: MarketWorldState) -> str:
|
||||
"""Deterministic hash of a MarketWorldState for trajectory recording."""
|
||||
parts = [
|
||||
str(state.ts_ns),
|
||||
state.venue.symbol,
|
||||
str(state.book.best_bid) if state.book.bids else "0",
|
||||
str(state.book.best_ask) if state.book.asks else "0",
|
||||
str(state.account.equity),
|
||||
str(len(state.open_orders)),
|
||||
]
|
||||
return hashlib.sha256(":".join(parts).encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
def _hash_action(action: Any) -> str:
|
||||
"""Deterministic hash of an action."""
|
||||
return hashlib.sha256(str(action).encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
def _compare_deep(
|
||||
i: int,
|
||||
expected: MarketWorldState,
|
||||
actual: MarketWorldState,
|
||||
tolerances: Optional[Mapping[str, float]] = None,
|
||||
) -> List[ReplayMismatch]:
|
||||
"""
|
||||
Deep comparison of two MarketWorldStates.
|
||||
|
||||
Compares all fields with appropriate tolerances:
|
||||
- ts_ns: exact match
|
||||
- venue: exact match
|
||||
- book prices: float tolerance (default 1e-6)
|
||||
- book quantities: float tolerance
|
||||
- account equity: float tolerance
|
||||
- open orders: count + individual comparison
|
||||
- positions: per-symbol comparison
|
||||
"""
|
||||
tol = tolerances or {}
|
||||
diffs: List[ReplayMismatch] = []
|
||||
|
||||
def _cmp(field: str, exp_val: Any, act_val: Any, tolerance: float = 1e-9) -> None:
|
||||
if isinstance(exp_val, float):
|
||||
if abs(exp_val - act_val) > tolerance:
|
||||
diffs.append(ReplayMismatch(i, field, exp_val, act_val, "warning", tolerance))
|
||||
elif isinstance(exp_val, int):
|
||||
if exp_val != act_val:
|
||||
diffs.append(ReplayMismatch(i, field, exp_val, act_val, "info"))
|
||||
elif exp_val != act_val:
|
||||
diffs.append(ReplayMismatch(i, field, str(exp_val), str(act_val), "info"))
|
||||
|
||||
# Timestamp
|
||||
_cmp("ts_ns", expected.ts_ns, actual.ts_ns)
|
||||
|
||||
# Venue
|
||||
_cmp("venue.symbol", expected.venue.symbol, actual.venue.symbol)
|
||||
_cmp("venue.exchange", expected.venue.exchange, actual.venue.exchange)
|
||||
_cmp("venue.tick_size", expected.venue.tick_size, actual.venue.tick_size)
|
||||
|
||||
# Book
|
||||
if expected.book and actual.book:
|
||||
book_tol = tol.get("book_price", 1e-6)
|
||||
_cmp("book.best_bid", expected.book.best_bid, actual.book.best_bid, book_tol)
|
||||
_cmp("book.best_ask", expected.book.best_ask, actual.book.best_ask, book_tol)
|
||||
_cmp("book.bid_depth", len(expected.book.bids), len(actual.book.bids))
|
||||
_cmp("book.ask_depth", len(expected.book.asks), len(actual.book.asks))
|
||||
|
||||
# Compare top N levels
|
||||
for n in range(min(5, len(expected.book.bids), len(actual.book.bids))):
|
||||
_cmp(f"book.bid[{n}].price", expected.book.bids[n].price, actual.book.bids[n].price, book_tol)
|
||||
_cmp(f"book.bid[{n}].qty", expected.book.bids[n].qty, actual.book.bids[n].qty, book_tol)
|
||||
for n in range(min(5, len(expected.book.asks), len(actual.book.asks))):
|
||||
_cmp(f"book.ask[{n}].price", expected.book.asks[n].price, actual.book.asks[n].price, book_tol)
|
||||
_cmp(f"book.ask[{n}].qty", expected.book.asks[n].qty, actual.book.asks[n].qty, book_tol)
|
||||
|
||||
# Account
|
||||
if expected.account and actual.account:
|
||||
acct_tol = tol.get("account_equity", 1e-6)
|
||||
_cmp("account.equity", expected.account.equity, actual.account.equity, acct_tol)
|
||||
_cmp("account.wallet_balance", expected.account.wallet_balance, actual.account.wallet_balance, acct_tol)
|
||||
_cmp("account.available_balance", expected.account.available_balance, actual.account.available_balance, acct_tol)
|
||||
_cmp("account.total_notional", expected.account.total_notional, actual.account.total_notional, acct_tol)
|
||||
|
||||
# Open orders
|
||||
_cmp("open_orders.count", len(expected.open_orders), len(actual.open_orders))
|
||||
for n in range(min(len(expected.open_orders), len(actual.open_orders))):
|
||||
eo = expected.open_orders[n]
|
||||
ao = actual.open_orders[n]
|
||||
_cmp(f"open_orders[{n}].price", eo.price, ao.price, tol.get("order_price", 1e-6))
|
||||
_cmp(f"open_orders[{n}].qty", eo.qty, ao.qty, tol.get("order_qty", 1e-9))
|
||||
_cmp(f"open_orders[{n}].side", eo.side.value, ao.side.value)
|
||||
|
||||
# Positions
|
||||
exp_pos = expected.account.positions if expected.account else {}
|
||||
act_pos = actual.account.positions if actual.account else {}
|
||||
_cmp("positions.count", len(exp_pos), len(act_pos))
|
||||
for sym in set(list(exp_pos.keys()) + list(act_pos.keys())):
|
||||
ep = exp_pos.get(sym)
|
||||
ap = act_pos.get(sym)
|
||||
if ep and ap:
|
||||
pos_tol = tol.get("position_qty", 1e-9)
|
||||
_cmp(f"positions[{sym}].qty", ep.qty, ap.qty, pos_tol)
|
||||
_cmp(f"positions[{sym}].avg_entry", ep.avg_entry, ap.avg_entry, pos_tol)
|
||||
_cmp(f"positions[{sym}].side", ep.side.value if ep.side else None, ap.side.value if ap.side else None)
|
||||
elif ep and not ap:
|
||||
diffs.append(ReplayMismatch(i, f"positions[{sym}]", "present", "missing", "critical"))
|
||||
elif not ep and ap:
|
||||
diffs.append(ReplayMismatch(i, f"positions[{sym}]", "missing", "present", "critical"))
|
||||
|
||||
# Trade path
|
||||
if expected.trade_path and actual.trade_path:
|
||||
ep = expected.trade_path
|
||||
ap = actual.trade_path
|
||||
_cmp("trade_path.pnl_bps", ep.pnl_bps, ap.pnl_bps, tol.get("pnl_bps", 0.1))
|
||||
_cmp("trade_path.mae_bps", ep.mae_bps, ap.mae_bps, tol.get("mae_bps", 0.1))
|
||||
_cmp("trade_path.mfe_bps", ep.mfe_bps, ap.mfe_bps, tol.get("mfe_bps", 0.1))
|
||||
|
||||
return diffs
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Binary search for first mismatch
|
||||
# ==============================================================================
|
||||
|
||||
def bisect_first_mismatch(
|
||||
cwm: CodeWorldModel,
|
||||
replay: Sequence[ReplayStep],
|
||||
lo: int = 0,
|
||||
hi: Optional[int] = None,
|
||||
tolerances: Optional[Mapping[str, float]] = None,
|
||||
) -> Optional[ReplayMismatch]:
|
||||
"""
|
||||
Binary search for the first mismatch in a replay trajectory.
|
||||
|
||||
Uses the CWM to re-simulate from known-good prefix, narrowing to the
|
||||
first divergence point. Much faster than linear scan for long trajectories.
|
||||
"""
|
||||
if hi is None:
|
||||
hi = len(replay) - 1
|
||||
|
||||
if lo > hi:
|
||||
return None
|
||||
|
||||
# Find any mismatch in the range
|
||||
mid = (lo + hi) // 2
|
||||
mismatches = _compare_deep(
|
||||
mid,
|
||||
replay[mid].after_ground_truth,
|
||||
cwm.transition(replay[mid].before, replay[mid].joint_action),
|
||||
tolerances,
|
||||
)
|
||||
|
||||
if mismatches:
|
||||
# Check if earlier steps also mismatch
|
||||
if mid > lo:
|
||||
earlier = bisect_first_mismatch(cwm, replay, lo, mid - 1, tolerances)
|
||||
if earlier:
|
||||
return earlier
|
||||
return mismatches[0]
|
||||
|
||||
# No mismatch at mid, check right half
|
||||
return bisect_first_mismatch(cwm, replay, mid + 1, hi, tolerances)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Trajectory recording for self-play verification
|
||||
# ==============================================================================
|
||||
|
||||
class TrajectoryRecorder:
|
||||
"""
|
||||
Records every state/action/next_state for deterministic re-run verification.
|
||||
|
||||
For self-play: persist trajectory → re-run must produce exact same states.
|
||||
For historical: persist trajectory → CWM prediction must match ground truth.
|
||||
"""
|
||||
|
||||
def __init__(self, max_steps: int = 10_000) -> None:
|
||||
self._max_steps = max_steps
|
||||
self._steps: list[TrajectoryRecord] = []
|
||||
self._full_states: list[Tuple[MarketWorldState, Any, MarketWorldState]] = []
|
||||
|
||||
def record(
|
||||
self,
|
||||
step_index: int,
|
||||
before: MarketWorldState,
|
||||
action: Any,
|
||||
after: MarketWorldState,
|
||||
) -> None:
|
||||
"""Record one step. Keeps full states for detailed comparison."""
|
||||
if len(self._steps) >= self._max_steps:
|
||||
return
|
||||
|
||||
self._steps.append(TrajectoryRecord(
|
||||
step_index=step_index,
|
||||
before_hash=_hash_state(before),
|
||||
action_hash=_hash_action(action),
|
||||
after_hash=_hash_state(after),
|
||||
ts_ns=after.ts_ns,
|
||||
symbol=before.venue.symbol,
|
||||
))
|
||||
self._full_states.append((before, action, after))
|
||||
|
||||
def verify_deterministic(
|
||||
self,
|
||||
cwm: CodeWorldModel,
|
||||
) -> Tuple[bool, List[ReplayMismatch]]:
|
||||
"""
|
||||
Re-run the trajectory through CWM and verify exact match.
|
||||
|
||||
Must produce identical states for same inputs.
|
||||
"""
|
||||
mismatches: List[ReplayMismatch] = []
|
||||
|
||||
for idx, (before, action, expected_after) in enumerate(self._full_states):
|
||||
actual_after = cwm.transition(before, action if isinstance(action, tuple) else (action,))
|
||||
step_mismatches = _compare_deep(idx, expected_after, actual_after)
|
||||
mismatches.extend(step_mismatches)
|
||||
if any(m.is_critical for m in step_mismatches):
|
||||
break
|
||||
|
||||
return (len(mismatches) == 0, mismatches)
|
||||
|
||||
def trajectory_hash(self) -> str:
|
||||
"""Hash of the entire trajectory for quick comparison."""
|
||||
parts = [s.before_hash + s.action_hash + s.after_hash for s in self._steps]
|
||||
return hashlib.sha256("".join(parts).encode()).hexdigest()[:16]
|
||||
|
||||
@property
|
||||
def step_count(self) -> int:
|
||||
return len(self._steps)
|
||||
|
||||
@property
|
||||
def steps(self) -> List[TrajectoryRecord]:
|
||||
return list(self._steps)
|
||||
|
||||
def to_replay_steps(self) -> List[ReplayStep]:
|
||||
"""Convert recorded trajectory to ReplayStep list."""
|
||||
return [
|
||||
ReplayStep(
|
||||
before=before,
|
||||
joint_action=action if isinstance(action, tuple) else (action,),
|
||||
after_ground_truth=after,
|
||||
step_index=idx,
|
||||
)
|
||||
for idx, (before, action, after) in enumerate(self._full_states)
|
||||
]
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# ReplayVerifier — main interface
|
||||
# ==============================================================================
|
||||
|
||||
class ReplayVerifier:
|
||||
"""
|
||||
Replay matching is mandatory.
|
||||
|
||||
A fast wrong CWM is worse than a slow correct one.
|
||||
|
||||
Three verification modes:
|
||||
1. verify(): compare CWM predictions against ground truth steps
|
||||
2. verify_deterministic(): re-run trajectory, check exact match
|
||||
3. bisect(): binary search for first mismatch
|
||||
|
||||
Tolerances:
|
||||
- Historical replay: exchange-data tolerances (feeds can drop)
|
||||
- Self-play replay: tight tolerances (deterministic)
|
||||
"""
|
||||
|
||||
HISTORICAL_TOLERANCES = {
|
||||
"book_price": 0.01, # 1 cent
|
||||
"book_qty": 0.001,
|
||||
"account_equity": 0.01,
|
||||
"position_qty": 0.0001,
|
||||
"pnl_bps": 0.5,
|
||||
}
|
||||
|
||||
SELF_PLAY_TOLERANCES = {
|
||||
"book_price": 1e-9,
|
||||
"book_qty": 1e-12,
|
||||
"account_equity": 1e-9,
|
||||
"position_qty": 1e-12,
|
||||
"pnl_bps": 1e-6,
|
||||
}
|
||||
|
||||
def verify(
|
||||
self,
|
||||
cwm: CodeWorldModel,
|
||||
replay: Sequence[ReplayStep],
|
||||
tolerances: Optional[Mapping[str, float]] = None,
|
||||
) -> ReplayResult:
|
||||
"""Verify CWM predictions against ground truth steps."""
|
||||
t0 = time.perf_counter_ns()
|
||||
tol = tolerances or self.HISTORICAL_TOLERANCES
|
||||
all_mismatches: List[ReplayMismatch] = []
|
||||
steps_verified = 0
|
||||
first_mismatch_idx = None
|
||||
|
||||
for i, step in enumerate(replay):
|
||||
pred = cwm.transition(step.before, step.joint_action)
|
||||
mismatches = _compare_deep(i, step.after_ground_truth, pred, tol)
|
||||
steps_verified += 1
|
||||
|
||||
if mismatches:
|
||||
all_mismatches.extend(mismatches)
|
||||
if first_mismatch_idx is None:
|
||||
first_mismatch_idx = i
|
||||
break
|
||||
|
||||
# Hash trajectory for caching
|
||||
traj_hash = hashlib.sha256(
|
||||
"".join(_hash_state(s.before) for s in replay).encode()
|
||||
).hexdigest()[:16]
|
||||
|
||||
return ReplayResult(
|
||||
passed=len(all_mismatches) == 0,
|
||||
mismatches=all_mismatches,
|
||||
steps_verified=steps_verified,
|
||||
total_steps=len(replay),
|
||||
first_mismatch_index=first_mismatch_idx,
|
||||
duration_ns=time.perf_counter_ns() - t0,
|
||||
trajectory_hash=traj_hash,
|
||||
)
|
||||
|
||||
def verify_determinism(
|
||||
self,
|
||||
cwm: CodeWorldModel,
|
||||
replay: Sequence[ReplayStep],
|
||||
) -> ReplayResult:
|
||||
"""Verify that re-running produces identical results."""
|
||||
t0 = time.perf_counter_ns()
|
||||
mismatches: List[ReplayMismatch] = []
|
||||
|
||||
# Run once, collect results
|
||||
first_results: list = []
|
||||
for step in replay:
|
||||
first_results.append(cwm.transition(step.before, step.joint_action))
|
||||
|
||||
# Run again, compare
|
||||
for i, step in enumerate(replay):
|
||||
second = cwm.transition(step.before, step.joint_action)
|
||||
step_mismatches = _compare_deep(i, first_results[i], second, self.SELF_PLAY_TOLERANCES)
|
||||
mismatches.extend(step_mismatches)
|
||||
if any(m.is_critical for m in step_mismatches):
|
||||
break
|
||||
|
||||
traj_hash = hashlib.sha256(
|
||||
"".join(_hash_state(s.before) for s in replay).encode()
|
||||
).hexdigest()[:16]
|
||||
|
||||
return ReplayResult(
|
||||
passed=len(mismatches) == 0,
|
||||
mismatches=mismatches,
|
||||
steps_verified=len(replay),
|
||||
total_steps=len(replay),
|
||||
first_mismatch_index=mismatches[0].index if mismatches else None,
|
||||
duration_ns=time.perf_counter_ns() - t0,
|
||||
trajectory_hash=traj_hash,
|
||||
)
|
||||
|
||||
def bisect(
|
||||
self,
|
||||
cwm: CodeWorldModel,
|
||||
replay: Sequence[ReplayStep],
|
||||
tolerances: Optional[Mapping[str, float]] = None,
|
||||
) -> Optional[ReplayMismatch]:
|
||||
"""Binary search for first mismatch."""
|
||||
return bisect_first_mismatch(cwm, replay, tolerances=tolerances)
|
||||
Reference in New Issue
Block a user