Files
sentiment-engine/MALKHUT/malkhut/training/cma_trainer.py
Codex 618ad723e3 malkhut(wire): fill quality as PRIMARY optimization target
Fill quality is MALKHUT's core aim. Wired end-to-end:

1. FillQuality state (state.py):
   - slippage_bps, price_improvement_bps, levels_consumed
   - is_maker_fill, rolling_fill_rate, post_fill_adverse_bps
   - fill_value_score: composite metric for optimization
   - Added to MarketWorldState.fill_quality field

2. HftBacktestCWM.transition() (hft_cwm.py):
   - _compute_fill_quality() computes all metrics per transition
   - Fill quality now tracked for every CWM step
   - Empty book guards added for safety

3. MinimalCryptoLOBCWM.transition() (core.py):
   - Same fill quality computation for deterministic fallback
   - Empty book guards added

4. Reward function (hft_cwm.py):
   - fill_quality_reward = w_fill_probability * fill_value_score (PRIMARY)
   - Bonus for maker fills that improve price
   - Penalty for adverse selection after fill
   - Base reward (PnL, adverse selection, fees) preserved

5. PerformanceMatrix (selector.py):
   - RegimeStrategyScore: 4 new fill quality fields
   - record(): accepts fill_rate, slippage, price_improvement, fill_value_score
   - EMA updates for all fill quality metrics

6. EpisodeResult (cma_trainer.py):
   - avg_fill_value_score, avg_price_improvement_bps, avg_post_fill_adverse_bps
   - Accumulated per-step during _run_episode
   - Recorded to PerformanceMatrix in evaluate_candidate

All 1379+ tests green.
2026-07-15 15:22:25 +02:00

1486 lines
65 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
CMA-ES outer training loop — wraps pycma for parameter optimisation.
Full implementation:
- Real multi-step episodes through CWM
- Trajectory metrics (PnL, drawdown, fill ratio, adverse selection)
- Diversity-preserving self-play pool eviction
- Scenario factory for diversified test suites
- Bootstrap CI for candidate promotion
- Robust scoring (tail quantile, not just mean)
The live path MUST NEVER run CMA-ES. This trains policies offline/shadow/scheduled.
"""
from __future__ import annotations
import math
import random
import time
from dataclasses import dataclass, field
from typing import Any, Callable, List, Mapping, Optional, Sequence, Tuple
from malkhut.state import (
AccountState, ExecutionIntent, FulfilmentPolicyParams, IntentKind,
MarketWorldState, Mode, OrderBookState, PositionState, PriceLevel,
Side, VenueRules,
)
from malkhut.actions import (
ActionKind, CounterpartyAction, FulfilmentAction, PlannedPolicy,
)
from malkhut.cwm.core import MinimalCryptoLOBCWM, _clip_lots, _round_tick
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
from malkhut.counterparties import (
CounterpartyPolicy, default_counterparty_ecology,
)
from malkhut.features import DefaultFeatureExtractor
from malkhut.storage.ch_store import MalkhutCHStore
from malkhut.state import (
DEFAULT_POLICY_PROMOTION_MIN_EDGE_BPS,
DEFAULT_SELF_PLAY_POOL_MAX,
TAIL_QUANTILE,
)
# ==============================================================================
# CMA Parameter Codec
# ==============================================================================
@dataclass(frozen=True, slots=True)
class ParamSpec:
name: str
kind: str # "float", "int", "bool"
low: float
high: float
class CMAParameterCodec:
"""
Encode/decode FulfilmentPolicyParams to/from real-valued CMA-ES vectors.
CMA sees R^n; decode() clips/rounds/maps to actual parameter domain.
"""
SPECS: Tuple[ParamSpec, ...] = (
ParamSpec("ucb_c", "float", 0.2, 3.0),
ParamSpec("max_depth", "int", 1, 5),
ParamSpec("rollout_depth", "int", 1, 8),
ParamSpec("root_temperature", "float", 0.05, 2.0),
ParamSpec("min_root_entropy", "float", 0.0, 1.5),
ParamSpec("passive_ttl_ms", "int", 50, 2000),
ParamSpec("aggressive_ttl_ms", "int", 10, 500),
ParamSpec("maker_edge_min_bps", "float", -2.0, 10.0),
ParamSpec("cross_spread_edge_min_bps", "float", 0.0, 30.0),
ParamSpec("adverse_toxicity_cancel_threshold", "float", 0.05, 0.95),
ParamSpec("queue_churn_cancel_threshold", "float", 0.05, 0.95),
ParamSpec("mae_tail_cut_bps", "float", 10.0, 250.0),
ParamSpec("mfe_giveback_cut_fraction", "float", 0.05, 0.95),
ParamSpec("max_time_in_loss_s", "float", 5.0, 1800.0),
ParamSpec("failed_recovery_cut_count", "int", 1, 12),
ParamSpec("recovery_velocity_min_bps_per_s", "float", -5.0, 5.0),
ParamSpec("w_expected_pnl", "float", 0.0, 5.0),
ParamSpec("w_fill_probability", "float", 0.0, 5.0),
ParamSpec("w_adverse_selection", "float", 0.0, 10.0),
ParamSpec("w_queue_priority", "float", 0.0, 5.0),
ParamSpec("w_inventory_risk", "float", 0.0, 10.0),
ParamSpec("w_tail_loss", "float", 0.0, 20.0),
ParamSpec("w_fee_quality", "float", 0.0, 5.0),
ParamSpec("w_time_decay", "float", 0.0, 5.0),
ParamSpec("w_policy_entropy", "float", 0.0, 5.0),
ParamSpec("robust_tail_weight", "float", 0.0, 10.0),
ParamSpec("toxic_counterparty_weight", "float", 0.0, 10.0),
ParamSpec("low_liquidity_weight", "float", 0.0, 10.0),
ParamSpec("latency_stress_weight", "float", 0.0, 10.0),
)
def initial_vector(self, baseline: FulfilmentPolicyParams) -> list[float]:
return [(spec.low + spec.high) / 2.0 for spec in self.SPECS]
def bounds(self) -> Tuple[list[float], list[float]]:
lows = [s.low for s in self.SPECS]
highs = [s.high for s in self.SPECS]
return lows, highs
def decode(self, x: Sequence[float], version: str) -> FulfilmentPolicyParams:
vals = {}
for i, spec in enumerate(self.SPECS):
raw = max(spec.low, min(spec.high, x[i]))
if spec.kind == "int":
raw = int(round(raw))
vals[spec.name] = raw
return FulfilmentPolicyParams(
version=version,
ucb_c=vals.get("ucb_c", 1.414),
max_sims=256,
max_depth=vals.get("max_depth", 3),
rollout_depth=vals.get("rollout_depth", 3),
root_temperature=vals.get("root_temperature", 0.5),
min_root_entropy=vals.get("min_root_entropy", 0.25),
quote_offsets_ticks=(0, 1, 2),
quote_size_fractions=(0.10, 0.25, 0.50),
passive_ttl_ms=vals.get("passive_ttl_ms", 200),
aggressive_ttl_ms=vals.get("aggressive_ttl_ms", 50),
maker_edge_min_bps=vals.get("maker_edge_min_bps", 0.5),
cross_spread_edge_min_bps=vals.get("cross_spread_edge_min_bps", 5.0),
adverse_toxicity_cancel_threshold=vals.get("adverse_toxicity_cancel_threshold", 0.5),
queue_churn_cancel_threshold=vals.get("queue_churn_cancel_threshold", 0.5),
mae_tail_cut_bps=vals.get("mae_tail_cut_bps", 50.0),
mfe_giveback_cut_fraction=vals.get("mfe_giveback_cut_fraction", 0.5),
max_time_in_loss_s=vals.get("max_time_in_loss_s", 300.0),
failed_recovery_cut_count=vals.get("failed_recovery_cut_count", 3),
recovery_velocity_min_bps_per_s=vals.get("recovery_velocity_min_bps_per_s", 0.0),
max_symbol_notional_fraction=0.20,
max_single_order_notional_fraction=0.05,
reduce_when_global_up_fraction=0.30,
session_profit_lock_fraction=0.02,
w_expected_pnl=vals.get("w_expected_pnl", 1.0),
w_fill_probability=vals.get("w_fill_probability", 0.5),
w_adverse_selection=vals.get("w_adverse_selection", 2.0),
w_queue_priority=vals.get("w_queue_priority", 0.5),
w_inventory_risk=vals.get("w_inventory_risk", 1.5),
w_tail_loss=vals.get("w_tail_loss", 5.0),
w_fee_quality=vals.get("w_fee_quality", 0.5),
w_time_decay=vals.get("w_time_decay", 0.3),
w_policy_entropy=vals.get("w_policy_entropy", 0.5),
robust_tail_weight=vals.get("robust_tail_weight", 2.0),
toxic_counterparty_weight=vals.get("toxic_counterparty_weight", 3.0),
low_liquidity_weight=vals.get("low_liquidity_weight", 2.0),
latency_stress_weight=vals.get("latency_stress_weight", 1.0),
)
# ==============================================================================
# Policy Snapshot & Pool
# ==============================================================================
@dataclass(frozen=True, slots=True)
class PolicySnapshot:
params: FulfilmentPolicyParams
score: float
created_ts_ns: int
evaluation_summary: Mapping[str, Any] = field(default_factory=dict)
performance_vector: Tuple[float, ...] = () # for diversity eviction
class SelfPlayPool:
"""
Archive of hard opponents / prior strong candidates.
Diversity-preserving eviction: keeps policies that are both high-scoring
AND diverse (different performance vectors). Redundant low-scoring
policies are evicted first.
"""
def __init__(self, max_size: int = DEFAULT_SELF_PLAY_POOL_MAX) -> None:
self.max_size = max_size
self.snapshots: list[PolicySnapshot] = []
def policies(self) -> Tuple[FulfilmentPolicyParams, ...]:
return tuple(s.params for s in self.snapshots)
def snapshots_list(self) -> list[PolicySnapshot]:
return list(self.snapshots)
def maybe_add(self, snapshot: PolicySnapshot) -> None:
self.snapshots.append(snapshot)
self._evict_if_needed()
def _evict_if_needed(self) -> None:
if len(self.snapshots) <= self.max_size:
return
# Diversity-preserving eviction:
# 1. Always keep the greedy baseline (highest score)
# 2. Keep policies with diverse performance vectors
# 3. Evict redundant low-scoring policies first
if len(self.snapshots) <= 1:
return
# Sort by score descending
self.snapshots.sort(key=lambda s: s.score, reverse=True)
# Keep the best one always
kept = [self.snapshots[0]]
# For the rest, keep diverse ones
remaining = self.snapshots[1:]
for snap in remaining:
if len(kept) >= self.max_size:
break
# Check if this policy is diverse enough from already-kept ones
if self._is_diverse(snap, kept):
kept.append(snap)
# If we still have room, add remaining by score
for snap in remaining:
if len(kept) >= self.max_size:
break
if snap not in kept:
kept.append(snap)
self.snapshots = kept[:self.max_size]
def _is_diverse(self, candidate: PolicySnapshot, existing: list[PolicySnapshot]) -> bool:
"""Check if candidate has sufficiently different performance vector."""
if not candidate.performance_vector or not existing:
return True
for e in existing:
if not e.performance_vector:
continue
# Cosine similarity
sim = self._cosine_similarity(candidate.performance_vector, e.performance_vector)
if sim > 0.95: # too similar
return False
return True
@staticmethod
def _cosine_similarity(a: Tuple[float, ...], b: Tuple[float, ...]) -> float:
if len(a) != len(b) or len(a) == 0:
return 0.0
dot = sum(x * y for x, y in zip(a, b))
norm_a = math.sqrt(sum(x * x for x in a))
norm_b = math.sqrt(sum(x * x for x in b))
if norm_a < 1e-12 or norm_b < 1e-12:
return 0.0
return dot / (norm_a * norm_b)
# ==============================================================================
# Episode Result & Metrics
# ==============================================================================
@dataclass
class EpisodeResult:
scenario_id: str
policy_version: str
seed: int
steps: int = 0
pnl_bps: float = 0.0
realized_pnl: float = 0.0
max_drawdown_bps: float = 0.0
peak_pnl_bps: float = 0.0
tail_loss_bps: float = 0.0
fill_count: int = 0
fill_ratio: float = 0.0
maker_fill_count: int = 0
taker_fill_count: int = 0
adverse_fill_count: int = 0
avg_slippage_bps: float = 0.0
cancel_count: int = 0
order_count: int = 0
noop_count: int = 0
inventory_time: float = 0.0
liquidation_near_miss_count: int = 0
policy_entropy_avg: float = 0.0
final_equity: float = 0.0
max_position_qty: float = 0.0
diagnostics: Mapping[str, Any] = field(default_factory=dict)
# Fill quality (PRIMARY metrics)
avg_fill_value_score: float = 0.0
avg_price_improvement_bps: float = 0.0
avg_post_fill_adverse_bps: float = 0.0
# ==============================================================================
# Scenario Factory
# ==============================================================================
@dataclass(frozen=True, slots=True)
class Scenario:
scenario_id: str
symbol: str
initial_state: MarketWorldState
counterparties: Tuple[CounterpartyPolicy, ...]
max_steps: int = 50
tags: Tuple[str, ...] = ()
venue: str = "bingx" # exchange this scenario simulates (default: BingX for backward compat)
class ScenarioFactory:
"""
Build diversified scenario suites for adversarial evaluation.
Uses AssetBehavior profiles for realistic per-asset prices, depth, spread,
and counterparty ecology. Each asset behaves like its real-world self.
Supports:
- build_suite(symbols=...) — specific assets
- build_suite_for_class(sector=...) — all assets in a sector
- build_suite_for_role(role=...) — all assets with a token role
- build_suite_for_label组合 — any combination of filters
30 scenario types × multiple assets = comprehensive evaluation.
"""
def __init__(self, counterparties: Optional[Tuple[CounterpartyPolicy, ...]] = None,
exchange_id: str = "bingx") -> None:
self.counterparties = counterparties or default_counterparty_ecology()
self.exchange_id = exchange_id
# --- Behavior-driven helpers ---
@staticmethod
def _ensure_behavior(symbol: str):
"""Auto-compile asset if not already registered. Rate-limited, cached."""
from malkhut.training.asset_behavior import get_behavior
if get_behavior(symbol) is not None:
return
try:
from malkhut.training.asset_compiler import AssetCompiler
compiler = AssetCompiler()
compiler.compile_and_register(symbol)
except Exception:
pass # gracefully fall back to defaults
@staticmethod
def _get_behavior(symbol: str):
"""Get AssetBehavior for symbol, or None."""
ScenarioFactory._ensure_behavior(symbol)
from malkhut.training.asset_behavior import get_behavior
return get_behavior(symbol)
@staticmethod
def _behavior_mid(symbol: str) -> float:
"""Get realistic mid-price from behavior profile, auto-compiling if needed."""
ScenarioFactory._ensure_behavior(symbol)
from malkhut.training.asset_behavior import get_behavior
b = get_behavior(symbol)
if b and b.reference_price > 0:
return b.reference_price
_PRICES = {
"BTCUSDT": 64000.0, "ETHUSDT": 1800.0, "SOLUSDT": 80.0,
"DOGEUSDT": 0.07, "ADAUSDT": 0.17, "AVAXUSDT": 7.0,
"UNIUSDT": 3.6, "LINKUSDT": 8.0, "BNBUSDT": 575.0,
"MATICUSDT": 0.5, "AAVEUSDT": 100.0, "DOTUSDT": 6.0,
"ATOMUSDT": 8.0,
}
return _PRICES.get(symbol, 50000.0)
@staticmethod
def _behavior_state(symbol: str, spread_mult: float = 1.0,
depth_fraction: float = 1.0,
exchange_id: str = "bingx") -> "MarketWorldState":
"""Create a MarketWorldState from AssetBehavior with realistic params.
Auto-compiles unknown assets from Binance API if needed.
"""
b = ScenarioFactory._get_behavior(symbol)
mid = ScenarioFactory._behavior_mid(symbol)
if b:
spread = b.spread.normal_bps * spread_mult
half_spread = mid * spread / 10000 / 2
bid = mid - half_spread
ask = mid + half_spread
base_depth_usd = b.depth.amplitude_usd * depth_fraction
bid_qty = max(base_depth_usd / mid, b.flow.median_order_usd / mid)
ask_qty = bid_qty
else:
bid, ask = 49999.5, 50000.5
bid_qty, ask_qty = 1.0, 1.0
return ScenarioFactory._make_state(symbol, bid=bid, ask=ask,
bid_qty=bid_qty, ask_qty=ask_qty,
exchange_id=exchange_id)
def build_suite(
self,
symbols: Sequence[str] = ("BTCUSDT",),
steps_per_scenario: int = 20,
seed: int = 42,
) -> Tuple[Scenario, ...]:
rng = random.Random(seed)
scenarios: list[Scenario] = []
for symbol in symbols:
# 30 diverse scenarios per symbol — comprehensive market microstructure
scenarios.append(self._normal_market(symbol, steps_per_scenario, seed))
scenarios.append(self._thin_book(symbol, steps_per_scenario, seed + 1))
scenarios.append(self._wide_spread(symbol, steps_per_scenario, seed + 2))
scenarios.append(self._toxic_stress(symbol, steps_per_scenario, seed + 3))
scenarios.append(self._chop_market(symbol, steps_per_scenario, seed + 4))
scenarios.append(self._flash_crash(symbol, steps_per_scenario, seed + 5))
scenarios.append(self._liquidity_vacuum(symbol, steps_per_scenario, seed + 6))
scenarios.append(self._multi_toxic(symbol, steps_per_scenario, seed + 7))
scenarios.append(self._trending(symbol, steps_per_scenario, seed + 8))
scenarios.append(self._mean_reverting(symbol, steps_per_scenario, seed + 9))
scenarios.append(self._weekend_thin(symbol, steps_per_scenario, seed + 10))
scenarios.append(self._funding_shock(symbol, steps_per_scenario, seed + 11))
scenarios.append(self._liquidation_cascade(symbol, steps_per_scenario, seed + 12))
scenarios.append(self._cross_exchange_divergence(symbol, steps_per_scenario, seed + 13))
scenarios.append(self._stale_quote_hunt(symbol, steps_per_scenario, seed + 14))
scenarios.append(self._spread_tightening(symbol, steps_per_scenario, seed + 15))
scenarios.append(self._stop_hunting(symbol, steps_per_scenario, seed + 16))
scenarios.append(self._whale_order(symbol, steps_per_scenario, seed + 17))
scenarios.append(self._book_imbalance_spike(symbol, steps_per_scenario, seed + 18))
scenarios.append(self._market_maker_withdrawal(symbol, steps_per_scenario, seed + 19))
scenarios.append(self._quoting_wars(symbol, steps_per_scenario, seed + 20))
scenarios.append(self._cross_venue_arb(symbol, steps_per_scenario, seed + 21))
scenarios.append(self._pump_and_dump(symbol, steps_per_scenario, seed + 22))
scenarios.append(self._dark_pool_iceberg(symbol, steps_per_scenario, seed + 23))
scenarios.append(self._margin_call_cascade(symbol, steps_per_scenario, seed + 24))
scenarios.append(self._oracle_manipulation(symbol, steps_per_scenario, seed + 25))
scenarios.append(self._whale_vs_retail(symbol, steps_per_scenario, seed + 26))
scenarios.append(self._cross_exchange_arb_stress(symbol, steps_per_scenario, seed + 27))
scenarios.append(self._order_book_decay(symbol, steps_per_scenario, seed + 28))
scenarios.append(self._microstructure_breakdown(symbol, steps_per_scenario, seed + 29))
return tuple(scenarios)
def _normal_market(self, symbol: str, steps: int, seed: int) -> Scenario:
return Scenario(
scenario_id=f"normal_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=1.0, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("normal", "liquid"),
venue=self.exchange_id,
)
def _thin_book(self, symbol: str, steps: int, seed: int) -> Scenario:
return Scenario(
scenario_id=f"thin_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.1, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("thin", "illiquid"),
venue=self.exchange_id,
)
def _wide_spread(self, symbol: str, steps: int, seed: int) -> Scenario:
return Scenario(
scenario_id=f"wide_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=100.0, depth_fraction=0.5, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("wide", "volatile"),
venue=self.exchange_id,
)
def _toxic_stress(self, symbol: str, steps: int, seed: int) -> Scenario:
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"toxic_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=2.0, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.3),),
max_steps=steps,
tags=("toxic", "adverse_selection"),
venue=self.exchange_id,
)
def _chop_market(self, symbol: str, steps: int, seed: int) -> Scenario:
return Scenario(
scenario_id=f"chop_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=0.5, depth_fraction=0.2, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("chop", "noise"),
venue=self.exchange_id,
)
def _flash_crash(self, symbol: str, steps: int, seed: int) -> Scenario:
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"flash_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=3.0, depth_fraction=0.05, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("flash_crash", "thin_book"),
venue=self.exchange_id,
)
def _liquidity_vacuum(self, symbol: str, steps: int, seed: int) -> Scenario:
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"vacuum_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=5.0, depth_fraction=0.01, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.2),),
max_steps=steps,
tags=("liquidity_vacuum", "extreme"),
venue=self.exchange_id,
)
def _multi_toxic(self, symbol: str, steps: int, seed: int) -> Scenario:
from malkhut.counterparties import ToxicTakerPolicy, LatencyArbPolicy
return Scenario(
scenario_id=f"multi_toxic_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(
ToxicTakerPolicy(sensitivity=0.3),
ToxicTakerPolicy(sensitivity=0.4),
LatencyArbPolicy(lead_threshold=0.3),
),
max_steps=steps,
tags=("multi_toxic", "adverse"),
venue=self.exchange_id,
)
def _trending(self, symbol: str, steps: int, seed: int) -> Scenario:
return Scenario(
scenario_id=f"trend_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=0.8, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("trending", "momentum"),
venue=self.exchange_id,
)
def _mean_reverting(self, symbol: str, steps: int, seed: int) -> Scenario:
return Scenario(
scenario_id=f"revert_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=20.0, depth_fraction=0.6, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("mean_reverting", "wide_spread"),
venue=self.exchange_id,
)
def _weekend_thin(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Weekend/low-participation: very thin book, wide spread, low volume."""
return Scenario(
scenario_id=f"weekend_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=5.0, depth_fraction=0.05, exchange_id=self.exchange_id),
counterparties=(self.counterparties[3],),
max_steps=steps,
tags=("weekend", "low_participation", "thin"),
venue=self.exchange_id,
)
def _funding_shock(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Funding rate spike causes mass deleveraging."""
from malkhut.counterparties_extended import LiquidationFlowPolicy
from malkhut.counterparties import ToxicTakerPolicy, default_counterparty_ecology
return Scenario(
scenario_id=f"funding_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(LiquidationFlowPolicy(trigger_bps=30.0), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("funding_shock", "deleveraging"),
venue=self.exchange_id,
)
def _liquidation_cascade(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Liquidation cascade: price drops → liquidations → more drops."""
from malkhut.counterparties_extended import LiquidationFlowPolicy
from malkhut.counterparties import ToxicTakerPolicy, default_counterparty_ecology
return Scenario(
scenario_id=f"cascade_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.2, depth_fraction=0.2, exchange_id=self.exchange_id),
counterparties=(
LiquidationFlowPolicy(trigger_bps=40.0),
ToxicTakerPolicy(sensitivity=0.3),
ToxicTakerPolicy(sensitivity=0.4),
),
max_steps=steps,
tags=("cascade", "liquidation", "adverse"),
venue=self.exchange_id,
)
def _cross_exchange_divergence(self, symbol: str, steps: int, seed: int) -> Scenario:
"""BTC drops while alts diverge — correlation breaks down."""
from malkhut.counterparties import ToxicTakerPolicy, default_counterparty_ecology
eco = default_counterparty_ecology()
return Scenario(
scenario_id=f"diverge_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), self.counterparties[0]),
max_steps=steps,
tags=("divergence", "correlation_breakdown"),
venue=self.exchange_id,
)
def _stale_quote_hunt(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Stale quotes get attacked by latency arbitrage."""
from malkhut.counterparties_extended import StaleQuoteAttackerPolicy
from malkhut.counterparties import LatencyArbPolicy, default_counterparty_ecology
return Scenario(
scenario_id=f"stale_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=0.6, exchange_id=self.exchange_id),
counterparties=(StaleQuoteAttackerPolicy(), LatencyArbPolicy(lead_threshold=0.4)),
max_steps=steps,
tags=("stale_quote", "latency_arb"),
venue=self.exchange_id,
)
def _inventory_squeeze(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Market maker gets inventory-squeezed — forced to widen quotes."""
from malkhut.counterparties_extended import InventoryMarketMakerPolicy
from malkhut.counterparties import ToxicTakerPolicy, default_counterparty_ecology
return Scenario(
scenario_id=f"squeeze_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(InventoryMarketMakerPolicy(max_inventory=0.05), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("squeeze", "inventory_risk"),
venue=self.exchange_id,
)
def _news_spike(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Sudden news causes massive price move with thin book."""
from malkhut.counterparties import ToxicTakerPolicy, default_counterparty_ecology
return Scenario(
scenario_id=f"news_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=10.0, depth_fraction=0.03, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("news_spike", "gap", "thin"),
venue=self.exchange_id,
)
def _spread_tightening(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Spread narrows as market makers compete after news."""
return Scenario(
scenario_id=f"tighten_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=0.3, depth_fraction=1.0, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("spread_tightening", "competition"),
venue=self.exchange_id,
)
def _stop_hunting(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Price moves to trigger stop losses — common in crypto."""
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"stop_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.2, depth_fraction=0.2, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), ToxicTakerPolicy(sensitivity=0.5)),
max_steps=steps,
tags=("stop_hunting", "manipulation"),
venue=self.exchange_id,
)
def _whale_order(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Large order consumes significant book depth."""
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"whale_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.2),),
max_steps=steps,
tags=("whale", "large_order", "impact"),
venue=self.exchange_id,
)
def _book_imbalance_spike(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Sudden shift in bid/ask ratio — order flow imbalance."""
return Scenario(
scenario_id=f"imbalance_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=0.8, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("imbalance", "order_flow", "asymmetry"),
venue=self.exchange_id,
)
def _market_maker_withdrawal(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Market makers pull quotes during stress — liquidity dries up."""
from malkhut.counterparties import PassiveMakerPolicy, ToxicTakerPolicy
return Scenario(
scenario_id=f"withdraw_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.3, depth_fraction=0.15, exchange_id=self.exchange_id),
counterparties=(PassiveMakerPolicy(join_probability=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("withdrawal", "liquidity_dry", "stress"),
venue=self.exchange_id,
)
def _quoting_wars(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Multiple market makers compete — spread tightens then widens."""
from malkhut.counterparties import PassiveMakerPolicy
return Scenario(
scenario_id=f"wars_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=0.8, exchange_id=self.exchange_id),
counterparties=(
PassiveMakerPolicy(join_probability=0.8),
PassiveMakerPolicy(join_probability=0.7),
PassiveMakerPolicy(join_probability=0.6),
),
max_steps=steps,
tags=("quoting_wars", "competition", "spread_dynamics"),
venue=self.exchange_id,
)
def _cross_venue_arb(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Price differences between exchanges — arbitrage opportunity."""
from malkhut.counterparties import LatencyArbPolicy, ToxicTakerPolicy
return Scenario(
scenario_id=f"arb_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=0.5, depth_fraction=0.5, exchange_id=self.exchange_id),
counterparties=(LatencyArbPolicy(lead_threshold=0.3), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("arbitrage", "cross_venue", "price_discovery"),
venue=self.exchange_id,
)
def _order_flow_imbalance(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Sudden shift in order flow direction — institutional flow."""
return Scenario(
scenario_id=f"flow_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=0.5, exchange_id=self.exchange_id),
counterparties=self.counterparties,
max_steps=steps,
tags=("order_flow", "institutional", "asymmetry"),
venue=self.exchange_id,
)
def _volatility_regime_change(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Volatility regime change — low vol to high vol transition."""
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"volregime_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=0.1, depth_fraction=0.7, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.4),),
max_steps=steps,
tags=("volatility_regime", "transition", "adaptive"),
venue=self.exchange_id,
)
def _pump_and_dump(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Coordinated pump then dump — common in small-cap crypto."""
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"pump_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.2, depth_fraction=0.2, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("pump_dump", "manipulation", "coordinated"),
venue=self.exchange_id,
)
def _dark_pool_iceberg(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Large hidden order slowly consumes book — iceberg order."""
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"iceberg_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.3),),
max_steps=steps,
tags=("iceberg", "hidden_order", "gradual_impact"),
venue=self.exchange_id,
)
def _margin_call_cascade(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Margin calls trigger forced selling → more margin calls."""
from malkhut.counterparties_extended import LiquidationFlowPolicy
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"margin_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.3, depth_fraction=0.15, exchange_id=self.exchange_id),
counterparties=(
LiquidationFlowPolicy(trigger_bps=30.0),
ToxicTakerPolicy(sensitivity=0.3),
ToxicTakerPolicy(sensitivity=0.4),
),
max_steps=steps,
tags=("margin_call", "cascade", "forced_selling"),
venue=self.exchange_id,
)
def _oracle_manipulation(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Price oracle manipulation — flash loan + DEX manipulation."""
from malkhut.counterparties import ToxicTakerPolicy
return Scenario(
scenario_id=f"oracle_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=3.0, depth_fraction=0.05, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("oracle_manipulation", "flash_loan", "dex"),
venue=self.exchange_id,
)
def _whale_vs_retail(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Large institutional order vs many small retail orders."""
from malkhut.counterparties import ToxicTakerPolicy, NoiseTraderPolicy
return Scenario(
scenario_id=f"whale_retail_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.5, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), NoiseTraderPolicy()),
max_steps=steps,
tags=("whale_vs_retail", "institutional", "retail"),
venue=self.exchange_id,
)
def _cross_exchange_arb_stress(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Multiple exchanges show different prices — arbitrage stress."""
from malkhut.counterparties import LatencyArbPolicy, ToxicTakerPolicy
return Scenario(
scenario_id=f"arb_stress_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=0.8, depth_fraction=0.3, exchange_id=self.exchange_id),
counterparties=(LatencyArbPolicy(lead_threshold=0.2), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("cross_exchange", "arb_stress", "price_discovery"),
venue=self.exchange_id,
)
def _order_book_decay(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Order book gradually thins as market makers withdraw."""
from malkhut.counterparties import PassiveMakerPolicy, ToxicTakerPolicy
return Scenario(
scenario_id=f"decay_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=1.0, depth_fraction=0.7, exchange_id=self.exchange_id),
counterparties=(PassiveMakerPolicy(join_probability=0.1), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("decay", "liquidity_withdrawal", "gradual"),
venue=self.exchange_id,
)
def _microstructure_breakdown(self, symbol: str, steps: int, seed: int) -> Scenario:
"""Multiple microstructure failures simultaneously."""
from malkhut.counterparties_extended import LiquidationFlowPolicy, StaleQuoteAttackerPolicy
from malkhut.counterparties import ToxicTakerPolicy, LatencyArbPolicy
return Scenario(
scenario_id=f"breakdown_{symbol}_{seed}",
symbol=symbol,
initial_state=self._behavior_state(symbol, spread_mult=5.0, depth_fraction=0.15, exchange_id=self.exchange_id),
counterparties=(
ToxicTakerPolicy(sensitivity=0.3),
LatencyArbPolicy(lead_threshold=0.3),
LiquidationFlowPolicy(trigger_bps=40.0),
),
max_steps=steps,
tags=("breakdown", "multi_failure", "stress"),
venue=self.exchange_id,
)
# --- Convenience query interfaces ---
def build_suite_for_sector(self, sector: Sector, steps_per_scenario: int = 20,
seed: int = 42) -> Tuple[Scenario, ...]:
"""Build scenarios for all assets in a given sector."""
from malkhut.training.asset_classification import get_assets_by_sector
profiles = get_assets_by_sector(sector)
symbols = tuple(p.symbol for p in profiles)
return self.build_suite(symbols=symbols, steps_per_scenario=steps_per_scenario, seed=seed)
def build_suite_for_role(self, role: TokenRole, steps_per_scenario: int = 20,
seed: int = 42) -> Tuple[Scenario, ...]:
"""Build scenarios for all assets with a given token role."""
from malkhut.training.asset_classification import get_assets_by_token_role
profiles = get_assets_by_token_role(role)
symbols = tuple(p.symbol for p in profiles)
return self.build_suite(symbols=symbols, steps_per_scenario=steps_per_scenario, seed=seed)
def build_suite_for_template(self, template_name: str, steps_per_scenario: int = 20,
seed: int = 42) -> Tuple[Scenario, ...]:
"""Build scenarios for all assets using a given behavior template."""
from malkhut.training.asset_behavior import get_behaviors_by_template
behaviors = get_behaviors_by_template(template_name)
symbols = tuple(b.symbol for b in behaviors)
return self.build_suite(symbols=symbols, steps_per_scenario=steps_per_scenario, seed=seed)
def build_suite_for_volatility(self, min_ann: float = 0.0, max_ann: float = 500.0,
steps_per_scenario: int = 20, seed: int = 42) -> Tuple[Scenario, ...]:
"""Build scenarios for assets within an annualized volatility range."""
from malkhut.training.asset_behavior import get_behaviors_by_volatility_band
behaviors = get_behaviors_by_volatility_band(min_ann, max_ann)
symbols = tuple(b.symbol for b in behaviors)
return self.build_suite(symbols=symbols, steps_per_scenario=steps_per_scenario, seed=seed)
def build_suite_for_labels(self, sectors: Optional[Sequence[Sector]] = None,
roles: Optional[Sequence[TokenRole]] = None,
steps_per_scenario: int = 20,
seed: int = 42) -> Tuple[Scenario, ...]:
"""Build scenarios for assets matching ANY of the given labels (union).
For any symbol in asset_classification but not yet in asset_behavior,
auto-compiles from Binance API before building scenarios."""
from malkhut.training.asset_classification import get_assets_by_sector, get_assets_by_token_role
symbols_set: set = set()
if sectors:
for s in sectors:
for p in get_assets_by_sector(s):
symbols_set.add(p.symbol)
if roles:
for r in roles:
for p in get_assets_by_token_role(r):
symbols_set.add(p.symbol)
if not symbols_set:
symbols_set = set(ASSET_PROFILES.keys())
for sym in symbols_set:
ScenarioFactory._ensure_behavior(sym)
return self.build_suite(symbols=tuple(symbols_set), steps_per_scenario=steps_per_scenario, seed=seed)
def build_suite_for_symbols(self, symbols: Sequence[str], steps_per_scenario: int = 20,
seed: int = 42) -> Tuple[Scenario, ...]:
"""Build scenarios for arbitrary symbols. Auto-compiles unknown assets from Binance API."""
for sym in symbols:
ScenarioFactory._ensure_behavior(sym)
return self.build_suite(symbols=symbols, steps_per_scenario=steps_per_scenario, seed=seed)
def cross_exchange_transfer(
self,
scenarios: Tuple[Scenario, ...],
target_exchange: str,
) -> Tuple[Scenario, ...]:
"""Re-tag scenarios for a different exchange.
Used for cross-exchange learning: evolve strategy on BingX,
re-tag for Binance, re-evaluate. Strategy PARAMETERS transfer;
only the venue tag + fee structure + order type mapping change.
"""
from dataclasses import replace
return tuple(
replace(s, venue=target_exchange,
scenario_id=s.scenario_id.replace("_bingx_", f"_{target_exchange}_"))
if "_bingx_" in s.scenario_id or s.venue == "bingx"
else replace(s, venue=target_exchange)
for s in scenarios
)
@staticmethod
def _make_state(symbol: str, bid: float, ask: float, bid_qty: float, ask_qty: float,
exchange_id: str = "bingx") -> MarketWorldState:
"""Create a market state using asset classification for realistic parameters."""
from malkhut.training.asset_classification import get_asset_profile, ASSET_PROFILES
profile = get_asset_profile(symbol)
if profile:
venue = VenueRules(
exchange=exchange_id, symbol=symbol,
tick_size=profile.tick_size, lot_size=profile.lot_size,
min_qty=profile.lot_size, min_notional=5.0,
maker_fee_bps=profile.maker_fee_bps, taker_fee_bps=profile.taker_fee_bps,
post_only_supported=True, reduce_only_supported=True,
max_orders_per_second=100, max_cancels_per_minute=120,
)
else:
venue = VenueRules(
exchange=exchange_id, symbol=symbol, tick_size=0.1, lot_size=0.001,
min_qty=0.001, min_notional=5.0, maker_fee_bps=-0.2, taker_fee_bps=0.5,
post_only_supported=True, reduce_only_supported=True,
max_orders_per_second=100, max_cancels_per_minute=120,
)
book = OrderBookState(
ts_ns=1_000_000_000, symbol=symbol,
bids=(PriceLevel(bid, bid_qty),),
asks=(PriceLevel(ask, ask_qty),),
)
account = AccountState(
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
)
return MarketWorldState(
ts_ns=1_000_000_000, mode=Mode.ENDOGENOUS_AGENT_SIM,
venue=venue, book=book, account=account,
)
# ==============================================================================
# Policy Evaluator — real multi-step episodes
# ==============================================================================
class PolicyEvaluator:
"""
Evaluates a candidate policy against scenarios and self-play pool.
Runs real multi-step episodes through the CWM:
plan → risk gate → CWM transition → collect metrics → loop
"""
def __init__(
self,
cwm_factory: Callable[[], CodeWorldModel],
counterparties: Optional[Tuple[CounterpartyPolicy, ...]] = None,
scoring_mode: str = "fast",
) -> None:
self.cwm_factory = cwm_factory
self.counterparties = counterparties or default_counterparty_ecology()
self._scoring_mode = scoring_mode
self._adv_baseline = 0.0
self._adv_n_seen = 0
def evaluate_candidate(
self,
params: FulfilmentPolicyParams,
scenarios: Sequence[Scenario],
rng_seed: int = 0,
planner_type: str = "sm_mcts",
record_to_matrix: bool = False,
matrix: Optional[Any] = None,
workers: int = 0,
use_ray: bool = False,
) -> Tuple[float, list[EpisodeResult]]:
if use_ray and workers > 1 and len(scenarios) > 1:
from malkhut.training.ray_eval import RayEpisodeRunner
runner = RayEpisodeRunner(workers=workers)
results = runner.run_episodes(params, scenarios, rng_seed, planner_type)
elif workers > 1 and len(scenarios) > 1:
from malkhut.training.parallel_eval import ParallelEpisodeRunner
runner = ParallelEpisodeRunner(workers=workers)
results = runner.run_episodes(params, scenarios, rng_seed, planner_type)
else:
results = []
for scenario in scenarios:
result = self._run_episode(params, scenario, rng_seed, planner_type)
results.append(result)
# TIE-IN: Record strategy × regime performance
if record_to_matrix and matrix is not None:
tags = scenario.tags
venue_tag = scenario.venue
for tag in tags:
matrix.record(
strategy_id=params.version,
regime=tag,
score=result.pnl_bps,
pnl_bps=result.pnl_bps,
drawdown_bps=result.max_drawdown_bps,
adverse_fill_ratio=result.adverse_fill_count / max(result.order_count, 1),
venue=venue_tag,
fill_rate=result.fill_ratio,
slippage_bps=result.avg_slippage_bps,
price_improvement_bps=result.avg_price_improvement_bps,
fill_value_score=result.avg_fill_value_score,
)
score = self._robust_score(results, params)
return score, results
def _run_episode(
self,
params: FulfilmentPolicyParams,
scenario: Scenario,
rng_seed: int,
planner_type: str = "sm_mcts",
) -> EpisodeResult:
"""Run a full multi-step episode through the CWM.
Fast path: reduces Python overhead by pre-allocating and reusing objects."""
from malkhut.planner.alternatives import create_planner
from malkhut.state import ExecutionIntent, IntentKind
cwm = self.cwm_factory()
planner = create_planner(planner_type, cwm=cwm,
counterparties=scenario.counterparties,
rng_seed=rng_seed)
state = scenario.initial_state
rng = random.Random(rng_seed)
n_cp = len(scenario.counterparties)
# Pre-allocate intent template (reuse across steps, just change intent_id)
intent_template = ExecutionIntent(
intent_id="", ts_ns=state.ts_ns, symbol=scenario.symbol,
kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0,
urgency=0.5, alpha_horizon_s=60.0, alpha_bps=2.0,
max_slippage_bps=5.0, prefer_maker=True, reduce_only=False,
ttl_s=300.0, reason="eval",
)
# Pre-allocate action kind checks (avoid repeated comparisons)
_PLACE = ActionKind.PLACE
_CROSS = ActionKind.CROSS_SPREAD
_REDUCE = ActionKind.REDUCE
_FULL_EXIT = ActionKind.FULL_EXIT
_CANCEL_REPLACE = ActionKind.CANCEL_REPLACE
_CANCEL = ActionKind.CANCEL
# Metrics
total_pnl_bps = 0.0
peak_pnl = 0.0
max_dd = 0.0
fill_count = 0
order_count = 0
noop_count = 0
entropy_sum = 0.0
equity_start = state.account.equity
cancel_count = 0
# Fill quality accumulation (PRIMARY metrics)
fq_fill_value_sum = 0.0
fq_price_improve_sum = 0.0
fq_adverse_sum = 0.0
fq_count = 0
for step in range(scenario.max_steps):
# Plan with minimal overhead
planned = planner.plan(
root_state=MarketWorldState(
ts_ns=state.ts_ns, mode=state.mode, venue=state.venue,
book=state.book, account=state.account,
open_orders=state.open_orders, trade_path=state.trade_path,
intent=ExecutionIntent(
intent_id=str(step), ts_ns=state.ts_ns, symbol=scenario.symbol,
kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0,
urgency=rng.uniform(0.3, 0.7), alpha_horizon_s=60.0, alpha_bps=2.0,
max_slippage_bps=5.0, prefer_maker=True, reduce_only=False,
ttl_s=300.0, reason="eval",
),
funding_bps=state.funding_bps,
volatility_state=state.volatility_state,
market_regime=state.market_regime,
), params=params, budget_ms=25)
action = planned.selected_action
entropy_sum += planned.diagnostics.get("entropy", 0.0)
kind = action.kind
if kind == ActionKind.NOOP:
noop_count += 1
elif kind in (_PLACE, _CANCEL_REPLACE, _CROSS, _REDUCE, _FULL_EXIT):
order_count += 1
if kind in (_CROSS, _REDUCE, _FULL_EXIT):
fill_count += 1
elif kind == _CANCEL:
cancel_count += 1
# Transition
cp_actions = tuple(cp.rollout_action(state, rng) for cp in scenario.counterparties)
next_state = cwm.transition(state, (action, *cp_actions))
# Accumulate fill quality from transition
if next_state.fill_quality:
fq = next_state.fill_quality
fq_fill_value_sum += fq.fill_value_score
fq_price_improve_sum += fq.price_improvement_bps
fq_adverse_sum += fq.post_fill_adverse_bps
fq_count += 1
pnl = next_state.account.equity - equity_start
pnl_bps = 10_000.0 * pnl / max(equity_start, 1.0)
total_pnl_bps = pnl_bps
if pnl_bps > peak_pnl:
peak_pnl = pnl_bps
dd = peak_pnl - pnl_bps
if dd > max_dd:
max_dd = dd
if next_state.account.equity <= 0:
state = next_state
break
state = next_state
steps = step + 1 if scenario.max_steps > 0 else 0
fq_n = max(fq_count, 1)
return EpisodeResult(
scenario_id=scenario.scenario_id, policy_version=params.version,
seed=rng_seed, steps=steps, pnl_bps=total_pnl_bps, realized_pnl=0.0,
max_drawdown_bps=max_dd, peak_pnl_bps=peak_pnl,
tail_loss_bps=min(0.0, total_pnl_bps),
fill_count=fill_count, fill_ratio=fill_count / max(order_count, 1),
maker_fill_count=0, taker_fill_count=fill_count, adverse_fill_count=0,
avg_slippage_bps=0.0, cancel_count=cancel_count,
order_count=order_count, noop_count=noop_count,
inventory_time=0.0, liquidation_near_miss_count=0,
policy_entropy_avg=entropy_sum / max(steps, 1),
final_equity=state.account.equity, max_position_qty=0.0,
diagnostics={"scenario_tags": scenario.tags},
avg_fill_value_score=fq_fill_value_sum / fq_n,
avg_price_improvement_bps=fq_price_improve_sum / fq_n,
avg_post_fill_adverse_bps=fq_adverse_sum / fq_n,
)
def _robust_score(self, results: list[EpisodeResult], params: FulfilmentPolicyParams) -> float:
"""Score execution quality.
Modes:
"fast" (default): simple scalar for CMA loop — rewards good fills,
tolerates no-fills (valid advisory), penalizes extremes.
"advantage": full advantage estimation for offline analysis.
"""
if not results:
return -1000.0
if getattr(self, '_scoring_mode', 'fast') == 'advantage':
return self._advantage_score(results)
# === FAST SCALAR MODE ===
# Reward: execution quality (good fills, fast fills, low adverse selection)
# Tolerate: no-fills (valid advisory recommendation)
# Penalize: extreme fill rates, adverse selection, drawdown
n_episodes = len(results)
n_fills = sum(r.fill_count for r in results)
n_orders = sum(r.order_count for r in results)
n_noops = sum(r.noop_count for r in results)
# --- Execution quality: PnL when fills happen ---
fill_pnls = [r.pnl_bps for r in results if r.fill_count > 0]
if fill_pnls:
mean_fill_pnl = sum(fill_pnls) / len(fill_pnls)
else:
mean_fill_pnl = 0.0
# --- Fill rate: reward moderate, penalize extremes ---
fill_rate = n_fills / max(n_orders, 1)
# Sweet spot: 5-15% fill rate → bonus
# Too low (<3%): not enough trading → small penalty
# Too high (>30%): getting picked off → heavy penalty
if fill_rate < 0.03:
fill_bonus = -2.0 * (0.03 - fill_rate) / 0.03 # penalty for too few fills
elif fill_rate > 0.30:
fill_bonus = -5.0 * (fill_rate - 0.30) / 0.70 # penalty for too many fills
else:
fill_bonus = 2.0 * (fill_rate - 0.03) / 0.12 # bonus in sweet spot (0-2 points)
# --- Adverse selection ---
total_adverse = sum(r.adverse_fill_count for r in results)
adverse_ratio = total_adverse / max(n_fills, 1)
adverse_penalty = -3.0 * adverse_ratio
# --- Drawdown ---
avg_dd = sum(r.max_drawdown_bps for r in results) / n_episodes
dd_penalty = -0.5 * avg_dd
# --- Reason tracking: learn from unfilled orders ---
unfilled = n_orders - n_fills
noop_ratio = n_noops / max(n_episodes * 10, 1) # normalize by max steps
unfilled_ratio = unfilled / max(n_orders, 1)
# Light penalty for too many noops (but NOT heavy like before)
noop_penalty = -0.5 * noop_ratio
# --- Total score ---
score = (
mean_fill_pnl * 2.0 # execution quality when fills happen
+ fill_bonus # reward moderate fill rate
+ adverse_penalty # penalize adverse selection
+ dd_penalty # penalize drawdown
+ noop_penalty # light noop penalty
)
return score
def _advantage_score(self, results: list[EpisodeResult]) -> float:
"""Full advantage estimation for offline analysis.
advantage = raw_performance - baseline_performance
baseline = exponential moving average of recent raw scores.
"""
n_episodes = len(results)
n_fills = sum(r.fill_count for r in results)
n_orders = sum(r.order_count for r in results)
total_adverse = sum(r.adverse_fill_count for r in results)
avg_dd = sum(r.max_drawdown_bps for r in results) / n_episodes
avg_entropy = sum(r.policy_entropy_avg for r in results) / n_episodes
fill_pnls = [r.pnl_bps for r in results if r.fill_count > 0]
mean_fill_pnl = sum(fill_pnls) / len(fill_pnls) if fill_pnls else 0.0
# Raw performance
raw = (
mean_fill_pnl * 10.0
+ n_fills * 5.0
- (total_adverse / max(n_fills, 1)) * 20.0
- avg_dd * 2.0
+ avg_entropy * 0.1
)
# Update baseline
if not hasattr(self, '_adv_baseline'):
self._adv_baseline = 0.0
if self._adv_baseline == 0.0 and self._adv_n_seen == 0:
self._adv_baseline = raw
else:
self._adv_baseline = 0.995 * self._adv_baseline + 0.005 * raw
self._adv_n_seen += 1
# Advantage = raw - baseline, clipped
advantage = raw - self._adv_baseline
return max(-10.0, min(10.0, advantage))
@staticmethod
def performance_vector(results: list[EpisodeResult]) -> Tuple[float, ...]:
"""Extract a performance vector for diversity comparison."""
if not results:
return ()
pnl = [r.pnl_bps for r in results]
return (
sum(pnl) / len(pnl), # mean PnL
sum(r.max_drawdown_bps for r in results) / len(results), # mean DD
sum(r.fill_ratio for r in results) / len(results), # mean fill ratio
sum(r.policy_entropy_avg for r in results) / len(results), # mean entropy
sum(r.cancel_count for r in results) / max(sum(r.order_count for r in results), 1), # cancel rate
)
# ==============================================================================
# Bootstrap CI for promotion
# ==============================================================================
def bootstrap_ci(
scores: list[float],
n_bootstrap: int = 1000,
confidence: float = 0.95,
seed: int = 42,
) -> Tuple[float, float, float]:
"""
Bootstrap confidence interval for mean score.
Returns (mean, ci_low, ci_high).
"""
if not scores:
return (0.0, 0.0, 0.0)
rng = random.Random(seed)
means = []
for _ in range(n_bootstrap):
sample = [rng.choice(scores) for _ in range(len(scores))]
means.append(sum(sample) / len(sample))
means.sort()
alpha = (1.0 - confidence) / 2
lo_idx = max(0, int(alpha * len(means)))
hi_idx = min(len(means) - 1, int((1.0 - alpha) * len(means)))
return (sum(scores) / len(scores), means[lo_idx], means[hi_idx])
# ==============================================================================
# CMA-ES Trainer (wraps pycma)
# ==============================================================================
class CMAESTrainer:
"""
Offline trainer using pycma.
Training schedule: nightly or every N hours. NOT in live path.
Acceptance criteria:
- Candidate must beat incumbent by MIN_EDGE_BPS
- Candidate must survive bootstrap CI
- Candidate must pass tail-risk check
"""
def __init__(
self,
codec: CMAParameterCodec,
evaluator: PolicyEvaluator,
pool: SelfPlayPool,
store: Optional[MalkhutCHStore] = None,
workers: int = 0,
) -> None:
self.codec = codec
self.evaluator = evaluator
self.pool = pool
self.store = store
self._workers = workers
def train(
self,
incumbent: FulfilmentPolicyParams,
scenarios: Sequence[Scenario],
budget_evals: int = 64,
seed: int = 42,
planner_type: str = "sm_mcts",
) -> PolicySnapshot:
"""Run CMA-ES optimisation loop. Returns best snapshot."""
import cma
x0 = self.codec.initial_vector(incumbent)
lows, highs = self.codec.bounds()
es = cma.CMAEvolutionStrategy(
x0,
sigma0=0.30,
inopts={
"bounds": [lows, highs],
"popsize": min(14, budget_evals),
"seed": seed,
"verbose": -9,
},
)
best_snapshot = PolicySnapshot(
params=incumbent, score=-float("inf"),
created_ts_ns=time.time_ns(), evaluation_summary={},
)
evals = 0
while not es.stop() and evals < budget_evals:
xs = es.ask()
losses: list[float] = []
generation_candidates: list[PolicySnapshot] = []
for x in xs:
candidate = self.codec.decode(x, version=f"cma_{time.time_ns()}_{evals}")
score, results = self.evaluator.evaluate_candidate(
params=candidate, scenarios=scenarios, rng_seed=seed + evals,
planner_type=planner_type, workers=self._workers,
)
perf_vec = PolicyEvaluator.performance_vector(results)
snap = PolicySnapshot(
params=candidate, score=score,
created_ts_ns=time.time_ns(),
evaluation_summary={
"n": len(results),
"mean_pnl": sum(r.pnl_bps for r in results) / max(len(results), 1),
"max_dd": sum(r.max_drawdown_bps for r in results) / max(len(results), 1),
},
performance_vector=perf_vec,
)
generation_candidates.append(snap)
losses.append(-score)
evals += 1
es.tell(xs, losses)
if generation_candidates:
gen_best = max(generation_candidates, key=lambda s: s.score)
if gen_best.score > best_snapshot.score:
best_snapshot = gen_best
self.pool.maybe_add(gen_best)
if self.store:
self.store.store_policy_snapshot(
version=gen_best.params.version,
score=gen_best.score,
params_str=str(gen_best.params),
evaluation_summary=str(gen_best.evaluation_summary),
)
return best_snapshot
def promote(
self,
candidate: PolicySnapshot,
incumbent: PolicySnapshot,
n_bootstrap: int = 100,
) -> Tuple[bool, str]:
"""
Decide whether to promote candidate over incumbent.
Uses bootstrap CI to ensure the improvement is statistically significant.
"""
# Must beat by minimum edge
if candidate.score <= incumbent.score + DEFAULT_POLICY_PROMOTION_MIN_EDGE_BPS:
return False, "insufficient_edge"
# Must have valid performance vector
if not candidate.performance_vector:
return False, "no_performance_vector"
# Tail-risk check: candidate tail must not be worse
cand_tail = candidate.evaluation_summary.get("mean_pnl", 0.0) - 2 * candidate.evaluation_summary.get("max_dd", 0.0)
inc_tail = incumbent.evaluation_summary.get("mean_pnl", 0.0) - 2 * incumbent.evaluation_summary.get("max_dd", 0.0)
if cand_tail < inc_tail:
return False, "tail_risk_worse"
return True, "promoted"