Files
sentiment-engine/MALKHUT/malkhut/training/cma_trainer.py
Codex 459215b7d8 malkhut(fix): wire workers into CMA training loop
CMAESTrainer.train() now accepts workers parameter and passes it to
evaluate_candidate(), enabling parallel episode evaluation during
actual training (not just in tests/benchmarks).

Benchmark result: ProcessPoolExecutor is optimal (4.76x speedup).
Ray is slower (0.36x) due to head init + plasma overhead for 90 scenarios.
2026-07-13 03:25:55 +02:00

1348 lines
58 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)
# ==============================================================================
# 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, ...] = ()
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) -> None:
self.counterparties = counterparties or default_counterparty_ecology()
# --- 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) -> "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)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("normal", "liquid"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("thin", "illiquid"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("wide", "volatile"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.3),),
max_steps=steps,
tags=("toxic", "adverse_selection"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("chop", "noise"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("flash_crash", "thin_book"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.2),),
max_steps=steps,
tags=("liquidity_vacuum", "extreme"),
)
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),
counterparties=(
ToxicTakerPolicy(sensitivity=0.3),
ToxicTakerPolicy(sensitivity=0.4),
LatencyArbPolicy(lead_threshold=0.3),
),
max_steps=steps,
tags=("multi_toxic", "adverse"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("trending", "momentum"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("mean_reverting", "wide_spread"),
)
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),
counterparties=(self.counterparties[3],),
max_steps=steps,
tags=("weekend", "low_participation", "thin"),
)
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),
counterparties=(LiquidationFlowPolicy(trigger_bps=30.0), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("funding_shock", "deleveraging"),
)
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),
counterparties=(
LiquidationFlowPolicy(trigger_bps=40.0),
ToxicTakerPolicy(sensitivity=0.3),
ToxicTakerPolicy(sensitivity=0.4),
),
max_steps=steps,
tags=("cascade", "liquidation", "adverse"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), self.counterparties[0]),
max_steps=steps,
tags=("divergence", "correlation_breakdown"),
)
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),
counterparties=(StaleQuoteAttackerPolicy(), LatencyArbPolicy(lead_threshold=0.4)),
max_steps=steps,
tags=("stale_quote", "latency_arb"),
)
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),
counterparties=(InventoryMarketMakerPolicy(max_inventory=0.05), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("squeeze", "inventory_risk"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("news_spike", "gap", "thin"),
)
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._make_state(symbol, bid=49950.0, ask=50050.0, bid_qty=2.0, ask_qty=2.0),
counterparties=self.counterparties,
max_steps=steps,
tags=("spread_tightening", "competition"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), ToxicTakerPolicy(sensitivity=0.5)),
max_steps=steps,
tags=("stop_hunting", "manipulation"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.2),),
max_steps=steps,
tags=("whale", "large_order", "impact"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("imbalance", "order_flow", "asymmetry"),
)
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),
counterparties=(PassiveMakerPolicy(join_probability=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("withdrawal", "liquidity_dry", "stress"),
)
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),
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"),
)
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._make_state(symbol, bid=49990.0, ask=50010.0, bid_qty=0.5, ask_qty=0.5),
counterparties=(LatencyArbPolicy(lead_threshold=0.3), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("arbitrage", "cross_venue", "price_discovery"),
)
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),
counterparties=self.counterparties,
max_steps=steps,
tags=("order_flow", "institutional", "asymmetry"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.4),),
max_steps=steps,
tags=("volatility_regime", "transition", "adaptive"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("pump_dump", "manipulation", "coordinated"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.3),),
max_steps=steps,
tags=("iceberg", "hidden_order", "gradual_impact"),
)
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),
counterparties=(
LiquidationFlowPolicy(trigger_bps=30.0),
ToxicTakerPolicy(sensitivity=0.3),
ToxicTakerPolicy(sensitivity=0.4),
),
max_steps=steps,
tags=("margin_call", "cascade", "forced_selling"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.2), ToxicTakerPolicy(sensitivity=0.3)),
max_steps=steps,
tags=("oracle_manipulation", "flash_loan", "dex"),
)
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),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), NoiseTraderPolicy()),
max_steps=steps,
tags=("whale_vs_retail", "institutional", "retail"),
)
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._make_state(symbol, bid=49980.0, ask=50020.0, bid_qty=0.3, ask_qty=0.3),
counterparties=(LatencyArbPolicy(lead_threshold=0.2), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("cross_exchange", "arb_stress", "price_discovery"),
)
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),
counterparties=(PassiveMakerPolicy(join_probability=0.1), ToxicTakerPolicy(sensitivity=0.4)),
max_steps=steps,
tags=("decay", "liquidity_withdrawal", "gradual"),
)
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),
counterparties=(
ToxicTakerPolicy(sensitivity=0.3),
LatencyArbPolicy(lead_threshold=0.3),
LiquidationFlowPolicy(trigger_bps=40.0),
),
max_steps=steps,
tags=("breakdown", "multi_failure", "stress"),
)
# --- 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)
@staticmethod
def _make_state(symbol: str, bid: float, ask: float, bid_qty: float, ask_qty: float) -> 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="bingx", 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="bingx", 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,
) -> None:
self.cwm_factory = cwm_factory
self.counterparties = counterparties or default_counterparty_ecology()
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
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),
)
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."""
from malkhut.planner.alternatives import create_planner
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)
# Trajectory metrics
total_pnl_bps = 0.0
peak_pnl = 0.0
max_dd = 0.0
fill_count = 0
maker_fills = 0
taker_fills = 0
adverse_fills = 0
cancel_count = 0
order_count = 0
noop_count = 0
entropy_sum = 0.0
max_pos_qty = 0.0
equity_start = state.account.equity
for step in range(scenario.max_steps):
# 1. Plan
intent = ExecutionIntent(
intent_id=f"ep_{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",
)
state_with_intent = 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=intent, funding_bps=state.funding_bps,
volatility_state=state.volatility_state,
market_regime=state.market_regime,
)
planned = planner.plan(root_state=state_with_intent, params=params, budget_ms=25)
# 2. Collect metrics from planned action
action = planned.selected_action
entropy_sum += planned.diagnostics.get("entropy", 0.0)
if action.kind == ActionKind.NOOP:
noop_count += 1
elif action.kind in (ActionKind.PLACE, ActionKind.CANCEL_REPLACE):
order_count += 1
elif action.kind == ActionKind.CROSS_SPREAD:
order_count += 1
fill_count += 1
taker_fills += 1
elif action.kind == ActionKind.REDUCE:
order_count += 1
fill_count += 1
elif action.kind == ActionKind.FULL_EXIT:
order_count += 1
fill_count += 1
elif action.kind == ActionKind.CANCEL:
cancel_count += 1
# 3. Transition through CWM
cp_actions = tuple(
cp.rollout_action(state, rng) for cp in scenario.counterparties
)
next_state = cwm.transition(state, (action, *cp_actions))
# 4. Track metrics
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
pos = next_state.account.positions.get(scenario.symbol)
if pos and abs(pos.qty) > max_pos_qty:
max_pos_qty = abs(pos.qty)
# 5. Check terminal
if next_state.account.equity <= 0:
state = next_state
break
state = next_state
steps = step + 1 if scenario.max_steps > 0 else 0
fill_ratio = fill_count / max(order_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_ratio,
maker_fill_count=maker_fills,
taker_fill_count=taker_fills,
adverse_fill_count=adverse_fills,
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=max_pos_qty,
diagnostics={"scenario_tags": scenario.tags},
)
def _robust_score(self, results: list[EpisodeResult], params: FulfilmentPolicyParams) -> float:
if not results:
return -float("inf")
pnl = [r.pnl_bps for r in results]
pnl_sorted = sorted(pnl)
tail_idx = max(0, int(TAIL_QUANTILE * (len(pnl_sorted) - 1)))
p05 = pnl_sorted[tail_idx]
mean = sum(pnl) / len(pnl)
adverse = sum(r.adverse_fill_count for r in results) / max(sum(r.order_count for r in results), 1)
slippage = sum(r.avg_slippage_bps for r in results) / len(results)
liq = sum(r.liquidation_near_miss_count for r in results)
dd = sum(r.max_drawdown_bps for r in results) / len(results)
entropy = sum(r.policy_entropy_avg for r in results) / len(results)
score = 0.0
score += mean * 10.0 # HEAVY PnL weight
score += params.robust_tail_weight * p05 * 5.0
score -= params.toxic_counterparty_weight * adverse * 100.0
score -= slippage
score -= 10.0 * liq
score -= 2.0 * dd # penalize drawdown
score += params.w_policy_entropy * entropy * 0.1 # reduced entropy weight
# NOOP penalty: penalize strategies that don't trade
noop_ratios = [r.noop_count / max(r.steps, 1) for r in results]
avg_noop_ratio = sum(noop_ratios) / len(noop_ratios) if noop_ratios else 0.0
score -= avg_noop_ratio * 50.0 # heavy penalty for not trading
# Fill reward: reward strategies that actually get fills
fill_ratios = [r.fill_count / max(r.order_count, 1) for r in results]
avg_fill_ratio = sum(fill_ratios) / len(fill_ratios) if fill_ratios else 0.0
score += avg_fill_ratio * 20.0 # reward fills
return score
return score
@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"