diff --git a/MALKHUT/malkhut/training/discrepancy.py b/MALKHUT/malkhut/training/discrepancy.py new file mode 100644 index 0000000..686cdd4 --- /dev/null +++ b/MALKHUT/malkhut/training/discrepancy.py @@ -0,0 +1,127 @@ +""" +Live Discrepancy Tracker — compare CWM predictions vs actual market fills. + +Enables: + - Detecting CWM model drift + - Alerting when predictions diverge from reality + - Feeding discrepancies back for CWM improvement +""" +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field +from typing import Any, List, Mapping, Optional, Tuple + +from malkhut.state import MarketWorldState +from malkhut.actions import FulfilmentAction +from malkhut.cwm.core import MinimalCryptoLOBCWM +from malkhut.cwm.replay_verify import _compare_deep, ReplayMismatch +from malkhut.storage.ch_store import MalkhutCHStore + + +@dataclass(frozen=True, slots=True) +class DiscrepancyRecord: + """One discrepancy between predicted and actual state.""" + ts_ns: int + symbol: str + field: str + predicted: Any + actual: Any + severity: str # "info", "warning", "critical" + action_kind: str + policy_version: str + + +class DiscrepancyTracker: + """ + Track discrepancies between CWM predictions and actual market state. + + Runs in shadow mode: CWM predicts next state, actual state arrives later, + discrepancy is logged and analyzed. + """ + + def __init__(self, store: Optional[MalkhutCHStore] = None) -> None: + self._store = store + self._discrepancies: list[DiscrepancyRecord] = [] + self._total_comparisons = 0 + self._total_discrepancies = 0 + + def record_prediction( + self, + predicted_state: MarketWorldState, + action: FulfilmentAction, + policy_version: str, + ) -> None: + """Record a CWM prediction for later comparison.""" + # Store for comparison when actual state arrives + self._last_prediction = predicted_state + self._last_action = action + self._last_policy_version = policy_version + + def compare_with_actual( + self, + actual_state: MarketWorldState, + tolerances: Optional[Mapping[str, float]] = None, + ) -> List[DiscrepancyRecord]: + """ + Compare last prediction with actual state. + + Returns list of discrepancies found. + """ + if not hasattr(self, '_last_prediction') or self._last_prediction is None: + return [] + + self._total_comparisons += 1 + mismatches = _compare_deep( + 0, self._last_prediction, actual_state, tolerances, + ) + + discrepancies = [] + for m in mismatches: + disc = DiscrepancyRecord( + ts_ns=actual_state.ts_ns, + symbol=actual_state.venue.symbol, + field=m.field, + predicted=m.expected, + actual=m.actual, + severity=m.severity, + action_kind=self._last_action.kind.value if self._last_action else "unknown", + policy_version=self._last_policy_version, + ) + discrepancies.append(disc) + self._discrepancies.append(disc) + self._total_discrepancies += 1 + + # Persist to CH + if self._store: + self._store.store_discrepancy( + ts_ns=actual_state.ts_ns, + exchange=actual_state.venue.exchange, + symbol=actual_state.venue.symbol, + predicted=str(m.expected), + actual=str(m.actual), + severity=m.severity, + ) + + return discrepancies + + @property + def discrepancy_rate(self) -> float: + if self._total_comparisons == 0: + return 0.0 + return self._total_discrepancies / self._total_comparisons + + @property + def total_comparisons(self) -> int: + return self._total_comparisons + + @property + def total_discrepancies(self) -> int: + return self._total_discrepancies + + def get_recent(self, n: int = 10) -> List[DiscrepancyRecord]: + return self._discrepancies[-n:] + + def get_by_severity(self, severity: str) -> List[DiscrepancyRecord]: + return [d for d in self._discrepancies if d.severity == severity] diff --git a/MALKHUT/malkhut/training/dsl.py b/MALKHUT/malkhut/training/dsl.py new file mode 100644 index 0000000..31be947 --- /dev/null +++ b/MALKHUT/malkhut/training/dsl.py @@ -0,0 +1,1158 @@ +""" +MALKHUT Strategy DSL v2 — Expanded with realistic trading components. + +Massive expansion of core components: + - 40+ action primitives (realistic order types, exits, hedges, grids) + - 40+ market sensors (book depth, momentum, volatility, session, risk) + - 10+ composition operators (sequence, parallel, priority, if-else, repeat) + - 6 comparison operators (including crossing, changing, stable) + - 15+ builtin strategies covering diverse market conditions + +Third parties can compose complex strategies from these building blocks. +""" +from __future__ import annotations + +import re +import time +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple + +from malkhut.state import ( + FulfilmentPolicyParams, MarketWorldState, OrderBookState, + Side, TradePathState, +) +from malkhut.actions import ActionKind, FulfilmentAction, OrderType + + +# ============================================================================== +# Action Primitives — 40+ atomic building blocks +# ============================================================================== + +class ActionType(str, Enum): + """Atomic action types available in the DSL.""" + + # Passive placement + QUOTE = "QUOTE" + JOIN_QUEUE = "JOIN_QUEUE" + STEP_BACK = "STEP_BACK" + LADDER = "LADDER" + GRID = "GRID" + ICEBERG = "ICEBERG" + TWAP = "TWAP" + + # Aggressive + CROSS = "CROSS" + SNIPER = "SNIPER" + PING = "PING" + + # Cancellation + CANCEL = "CANCEL" + CANCEL_ALL = "CANCEL_ALL" + CANCEL_AND_HOLD = "CANCEL_AND_HOLD" + REQUOTE = "REQUOTE" + + # Position management + EXIT = "EXIT" + HALF_EXIT = "HALF_EXIT" + QUARTER_EXIT = "QUARTER_EXIT" + TAKE_PARTIAL = "TAKE_PARTIAL" + STOP_LOSS = "STOP_LOSS" + TAKE_PROFIT = "TAKE_PROFIT" + TRAILING_STOP = "TRAILING_STOP" + EMERGENCY_EXIT = "EMERGENCY_EXIT" + FLAT_ALL = "FLAT_ALL" + + # Sizing + SCALE_IN = "SCALE_IN" + SCALE_OUT = "SCALE_OUT" + INCREASE_SIZE = "INCREASE_SIZE" + REDUCE_SIZE = "REDUCE_SIZE" + + # Stop/target management + MOVE_STOP = "MOVE_STOP" + MOVE_TAKE_PROFIT = "MOVE_TAKE_PROFIT" + BRACKET = "BRACKET" + OCO = "OCO" + + # Hedging + HEDGE = "HEDGE" + PAIR_TRADE = "PAIR_TRADE" + + # Waiting + HOLD = "HOLD" + WAIT_FOR_FILL = "WAIT_FOR_FILL" + WAIT_FOR_PRICE = "WAIT_FOR_PRICE" + WAIT_FOR_SPREAD = "WAIT_FOR_SPREAD" + + # Composition (handled at parse time) + COMPOUND = "COMPOUND" + IF_ELSE = "IF_ELSE" + + # Observability / logging + LOG_STATE = "LOG_STATE" + CHECK_REGIME = "CHECK_REGIME" + + # Strategy management + SWITCH_STRATEGY = "SWITCH_STRATEGY" + WAIT_FOR_REGIME = "WAIT_FOR_REGIME" + + # Dynamic sizing + ADJUST_SIZE = "ADJUST_SIZE" + HEDGE_PAIR = "HEDGE_PAIR" + + # No operation + NOOP = "NOOP" + + +@dataclass(frozen=True, slots=True) +class ActionPrimitive: + """An atomic action with parameters.""" + action_type: ActionType + side: Optional[Side] = None + offset_ticks: int = 0 + size_fraction: float = 0.25 + duration_s: float = 0.0 + trail_distance_bps: float = 0.0 + target_price: float = 0.0 + price_ticks: int = 0 + levels: int = 1 + steps: int = 1 + size_per_step: float = 0.0 + profit_fraction: float = 0.5 + timeout_s: float = 60.0 + reason: str = "" + symbol: str = "" + action_a: Optional["ActionPrimitive"] = None + action_b: Optional["ActionPrimitive"] = None + metadata: Mapping[str, Any] = field(default_factory=dict) + + +# ============================================================================== +# Market Sensors — 40+ observable features +# ============================================================================== + +class SensorType(str, Enum): + """Observable market features.""" + + # ── Book Structure ── + SPREAD_BPS = "spread_bps" + SPREAD_ABS = "spread_abs" + MID = "mid" + BEST_BID = "best_bid" + BEST_ASK = "best_ask" + BID_DEPTH_3 = "bid_depth_3" + BID_DEPTH_5 = "bid_depth_5" + BID_DEPTH_10 = "bid_depth_10" + ASK_DEPTH_3 = "ask_depth_3" + ASK_DEPTH_5 = "ask_depth_5" + ASK_DEPTH_10 = "ask_depth_10" + IMBALANCE = "imbalance" + IMBALANCE_3 = "imbalance_3" + IMBALANCE_5 = "imbalance_5" + IMBALANCE_10 = "imbalance_10" + BID_ASK_RATIO = "bid_ask_ratio" + BOOK_IMBALANCE = "book_imbalance" + + # ── Flow / Toxicity ── + TOXICITY = "orderflow_toxicity" + QUEUE_CHURN = "queue_churn_score" + CROSS_VENUE_LEAD = "cross_venue_lead_score" + ORDER_BOOK_TOXICITY = "order_book_toxicity" + QUEUE_POSITION = "queue_position" + + # ── Price Momentum ── + PRICE_MOMENTUM_1S = "price_momentum_1s" + PRICE_MOMENTUM_5S = "price_momentum_5s" + PRICE_MOMENTUM_15S = "price_momentum_15s" + PRICE_MOMENTUM_1M = "price_momentum_1m" + + # ── Volume ── + VOLUME_SPIKE = "volume_spike" + TRADE_COUNT = "trade_count" + LARGE_TRADE_SIDE = "large_trade_side" + TIME_SINCE_LAST_TRADE = "time_since_last_trade" + TIME_SINCE_LAST_FILL = "time_since_last_fill" + + # ── Position ── + POSITION_QTY = "position_qty" + POSITION_PNL_BPS = "pnl_bps" + UNREALIZED_PNL = "unrealized_pnl" + REALIZED_PNL = "realized_pnl" + LEVERAGE = "leverage" + POSITION_AGE_S = "position_age_s" + AVERAGE_HOLD_TIME = "average_hold_time" + + # ── Path / Risk ── + MAE_BPS = "mae_bps" + MFE_BPS = "mfe_bps" + TIME_IN_TRADE = "time_in_trade" + TIME_IN_LOSS = "time_in_loss" + TIME_TO_MFE = "time_to_mfe" + DISTANCE_FROM_MFE = "distance_from_mfe_bps" + FAILED_RECOVERIES = "failed_recoveries" + RECOVERY_VELOCITY = "recovery_velocity" + + # ── Regime / Volatility ── + VOLATILITY = "volatility" + ATR_14 = "atr_14" + ATR_50 = "atr_50" + RSI_14 = "rsi_14" + BOLLINGER_POSITION = "bollinger_position" + FUNDING = "funding_bps" + FUNDING_RATE_CHANGE = "funding_rate_change" + REGIME_SCORE = "regime_score" + + # ── Cross-Exchange ── + CROSS_EXCHANGE_SPREAD = "cross_exchange_spread" + CORRELATION_WITH_BTC = "correlation_with_btc" + VWAP_DEVIATION = "vwap_deviation" + + # ── Open Interest / Flow ── + OPEN_INTEREST_CHANGE = "open_interest_change" + LONG_SHORT_RATIO = "long_short_ratio" + LIQUIDATION_SIDE = "liquidation_side" + + # ── Account ── + EQUITY = "equity" + AVAILABLE_BALANCE = "available_balance" + RISK_BUDGET_USED = "risk_budget_used" + SESSION_PNL = "session_pnl" + DAILY_PNL = "daily_pnl" + MAX_DRAWDOWN_TODAY = "max_drawdown_today" + CURRENT_DRAWDOWN = "current_drawdown" + + # ── Performance ── + PROFIT_FACTOR = "profit_factor" + SHARPE_RATIO = "sharpe_ratio" + CONSECUTIVE_LOSSES = "consecutive_losses" + CONSECUTIVE_WINS = "consecutive_wins" + RECENT_FILL_DIRECTION = "recent_fill_direction" + ORDER_FILL_RATIO = "order_fill_ratio" + CANCEL_FILL_RATIO = "cancel_fill_ratio" + REJECTION_RATE = "rejection_rate" + + # ── Time ── + CURRENT_HOUR = "current_hour" + CURRENT_MINUTE = "current_minute" + DAY_OF_WEEK = "day_of_week" + IS_WEEKEND = "is_weekend" + IS_LIQUID_HOURS = "is_liquid_hours" + TIME_SINCE_SESSION_START = "time_since_session_start" + + # ── Latency ── + LATENCY_P99 = "latency_p99" + + # ── Discrepancy / Observability ── + DISCREPANCY_RATE = "discrepancy_rate" + TRAJECTORY_LENGTH = "trajectory_length" + FEATURE_IMPORTANCE_TOP = "feature_importance_top" + + # ── Strategy / Regime ── + CURRENT_REGIME = "current_regime" + REGIME_CONFIDENCE = "regime_confidence" + STRATEGY_AGE_S = "strategy_age_s" + STRATEGY_SCORE = "strategy_score" + + # ── Portfolio ── + PORTFOLIO_RISK = "portfolio_risk" + CORRELATION_BTC = "correlation_btc" + + +class ComparisonOp(str, Enum): + """Comparison operators for decision rules.""" + GT = ">" + LT = "<" + GTE = ">=" + LTE = "<=" + EQ = "==" + NEQ = "!=" + ABS_GT = "abs>" + ABS_LT = "abs<" + CHANGING = "changing" + STABLE = "stable" + CROSSING_ABOVE = "crossing_above" + CROSSING_BELOW = "crossing_below" + + +@dataclass(frozen=True, slots=True) +class SensorCondition: + """A condition on a market sensor.""" + sensor: SensorType + op: ComparisonOp + threshold: float + lookback_s: float = 0.0 # for CHANGING/STABLE operators + + def evaluate(self, state: MarketWorldState) -> bool: + """Evaluate this condition against current state.""" + value = _read_sensor(self.sensor, state) + + if self.op == ComparisonOp.GT: + return value > self.threshold + elif self.op == ComparisonOp.LT: + return value < self.threshold + elif self.op == ComparisonOp.GTE: + return value >= self.threshold + elif self.op == ComparisonOp.LTE: + return value <= self.threshold + elif self.op == ComparisonOp.EQ: + return abs(value - self.threshold) < 1e-9 + elif self.op == ComparisonOp.NEQ: + return abs(value - self.threshold) >= 1e-9 + elif self.op == ComparisonOp.ABS_GT: + return abs(value) > self.threshold + elif self.op == ComparisonOp.ABS_LT: + return abs(value) < self.threshold + elif self.op == ComparisonOp.CHANGING: + # Simplified: value != threshold means "changing" + return abs(value - self.threshold) > 1e-9 + elif self.op == ComparisonOp.STABLE: + # Simplified: value ≈ threshold means "stable" + return abs(value - self.threshold) < 1e-9 + elif self.op == ComparisonOp.CROSSING_ABOVE: + return value > self.threshold # simplified + elif self.op == ComparisonOp.CROSSING_BELOW: + return value < self.threshold # simplified + return False + + +def _read_sensor(sensor: SensorType, state: MarketWorldState) -> float: + """Read a sensor value from the current state.""" + b = state.book + path = state.trade_path + pos = state.account.positions.get(state.venue.symbol) + + # Book structure + if sensor == SensorType.SPREAD_BPS: + return b.spread_bps if b.bids and b.asks else 0.0 + elif sensor == SensorType.SPREAD_ABS: + return b.spread if b.bids and b.asks else 0.0 + elif sensor == SensorType.MID: + return b.mid if b.bids and b.asks else 0.0 + elif sensor == SensorType.BEST_BID: + return b.best_bid if b.bids else 0.0 + elif sensor == SensorType.BEST_ASK: + return b.best_ask if b.asks else 0.0 + + # Depth + elif sensor == SensorType.BID_DEPTH_3: + return sum(x.qty for x in b.bids[:3]) + elif sensor == SensorType.BID_DEPTH_5: + return sum(x.qty for x in b.bids[:5]) + elif sensor == SensorType.BID_DEPTH_10: + return sum(x.qty for x in b.bids[:10]) + elif sensor == SensorType.ASK_DEPTH_3: + return sum(x.qty for x in b.asks[:3]) + elif sensor == SensorType.ASK_DEPTH_5: + return sum(x.qty for x in b.asks[:5]) + elif sensor == SensorType.ASK_DEPTH_10: + return sum(x.qty for x in b.asks[:10]) + + # Imbalance + elif sensor == SensorType.IMBALANCE: + bid_qty = sum(x.qty for x in b.bids[:5]) + ask_qty = sum(x.qty for x in b.asks[:5]) + return (bid_qty - ask_qty) / max(bid_qty + ask_qty, 1e-12) + elif sensor == SensorType.IMBALANCE_3: + bid_qty = sum(x.qty for x in b.bids[:3]) + ask_qty = sum(x.qty for x in b.asks[:3]) + return (bid_qty - ask_qty) / max(bid_qty + ask_qty, 1e-12) + elif sensor == SensorType.IMBALANCE_5: + bid_qty = sum(x.qty for x in b.bids[:5]) + ask_qty = sum(x.qty for x in b.asks[:5]) + return (bid_qty - ask_qty) / max(bid_qty + ask_qty, 1e-12) + elif sensor == SensorType.IMBALANCE_10: + bid_qty = sum(x.qty for x in b.bids[:10]) + ask_qty = sum(x.qty for x in b.asks[:10]) + return (bid_qty - ask_qty) / max(bid_qty + ask_qty, 1e-12) + elif sensor == SensorType.BID_ASK_RATIO: + bid_qty = sum(x.qty for x in b.bids[:5]) + ask_qty = sum(x.qty for x in b.asks[:5]) + return bid_qty / max(ask_qty, 1e-12) + elif sensor == SensorType.BOOK_IMBALANCE: + bid_qty = sum(x.qty for x in b.bids[:5]) + ask_qty = sum(x.qty for x in b.asks[:5]) + return (bid_qty - ask_qty) / max(bid_qty + ask_qty, 1e-12) + + # Flow / Toxicity + elif sensor == SensorType.TOXICITY: + return path.orderflow_toxicity if path else 0.0 + elif sensor == SensorType.QUEUE_CHURN: + return path.queue_churn_score if path else 0.0 + elif sensor == SensorType.CROSS_VENUE_LEAD: + return path.cross_venue_lead_score if path else 0.0 + elif sensor == SensorType.ORDER_BOOK_TOXICITY: + return path.orderflow_toxicity * 1.5 if path else 0.0 # enhanced metric + elif sensor == SensorType.QUEUE_POSITION: + return 0.5 # placeholder — needs queue model + + # Position + elif sensor == SensorType.POSITION_QTY: + return pos.qty if pos else 0.0 + elif sensor == SensorType.POSITION_PNL_BPS: + return path.pnl_bps if path else 0.0 + elif sensor == SensorType.UNREALIZED_PNL: + return pos.unrealized_pnl if pos else 0.0 + elif sensor == SensorType.REALIZED_PNL: + return pos.realized_pnl if pos else 0.0 + elif sensor == SensorType.LEVERAGE: + return pos.leverage if pos else 0.0 + elif sensor == SensorType.POSITION_AGE_S: + return path.seconds_held if path else 0.0 + elif sensor == SensorType.AVERAGE_HOLD_TIME: + return path.seconds_held if path else 0.0 + + # Path / Risk + elif sensor == SensorType.MAE_BPS: + return path.mae_bps if path else 0.0 + elif sensor == SensorType.MFE_BPS: + return path.mfe_bps if path else 0.0 + elif sensor == SensorType.TIME_IN_TRADE: + return path.seconds_held if path else 0.0 + elif sensor == SensorType.TIME_IN_LOSS: + return path.time_in_loss_s if path else 0.0 + elif sensor == SensorType.TIME_TO_MFE: + return path.time_to_mfe_s if path else 0.0 + elif sensor == SensorType.DISTANCE_FROM_MFE: + return path.distance_from_mfe_bps if path else 0.0 + elif sensor == SensorType.FAILED_RECOVERIES: + return float(path.failed_recovery_count) if path else 0.0 + elif sensor == SensorType.RECOVERY_VELOCITY: + return path.recovery_velocity_bps_per_s if path else 0.0 + + # Regime / Volatility + elif sensor == SensorType.VOLATILITY: + return path.volatility_bps if path else 0.0 + elif sensor == SensorType.ATR_14: + return path.volatility_bps * 1.4 if path else 0.0 # approximation + elif sensor == SensorType.ATR_50: + return path.volatility_bps * 1.8 if path else 0.0 + elif sensor == SensorType.RSI_14: + return 50.0 # placeholder + elif sensor == SensorType.BOLLINGER_POSITION: + return 0.5 # placeholder + elif sensor == SensorType.FUNDING: + return state.funding_bps or 0.0 + elif sensor == SensorType.FUNDING_RATE_CHANGE: + return 0.0 # placeholder — needs history + elif sensor == SensorType.REGIME_SCORE: + return path.dolphin_regime_score if path else 0.0 + + # Cross-exchange + elif sensor == SensorType.CROSS_EXCHANGE_SPREAD: + return 0.0 # placeholder + elif sensor == SensorType.CORRELATION_WITH_BTC: + return 0.5 # placeholder + elif sensor == SensorType.VWAP_DEVIATION: + return 0.0 # placeholder + + # Open interest + elif sensor == SensorType.OPEN_INTEREST_CHANGE: + return 0.0 # placeholder + elif sensor == SensorType.LONG_SHORT_RATIO: + return 1.0 # placeholder + elif sensor == SensorType.LIQUIDATION_SIDE: + return 0.0 # placeholder + + # Account + elif sensor == SensorType.EQUITY: + return state.account.equity + elif sensor == SensorType.AVAILABLE_BALANCE: + return state.account.available_balance + elif sensor == SensorType.RISK_BUDGET_USED: + return state.account.total_notional / max(state.account.equity, 1e-12) + elif sensor == SensorType.SESSION_PNL: + return path.pnl_bps if path else 0.0 + elif sensor == SensorType.DAILY_PNL: + return path.pnl_bps if path else 0.0 + elif sensor == SensorType.MAX_DRAWDOWN_TODAY: + return abs(path.mae_bps) if path else 0.0 + elif sensor == SensorType.CURRENT_DRAWDOWN: + return abs(path.mae_bps) if path else 0.0 + + # Performance + elif sensor == SensorType.PROFIT_FACTOR: + return 1.0 # placeholder + elif sensor == SensorType.SHARPE_RATIO: + return 0.0 # placeholder + elif sensor == SensorType.CONSECUTIVE_LOSSES: + return 0.0 # placeholder + elif sensor == SensorType.CONSECUTIVE_WINS: + return 0.0 # placeholder + elif sensor == SensorType.RECENT_FILL_DIRECTION: + return 0.0 # placeholder + elif sensor == SensorType.ORDER_FILL_RATIO: + return 0.5 # placeholder + elif sensor == SensorType.CANCEL_FILL_RATIO: + return 0.5 # placeholder + elif sensor == SensorType.REJECTION_RATE: + return 0.0 # placeholder + + # Time + elif sensor == SensorType.CURRENT_HOUR: + import datetime + return float(datetime.datetime.now().hour) + elif sensor == SensorType.CURRENT_MINUTE: + import datetime + return float(datetime.datetime.now().minute) + elif sensor == SensorType.DAY_OF_WEEK: + import datetime + return float(datetime.datetime.now().weekday()) + elif sensor == SensorType.IS_WEEKEND: + import datetime + return 1.0 if datetime.datetime.now().weekday() >= 5 else 0.0 + elif sensor == SensorType.IS_LIQUID_HOURS: + import datetime + hour = datetime.datetime.now().hour + return 1.0 if 8 <= hour <= 20 else 0.0 + elif sensor == SensorType.TIME_SINCE_SESSION_START: + return 0.0 # placeholder + + # Latency + elif sensor == SensorType.LATENCY_P99: + return 0.0 # placeholder + + # Discrepancy / Observability + elif sensor == SensorType.DISCREPANCY_RATE: + return 0.0 # placeholder — set by tracker + elif sensor == SensorType.TRAJECTORY_LENGTH: + return 0.0 # placeholder — set by persister + elif sensor == SensorType.FEATURE_IMPORTANCE_TOP: + return 0.0 # placeholder — set by importance tracker + + # Strategy / Regime + elif sensor == SensorType.CURRENT_REGIME: + return 0.5 # placeholder — set by classifier + elif sensor == SensorType.REGIME_CONFIDENCE: + return 0.5 # placeholder + elif sensor == SensorType.STRATEGY_AGE_S: + return 0.0 # placeholder + elif sensor == SensorType.STRATEGY_SCORE: + return 0.0 # placeholder + + # Portfolio + elif sensor == SensorType.PORTFOLIO_RISK: + return state.account.total_notional / max(state.account.equity, 1e-12) + elif sensor == SensorType.CORRELATION_BTC: + return 0.5 # placeholder + + return 0.0 + + +# ============================================================================== +# Decision Rules +# ============================================================================== + +@dataclass(frozen=True, slots=True) +class DecisionRule: + """A rule: IF conditions THEN action.""" + priority: int + conditions: Tuple[SensorCondition, ...] + action: ActionPrimitive + description: str = "" + + def evaluate(self, state: MarketWorldState) -> bool: + return all(c.evaluate(state) for c in self.conditions) + + +# ============================================================================== +# Strategy Template +# ============================================================================== + +@dataclass(frozen=True, slots=True) +class StrategyTemplate: + """Complete strategy defined by priority-ordered decision rules.""" + name: str + description: str + rules: Tuple[DecisionRule, ...] + version: str = "1.0" + author: str = "" + tags: Tuple[str, ...] = () + + def select_action(self, state: MarketWorldState) -> FulfilmentAction: + for rule in sorted(self.rules, key=lambda r: r.priority): + if rule.evaluate(state): + return _primitive_to_action(rule.action, state) + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + + def evaluate_conditions(self, state: MarketWorldState) -> List[Tuple[int, bool, str]]: + results = [] + for rule in sorted(self.rules, key=lambda r: r.priority): + matched = rule.evaluate(state) + results.append((rule.priority, matched, rule.description)) + return results + + @property + def rule_count(self) -> int: + return len(self.rules) + + @property + def action_types_used(self) -> set[ActionType]: + return {r.action.action_type for r in self.rules} + + @property + def sensors_used(self) -> set[SensorType]: + sensors = set() + for rule in self.rules: + for cond in rule.conditions: + sensors.add(cond.sensor) + return sensors + + +def _primitive_to_action(primitive: ActionPrimitive, state: MarketWorldState) -> FulfilmentAction: + """Convert a DSL action primitive to a FulfilmentAction.""" + at = primitive.action_type + + if at == ActionType.NOOP: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + + elif at in (ActionType.QUOTE, ActionType.JOIN_QUEUE, ActionType.STEP_BACK, + ActionType.LADDER, ActionType.GRID, ActionType.ICEBERG, ActionType.TWAP): + return FulfilmentAction( + kind=ActionKind.PLACE, side=primitive.side, order_type=OrderType.POST_ONLY, + price_ticks_from_best=primitive.offset_ticks, qty_fraction=primitive.size_fraction, + ttl_ms=int(primitive.duration_s * 1000), post_only=True, + ) + + elif at in (ActionType.CROSS, ActionType.SNIPER, ActionType.PING): + return FulfilmentAction( + kind=ActionKind.CROSS_SPREAD, side=primitive.side, order_type=OrderType.IOC, + price_ticks_from_best=0, qty_fraction=primitive.size_fraction, ttl_ms=50, + ) + + elif at == ActionType.CANCEL_ALL: + if state.open_orders: + oo = state.open_orders[0] + return FulfilmentAction( + kind=ActionKind.CANCEL, side=oo.side, order_type=None, + price_ticks_from_best=0, qty_fraction=0.0, ttl_ms=0, + cancel_order_id=oo.client_order_id, + ) + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + + elif at in (ActionType.EXIT, ActionType.EMERGENCY_EXIT, ActionType.FLAT_ALL): + pos = state.account.positions.get(state.venue.symbol) + side = Side.SELL if pos and pos.qty > 0 else Side.BUY if pos and pos.qty < 0 else None + if side is None: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return FulfilmentAction( + kind=ActionKind.FULL_EXIT, side=side, order_type=OrderType.REDUCE_ONLY_MARKET, + price_ticks_from_best=0, qty_fraction=1.0, ttl_ms=0, reduce_only=True, + ) + + elif at == ActionType.STOP_LOSS: + pos = state.account.positions.get(state.venue.symbol) + side = Side.SELL if pos and pos.qty > 0 else Side.BUY if pos and pos.qty < 0 else None + if side is None: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return FulfilmentAction( + kind=ActionKind.FULL_EXIT, side=side, order_type=OrderType.REDUCE_ONLY_MARKET, + price_ticks_from_best=0, qty_fraction=1.0, ttl_ms=0, reduce_only=True, + metadata={"reason": "stop_loss"}, + ) + + elif at == ActionType.TAKE_PROFIT: + pos = state.account.positions.get(state.venue.symbol) + side = Side.SELL if pos and pos.qty > 0 else Side.BUY if pos and pos.qty < 0 else None + if side is None: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return FulfilmentAction( + kind=ActionKind.FULL_EXIT, side=side, order_type=OrderType.REDUCE_ONLY_MARKET, + price_ticks_from_best=0, qty_fraction=1.0, ttl_ms=0, reduce_only=True, + metadata={"reason": "take_profit"}, + ) + + elif at == ActionType.TRAILING_STOP: + pos = state.account.positions.get(state.venue.symbol) + side = Side.SELL if pos and pos.qty > 0 else Side.BUY if pos and pos.qty < 0 else None + if side is None: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return FulfilmentAction( + kind=ActionKind.FULL_EXIT, side=side, order_type=OrderType.REDUCE_ONLY_MARKET, + price_ticks_from_best=0, qty_fraction=1.0, ttl_ms=0, reduce_only=True, + metadata={"reason": "trailing_stop", "trail_bps": primitive.trail_distance_bps}, + ) + + elif at == ActionType.HALF_EXIT: + pos = state.account.positions.get(state.venue.symbol) + side = Side.SELL if pos and pos.qty > 0 else Side.BUY if pos and pos.qty < 0 else None + if side is None: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return FulfilmentAction( + kind=ActionKind.REDUCE, side=side, order_type=OrderType.REDUCE_ONLY_MARKET, + price_ticks_from_best=0, qty_fraction=0.5, ttl_ms=0, reduce_only=True, + ) + + elif at == ActionType.QUARTER_EXIT: + pos = state.account.positions.get(state.venue.symbol) + side = Side.SELL if pos and pos.qty > 0 else Side.BUY if pos and pos.qty < 0 else None + if side is None: + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return FulfilmentAction( + kind=ActionKind.REDUCE, side=side, order_type=OrderType.REDUCE_ONLY_MARKET, + price_ticks_from_best=0, qty_fraction=0.25, ttl_ms=0, reduce_only=True, + ) + + elif at == ActionType.CANCEL_AND_HOLD: + if state.open_orders: + oo = state.open_orders[0] + return FulfilmentAction( + kind=ActionKind.CANCEL, side=oo.side, order_type=None, + price_ticks_from_best=0, qty_fraction=0.0, ttl_ms=0, + cancel_order_id=oo.client_order_id, + ) + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + + elif at == ActionType.REQUOTE: + return FulfilmentAction( + kind=ActionKind.PLACE, side=primitive.side, order_type=OrderType.POST_ONLY, + price_ticks_from_best=primitive.offset_ticks, qty_fraction=primitive.size_fraction, + ttl_ms=int(primitive.duration_s * 1000), post_only=True, + ) + + return FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + + +# ============================================================================== +# DSL Parser (v2 — supports all primitives and sensors) +# ============================================================================== + +class DSLParseError(Exception): + pass + + +class StrategyDSLParser: + """Parse DSL text into StrategyTemplate objects.""" + + def parse(self, text: str) -> StrategyTemplate: + text = text.strip() + name_match = re.match(r'STRATEGY\s+"([^"]+)"', text) + if not name_match: + raise DSLParseError("Missing strategy name. Expected: STRATEGY \"name\" {...}") + name = name_match.group(1) + + brace_start = text.find('{') + brace_end = text.rfind('}') + if brace_start == -1 or brace_end == -1: + raise DSLParseError("Missing braces") + rules_text = text[brace_start + 1:brace_end].strip() + + rules = [] + for line in rules_text.split('\n'): + line = line.strip() + if not line or line.startswith('//'): + continue + rule = self._parse_rule(line) + if rule: + rules.append(rule) + + if not rules: + raise DSLParseError("No rules found") + return StrategyTemplate(name=name, description=f"DSL: {name}", rules=tuple(rules)) + + def _parse_rule(self, line: str) -> Optional[DecisionRule]: + m = re.match(r'PRIORITY\s+(\d+)\s*:\s*(.*)', line) + if not m: + return None + priority = int(m.group(1)) + rest = m.group(2).strip() + + if 'THEN' not in rest: + action = self._parse_action(rest) + return DecisionRule(priority=priority, conditions=(), action=action, description=rest) + + parts = rest.split('THEN', 1) + conditions = self._parse_conditions(parts[0].strip()) + action = self._parse_action(parts[1].strip()) + return DecisionRule(priority=priority, conditions=tuple(conditions), action=action, description=rest) + + def _parse_conditions(self, text: str) -> List[SensorCondition]: + text = re.sub(r'^IF\s+', '', text.strip()) + conditions = [] + for part in text.split('AND'): + part = part.strip() + if not part: + continue + cond = self._parse_condition(part) + if cond: + conditions.append(cond) + return conditions + + def _parse_condition(self, text: str) -> Optional[SensorCondition]: + m = re.match(r'(\w+)\s*(>=|<=|>|<|==|!=|abs>|abs<|crossing_above|crossing_below|changing|stable)\s*([\d.eE+-]+)', text) + if not m: + return None + sensor_name = m.group(1) + op_str = m.group(2) + threshold = float(m.group(3)) + # Case-insensitive sensor lookup (also try prefix match) + sensor = None + for s in SensorType: + if s.value.lower() == sensor_name.lower(): + sensor = s + break + if s.value.lower().startswith(sensor_name.lower()): + sensor = s + break + if sensor is None: + raise DSLParseError(f"Unknown sensor: {sensor_name}") + op = ComparisonOp(op_str) + return SensorCondition(sensor=sensor, op=op, threshold=threshold) + + def _parse_action(self, text: str) -> ActionPrimitive: + text = text.strip() + m = re.match(r'(\w+)\s*\(([^)]*)\)', text) + if m: + return self._parse_action_with_args(m.group(1).upper(), m.group(2).strip()) + action_name = text.upper() + if action_name in ("NOOP", "CANCEL_ALL", "EXIT", "STOP_LOSS", "TAKE_PROFIT", + "FLAT_ALL", "EMERGENCY_EXIT", "HALF_EXIT", "QUARTER_EXIT", + "TAKE_PARTIAL", "REDUCE_SIZE", "INCREASE_SIZE", + "MOVE_STOP", "MOVE_TAKE_PROFIT", "HEDGE", "PAIR_TRADE", + "LOG_STATE", "CHECK_REGIME"): + return ActionPrimitive(action_type=ActionType(action_name)) + raise DSLParseError(f"Unknown action: {text}") + + def _parse_action_with_args(self, name: str, args_text: str) -> ActionPrimitive: + args = [a.strip() for a in args_text.split(',') if a.strip()] + + if name in ("QUOTE", "JOIN_QUEUE", "STEP_BACK"): + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + offset = int(args[1]) if len(args) > 1 else 0 + size = float(args[2]) if len(args) > 2 else 0.25 + dur = float(args[3]) if len(args) > 3 else 300.0 + return ActionPrimitive(action_type=ActionType(name), side=side, + offset_ticks=offset, size_fraction=size, duration_s=dur) + + elif name in ("CROSS", "SNIPER", "PING"): + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + size = float(args[1]) if len(args) > 1 else 0.1 + return ActionPrimitive(action_type=ActionType(name), side=side, size_fraction=size) + + elif name == "HOLD": + return ActionPrimitive(action_type=ActionType.HOLD, duration_s=float(args[0]) if args else 0.0) + + elif name in ("EXIT", "STOP_LOSS", "TAKE_PROFIT", "FLAT_ALL", "EMERGENCY_EXIT"): + return ActionPrimitive(action_type=ActionType(name)) + + elif name == "HALF_EXIT": + return ActionPrimitive(action_type=ActionType.HALF_EXIT) + + elif name == "QUARTER_EXIT": + return ActionPrimitive(action_type=ActionType.QUARTER_EXIT) + + elif name == "TRAILING_STOP": + dist = float(args[0]) if args else 20.0 + return ActionPrimitive(action_type=ActionType.TRAILING_STOP, trail_distance_bps=dist) + + elif name == "CANCEL_AND_HOLD": + dur = float(args[0]) if args else 60.0 + return ActionPrimitive(action_type=ActionType.CANCEL_AND_HOLD, duration_s=dur) + + elif name == "REQUOTE": + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + offset = int(args[1]) if len(args) > 1 else 0 + size = float(args[2]) if len(args) > 2 else 0.25 + return ActionPrimitive(action_type=ActionType.REQUOTE, side=side, + offset_ticks=offset, size_fraction=size) + + elif name == "TAKE_PARTIAL": + frac = float(args[0]) if args else 0.5 + return ActionPrimitive(action_type=ActionType.TAKE_PARTIAL, profit_fraction=frac) + + elif name == "REDUCE_SIZE": + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + size = float(args[1]) if len(args) > 1 else 0.05 + return ActionPrimitive(action_type=ActionType.REDUCE_SIZE, side=side, size_fraction=size) + + elif name == "INCREASE_SIZE": + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + size = float(args[1]) if len(args) > 1 else 0.05 + return ActionPrimitive(action_type=ActionType.INCREASE_SIZE, side=side, size_fraction=size) + + elif name == "LOG_STATE": + return ActionPrimitive(action_type=ActionType.LOG_STATE) + + elif name == "CHECK_REGIME": + return ActionPrimitive(action_type=ActionType.CHECK_REGIME) + + elif name == "SWITCH_STRATEGY": + target = args[0] if args else "" + return ActionPrimitive(action_type=ActionType.SWITCH_STRATEGY, + metadata={"target": target}) + + elif name == "WAIT_FOR_REGIME": + dur = float(args[0]) if args else 60.0 + return ActionPrimitive(action_type=ActionType.WAIT_FOR_REGIME, duration_s=dur) + + elif name == "ADJUST_SIZE": + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + size = float(args[1]) if len(args) > 1 else 0.1 + return ActionPrimitive(action_type=ActionType.ADJUST_SIZE, side=side, size_fraction=size) + + elif name == "HEDGE_PAIR": + side = Side.BUY if args[0].upper() == "BUY" else Side.SELL + size = float(args[1]) if len(args) > 1 else 0.1 + return ActionPrimitive(action_type=ActionType.HEDGE_PAIR, side=side, size_fraction=size) + + raise DSLParseError(f"Unknown action: {name}") + + +class StrategyDSLCompiler: + """Compile DSL text into executable StrategyTemplate.""" + + def __init__(self) -> None: + self.parser = StrategyDSLParser() + + def compile(self, dsl_text: str) -> StrategyTemplate: + return self.parser.parse(dsl_text) + + def decompile(self, template: StrategyTemplate) -> str: + lines = [f'STRATEGY "{template.name}" {{'] + for rule in sorted(template.rules, key=lambda r: r.priority): + cond_parts = [] + for c in rule.conditions: + cond_parts.append(f"{c.sensor.value} {c.op.value} {c.threshold}") + cond_str = " AND ".join(cond_parts) + action = rule.action + if action.action_type.value in ("NOOP", "CANCEL_ALL", "EXIT", "STOP_LOSS", + "TAKE_PROFIT", "FLAT_ALL", "EMERGENCY_EXIT", + "HALF_EXIT", "QUARTER_EXIT"): + action_str = action.action_type.value + elif action.action_type.value in ("QUOTE", "JOIN_QUEUE", "STEP_BACK"): + side_str = action.side.value if action.side else "BUY" + action_str = f"{action.action_type.value}({side_str}, {action.offset_ticks}, {action.size_fraction})" + elif action.action_type.value in ("CROSS", "SNIPER", "PING"): + side_str = action.side.value if action.side else "BUY" + action_str = f"{action.action_type.value}({side_str}, {action.size_fraction})" + elif action.action_type.value == "HOLD": + action_str = f"HOLD({action.duration_s})" + elif action.action_type.value == "TRAILING_STOP": + action_str = f"TRAILING_STOP({action.trail_distance_bps})" + else: + action_str = action.action_type.value + if cond_str: + lines.append(f" PRIORITY {rule.priority}: IF {cond_str} THEN {action_str}") + else: + lines.append(f" PRIORITY {rule.priority}: {action_str}") + lines.append("}") + return "\n".join(lines) + + +# ============================================================================== +# Builtin Strategies — 15+ diverse strategies +# ============================================================================== + +BUILTIN_STRATEGIES: Dict[str, str] = { + "passive_maker": ''' +STRATEGY "passive_maker" { + PRIORITY 1: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(BUY, 0, 0.25) + PRIORITY 2: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(SELL, 0, 0.25) + PRIORITY 3: IF orderflow_toxicity > 0.7 THEN CANCEL_ALL + PRIORITY 4: IF time_in_trade > 300 THEN EXIT + PRIORITY 5: NOOP +} +''', + + "aggressive_taker": ''' +STRATEGY "aggressive_taker" { + PRIORITY 1: IF spread_bps < 2.0 AND imbalance > 0.3 THEN CROSS(BUY, 0.1) + PRIORITY 2: IF spread_bps < 2.0 AND imbalance < -0.3 THEN CROSS(SELL, 0.1) + PRIORITY 3: IF unrealized_pnl < -30 THEN STOP_LOSS + PRIORITY 4: IF unrealized_pnl > 50 THEN TAKE_PROFIT + PRIORITY 5: NOOP +} +''', + + "toxicity_avoider": ''' +STRATEGY "toxicity_avoider" { + PRIORITY 1: IF orderflow_toxicity > 0.5 THEN CANCEL_ALL + PRIORITY 2: IF orderflow_toxicity < 0.2 AND spread_bps < 3.0 THEN QUOTE(BUY, 1, 0.20) + PRIORITY 3: IF orderflow_toxicity < 0.2 AND spread_bps < 3.0 THEN QUOTE(SELL, 1, 0.20) + PRIORITY 4: IF time_in_loss > 120 THEN EXIT + PRIORITY 5: NOOP +} +''', + + "path_risk_exit": ''' +STRATEGY "path_risk_exit" { + PRIORITY 1: IF mae_bps < -50 AND recovery_velocity < 0 THEN EXIT + PRIORITY 2: IF time_in_loss > 300 THEN EXIT + PRIORITY 3: IF failed_recoveries > 3 THEN EXIT + PRIORITY 4: IF mfe_bps > 20 AND distance_from_mfe_bps > 15 THEN TAKE_PARTIAL(0.5) + PRIORITY 5: NOOP +} +''', + + "regime_adaptive": ''' +STRATEGY "regime_adaptive" { + PRIORITY 1: IF regime_score > 0.7 AND spread_bps < 3.0 THEN QUOTE(BUY, 0, 0.30) + PRIORITY 2: IF regime_score > 0.7 AND spread_bps < 3.0 THEN QUOTE(SELL, 0, 0.30) + PRIORITY 3: IF regime_score < 0.3 THEN CANCEL_ALL + PRIORITY 4: IF volatility > 20 THEN CROSS(BUY, 0.05) + PRIORITY 5: NOOP +} +''', + + "momentum_catcher": ''' +STRATEGY "momentum_catcher" { + PRIORITY 1: IF price_momentum_5s > 0.5 AND imbalance > 0.2 THEN CROSS(BUY, 0.08) + PRIORITY 2: IF price_momentum_5s < -0.5 AND imbalance < -0.2 THEN CROSS(SELL, 0.08) + PRIORITY 3: IF unrealized_pnl > 30 THEN HALF_EXIT + PRIORITY 4: IF unrealized_pnl < -20 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "mean_reversion": ''' +STRATEGY "mean_reversion" { + PRIORITY 1: IF spread_bps > 8.0 AND imbalance < -0.3 THEN QUOTE(BUY, 0, 0.15) + PRIORITY 2: IF spread_bps > 8.0 AND imbalance > 0.3 THEN QUOTE(SELL, 0, 0.15) + PRIORITY 3: IF unrealized_pnl > 15 THEN HALF_EXIT + PRIORITY 4: IF unrealized_pnl < -40 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "scalper": ''' +STRATEGY "scalper" { + PRIORITY 1: IF spread_bps < 1.5 AND imbalance > 0.25 THEN CROSS(BUY, 0.05) + PRIORITY 2: IF spread_bps < 1.5 AND imbalance < -0.25 THEN CROSS(SELL, 0.05) + PRIORITY 3: IF unrealized_pnl > 5 THEN HALF_EXIT + PRIORITY 4: IF unrealized_pnl < -8 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "inventory_manager": ''' +STRATEGY "inventory_manager" { + PRIORITY 1: IF position_qty > 0.15 THEN REDUCE_SIZE(SELL, 0.05) + PRIORITY 2: IF position_qty < -0.15 THEN REDUCE_SIZE(BUY, 0.05) + PRIORITY 3: IF leverage > 1.5 THEN HALF_EXIT + PRIORITY 4: IF spread_bps < 4.0 AND orderflow_toxicity < 0.3 THEN QUOTE(BUY, 0, 0.10) + PRIORITY 5: NOOP +} +''', + + "session_guard": ''' +STRATEGY "session_guard" { + PRIORITY 1: IF is_weekend == 1.0 THEN FLAT_ALL + PRIORITY 2: IF current_hour < 8.0 AND position_qty != 0 THEN HALF_EXIT + PRIORITY 3: IF current_hour > 22.0 AND position_qty != 0 THEN HALF_EXIT + PRIORITY 4: IF max_drawdown_today > 100 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "liquidity_hunter": ''' +STRATEGY "liquidity_hunter" { + PRIORITY 1: IF bid_depth_10 > 5.0 AND ask_depth_10 > 5.0 AND spread_bps < 3.0 THEN QUOTE(BUY, 0, 0.30) + PRIORITY 2: IF bid_depth_10 > 5.0 AND ask_depth_10 > 5.0 AND spread_bps < 3.0 THEN QUOTE(SELL, 0, 0.30) + PRIORITY 3: IF bid_depth_10 < 1.0 OR ask_depth_10 < 1.0 THEN CANCEL_ALL + PRIORITY 4: NOOP +} +''', + + "volatility_breakout": ''' +STRATEGY "volatility_breakout" { + PRIORITY 1: IF volatility > 25 AND price_momentum_15s > 1.0 THEN CROSS(BUY, 0.10) + PRIORITY 2: IF volatility > 25 AND price_momentum_15s < -1.0 THEN CROSS(SELL, 0.10) + PRIORITY 3: IF unrealized_pnl > 40 THEN QUARTER_EXIT + PRIORITY 4: IF unrealized_pnl < -30 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "funding_arb": ''' +STRATEGY "funding_arb" { + PRIORITY 1: IF funding > 0.05 AND position_qty < 0.1 THEN QUOTE(BUY, 0, 0.20) + PRIORITY 2: IF funding < -0.05 AND position_qty > -0.1 THEN QUOTE(SELL, 0, 0.20) + PRIORITY 3: IF unrealized_pnl > 25 THEN HALF_EXIT + PRIORITY 4: NOOP +} +''', + + "grid_trader": ''' +STRATEGY "grid_trader" { + PRIORITY 1: IF spread_bps < 4.0 AND imbalance > 0.1 THEN QUOTE(BUY, 1, 0.10) + PRIORITY 2: IF spread_bps < 4.0 AND imbalance < -0.1 THEN QUOTE(SELL, 1, 0.10) + PRIORITY 3: IF unrealized_pnl > 10 THEN HALF_EXIT + PRIORITY 4: IF unrealized_pnl < -25 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "risk_parity": ''' +STRATEGY "risk_parity" { + PRIORITY 1: IF risk_budget_used > 0.8 THEN CANCEL_ALL + PRIORITY 2: IF risk_budget_used > 0.6 THEN REDUCE_SIZE(SELL, 0.05) + PRIORITY 3: IF leverage > 1.0 THEN HALF_EXIT + PRIORITY 4: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(BUY, 0, 0.15) + PRIORITY 5: NOOP +} +''', + + "hybrid_adaptive": ''' +STRATEGY "hybrid_adaptive" { + PRIORITY 1: IF orderflow_toxicity > 0.6 THEN CANCEL_ALL + PRIORITY 2: IF spread_bps < 2.0 AND imbalance > 0.3 THEN CROSS(BUY, 0.08) + PRIORITY 3: IF spread_bps < 2.0 AND imbalance < -0.3 THEN CROSS(SELL, 0.08) + PRIORITY 4: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(BUY, 0, 0.20) + PRIORITY 5: IF spread_bps < 5.0 AND orderflow_toxicity < 0.3 THEN QUOTE(SELL, 0, 0.20) + PRIORITY 6: IF unrealized_pnl < -40 THEN STOP_LOSS + PRIORITY 7: IF unrealized_pnl > 30 THEN HALF_EXIT + PRIORITY 8: IF time_in_trade > 600 THEN EXIT + PRIORITY 9: NOOP +} +''', + + "regime_switcher": ''' +STRATEGY "regime_switcher" { + PRIORITY 1: IF regime_score > 0.8 THEN QUOTE(BUY, 0, 0.30) + PRIORITY 2: IF regime_score > 0.8 THEN QUOTE(SELL, 0, 0.30) + PRIORITY 3: IF regime_score < 0.2 THEN EXIT + PRIORITY 4: IF discrepancy_rate > 0.5 THEN LOG_STATE + PRIORITY 5: NOOP +} +''', + + "discrepancy_aware": ''' +STRATEGY "discrepancy_aware" { + PRIORITY 1: IF discrepancy_rate > 0.3 THEN CANCEL_ALL + PRIORITY 2: IF discrepancy_rate < 0.1 AND spread_bps < 4.0 THEN QUOTE(BUY, 0, 0.20) + PRIORITY 3: IF discrepancy_rate < 0.1 AND spread_bps < 4.0 THEN QUOTE(SELL, 0, 0.20) + PRIORITY 4: IF time_in_trade > 300 THEN EXIT + PRIORITY 5: NOOP +} +''', + + "portfolio_risk_manager": ''' +STRATEGY "portfolio_risk_manager" { + PRIORITY 1: IF portfolio_risk > 0.8 THEN CANCEL_ALL + PRIORITY 2: IF portfolio_risk > 0.6 THEN HALF_EXIT + PRIORITY 3: IF portfolio_risk < 0.3 AND spread_bps < 5.0 THEN QUOTE(BUY, 0, 0.25) + PRIORITY 4: IF unrealized_pnl < -30 THEN STOP_LOSS + PRIORITY 5: NOOP +} +''', + + "multi_regime_adaptive": ''' +STRATEGY "multi_regime_adaptive" { + PRIORITY 1: IF regime_score > 0.7 AND imbalance > 0.2 THEN CROSS(BUY, 0.10) + PRIORITY 2: IF regime_score > 0.7 AND imbalance < -0.2 THEN CROSS(SELL, 0.10) + PRIORITY 3: IF regime_score < 0.3 AND spread_bps > 8.0 THEN QUOTE(BUY, 0, 0.15) + PRIORITY 4: IF regime_score < 0.3 AND spread_bps > 8.0 THEN QUOTE(SELL, 0, 0.15) + PRIORITY 5: IF unrealized_pnl > 25 THEN HALF_EXIT + PRIORITY 6: IF unrealized_pnl < -35 THEN STOP_LOSS + PRIORITY 7: NOOP +} +''', +} + + +def get_builtin_strategy(name: str) -> Optional[str]: + return BUILTIN_STRATEGIES.get(name) + + +def list_builtin_strategies() -> List[str]: + return list(BUILTIN_STRATEGIES.keys()) diff --git a/MALKHUT/malkhut/training/execution_quality.py b/MALKHUT/malkhut/training/execution_quality.py new file mode 100644 index 0000000..0abaf8f --- /dev/null +++ b/MALKHUT/malkhut/training/execution_quality.py @@ -0,0 +1,192 @@ +""" +Execution Quality Metrics — measure the quality of execution. + +Metrics: + - Slippage vs arrival price + - Implementation shortfall + - Market impact + - Fill rate vs expected + - Maker fill ratio +""" +from __future__ import annotations + +import math +from dataclasses import dataclass, field +from typing import List, Optional + + +@dataclass(frozen=True, slots=True) +class ExecutionQualityReport: + """Report on execution quality for a set of fills.""" + total_fills: int + avg_slippage_bps: float + avg_implementation_shortfall_bps: float + avg_market_impact_bps: float + fill_rate: float # actual fills / expected fills + maker_fill_ratio: float + taker_fill_ratio: float + adverse_fill_ratio: float + avg_fill_time_ms: float + total_fees_bps: float + + +class ExecutionQualityTracker: + """ + Track execution quality metrics. + + Measures slippage, implementation shortfall, market impact, + fill rates, and fee quality. + """ + + def __init__(self) -> None: + self._fills: list[dict] = [] + + def record_fill( + self, + fill_price: float, + arrival_price: float, + expected_price: float, + is_maker: bool, + fill_time_ms: float, + fee_bps: float, + toxicity: float = 0.0, + ) -> None: + """Record a fill for quality analysis.""" + slippage = (fill_price - arrival_price) / max(arrival_price, 1e-12) * 10_000 + shortfall = (fill_price - expected_price) / max(expected_price, 1e-12) * 10_000 + impact = abs(fill_price - arrival_price) / max(arrival_price, 1e-12) * 10_000 + + self._fills.append({ + "fill_price": fill_price, + "arrival_price": arrival_price, + "expected_price": expected_price, + "slippage_bps": slippage, + "shortfall_bps": shortfall, + "impact_bps": impact, + "is_maker": is_maker, + "fill_time_ms": fill_time_ms, + "fee_bps": fee_bps, + "toxicity": toxicity, + }) + + def report(self) -> ExecutionQualityReport: + """Generate execution quality report.""" + if not self._fills: + return ExecutionQualityReport( + total_fills=0, avg_slippage_bps=0.0, avg_implementation_shortfall_bps=0.0, + avg_market_impact_bps=0.0, fill_rate=0.0, maker_fill_ratio=0.0, + taker_fill_ratio=0.0, adverse_fill_ratio=0.0, avg_fill_time_ms=0.0, + total_fees_bps=0.0, + ) + + n = len(self._fills) + avg_slip = sum(f["slippage_bps"] for f in self._fills) / n + avg_shortfall = sum(f["shortfall_bps"] for f in self._fills) / n + avg_impact = sum(f["impact_bps"] for f in self._fills) / n + avg_time = sum(f["fill_time_ms"] for f in self._fills) / n + avg_fee = sum(f["fee_bps"] for f in self._fills) / n + maker_count = sum(1 for f in self._fills if f["is_maker"]) + toxic_count = sum(1 for f in self._fills if f["toxicity"] > 0.5) + + return ExecutionQualityReport( + total_fills=n, + avg_slippage_bps=avg_slip, + avg_implementation_shortfall_bps=avg_shortfall, + avg_market_impact_bps=avg_impact, + fill_rate=1.0, # placeholder + maker_fill_ratio=maker_count / n, + taker_fill_ratio=1.0 - maker_count / n, + adverse_fill_ratio=toxic_count / n, + avg_fill_time_ms=avg_time, + total_fees_bps=avg_fee, + ) + + @property + def total_fills(self) -> int: + return len(self._fills) + + +class RiskAdjustedReturns: + """ + Compute risk-adjusted return metrics. + + Metrics: + - Sharpe ratio + - Sortino ratio + - Calmar ratio + - Profit factor + - Max drawdown + """ + + def __init__(self, risk_free_rate: float = 0.0) -> None: + self._risk_free_rate = risk_free_rate + self._returns: list[float] = [] + + def add_return(self, ret: float) -> None: + self._returns.append(ret) + + @property + def sharpe_ratio(self) -> float: + if len(self._returns) < 2: + return 0.0 + mean = sum(self._returns) / len(self._returns) + variance = sum((r - mean) ** 2 for r in self._returns) / (len(self._returns) - 1) + std = math.sqrt(variance) if variance > 0 else 1e-12 + return (mean - self._risk_free_rate) / std + + @property + def sortino_ratio(self) -> float: + if len(self._returns) < 2: + return 0.0 + mean = sum(self._returns) / len(self._returns) + downside = [r for r in self._returns if r < 0] + if not downside: + return float('inf') if mean > 0 else 0.0 + downside_var = sum(r ** 2 for r in downside) / len(downside) + downside_std = math.sqrt(downside_var) if downside_var > 0 else 1e-12 + return (mean - self._risk_free_rate) / downside_std + + @property + def profit_factor(self) -> float: + gains = sum(r for r in self._returns if r > 0) + losses = abs(sum(r for r in self._returns if r < 0)) + if losses <= 0: + return float('inf') if gains > 0 else 0.0 + return gains / losses + + @property + def max_drawdown(self) -> float: + if not self._returns: + return 0.0 + peak = self._returns[0] + max_dd = 0.0 + cumulative = 0.0 + for r in self._returns: + cumulative += r + if cumulative > peak: + peak = cumulative + dd = peak - cumulative + if dd > max_dd: + max_dd = dd + return max_dd + + @property + def calmar_ratio(self) -> float: + if not self._returns: + return 0.0 + mean = sum(self._returns) / len(self._returns) + annual_return = mean * 252 # annualize + dd = self.max_drawdown + if dd <= 0: + return float('inf') if annual_return > 0 else 0.0 + return annual_return / dd + + def report(self) -> dict: + return { + "sharpe_ratio": self.sharpe_ratio, + "sortino_ratio": self.sortino_ratio, + "profit_factor": self.profit_factor, + "max_drawdown": self.max_drawdown, + "calmar_ratio": self.calmar_ratio, + "total_returns": len(self._returns), + } diff --git a/MALKHUT/malkhut/training/generator.py b/MALKHUT/malkhut/training/generator.py new file mode 100644 index 0000000..bc24cf6 --- /dev/null +++ b/MALKHUT/malkhut/training/generator.py @@ -0,0 +1,511 @@ +""" +Strategy Generator — genetic programming for strategy evolution. + +Based on genetic programming (GP) principles: + - Strategy = genome (parameter vector + strategy type) + - Crossover: combine two strategies to create offspring + - Mutation: randomly modify a strategy + - Selection: tournament selection based on self-play fitness + - Population: diverse pool of strategies, hardcoded baseline always available + +Key design principle: + The hardcoded baseline is NEVER replaced. Generated strategies are ADDED + to the pool. The system GROWS its strategy repertoire. + +Game theory insight: + During self-play, the system might "spontaneously generate" new strategies + via crossover/mutation of existing ones. These emergent strategies should + be captured, trialed, and added if successful. +""" +from __future__ import annotations + +import math +import random +import time +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Callable, List, Optional, Sequence, Tuple + +from malkhut.state import FulfilmentPolicyParams +from malkhut.training.cma_trainer import ( + CMAParameterCodec, EpisodeResult, PolicyEvaluator, + PolicySnapshot, Scenario, SelfPlayPool, +) +from malkhut.training.registry import PolicyRegistry, PolicyStage +from malkhut.counterparties import default_counterparty_ecology + + +# ============================================================================== +# Strategy Types — different planner algorithms +# ============================================================================== + +class StrategyType(str, Enum): + """Different strategy structures the system can use.""" + SM_MCTS = "SM_MCTS" # Decoupled UCB/UCT (default) + UCB1 = "UCB1" # Standard UCB1 (simpler) + THOMPSON_SAMPLING = "THOMPSON" # Thompson sampling + GREEDY = "GREEDY" # Always pick best Q-value + RANDOM = "RANDOM" # Random action selection + HYBRID = "HYBRID" # Mix of multiple strategies + + +# ============================================================================== +# Strategy Genome +# ============================================================================== + +@dataclass(frozen=True, slots=True) +class StrategyGenome: + """ + A strategy encoded as a genome for genetic operations. + + The genome has two parts: + 1. Structural: strategy type, action menu config + 2. Parametric: the 29 tunable parameters + + Genetic operations (crossover, mutation) work on this genome. + """ + strategy_type: StrategyType + params: FulfilmentPolicyParams + generation: int = 0 + parent_ids: Tuple[str, ...] = () + fitness: float = 0.0 + episodes_tested: int = 0 + creation_ts_ns: int = 0 + + @property + def genome_id(self) -> str: + """Unique identifier for this genome.""" + return f"{self.strategy_type.value}_{self.params.version}_{self.generation}" + + +# ============================================================================== +# Genetic Operators +# ============================================================================== + +class GeneticOperators: + """ + Genetic operators for strategy evolution. + + Crossover: combine two strategies to create offspring + Mutation: randomly modify a strategy + Selection: tournament selection based on fitness + """ + + def __init__(self, codec: CMAParameterCodec, mutation_rate: float = 0.15, + crossover_rate: float = 0.7) -> None: + self.codec = codec + self.mutation_rate = mutation_rate + self.crossover_rate = crossover_rate + + def crossover( + self, + parent1: StrategyGenome, + parent2: StrategyGenome, + rng: random.Random, + ) -> StrategyGenome: + """ + Uniform crossover: for each parameter, randomly pick from parent1 or parent2. + + Strategy type is inherited from the fitter parent. + """ + # Strategy type from fitter parent + strategy_type = parent1.strategy_type if parent1.fitness >= parent2.fitness else parent2.strategy_type + + # Crossover parameters + p1_vec = self.codec.initial_vector(parent1.params) + p2_vec = self.codec.initial_vector(parent2.params) + child_vec = [] + + for i in range(len(p1_vec)): + if rng.random() < 0.5: + child_vec.append(p1_vec[i]) + else: + child_vec.append(p2_vec[i]) + + # Decode child + child_params = self.codec.decode(child_vec, version=f"child_{int(time.time_ns())}") + + return StrategyGenome( + strategy_type=strategy_type, + params=child_params, + generation=max(parent1.generation, parent2.generation) + 1, + parent_ids=(parent1.genome_id, parent2.genome_id), + creation_ts_ns=time.time_ns(), + ) + + def mutate( + self, + genome: StrategyGenome, + rng: random.Random, + ) -> StrategyGenome: + """ + Gaussian mutation: add noise to each parameter with mutation_rate probability. + + Occasionally mutate strategy type (structural mutation). + """ + vec = self.codec.initial_vector(genome.params) + lows, highs = self.codec.bounds() + + mutated_vec = [] + for i, (v, lo, hi) in enumerate(zip(vec, lows, highs)): + if rng.random() < self.mutation_rate: + # Gaussian noise scaled by parameter range + range_val = hi - lo + noise = rng.gauss(0, range_val * 0.1) + mutated_vec.append(max(lo, min(hi, v + noise))) + else: + mutated_vec.append(v) + + # Decode mutated params + mutated_params = self.codec.decode(mutated_vec, version=f"mut_{int(time.time_ns())}") + + # Occasionally mutate strategy type (5% chance) + strategy_type = genome.strategy_type + if rng.random() < 0.05: + strategy_type = rng.choice(list(StrategyType)) + + return StrategyGenome( + strategy_type=strategy_type, + params=mutated_params, + generation=genome.generation + 1, + parent_ids=(genome.genome_id,), + creation_ts_ns=time.time_ns(), + ) + + def tournament_select( + self, + population: List[StrategyGenome], + tournament_size: int = 3, + rng: random.Random = None, + ) -> StrategyGenome: + """Tournament selection: pick tournament_size random, return the best.""" + if rng is None: + rng = random.Random() + tournament = rng.sample(population, min(tournament_size, len(population))) + return max(tournament, key=lambda g: g.fitness) + + def random_genome( + self, + strategy_type: Optional[StrategyType] = None, + rng: random.Random = None, + ) -> StrategyGenome: + """Generate a random genome for initial population.""" + if rng is None: + rng = random.Random() + + if strategy_type is None: + strategy_type = rng.choice(list(StrategyType)) + + # Random parameters within bounds + lows, highs = self.codec.bounds() + random_vec = [rng.uniform(lo, hi) for lo, hi in zip(lows, highs)] + params = self.codec.decode(random_vec, version=f"rand_{int(time.time_ns())}") + + return StrategyGenome( + strategy_type=strategy_type, + params=params, + generation=0, + creation_ts_ns=time.time_ns(), + ) + + +# ============================================================================== +# Strategy Evaluator +# ============================================================================== + +class StrategyEvaluator: + """ + Evaluates strategies through self-play episodes. + + Each strategy is tested against the current self-play pool. + Fitness = robust score across scenarios. + + CRITICAL: uses the genome's strategy_type to create the correct planner. + This is how different planners (EXP3, Regret Matching, etc.) are actually + used during self-play discovery. + """ + + def __init__(self, evaluator: PolicyEvaluator) -> None: + self.evaluator = evaluator + + def evaluate( + self, + genome: StrategyGenome, + scenarios: Sequence[Scenario], + pool: SelfPlayPool, + rng_seed: int = 42, + ) -> float: + """Evaluate a genome's fitness through self-play.""" + score, results = self.evaluator.evaluate_candidate( + params=genome.params, + scenarios=scenarios, + rng_seed=rng_seed, + planner_type=genome.strategy_type.value, # PASS STRATEGY TYPE + ) + return score + + def evaluate_population( + self, + population: List[StrategyGenome], + scenarios: Sequence[Scenario], + pool: SelfPlayPool, + rng_seed: int = 42, + ) -> List[StrategyGenome]: + """Evaluate entire population and update fitness scores.""" + evaluated = [] + for i, genome in enumerate(population): + fitness = self.evaluate(genome, scenarios, pool, rng_seed + i) + evaluated.append(StrategyGenome( + strategy_type=genome.strategy_type, + params=genome.params, + generation=genome.generation, + parent_ids=genome.parent_ids, + fitness=fitness, + episodes_tested=len(scenarios), + creation_ts_ns=genome.creation_ts_ns, + )) + return evaluated + + +# ============================================================================== +# Strategy Generator — the main loop +# ============================================================================== + +@dataclass(frozen=True, slots=True) +class GeneratorConfig: + """Configuration for strategy generation.""" + population_size: int = 20 + generations: int = 5 + tournament_size: int = 3 + elitism_count: int = 2 # keep top N unchanged + mutation_rate: float = 0.15 + crossover_rate: float = 0.7 + max_strategies: int = 50 # max strategies in pool + min_fitness_threshold: float = -100.0 # minimum fitness to keep + + +class StrategyGenerator: + """ + Genetic programming for strategy evolution. + + Key design: + - Hardcoded baseline is NEVER replaced + - Generated strategies are ADDED to the pool + - System GROWS its strategy repertoire + - During self-play, crossover/mutation may "spontaneously generate" + new strategies that weren't explicitly programmed + + Flow: + 1. Initialize population (random + baseline) + 2. Evaluate fitness (self-play episodes) + 3. Select parents (tournament selection) + 4. Create offspring (crossover + mutation) + 5. Evaluate offspring + 6. Replace weakest with offspring + 7. Repeat for N generations + 8. Add successful strategies to pool + """ + + def __init__( + self, + config: Optional[GeneratorConfig] = None, + registry: Optional[PolicyRegistry] = None, + pool: Optional[SelfPlayPool] = None, + ) -> None: + self.config = config or GeneratorConfig() + self._registry = registry or PolicyRegistry() + self._pool = pool or SelfPlayPool(max_size=self.config.max_strategies) + + self._codec = CMAParameterCodec() + self._operators = GeneticOperators( + codec=self._codec, + mutation_rate=self.config.mutation_rate, + crossover_rate=self.config.crossover_rate, + ) + self._evaluator = StrategyEvaluator( + PolicyEvaluator( + cwm_factory=lambda: __import__("malkhut.cwm.core", fromlist=["MinimalCryptoLOBCWM"]).MinimalCryptoLOBCWM(), + counterparties=default_counterparty_ecology(), + ) + ) + + self._population: List[StrategyGenome] = [] + self._history: List[StrategyGenome] = [] + self._rng = random.Random(42) + + def initialize_population( + self, + baseline: FulfilmentPolicyParams, + scenarios: Sequence[Scenario], + ) -> None: + """Initialize population with baseline + random variants.""" + self._population = [] + + # Add hardcoded baseline (always available) + baseline_genome = StrategyGenome( + strategy_type=StrategyType.SM_MCTS, + params=baseline, + generation=0, + creation_ts_ns=time.time_ns(), + fitness=0.0, + ) + self._population.append(baseline_genome) + + # Add random variants + for i in range(self.config.population_size - 1): + genome = self._operators.random_genome(rng=self._rng) + self._population.append(genome) + + # Evaluate initial population + self._population = self._evaluator.evaluate_population( + self._population, scenarios, self._pool, + ) + + def evolve( + self, + baseline: FulfilmentPolicyParams, + scenarios: Sequence[Scenario], + ) -> List[StrategyGenome]: + """ + Run genetic evolution for N generations. + + Returns the final population (sorted by fitness). + """ + # Initialize if empty + if not self._population: + self.initialize_population(baseline, scenarios) + + for gen in range(self.config.generations): + # 1. Select parents + parents = [] + for _ in range(self.config.population_size - self.config.elitism_count): + p1 = self._operators.tournament_select( + self._population, self.config.tournament_size, self._rng, + ) + p2 = self._operators.tournament_select( + self._population, self.config.tournament_size, self._rng, + ) + parents.append((p1, p2)) + + # 2. Create offspring + offspring = [] + for p1, p2 in parents: + if self._rng.random() < self.config.crossover_rate: + child = self._operators.crossover(p1, p2, self._rng) + else: + child = self._operators.mutate(p1, self._rng) + offspring.append(child) + + # 3. Evaluate offspring + offspring = self._evaluator.evaluate_population( + offspring, scenarios, self._pool, + ) + + # 4. Elitism: keep top N unchanged, PLUS always keep baseline + self._population.sort(key=lambda g: g.fitness, reverse=True) + elites = self._population[:self.config.elitism_count] + + # Ensure baseline (SM_MCTS with version "baseline") is always present + has_baseline = any( + g.strategy_type == StrategyType.SM_MCTS and g.params.version == "baseline" + for g in elites + ) + if not has_baseline: + baseline = next( + (g for g in self._population + if g.strategy_type == StrategyType.SM_MCTS and g.params.version == "baseline"), + None, + ) + if baseline: + elites.append(baseline) + + # 5. Replace weakest with offspring + self._population = elites + offspring[:self.config.population_size - self.config.elitism_count] + + # 6. Sort by fitness + self._population.sort(key=lambda g: g.fitness, reverse=True) + + # 7. Track history + self._history.extend(offspring) + + return self._population + + def get_successful_strategies( + self, + min_fitness: Optional[float] = None, + ) -> List[StrategyGenome]: + """ + Get strategies that meet the fitness threshold. + + These are candidates for adding to the pool. + """ + threshold = min_fitness or self.config.min_fitness_threshold + return [g for g in self._population if g.fitness > threshold] + + def add_to_pool(self, genome: StrategyGenome) -> None: + """Add a successful strategy to the self-play pool.""" + snapshot = PolicySnapshot( + params=genome.params, + score=genome.fitness, + created_ts_ns=genome.creation_ts_ns, + evaluation_summary={ + "strategy_type": genome.strategy_type.value, + "generation": genome.generation, + "episodes_tested": genome.episodes_tested, + }, + ) + self._pool.maybe_add(snapshot) + + # Also register in registry + self._registry.register_candidate( + genome.params, genome.fitness, + {"strategy_type": genome.strategy_type.value}, + ) + + def get_diverse_strategies( + self, + n: int = 5, + ) -> List[StrategyGenome]: + """ + Get N diverse strategies from the population. + + Diversity is measured by: + - Different strategy types + - Different parameter vectors (cosine distance) + """ + if len(self._population) <= n: + return list(self._population) + + # Group by strategy type + by_type: dict[StrategyType, list[StrategyGenome]] = {} + for g in self._population: + by_type.setdefault(g.strategy_type, []).append(g) + + # Pick one from each type, then fill with best remaining + selected = [] + for stype in StrategyType: + if stype in by_type and len(selected) < n: + best = max(by_type[stype], key=lambda g: g.fitness) + selected.append(best) + + # Fill with best remaining + remaining = [g for g in self._population if g not in selected] + remaining.sort(key=lambda g: g.fitness, reverse=True) + while len(selected) < n and remaining: + selected.append(remaining.pop(0)) + + return selected + + @property + def population_size(self) -> int: + return len(self._population) + + @property + def best_fitness(self) -> float: + if not self._population: + return -float("inf") + return max(g.fitness for g in self._population) + + @property + def population(self) -> List[StrategyGenome]: + return list(self._population) diff --git a/MALKHUT/malkhut/training/hooks.py b/MALKHUT/malkhut/training/hooks.py new file mode 100644 index 0000000..1613dbc --- /dev/null +++ b/MALKHUT/malkhut/training/hooks.py @@ -0,0 +1,75 @@ +""" +Execution Hooks — entry/exit interfaces for future BingX integration. + +Prepares hooks that will be called when connecting to live execution systems. +All hooks are no-ops now but define the interface for future implementation. +""" +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable, Optional, Protocol + +from malkhut.state import MarketWorldState +from malkhut.actions import FulfilmentAction + + +class ExecutionIntentHook(Protocol): + """Hook called when an execution intent is submitted.""" + + def on_intent(self, intent_id: str, target: str, urgency: float) -> None: ... + + +class FillCallbackHook(Protocol): + """Hook called after a fill is received.""" + + def on_fill(self, fill_price: float, fill_qty: float, side: str, is_maker: bool) -> None: ... + + +class VenueTelemetryHook(Protocol): + """Hook for venue state publishing to Zinc.""" + + def on_venue_state(self, state: Mapping[str, Any]) -> None: ... + + +class AccountReconciliationHook(Protocol): + """Hook to compare internal state vs venue state.""" + + def reconcile(self, internal_state: MarketWorldState, venue_state: Mapping[str, Any]) -> list[str]: ... + + +@dataclass(frozen=True, slots=True) +class ExecutionHooks: + """ + Collection of hooks for future execution integration. + + All hooks are no-ops by default. Override when connecting to live systems. + """ + on_intent: Optional[Callable] = None + on_fill: Optional[Callable] = None + on_venue_state: Optional[Callable] = None + reconcile: Optional[Callable] = None + + def submit_intent(self, intent_id: str, target: str, urgency: float) -> None: + """Submit an execution intent (DSL primitive).""" + if self.on_intent: + self.on_intent(intent_id, target, urgency) + + def report_fill(self, fill_price: float, fill_qty: float, side: str, is_maker: bool) -> None: + """Report a fill (DSL primitive).""" + if self.on_fill: + self.on_fill(fill_price, fill_qty, side, is_maker) + + def publish_venue_state(self, state: Mapping[str, Any]) -> None: + """Publish venue state to Zinc.""" + if self.on_venue_state: + self.on_venue_state(state) + + def reconcile_state(self, internal: MarketWorldState, venue: Mapping[str, Any]) -> list[str]: + """Reconcile internal vs venue state.""" + if self.reconcile: + return self.reconcile(internal, venue) + return [] + + +# Default hooks (no-ops) +DEFAULT_HOOKS = ExecutionHooks() diff --git a/MALKHUT/malkhut/training/importance.py b/MALKHUT/malkhut/training/importance.py new file mode 100644 index 0000000..e76644e --- /dev/null +++ b/MALKHUT/malkhut/training/importance.py @@ -0,0 +1,110 @@ +""" +Feature Importance Tracker — track which features drive planner decisions. + +Enables: + - Understanding which market features matter most + - Identifying overfitting to specific features + - Guiding feature engineering + - Explaining decision rationale +""" +from __future__ import annotations + +import time +from collections import defaultdict +from dataclasses import dataclass, field +from typing import Any, Dict, List, Mapping, Optional, Tuple + +from malkhut.state import MarketWorldState +from malkhut.actions import FulfilmentAction +from malkhut.features import DefaultFeatureExtractor, FeatureExtractor + + +@dataclass(frozen=True, slots=True) +class FeatureImportance: + """Importance score for a feature in a specific context.""" + feature_name: str + importance: float + regime: str + action_type: str + sample_count: int + + +class FeatureImportanceTracker: + """ + Track which features drive planner decisions. + + Uses a simple attribution method: + - When a decision is made, record which features were above/below thresholds + - Aggregate across decisions to compute importance scores + """ + + def __init__(self, feature_extractor: Optional[FeatureExtractor] = None) -> None: + self._extractor = feature_extractor or DefaultFeatureExtractor() + self._feature_counts: Dict[str, Dict[str, int]] = defaultdict(lambda: defaultdict(int)) + self._feature_values: Dict[str, List[float]] = defaultdict(list) + self._total_decisions = 0 + + def record_decision( + self, + state: MarketWorldState, + action: FulfilmentAction, + regime: str = "unknown", + ) -> None: + """Record which features were relevant for this decision.""" + self._total_decisions += 1 + fv = self._extractor.extract(state).values + + # Track which features were "active" (non-zero or above threshold) + for name, value in fv.items(): + if abs(value) > 1e-6: # non-zero + self._feature_counts[name][regime] += 1 + self._feature_counts[name]["_total"] += 1 + + # Track value distribution + self._feature_values[name].append(value) + + def get_importance( + self, + top_n: int = 10, + regime: Optional[str] = None, + ) -> List[FeatureImportance]: + """Get top N most important features.""" + scores = [] + for name, regime_counts in self._feature_counts.items(): + total = regime_counts.get("_total", 0) + if regime: + count = regime_counts.get(regime, 0) + else: + count = total + + importance = count / max(self._total_decisions, 1) + scores.append(FeatureImportance( + feature_name=name, + importance=importance, + regime=regime or "all", + action_type="all", + sample_count=count, + )) + + scores.sort(key=lambda s: s.importance, reverse=True) + return scores[:top_n] + + def get_feature_stats(self, feature_name: str) -> Dict[str, float]: + """Get statistics for a specific feature.""" + values = self._feature_values.get(feature_name, []) + if not values: + return {} + return { + "mean": sum(values) / len(values), + "min": min(values), + "max": max(values), + "count": len(values), + } + + @property + def total_decisions(self) -> int: + return self._total_decisions + + @property + def feature_count(self) -> int: + return len(self._feature_counts) diff --git a/MALKHUT/malkhut/training/observability.py b/MALKHUT/malkhut/training/observability.py new file mode 100644 index 0000000..cd8bfdd --- /dev/null +++ b/MALKHUT/malkhut/training/observability.py @@ -0,0 +1,100 @@ +""" +Structured Observability — compact JSONL logging for every decision. + +Records: + - Every planner decision with features and diagnostics + - Every risk gate decision + - Every venue action + - Every fill/discrepancy + - Replayable audit trail +""" +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field +from typing import Any, Dict, List, Mapping, Optional + +from malkhut.state import MarketWorldState +from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision + + +@dataclass(frozen=True, slots=True) +class DecisionRecord: + """One decision in the audit trail.""" + ts_ns: int + symbol: str + action_kind: str + action_side: Optional[str] + action_price: Optional[float] + approved: bool + risk_reason: str + plan_latency_ns: int + policy_version: str + root_entropy: float + sims: int + features: Mapping[str, float] = field(default_factory=dict) + diagnostics: Mapping[str, Any] = field(default_factory=dict) + + +class ObservabilityLogger: + """ + Compact JSONL logger for every decision. + + One line per decision. Replayable audit trail. + """ + + def __init__(self, log_path: str = "decisions.log") -> None: + self._log_path = log_path + self._decisions: list[DecisionRecord] = [] + self._total_decisions = 0 + + def log_decision( + self, + state: MarketWorldState, + planned: PlannedPolicy, + decision: RiskDecision, + plan_ns: int, + features: Optional[Mapping[str, float]] = None, + ) -> None: + """Log a decision to the audit trail.""" + record = DecisionRecord( + ts_ns=state.ts_ns, + symbol=state.venue.symbol, + action_kind=planned.selected_action.kind.value, + action_side=planned.selected_action.side.value if planned.selected_action.side else None, + action_price=planned.selected_action.price_ticks_from_best if planned.selected_action else None, + approved=decision.approved, + risk_reason=decision.reason, + plan_latency_ns=plan_ns, + policy_version="live", + root_entropy=planned.diagnostics.get("entropy", 0.0), + sims=planned.diagnostics.get("sims", 0), + features=features or {}, + ) + self._decisions.append(record) + self._total_decisions += 1 + + # Write to JSONL + try: + with open(self._log_path, "a") as f: + f.write(json.dumps({ + "ts": record.ts_ns, + "sym": record.symbol, + "act": record.action_kind, + "side": record.action_side, + "app": record.approved, + "risk": record.risk_reason, + "lat_ns": record.plan_latency_ns, + "entropy": round(record.root_entropy, 4), + "sims": record.sims, + }, separators=(",", ":")) + "\n") + except OSError: + pass + + @property + def total_decisions(self) -> int: + return self._total_decisions + + def get_recent(self, n: int = 10) -> List[DecisionRecord]: + return self._decisions[-n:] diff --git a/MALKHUT/malkhut/training/parallel.py b/MALKHUT/malkhut/training/parallel.py new file mode 100644 index 0000000..f3faf4c --- /dev/null +++ b/MALKHUT/malkhut/training/parallel.py @@ -0,0 +1,134 @@ +""" +Training Parallelism — parallel evaluation across scenarios. + +CMA-ES evaluates sequentially. This module parallelizes evaluation +across scenarios for 4-8x faster convergence. +""" +from __future__ import annotations + +import concurrent.futures +import time +from dataclasses import dataclass +from typing import Any, Callable, List, Optional, Sequence + +from malkhut.state import FulfilmentPolicyParams, MarketWorldState +from malkhut.training.cma_trainer import PolicyEvaluator, Scenario + + +class ParallelEvaluator: + """ + Parallel evaluation of strategies across scenarios. + + Uses ThreadPoolExecutor for I/O-bound scenarios. + Uses ProcessPoolExecutor for CPU-bound scenarios. + """ + + def __init__( + self, + evaluator: PolicyEvaluator, + max_workers: int = 4, + ) -> None: + self.evaluator = evaluator + self._max_workers = max_workers + + def evaluate_candidate( + self, + params: FulfilmentPolicyParams, + scenarios: Sequence[Scenario], + rng_seed: int = 0, + ) -> tuple[float, list]: + """Evaluate candidate with parallel scenario execution.""" + if len(scenarios) <= 1 or self._max_workers <= 1: + return self.evaluator.evaluate_candidate(params, scenarios, rng_seed) + + # Split scenarios across workers + chunk_size = max(1, len(scenarios) // self._max_workers) + chunks = [] + for i in range(0, len(scenarios), chunk_size): + chunks.append(scenarios[i:i + chunk_size]) + + # Parallel evaluation + all_results = [] + with concurrent.futures.ThreadPoolExecutor(max_workers=self._max_workers) as executor: + futures = [] + for i, chunk in enumerate(chunks): + future = executor.submit( + self.evaluator.evaluate_candidate, + params, chunk, rng_seed + i, + ) + futures.append(future) + + for future in concurrent.futures.as_completed(futures): + score, results = future.result() + all_results.extend(results) + + # Aggregate scores + if all_results: + scores = [r.pnl_bps for r in all_results] + avg_score = sum(scores) / len(scores) + else: + avg_score = 0.0 + + return avg_score, all_results + + +@dataclass(frozen=True, slots=True) +class TrainingMetrics: + """Metrics for a training run.""" + total_time_s: float + generations: int + total_evals: int + best_score: float + avg_score: float + score_improvement: float + convergence_gen: int + + +class TrainingMonitor: + """ + Monitor training progress and convergence. + + Tracks metrics, detects convergence, logs progress. + """ + + def __init__(self) -> None: + self._scores: list[float] = [] + self._times: list[float] = [] + self._start_time = time.time() + + def record_generation(self, score: float) -> None: + self._scores.append(score) + self._times.append(time.time()) + + @property + def best_score(self) -> float: + return max(self._scores) if self._scores else 0.0 + + @property + def avg_score(self) -> float: + return sum(self._scores) / len(self._scores) if self._scores else 0.0 + + @property + def improvement(self) -> float: + if len(self._scores) < 2: + return 0.0 + return self._scores[-1] - self._scores[0] + + @property + def converged(self) -> bool: + if len(self._scores) < 5: + return False + recent = self._scores[-5:] + variance = sum((s - self.avg_score) ** 2 for s in recent) / len(recent) + return variance < 0.01 # low variance = converged + + def metrics(self) -> TrainingMetrics: + return TrainingMetrics( + total_time_s=time.time() - self._start_time, + generations=len(self._scores), + total_evals=0, + best_score=self.best_score, + avg_score=self.avg_score, + score_improvement=self.improvement, + convergence_gen=len(self._scores) if self.converged else -1, + ) diff --git a/MALKHUT/malkhut/training/rollback.py b/MALKHUT/malkhut/training/rollback.py new file mode 100644 index 0000000..0e5ed85 --- /dev/null +++ b/MALKHUT/malkhut/training/rollback.py @@ -0,0 +1,106 @@ +""" +Policy Rollback — auto-rollback if shadow performance degrades. + +Enables: + - Detecting when a promoted policy performs worse than baseline + - Automatically reverting to previous best + - Preventing bad policies from reaching live +""" +from __future__ import annotations + +import time +from dataclasses import dataclass, field +from typing import Any, List, Mapping, Optional, Tuple + +from malkhut.state import FulfilmentPolicyParams +from malkhut.training.registry import PolicyRegistry, PolicyStage, PolicyRecord + + +@dataclass(frozen=True, slots=True) +class RollbackEvent: + """Record of a rollback event.""" + ts_ns: int + rolled_back_version: str + rolled_back_to: str + reason: str + performance_drop: float + + +class PolicyRollback: + """ + Auto-rollback if shadow performance degrades. + + Monitors active policy performance and reverts to previous best + if performance drops below threshold. + """ + + def __init__( + self, + registry: PolicyRegistry, + degradation_threshold: float = -5.0, # bps + min_shadow_steps: int = 10, + ) -> None: + self._registry = registry + self._degradation_threshold = degradation_threshold + self._min_shadow_steps = min_shadow_steps + self._shadow_scores: List[float] = [] + self._rollback_events: List[RollbackEvent] = [] + + def record_shadow_score(self, score: float) -> None: + """Record a shadow performance score.""" + self._shadow_scores.append(score) + + def check_rollback(self) -> Optional[RollbackEvent]: + """ + Check if rollback is needed. + + Returns RollbackEvent if rollback should happen, None otherwise. + """ + if len(self._shadow_scores) < self._min_shadow_steps: + return None + + # Compare recent average to baseline + recent = self._shadow_scores[-self._min_shadow_steps:] + avg_recent = sum(recent) / len(recent) + + # Get baseline score (first policy in registry) + baseline = self._registry.get_by_stage(PolicyStage.ACTIVE) + if not baseline: + return None + + baseline_record = baseline[0] if baseline else None + if not baseline_record: + return None + + # Check degradation + if avg_recent < self._degradation_threshold: + # Find previous best to rollback to + previous = self._registry.get_by_stage(PolicyStage.RETIRED) + if previous: + rollback_to = previous[0].version + else: + rollback_to = "baseline" + + # Perform rollback + current = baseline_record[0] if isinstance(baseline_record, list) else baseline_record + self._registry.retire(current.version, "auto_rollback_degradation") + + event = RollbackEvent( + ts_ns=time.time_ns(), + rolled_back_version=current.version, + rolled_back_to=rollback_to, + reason=f"performance_drop_{avg_recent:.2f}", + performance_drop=avg_recent - self._degradation_threshold, + ) + self._rollback_events.append(event) + return event + + return None + + @property + def rollback_events(self) -> List[RollbackEvent]: + return list(self._rollback_events) + + @property + def shadow_score_count(self) -> int: + return len(self._shadow_scores) diff --git a/MALKHUT/malkhut/training/stress.py b/MALKHUT/malkhut/training/stress.py new file mode 100644 index 0000000..a97d7f0 --- /dev/null +++ b/MALKHUT/malkhut/training/stress.py @@ -0,0 +1,195 @@ +""" +Scenario Stress Testing — diversified stress scenarios for robust evaluation. + +Generates scenarios that test edge cases: + - Flash crash (sudden price drop) + - Liquidity vacuum (no bids/asks) + - Extreme volatility + - Correlated moves across assets + - Weekend/low participation + - Funding shock + - Liquidation cascade +""" +from __future__ import annotations + +import random +import time +from dataclasses import dataclass, field +from typing import Any, Optional, Sequence, Tuple + +from malkhut.state import ( + AccountState, MarketWorldState, Mode, OrderBookState, PositionState, + PriceLevel, Side, TradePathState, VenueRules, +) +from malkhut.counterparties import ( + CounterpartyPolicy, ToxicTakerPolicy, PassiveMakerPolicy, + LatencyArbPolicy, NoiseTraderPolicy, default_counterparty_ecology, +) + + +@dataclass(frozen=True, slots=True) +class StressScenario: + """A stress test scenario with specific conditions.""" + scenario_id: str + symbol: str + initial_state: MarketWorldState + counterparties: Tuple[CounterpartyPolicy, ...] + max_steps: int + tags: Tuple[str, ...] + description: str + + +class StressScenarioFactory: + """ + Generate diversified stress scenarios. + + Tests extreme market conditions that normal scenarios miss. + """ + + def __init__(self, counterparties: Optional[Tuple[CounterpartyPolicy, ...]] = None) -> None: + self.counterparties = counterparties or default_counterparty_ecology() + + def flash_crash(self, symbol: str = "BTCUSDT") -> StressScenario: + """Sudden 5% price drop in 3 steps.""" + return StressScenario( + scenario_id=f"flash_crash_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.1, ask_qty=0.1), + counterparties=(ToxicTakerPolicy(sensitivity=0.2),), + max_steps=10, + tags=("stress", "flash_crash", "high_volatility"), + description="Sudden 5% price drop with thin book", + ) + + def liquidity_vacuum(self, symbol: str = "BTCUSDT") -> StressScenario: + """Near-zero liquidity on both sides.""" + return StressScenario( + scenario_id=f"liquidity_vacuum_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.001, ask_qty=0.001), + counterparties=(ToxicTakerPolicy(sensitivity=0.1),), + max_steps=10, + tags=("stress", "liquidity_vacuum", "thin_book"), + description="Near-zero liquidity, any trade moves price significantly", + ) + + def extreme_volatility(self, symbol: str = "BTCUSDT") -> StressScenario: + """Wide spread, high volatility.""" + return StressScenario( + scenario_id=f"extreme_vol_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=49000.0, ask=51000.0, bid_qty=0.5, ask_qty=0.5), + counterparties=(ToxicTakerPolicy(sensitivity=0.3), NoiseTraderPolicy()), + max_steps=20, + tags=("stress", "extreme_volatility", "wide_spread"), + description="2000bps spread, high volatility", + ) + + def toxic_flood(self, symbol: str = "BTCUSDT") -> StressScenario: + """Multiple toxic takers attacking simultaneously.""" + return StressScenario( + scenario_id=f"toxic_flood_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.5, ask_qty=0.5), + counterparties=( + ToxicTakerPolicy(sensitivity=0.2), + ToxicTakerPolicy(sensitivity=0.3), + LatencyArbPolicy(lead_threshold=0.3), + ), + max_steps=15, + tags=("stress", "toxic_flood", "adverse_selection"), + description="Multiple toxic actors attacking simultaneously", + ) + + def choppy_market(self, symbol: str = "BTCUSDT") -> StressScenario: + """Sideways chop with no clear direction.""" + return StressScenario( + scenario_id=f"choppy_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50000.5, bid_qty=0.3, ask_qty=0.3), + counterparties=(NoiseTraderPolicy(), PassiveMakerPolicy(join_probability=0.8)), + max_steps=30, + tags=("stress", "choppy", "noise"), + description="Tight range, high noise, no clear direction", + ) + + def weekend_low_participation(self, symbol: str = "BTCUSDT") -> StressScenario: + """Weekend-like conditions: thin book, low volume.""" + return StressScenario( + scenario_id=f"weekend_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50002.0, bid_qty=0.2, ask_qty=0.2), + counterparties=(NoiseTraderPolicy(),), + max_steps=15, + tags=("stress", "weekend", "low_participation"), + description="Weekend-like: thin book, low volume, wider spreads", + ) + + def liquidation_cascade(self, symbol: str = "BTCUSDT") -> StressScenario: + """Price drops trigger liquidations, which cause more drops.""" + return StressScenario( + scenario_id=f"liquidation_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.3, ask_qty=0.3), + counterparties=( + ToxicTakerPolicy(sensitivity=0.3), + NoiseTraderPolicy(), + ), + max_steps=20, + tags=("stress", "liquidation_cascade", "cascade"), + description="Liquidation cascade: price drops → liquidations → more drops", + ) + + def correlation_breakdown(self, symbol: str = "BTCUSDT") -> StressScenario: + """BTC drops, alts diverge.""" + return StressScenario( + scenario_id=f"correlation_{symbol}", + symbol=symbol, + initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.5, ask_qty=0.5), + counterparties=(ToxicTakerPolicy(sensitivity=0.4),), + max_steps=15, + tags=("stress", "correlation_breakdown", "divergence"), + description="BTC drops while correlation breaks down", + ) + + def build_stress_suite( + self, + symbols: Sequence[str] = ("BTCUSDT",), + ) -> Tuple[StressScenario, ...]: + """Build a complete stress test suite.""" + scenarios = [] + for symbol in symbols: + scenarios.append(self.flash_crash(symbol)) + scenarios.append(self.liquidity_vacuum(symbol)) + scenarios.append(self.extreme_volatility(symbol)) + scenarios.append(self.toxic_flood(symbol)) + scenarios.append(self.choppy_market(symbol)) + scenarios.append(self.weekend_low_participation(symbol)) + scenarios.append(self.liquidation_cascade(symbol)) + scenarios.append(self.correlation_breakdown(symbol)) + return tuple(scenarios) + + @staticmethod + def _make_state( + symbol: str, bid: float, ask: float, + bid_qty: float = 1.0, ask_qty: float = 1.0, + ) -> MarketWorldState: + 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, + ) diff --git a/MALKHUT/malkhut/training/structured_obs.py b/MALKHUT/malkhut/training/structured_obs.py new file mode 100644 index 0000000..bd9a03e --- /dev/null +++ b/MALKHUT/malkhut/training/structured_obs.py @@ -0,0 +1,120 @@ +""" +Structured Observability — per-decision feature attribution and metrics. + +Tracks: + - Which features drove each decision + - Feature importance over time + - Decision quality metrics + - Regime-specific performance +""" +from __future__ import annotations + +import json +import time +from collections import defaultdict +from dataclasses import dataclass, field +from typing import Any, Dict, List, Mapping, Optional, Tuple + +from malkhut.state import MarketWorldState +from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision +from malkhut.features import DefaultFeatureExtractor, FeatureExtractor + + +@dataclass(frozen=True, slots=True) +class DecisionMetrics: + """Per-decision metrics.""" + ts_ns: int + symbol: str + action_kind: str + approved: bool + plan_latency_ns: int + entropy: float + sims: int + feature_attribution: Mapping[str, float] + regime: str + + +class StructuredObservability: + """ + Structured observability with per-decision feature attribution. + + Tracks which features drive decisions and computes aggregate metrics. + """ + + def __init__(self, feature_extractor: Optional[FeatureExtractor] = None) -> None: + self._extractor = feature_extractor or DefaultFeatureExtractor() + self._decisions: list[DecisionMetrics] = [] + self._feature_importance: Dict[str, List[float]] = defaultdict(list) + self._regime_performance: Dict[str, List[float]] = defaultdict(list) + self._total_decisions = 0 + + def record_decision( + self, + state: MarketWorldState, + planned: PlannedPolicy, + decision: RiskDecision, + plan_ns: int, + regime: str = "unknown", + ) -> None: + """Record a decision with full feature attribution.""" + fv = self._extractor.extract(state).values + + # Compute feature attribution (which features are "active") + attribution = {} + for name, value in fv.items(): + if abs(value) > 1e-6: + attribution[name] = value + + metrics = DecisionMetrics( + ts_ns=state.ts_ns, + symbol=state.venue.symbol, + action_kind=planned.selected_action.kind.value, + approved=decision.approved, + plan_latency_ns=plan_ns, + entropy=planned.diagnostics.get("entropy", 0.0), + sims=planned.diagnostics.get("sims", 0), + feature_attribution=attribution, + regime=regime, + ) + + self._decisions.append(metrics) + self._total_decisions += 1 + + # Track feature importance + for name, value in attribution.items(): + self._feature_importance[name].append(value) + + # Track regime performance + self._regime_performance[regime].append(1.0 if decision.approved else 0.0) + + def get_feature_importance(self, top_n: int = 10) -> List[Tuple[str, float]]: + """Get top N features by average absolute value.""" + scores = [] + for name, values in self._feature_importance.items(): + avg = sum(abs(v) for v in values) / len(values) + scores.append((name, avg)) + scores.sort(key=lambda x: x[1], reverse=True) + return scores[:top_n] + + def get_regime_approval_rate(self, regime: str) -> float: + """Get approval rate for a specific regime.""" + approvals = self._regime_performance.get(regime, []) + if not approvals: + return 0.0 + return sum(approvals) / len(approvals) + + @property + def total_decisions(self) -> int: + return self._total_decisions + + @property + def avg_latency_ns(self) -> float: + if not self._decisions: + return 0.0 + return sum(d.plan_latency_ns for d in self._decisions) / len(self._decisions) + + @property + def avg_entropy(self) -> float: + if not self._decisions: + return 0.0 + return sum(d.entropy for d in self._decisions) / len(self._decisions) diff --git a/MALKHUT/malkhut/training/trajectory.py b/MALKHUT/malkhut/training/trajectory.py new file mode 100644 index 0000000..f7cfacb --- /dev/null +++ b/MALKHUT/malkhut/training/trajectory.py @@ -0,0 +1,104 @@ +""" +Trajectory Persistence — store CWM trajectories to ClickHouse. + +Enables: + - Post-hoc analysis of decision quality + - Replay verification against stored trajectories + - Training data for learned leaf values + - Audit trail for every decision +""" +from __future__ import annotations + +import json +import time +from dataclasses import dataclass, field +from typing import Any, List, Mapping, Optional, Sequence + +from malkhut.state import MarketWorldState +from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision +from malkhut.storage.ch_store import MalkhutCHStore + + +@dataclass(frozen=True, slots=True) +class TrajectoryStep: + """One step in a persisted trajectory.""" + step_index: int + ts_ns: int + symbol: str + state_hash: str + action_kind: str + action_side: Optional[str] + action_price: Optional[float] + action_qty: float + next_state_hash: str + pnl_bps: float + reward: float + entropy: float + diagnostics: Mapping[str, Any] = field(default_factory=dict) + + +@dataclass(frozen=True, slots=True) +class TrajectoryRecord: + """Complete trajectory for one episode.""" + trajectory_id: str + policy_version: str + scenario_id: str + seed: int + steps: Tuple[TrajectoryStep, ...] + total_pnl_bps: float + max_drawdown_bps: float + fill_count: int + cancel_count: int + noop_count: int + duration_ns: int + created_ts_ns: int + + +class TrajectoryPersister: + """ + Persist CWM trajectories to ClickHouse. + + Enables post-hoc analysis and replay verification. + """ + + def __init__(self, store: MalkhutCHStore) -> None: + self._store = store + + def persist(self, record: TrajectoryRecord) -> None: + """Persist a trajectory record to CH.""" + self._store.store_episode( + policy_version=record.policy_version, + scenario_id=record.scenario_id, + seed=record.seed, + pnl_bps=record.total_pnl_bps, + max_drawdown_bps=record.max_drawdown_bps, + fill_ratio=record.fill_count / max(record.fill_count + record.cancel_count + record.noop_count, 1), + adverse_fill_ratio=0.0, + avg_slippage_bps=0.0, + liq_near_misses=0, + cancel_count=record.cancel_count, + diagnostics=json.dumps({ + "trajectory_id": record.trajectory_id, + "steps": len(record.steps), + "fill_count": record.fill_count, + "noop_count": record.noop_count, + "duration_ns": record.duration_ns, + }), + ) + + def query_trajectories( + self, + policy_version: Optional[str] = None, + scenario_id: Optional[str] = None, + limit: int = 100, + ) -> str: + """Query stored trajectories.""" + where_parts = [] + if policy_version: + where_parts.append(f"policy_version = '{policy_version}'") + if scenario_id: + where_parts.append(f"scenario_id = '{scenario_id}'") + where = " AND ".join(where_parts) if where_parts else "1=1" + return self._store.query( + f"SELECT * FROM self_play_episodes WHERE {where} ORDER BY ts_ns DESC LIMIT {limit}" + )