Files
sentiment-engine/MALKHUT/malkhut/training/cma_trainer.py

1517 lines
69 KiB
Python
Raw Normal View History

"""
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),
ParamSpec("wait_to_retry_ms", "int", 0, 2000),
ParamSpec("chase_offset_ticks", "int", 0, 10),
ParamSpec("chase_max_retries", "int", 0, 5),
ParamSpec("urgency_taker_threshold", "float", 0.1, 0.9),
ParamSpec("urgency_taker_penalty_bps", "float", 0.0, 10.0),
ParamSpec("execution_friction_threshold_bps", "float", 0.5, 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),
wait_to_retry_ms=vals.get("wait_to_retry_ms", 0),
chase_enabled=vals.get("chase_max_retries", 0) > 0,
chase_offset_ticks=vals.get("chase_offset_ticks", 1),
chase_max_retries=vals.get("chase_max_retries", 0),
urgency_taker_threshold=vals.get("urgency_taker_threshold", 0.65),
urgency_taker_penalty_bps=vals.get("urgency_taker_penalty_bps", 2.0),
execution_friction_threshold_bps=vals.get("execution_friction_threshold_bps", 3.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
# Per-scenario friction overrides (None = use AssetProfile defaults)
maker_fee_bps: Optional[float] = None
taker_fee_bps: Optional[float] = None
adverse_cost_bps: Optional[float] = None # per-fill adverse selection cost
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",
maker_fee_bps: Optional[float] = None,
taker_fee_bps: Optional[float] = None,
adverse_cost_bps: Optional[float] = None) -> None:
self.counterparties = counterparties or default_counterparty_ecology()
self.exchange_id = exchange_id
self.maker_fee_bps = maker_fee_bps
self.taker_fee_bps = taker_fee_bps
self.adverse_cost_bps = adverse_cost_bps
# --- 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",
maker_fee_bps: Optional[float] = None,
taker_fee_bps: Optional[float] = None) -> "MarketWorldState":
"""Create a MarketWorldState from AssetBehavior with realistic params.
Friction overrides take precedence over defaults.
"""
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,
maker_fee_bps=maker_fee_bps,
taker_fee_bps=taker_fee_bps)
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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), maker_fee_bps=self.maker_fee_bps, taker_fee_bps=self.taker_fee_bps,
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",
maker_fee_bps: Optional[float] = None,
taker_fee_bps: Optional[float] = None) -> MarketWorldState:
"""Create a market state using asset classification for realistic parameters.
Friction overrides take precedence over AssetProfile defaults.
"""
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=maker_fee_bps if maker_fee_bps is not None else profile.maker_fee_bps,
taker_fee_bps=taker_fee_bps if taker_fee_bps is not None else 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"