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.
462 lines
16 KiB
Python
462 lines
16 KiB
Python
"""
|
|
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)
|