Files
sentiment-engine/MALKHUT/malkhut/cwm/replay_verify.py
Codex f943191d56 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

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)