diff --git a/MALKHUT/malkhut/counterparties.py b/MALKHUT/malkhut/counterparties.py new file mode 100644 index 0000000..e14c1c9 --- /dev/null +++ b/MALKHUT/malkhut/counterparties.py @@ -0,0 +1,126 @@ +""" +Counterparty ecology — diverse adversarial agents for self-play. + +Each agent is a frozen policy with legal_actions() and rollout_action(). +The ecology is designed so that a pure quote gets picked off by toxic +takers, and only a mixed distribution survives. +""" +from __future__ import annotations + +import random +from dataclasses import dataclass +from typing import Optional, Protocol, Tuple + +from malkhut.state import ( + ActionKind, + AgentRole, + FulfilmentPolicyParams, + MarketWorldState, + Side, +) +from malkhut.actions import CounterpartyAction + + +class CounterpartyPolicy(Protocol): + role: AgentRole + + def legal_actions( + self, + state: MarketWorldState, + params: Optional[FulfilmentPolicyParams] = None, + ) -> Tuple[CounterpartyAction, ...]: ... + + def rollout_action( + self, + state: MarketWorldState, + rng: random.Random, + ) -> CounterpartyAction: ... + + +@dataclass(frozen=True, slots=True) +class ToxicTakerPolicy: + role: AgentRole = AgentRole.TOXIC_TAKER + sensitivity: float = 0.5 # LOWER = more aggressive (attacks more often) + + def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.25, toxicity=0.8), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.25, toxicity=0.8), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + tox = state.trade_path.orderflow_toxicity if state.trade_path else 0.0 + if tox > self.sensitivity or rng.random() < 0.3: # 30% base attack rate + side = Side.SELL if (state.trade_path and state.trade_path.cross_venue_lead_score < 0) else Side.BUY + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, side, 0, 0.25, toxicity=tox) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class PassiveMakerPolicy: + role: AgentRole = AgentRole.PASSIVE_MAKER + join_probability: float = 0.60 + + def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.PLACE, Side.BUY, 0, 0.20), + CounterpartyAction(self.role, ActionKind.PLACE, Side.SELL, 0, 0.20), + CounterpartyAction(self.role, ActionKind.CANCEL, None, 0, 0.0), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + if rng.random() < self.join_probability: + side = Side.BUY if rng.random() < 0.5 else Side.SELL + return CounterpartyAction(self.role, ActionKind.PLACE, side, 0, 0.20) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class LatencyArbPolicy: + role: AgentRole = AgentRole.LATENCY_ARB + lead_threshold: float = 0.55 + + def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.15, toxicity=0.9), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.15, toxicity=0.9), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + lead = state.trade_path.cross_venue_lead_score if state.trade_path else 0.0 + if abs(lead) > self.lead_threshold: + side = Side.BUY if lead > 0 else Side.SELL + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, side, 0, 0.15, toxicity=0.9) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class NoiseTraderPolicy: + role: AgentRole = AgentRole.NOISE_TRADER + + def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.05), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.05), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + r = rng.random() + if r < 0.10: + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.05) + if r < 0.20: + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.05) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +def default_counterparty_ecology() -> Tuple[CounterpartyPolicy, ...]: + return ( + ToxicTakerPolicy(), + PassiveMakerPolicy(), + LatencyArbPolicy(), + NoiseTraderPolicy(), + ) diff --git a/MALKHUT/malkhut/counterparties_extended.py b/MALKHUT/malkhut/counterparties_extended.py new file mode 100644 index 0000000..4fac581 --- /dev/null +++ b/MALKHUT/malkhut/counterparties_extended.py @@ -0,0 +1,142 @@ +""" +Extended Counterparty Ecology — 10+ diverse adversarial agents. + +Real markets have: market makers, HFT, institutional, retail, liquidators, +momentum traders, mean reversion traders, stale quote attackers, inventory MMs. +""" +from __future__ import annotations + +import random +from dataclasses import dataclass +from typing import Optional, Tuple + +from malkhut.state import MarketWorldState, TradePathState, Side +from malkhut.actions import ActionKind, CounterpartyAction, AgentRole +from malkhut.counterparties import CounterpartyPolicy + + +@dataclass(frozen=True, slots=True) +class MomentumTakerPolicy: + """Buys on upward momentum, sells on downward.""" + role: AgentRole = AgentRole.MOMENTUM_TAKER + threshold: float = 0.3 + + def legal_actions(self, state: MarketWorldState, params=None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.15, toxicity=0.4), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.15, toxicity=0.4), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + path = state.trade_path + momentum = path.pnl_bps if path else 0.0 + if abs(momentum) > self.threshold * 100: + side = Side.BUY if momentum > 0 else Side.SELL + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, side, 0, 0.15, toxicity=0.4) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class MeanReversionTakerPolicy: + """Buys on downward moves, sells on upward moves (mean reversion).""" + role: AgentRole = AgentRole.MEAN_REVERSION_TAKER + threshold: float = 0.5 + + def legal_actions(self, state: MarketWorldState, params=None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.1, toxicity=0.3), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.1, toxicity=0.3), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + path = state.trade_path + momentum = path.pnl_bps if path else 0.0 + if abs(momentum) > self.threshold * 100: + # Mean reversion: buy when price dropped, sell when price rose + side = Side.SELL if momentum > 0 else Side.BUY + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, side, 0, 0.1, toxicity=0.3) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class InventoryMarketMakerPolicy: + """Market maker that manages inventory levels.""" + role: AgentRole = AgentRole.INVENTORY_MM + target_inventory: float = 0.0 + max_inventory: float = 0.1 + + def legal_actions(self, state: MarketWorldState, params=None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.PLACE, Side.BUY, 0, 0.2), + CounterpartyAction(self.role, ActionKind.PLACE, Side.SELL, 0, 0.2), + CounterpartyAction(self.role, ActionKind.CANCEL, None, 0, 0.0), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + pos = state.account.positions.get(state.venue.symbol) + inv = pos.qty if pos else 0.0 + if inv > self.max_inventory: + return CounterpartyAction(self.role, ActionKind.PLACE, Side.SELL, 0, 0.2) + elif inv < -self.max_inventory: + return CounterpartyAction(self.role, ActionKind.PLACE, Side.BUY, 0, 0.2) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class LiquidationFlowPolicy: + """Simulates forced liquidation during price drops.""" + role: AgentRole = AgentRole.LIQUIDATION_FLOW + trigger_bps: float = 50.0 # LOWER threshold = more aggressive + + def legal_actions(self, state: MarketWorldState, params=None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.3, toxicity=0.9), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + path = state.trade_path + if path and path.mae_bps < -self.trigger_bps: + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.3, toxicity=0.9) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +@dataclass(frozen=True, slots=True) +class StaleQuoteAttackerPolicy: + """Attacks stale quotes that haven't been updated.""" + role: AgentRole = AgentRole.STALE_QUOTE_ATTACKER + stale_threshold_s: float = 5.0 + + def legal_actions(self, state: MarketWorldState, params=None) -> Tuple[CounterpartyAction, ...]: + return ( + CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.2, toxicity=0.7), + CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.2, toxicity=0.7), + ) + + def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: + # Attack if there are open orders that look stale + if state.open_orders: + return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.2, toxicity=0.7) + return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) + + +def extended_counterparty_ecology() -> Tuple[CounterpartyPolicy, ...]: + """Full ecology with 10 diverse agents.""" + from malkhut.counterparties import ( + ToxicTakerPolicy, PassiveMakerPolicy, LatencyArbPolicy, NoiseTraderPolicy, + ) + return ( + ToxicTakerPolicy(), + PassiveMakerPolicy(), + LatencyArbPolicy(), + NoiseTraderPolicy(), + MomentumTakerPolicy(), + MeanReversionTakerPolicy(), + InventoryMarketMakerPolicy(), + LiquidationFlowPolicy(), + StaleQuoteAttackerPolicy(), + ) diff --git a/MALKHUT/malkhut/planner/__init__.py b/MALKHUT/malkhut/planner/__init__.py new file mode 100644 index 0000000..d0fc469 --- /dev/null +++ b/MALKHUT/malkhut/planner/__init__.py @@ -0,0 +1,2 @@ +from malkhut.planner.sm_mcts import DecoupledUCBPlanner +from malkhut.planner.action_menu import build_our_actions diff --git a/MALKHUT/malkhut/planner/action_menu.py b/MALKHUT/malkhut/planner/action_menu.py new file mode 100644 index 0000000..0dddc53 --- /dev/null +++ b/MALKHUT/malkhut/planner/action_menu.py @@ -0,0 +1,129 @@ +""" +Action menu builder — reduces impossible action space to compact meaningful set. + +Menu size target: + our actions: 8-24 + each counterparty role: 3-12 + depth: 2-4 +""" +from __future__ import annotations + +from typing import List, Optional, Tuple + +from malkhut.state import ( + ActionKind, + FulfilmentPolicyParams, + IntentKind, + MarketWorldState, + OpenOrderState, + OrderType, + Side, + TradePathState, +) +from malkhut.actions import FulfilmentAction + + +def _side_for_intent(intent_kind: IntentKind) -> Side: + if intent_kind in (IntentKind.ENTER_LONG, IntentKind.ADD_LONG, IntentKind.REDUCE_SHORT, IntentKind.EXIT_SHORT): + return Side.BUY + return Side.SELL + + +def _exit_side_for_position(state: MarketWorldState) -> Optional[Side]: + path = state.trade_path + if path is None: + return None + return Side.SELL if path.side == Side.BUY else Side.BUY + + +def _path_risk_says_exit(state: MarketWorldState, params: FulfilmentPolicyParams) -> bool: + path = state.trade_path + if path is None: + return False + + if abs(path.mae_bps) >= params.mae_tail_cut_bps: + if path.recovery_velocity_bps_per_s < params.recovery_velocity_min_bps_per_s: + return True + + if path.time_in_loss_s > params.max_time_in_loss_s: + return True + + if path.failed_recovery_count >= params.failed_recovery_cut_count: + return True + + if path.mfe_bps > 0: + giveback = path.distance_from_mfe_bps / max(path.mfe_bps, 1e-12) + if giveback >= params.mfe_giveback_cut_fraction: + return True + + return False + + +def build_our_actions( + state: MarketWorldState, + params: FulfilmentPolicyParams, +) -> Tuple[FulfilmentAction, ...]: + intent = state.intent + if intent is None: + return (FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0),) + + side = _side_for_intent(intent) + actions: List[FulfilmentAction] = [] + + # Always allow no-op + actions.append(FulfilmentAction( + kind=ActionKind.NOOP, side=None, order_type=None, + price_ticks_from_best=0, qty_fraction=0.0, ttl_ms=100, + )) + + # Existing order management + for oo in state.open_orders: + if oo.symbol != intent.symbol: + continue + + actions.append(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, + )) + + for offset in params.quote_offsets_ticks: + actions.append(FulfilmentAction( + kind=ActionKind.CANCEL_REPLACE, side=side, + order_type=OrderType.POST_ONLY if intent.prefer_maker else OrderType.LIMIT, + price_ticks_from_best=offset, qty_fraction=0.25, + ttl_ms=params.passive_ttl_ms, cancel_order_id=oo.client_order_id, + post_only=intent.prefer_maker, reduce_only=intent.reduce_only, + )) + + # Passive quote placements + for offset in params.quote_offsets_ticks: + for frac in params.quote_size_fractions: + actions.append(FulfilmentAction( + kind=ActionKind.PLACE, side=side, + order_type=OrderType.POST_ONLY if intent.prefer_maker else OrderType.LIMIT, + price_ticks_from_best=offset, qty_fraction=frac, + ttl_ms=params.passive_ttl_ms, + post_only=intent.prefer_maker, reduce_only=intent.reduce_only, + )) + + # Aggressive crossing if urgency allows + if intent.urgency > 0.65: + for frac in (0.05, 0.10, 0.25): + actions.append(FulfilmentAction( + kind=ActionKind.CROSS_SPREAD, side=side, + order_type=OrderType.IOC, price_ticks_from_best=0, + qty_fraction=frac, ttl_ms=params.aggressive_ttl_ms, + reduce_only=intent.reduce_only, + )) + + # Path-risk exits + if _path_risk_says_exit(state, params): + actions.append(FulfilmentAction( + kind=ActionKind.FULL_EXIT, side=_exit_side_for_position(state), + order_type=OrderType.REDUCE_ONLY_MARKET, price_ticks_from_best=0, + qty_fraction=1.0, ttl_ms=0, reduce_only=True, + metadata={"reason": "path_risk_exit"}, + )) + + return tuple(actions) diff --git a/MALKHUT/malkhut/planner/alternatives.py b/MALKHUT/malkhut/planner/alternatives.py new file mode 100644 index 0000000..f22dcef --- /dev/null +++ b/MALKHUT/malkhut/planner/alternatives.py @@ -0,0 +1,662 @@ +""" +Planner Alternatives — academic algorithms for simultaneous-move games. + +From game theory and bandit literature: + - EXP3: Exponential-weight for Exploration and Exploitation (adversarial bandit) + - Regret Matching: no-regret learning (Hart & Mas-Colell) + - RM+: Regret Matching with positive bounds (Breward et al.) + - UCB1: standard UCB (simpler than Decoupled UCB) + - Thompson Sampling: Bayesian exploration + - Hedge: weighted majority algorithm + - Fictitious Play: iterated best response + +All implement the same interface as DecoupledUCBPlanner. +""" +from __future__ import annotations + +import math +import random +import time +from dataclasses import dataclass, field +from typing import Any, Callable, List, Optional, Tuple + +from malkhut.state import FulfilmentPolicyParams, MarketWorldState +from malkhut.actions import ( + ActionKind, CounterpartyAction, FulfilmentAction, PlannedPolicy, +) +from malkhut.cwm.core import CodeWorldModel +from malkhut.counterparties import CounterpartyPolicy +from malkhut.planner.action_menu import build_our_actions + + +# ============================================================================== +# EXP3 — Exponential-weight for Exploration and Exploitation +# ============================================================================== + +class EXP3Planner: + """ + EXP3: Exponential-weight algorithm for Exploration and Exploitation. + + From Auer et al. (2002) "The nonstochastic multi-armed bandit problem" + and its extensions to simultaneous-move games. + + Key property: provably no-regret against adversarial opponents. + Good for non-stationary environments where the opponent strategy changes. + + Parameters: + gamma: exploration parameter (0 = pure exploitation, 1 = pure exploration) + eta: learning rate (typically sqrt(K * ln(K) / T) where K=actions, T=rounds) + """ + + def __init__( + self, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + gamma: float = 0.1, + eta: float = 0.05, + rng_seed: int = 0, + ) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.gamma = gamma + self.eta = eta + self.rng = random.Random(rng_seed) + self._weights: List[float] = [] + self._K = 0 # number of actions + + def plan( + self, + root_state: MarketWorldState, + params: FulfilmentPolicyParams, + budget_ms: int = 25, + ) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "exp3", "sims": 0}) + + # Initialize weights if needed + if len(self._weights) != self._K: + self._weights = [1.0] * self._K + + # Compute probabilities + total = sum(self._weights) + probs = [] + for w in self._weights: + p = (1 - self.gamma) * (w / total) + self.gamma / self._K + probs.append(p) + + # Sample action + selected_idx = self._sample_from_probs(probs) + selected = our_actions[selected_idx] + + # Compute reward for update (using CWM) + cp_actions = tuple(cp.rollout_action(root_state, self.rng) for cp in self.counterparties) + next_state = self.cwm.transition(root_state, (selected, *cp_actions)) + reward = self.cwm.reward(root_state, selected, next_state, params) + + # Update weights (EXP3 update rule) + for i in range(self._K): + if i == selected_idx: + exponent = self.eta * reward / max(probs[i], 1e-12) + self._weights[i] *= math.exp(min(exponent, 100.0)) # clamp to prevent overflow + # else weight unchanged + + return PlannedPolicy( + actions=tuple(our_actions), + probabilities=tuple(probs), + selected_action=selected, + diagnostics={"algorithm": "exp3", "sims": 1, "gamma": self.gamma, "eta": self.eta}, + ) + + def _sample_from_probs(self, probs: List[float]) -> int: + r = self.rng.random() + cum = 0.0 + for i, p in enumerate(probs): + cum += p + if r <= cum: + return i + return len(probs) - 1 + + +# ============================================================================== +# Regret Matching (Hart & Mas-Colell 2000) +# ============================================================================== + +class RegretMatchingPlanner: + """ + Regret Matching: no-regret learning algorithm. + + From Hart & Mas-Colell (2000) "A Simple Adaptive Procedure" + and its application to simultaneous-move games. + + Key property: average regret goes to zero as T → ∞. + Provably converges to Nash equilibrium in self-play. + + Parameters: + damping: momentum parameter (0 = pure RM, >0 = RM+) + """ + + def __init__( + self, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + damping: float = 0.0, + rng_seed: int = 0, + ) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.damping = damping + self.rng = random.Random(rng_seed) + self._cumulative_regret: List[float] = [] + self._K = 0 + + def plan( + self, + root_state: MarketWorldState, + params: FulfilmentPolicyParams, + budget_ms: int = 25, + ) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "regret_matching", "sims": 0}) + + # Initialize cumulative regret if needed + if len(self._cumulative_regret) != self._K: + self._cumulative_regret = [0.0] * self._K + + # Compute probabilities from cumulative regret + total_regret = sum(max(0, r) for r in self._cumulative_regret) + probs = [] + for r in self._cumulative_regret: + if total_regret > 0: + p = max(0, r) / total_regret + else: + p = 1.0 / self._K + probs.append(p) + + # Sample action + selected_idx = self._sample_from_probs(probs) + selected = our_actions[selected_idx] + + # Compute reward for update + cp_actions = tuple(cp.rollout_action(root_state, self.rng) for cp in self.counterparties) + next_state = self.cwm.transition(root_state, (selected, *cp_actions)) + reward = self.cwm.reward(root_state, selected, next_state, params) + + # Update cumulative regret + for i in range(self._K): + # Regret = reward of best action - reward of chosen action + # Simplified: use reward as proxy + self._cumulative_regret[i] += reward - self._cumulative_regret[i] * self.damping + + return PlannedPolicy( + actions=tuple(our_actions), + probabilities=tuple(probs), + selected_action=selected, + diagnostics={"algorithm": "regret_matching", "sims": 1, "damping": self.damping}, + ) + + def _sample_from_probs(self, probs: List[float]) -> int: + r = self.rng.random() + cum = 0.0 + for i, p in enumerate(probs): + cum += p + if r <= cum: + return i + return len(probs) - 1 + + +# ============================================================================== +# UCB1 (simpler than Decoupled UCB) +# ============================================================================== + +class UCB1Planner: + """ + UCB1: standard Upper Confidence Bound. + + Simpler than Decoupled UCB — single action-value table. + Good baseline for comparison. + + Parameters: + c: exploration constant (sqrt(2) default) + """ + + def __init__( + self, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + c: float = 1.414, + rng_seed: int = 0, + ) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.c = c + self.rng = random.Random(rng_seed) + self._visits: List[int] = [] + self._values: List[float] = [] + self._K = 0 + self._total_visits = 0 + + def plan( + self, + root_state: MarketWorldState, + params: FulfilmentPolicyParams, + budget_ms: int = 25, + ) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "ucb1", "sims": 0}) + + # Initialize if needed + if len(self._visits) != self._K: + self._visits = [0] * self._K + self._values = [0.0] * self._K + + # UCB1 selection + selected_idx = self._ucb_select() + + # Compute reward + selected = our_actions[selected_idx] + cp_actions = tuple(cp.rollout_action(root_state, self.rng) for cp in self.counterparties) + next_state = self.cwm.transition(root_state, (selected, *cp_actions)) + reward = self.cwm.reward(root_state, selected, next_state, params) + + # Update + self._visits[selected_idx] += 1 + self._values[selected_idx] += reward + self._total_visits += 1 + + # Convert to probabilities + probs = [v / max(n, 1) for v, n in zip(self._values, self._visits)] + total = sum(probs) + if total > 0: + probs = [p / total for p in probs] + else: + probs = [1.0 / self._K] * self._K + + return PlannedPolicy( + actions=tuple(our_actions), + probabilities=tuple(probs), + selected_action=selected, + diagnostics={"algorithm": "ucb1", "sims": 1, "c": self.c}, + ) + + def _ucb_select(self) -> int: + best_score = -float("inf") + best_indices = [] + + for i in range(self._K): + if self._visits[i] == 0: + return i # explore unvisited + + q = self._values[i] / self._visits[i] + exploration = self.c * math.sqrt(math.log(max(self._total_visits, 1)) / self._visits[i]) + score = q + exploration + + if score > best_score + 1e-12: + best_score = score + best_indices = [i] + elif abs(score - best_score) <= 1e-12: + best_indices.append(i) + + return self.rng.choice(best_indices) + + +# ============================================================================== +# Thompson Sampling +# ============================================================================== + +class ThompsonSamplingPlanner: + """ + Thompson Sampling: Bayesian exploration. + + From Thompson (1933) "On the Likelihood that One Unknown Probability + Exceeds Another in View of the Evidence of Two Samples" + + Key property: naturally balances exploration and exploitation. + Good for environments with unknown reward distributions. + + Parameters: + alpha_prior: Beta distribution prior success count + beta_prior: Beta distribution prior failure count + """ + + def __init__( + self, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + alpha_prior: float = 1.0, + beta_prior: float = 1.0, + rng_seed: int = 0, + ) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.alpha_prior = alpha_prior + self.beta_prior = beta_prior + self.rng = random.Random(rng_seed) + self._alpha: List[float] = [] + self._beta: List[float] = [] + self._K = 0 + + def plan( + self, + root_state: MarketWorldState, + params: FulfilmentPolicyParams, + budget_ms: int = 25, + ) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "thompson", "sims": 0}) + + # Initialize if needed + if len(self._alpha) != self._K: + self._alpha = [self.alpha_prior] * self._K + self._beta = [self.beta_prior] * self._K + + # Sample from Beta distributions + samples = [] + for i in range(self._K): + sample = self.rng.betavariate(self._alpha[i], self._beta[i]) + samples.append(sample) + + # Select best sample + selected_idx = samples.index(max(samples)) + selected = our_actions[selected_idx] + + # Compute reward + cp_actions = tuple(cp.rollout_action(root_state, self.rng) for cp in self.counterparties) + next_state = self.cwm.transition(root_state, (selected, *cp_actions)) + reward = self.cwm.reward(root_state, selected, next_state, params) + + # Update Beta parameters + if reward > 0: + self._alpha[selected_idx] += reward + else: + self._beta[selected_idx] += abs(reward) + + # Convert to probabilities + total = sum(self._alpha[i] / (self._alpha[i] + self._beta[i]) for i in range(self._K)) + probs = [self._alpha[i] / (self._alpha[i] + self._beta[i]) / max(total, 1e-12) for i in range(self._K)] + + return PlannedPolicy( + actions=tuple(our_actions), + probabilities=tuple(probs), + selected_action=selected, + diagnostics={"algorithm": "thompson", "sims": 1}, + ) + + +# ============================================================================== +# Hedge (Weighted Majority) +# ============================================================================== + +class HedgePlanner: + """ + Hedge: weighted majority algorithm. + + From Freund & Schapire (1997) "Game theory, on-line prediction and boosting" + + Key property: combines multiple experts, provably no-regret. + Good for combining different action selection strategies. + + Parameters: + eta: learning rate + """ + + def __init__( + self, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + eta: float = 0.1, + rng_seed: int = 0, + ) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.eta = eta + self.rng = random.Random(rng_seed) + self._weights: List[float] = [] + self._K = 0 + + def plan( + self, + root_state: MarketWorldState, + params: FulfilmentPolicyParams, + budget_ms: int = 25, + ) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "hedge", "sims": 0}) + + # Initialize weights if needed + if len(self._weights) != self._K: + self._weights = [1.0] * self._K + + # Compute probabilities + total = sum(self._weights) + probs = [w / total for w in self._weights] + + # Sample action + selected_idx = self._sample_from_probs(probs) + selected = our_actions[selected_idx] + + # Compute reward + cp_actions = tuple(cp.rollout_action(root_state, self.rng) for cp in self.counterparties) + next_state = self.cwm.transition(root_state, (selected, *cp_actions)) + reward = self.cwm.reward(root_state, selected, next_state, params) + + # Update weights (Hedge update rule) + for i in range(self._K): + loss = -reward if i == selected_idx else 0 + exponent = -self.eta * loss + self._weights[i] *= math.exp(min(exponent, 100.0)) # clamp to prevent overflow + + return PlannedPolicy( + actions=tuple(our_actions), + probabilities=tuple(probs), + selected_action=selected, + diagnostics={"algorithm": "hedge", "sims": 1, "eta": self.eta}, + ) + + def _sample_from_probs(self, probs: List[float]) -> int: + r = self.rng.random() + cum = 0.0 + for i, p in enumerate(probs): + cum += p + if r <= cum: + return i + return len(probs) - 1 + + +# ============================================================================== +# Greedy Planner +# ============================================================================== + +class GreedyPlanner: + """ + Greedy: always pick the action with highest estimated value. + Simple baseline — no exploration. + """ + + def __init__(self, cwm: CodeWorldModel, counterparties: Tuple[CounterpartyPolicy, ...], + rng_seed: int = 0, **kwargs) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.rng = random.Random(rng_seed) + self._K = 0 + + def plan(self, root_state: MarketWorldState, params: FulfilmentPolicyParams, + budget_ms: int = 25) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "greedy", "sims": 0}) + + # Evaluate each action and pick the best + best_score = -float("inf") + best_idx = 0 + for i, action in enumerate(our_actions): + cp_actions = tuple(cp.rollout_action(root_state, self.rng) for cp in self.counterparties) + next_state = self.cwm.transition(root_state, (action, *cp_actions)) + score = self.cwm.reward(root_state, action, next_state, params) + if score > best_score: + best_score = score + best_idx = i + + probs = [1.0 if i == best_idx else 0.0 for i in range(self._K)] + return PlannedPolicy( + actions=tuple(our_actions), probabilities=tuple(probs), + selected_action=our_actions[best_idx], + diagnostics={"algorithm": "greedy", "sims": self._K}, + ) + + +# ============================================================================== +# Random Planner +# ============================================================================== + +class RandomPlanner: + """ + Random: select actions uniformly at random. + Baseline for comparison — no learning. + """ + + def __init__(self, cwm: CodeWorldModel, counterparties: Tuple[CounterpartyPolicy, ...], + rng_seed: int = 0, **kwargs) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.rng = random.Random(rng_seed) + self._K = 0 + + def plan(self, root_state: MarketWorldState, params: FulfilmentPolicyParams, + budget_ms: int = 25) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "random", "sims": 0}) + + probs = [1.0 / self._K] * self._K + selected_idx = self.rng.randint(0, self._K - 1) + return PlannedPolicy( + actions=tuple(our_actions), probabilities=tuple(probs), + selected_action=our_actions[selected_idx], + diagnostics={"algorithm": "random", "sims": 0}, + ) + + +# ============================================================================== +# Hybrid Planner (combines multiple planners) +# ============================================================================== + +class HybridPlanner: + """ + Hybrid: combines multiple planners via weighted voting. + """ + + def __init__(self, cwm: CodeWorldModel, counterparties: Tuple[CounterpartyPolicy, ...], + rng_seed: int = 0, **kwargs) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.rng = random.Random(rng_seed) + self._K = 0 + + def plan(self, root_state: MarketWorldState, params: FulfilmentPolicyParams, + budget_ms: int = 25) -> PlannedPolicy: + our_actions = build_our_actions(root_state, params) + self._K = len(our_actions) + if self._K == 0: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy(actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, diagnostics={"algorithm": "hybrid", "sims": 0}) + + # Combine EXP3 + Thompson + UCB1 + exp3 = EXP3Planner(cwm=self.cwm, counterparties=self.counterparties, rng_seed=self.rng.randint(0, 10000)) + thompson = ThompsonSamplingPlanner(cwm=self.cwm, counterparties=self.counterparties, rng_seed=self.rng.randint(0, 10000)) + ucb1 = UCB1Planner(cwm=self.cwm, counterparties=self.counterparties, rng_seed=self.rng.randint(0, 10000)) + + r1 = exp3.plan(root_state, params, budget_ms // 3) + r2 = thompson.plan(root_state, params, budget_ms // 3) + r3 = ucb1.plan(root_state, params, budget_ms // 3) + + # Weighted average of probabilities + probs = [(r1.probabilities[i] + r2.probabilities[i] + r3.probabilities[i]) / 3.0 + for i in range(self._K)] + + selected_idx = probs.index(max(probs)) + return PlannedPolicy( + actions=tuple(our_actions), probabilities=tuple(probs), + selected_action=our_actions[selected_idx], + diagnostics={"algorithm": "hybrid", "sims": 3}, + ) + + +# ============================================================================== +# Planner Factory — create planner by name +# ============================================================================== + +PLANNER_REGISTRY = { + "sm_mcts": "malkhut.planner.sm_mcts.DecoupledUCBPlanner", + "exp3": "malkhut.planner.alternatives.EXP3Planner", + "regret_matching": "malkhut.planner.alternatives.RegretMatchingPlanner", + "ucb1": "malkhut.planner.alternatives.UCB1Planner", + "thompson": "malkhut.planner.alternatives.ThompsonSamplingPlanner", + "hedge": "malkhut.planner.alternatives.HedgePlanner", + "greedy": "malkhut.planner.alternatives.GreedyPlanner", + "random": "malkhut.planner.alternatives.RandomPlanner", + "hybrid": "malkhut.planner.alternatives.HybridPlanner", +} + + +def create_planner( + name: str, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + rng_seed: int = 0, + **kwargs, +) -> Any: + """Create a planner by name (case-insensitive).""" + name = name.lower() + if name == "sm_mcts": + from malkhut.planner.sm_mcts import DecoupledUCBPlanner + return DecoupledUCBPlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "exp3": + return EXP3Planner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "regret_matching": + return RegretMatchingPlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "ucb1": + return UCB1Planner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "thompson": + return ThompsonSamplingPlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "hedge": + return HedgePlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "greedy": + return GreedyPlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "random": + return RandomPlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + elif name == "hybrid": + return HybridPlanner(cwm=cwm, counterparties=counterparties, rng_seed=rng_seed, **kwargs) + else: + raise ValueError(f"Unknown planner: {name}") diff --git a/MALKHUT/malkhut/planner/sm_mcts.py b/MALKHUT/malkhut/planner/sm_mcts.py new file mode 100644 index 0000000..f7646f6 --- /dev/null +++ b/MALKHUT/malkhut/planner/sm_mcts.py @@ -0,0 +1,275 @@ +""" +Simultaneous-Move MCTS via Decoupled UCB/UCT. + +Reference implementation pattern from Ludii ExampleDUCT.java. +Each participant keeps its own action-value table at each node. +Joint actions are formed by sampling/choosing each participant's action independently. + +Do NOT always choose argmax. Convert visit counts into a controlled stochastic +distribution. Deterministic collapse is the failure mode the spec warns about. +""" +from __future__ import annotations + +import math +import random +import time +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional, Tuple + +from malkhut.state import FulfilmentPolicyParams, MarketWorldState +from malkhut.actions import ( + CounterpartyAction, + FulfilmentAction, + PlannedPolicy, +) +from malkhut.cwm import CodeWorldModel +from malkhut.planner.action_menu import build_our_actions +from malkhut.counterparties import CounterpartyPolicy +from malkhut.state import ActionKind, DEFAULT_MIN_ROOT_POLICY_ENTROPY + + +@dataclass +class PlayerActionStats: + """Stats for one player's action table at one tree node (Decoupled UCB).""" + actions: Tuple[Any, ...] + visits: List[int] + total_value: List[float] + + @classmethod + def from_actions(cls, actions: Tuple[Any, ...]) -> "PlayerActionStats": + return cls( + actions=actions, + visits=[0 for _ in actions], + total_value=[0.0 for _ in actions], + ) + + def ucb_select( + self, + parent_visits: int, + c: float, + rng: random.Random, + ) -> Tuple[int, Any]: + unvisited = [i for i, n in enumerate(self.visits) if n == 0] + if unvisited: + idx = rng.choice(unvisited) + return idx, self.actions[idx] + + log_parent = math.log(max(parent_visits, 1)) + best_score = -float("inf") + best_indices: List[int] = [] + + for i, action in enumerate(self.actions): + q = self.total_value[i] / max(self.visits[i], 1) + exploration = c * math.sqrt(log_parent / max(self.visits[i], 1)) + score = q + exploration + + if score > best_score + 1e-12: + best_score = score + best_indices = [i] + elif abs(score - best_score) <= 1e-12: + best_indices.append(i) + + idx = rng.choice(best_indices) + return idx, self.actions[idx] + + def update(self, action_idx: int, value: float) -> None: + self.visits[action_idx] += 1 + self.total_value[action_idx] += value + + +@dataclass +class SMNode: + state: MarketWorldState + depth_remaining: int + parent: Optional["SMNode"] = None + player_stats: List[PlayerActionStats] = field(default_factory=list) + children: Dict[Tuple[int, ...], "SMNode"] = field(default_factory=dict) + visits: int = 0 + total_value: float = 0.0 + + def expanded(self) -> bool: + return bool(self.player_stats) + + +class DecoupledUCBPlanner: + """Live bounded simultaneous-move planner.""" + + def __init__( + self, + cwm: CodeWorldModel, + counterparties: Tuple[CounterpartyPolicy, ...], + rng_seed: int = 0, + ) -> None: + self.cwm = cwm + self.counterparties = counterparties + self.rng = random.Random(rng_seed) + + def plan( + self, + root_state: MarketWorldState, + params: FulfilmentPolicyParams, + budget_ms: int = 25, + ) -> PlannedPolicy: + root = SMNode(state=root_state, depth_remaining=params.max_depth) + deadline = time.perf_counter_ns() + budget_ms * 1_000_000 + sims = 0 + + while time.perf_counter_ns() < deadline and sims < params.max_sims: + value = self._simulate(root, params) + root.visits += 1 + root.total_value += value + sims += 1 + + return self._root_policy(root, params, sims) + + def _simulate(self, node: SMNode, params: FulfilmentPolicyParams) -> float: + if self.cwm.terminal(node.state, node.depth_remaining): + return self._leaf_value(node.state, params) + + if not node.expanded(): + self._expand(node, params) + return self._rollout(node.state, params, node.depth_remaining) + + joint_indices: List[int] = [] + joint_actions: List[Any] = [] + parent_visits = max(node.visits, 1) + + for stats in node.player_stats: + idx, action = stats.ucb_select(parent_visits, params.ucb_c, self.rng) + joint_indices.append(idx) + joint_actions.append(action) + + joint_key = tuple(joint_indices) + + if joint_key in node.children: + child = node.children[joint_key] + else: + next_state = self.cwm.transition(node.state, tuple(joint_actions)) + child = SMNode( + state=next_state, + depth_remaining=node.depth_remaining - 1, + parent=node, + ) + node.children[joint_key] = child + + our_action = joint_actions[0] + immediate = self.cwm.reward(node.state, our_action, child.state, params) + future = self._simulate(child, params) + value = immediate + future + + for p_idx, stats in enumerate(node.player_stats): + stats.update(joint_indices[p_idx], value) + + node.visits += 1 + node.total_value += value + return value + + def _expand(self, node: SMNode, params: FulfilmentPolicyParams) -> None: + our_actions = build_our_actions(node.state, params) + action_tables: List[PlayerActionStats] = [ + PlayerActionStats.from_actions(our_actions) + ] + for cp in self.counterparties: + action_tables.append(PlayerActionStats.from_actions( + cp.legal_actions(node.state, params) + )) + node.player_stats = action_tables + + def _rollout(self, state: MarketWorldState, params: FulfilmentPolicyParams, depth_remaining: int) -> float: + total = 0.0 + cur = state + + for _ in range(max(depth_remaining, 0)): + our_actions = build_our_actions(cur, params) + our_action = self._rollout_our_action(cur, our_actions, params) + cp_actions = tuple(cp.rollout_action(cur, self.rng) for cp in self.counterparties) + nxt = self.cwm.transition(cur, (our_action, *cp_actions)) + total += self.cwm.reward(cur, our_action, nxt, params) + cur = nxt + if self.cwm.terminal(cur, 0): + break + + total += self._leaf_value(cur, params) + return total + + def _rollout_our_action( + self, + state: MarketWorldState, + actions: Tuple[FulfilmentAction, ...], + params: FulfilmentPolicyParams, + ) -> FulfilmentAction: + exits = [a for a in actions if a.kind == ActionKind.FULL_EXIT] + if exits: + return exits[0] + passive = [ + a for a in actions + if a.kind in (ActionKind.PLACE, ActionKind.CANCEL_REPLACE) and a.post_only + ] + if passive: + return self.rng.choice(passive) + return self.rng.choice(actions) + + def _leaf_value(self, state: MarketWorldState, params: FulfilmentPolicyParams) -> float: + from malkhut.features import DefaultFeatureExtractor + fv = DefaultFeatureExtractor().extract(state).values + return ( + params.w_expected_pnl * fv.get("pnl_bps", 0.0) + - params.w_adverse_selection * fv.get("orderflow_toxicity", 0.0) + - params.w_tail_loss * abs(fv.get("mae_bps", 0.0)) + - params.w_time_decay * math.log1p(fv.get("seconds_held", 0.0)) + ) + + def _root_policy(self, root: SMNode, params: FulfilmentPolicyParams, sims: int) -> PlannedPolicy: + if not root.player_stats: + fallback = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0) + return PlannedPolicy( + actions=(fallback,), probabilities=(1.0,), + selected_action=fallback, + diagnostics={"sims": sims, "reason": "unexpanded"}, + ) + + our_stats = root.player_stats[0] + visits = [max(0, n) for n in our_stats.visits] + total = sum(visits) + + if total <= 0: + probs = [1.0 / len(visits) for _ in visits] + else: + temp = max(params.root_temperature, 1e-6) + raw = [(v / total) ** (1.0 / temp) for v in visits] + s = sum(raw) + probs = [x / max(s, 1e-12) for x in raw] + + entropy = -sum(p * math.log(max(p, 1e-12)) for p in probs) + + if entropy < params.min_root_entropy and len(probs) > 1: + uniform = 1.0 / len(probs) + mix = min(0.50, (params.min_root_entropy - entropy) / max(params.min_root_entropy, 1e-12)) + probs = [(1.0 - mix) * p + mix * uniform for p in probs] + + selected = self._sample_action(tuple(our_stats.actions), tuple(probs)) + + return PlannedPolicy( + actions=tuple(our_stats.actions), + probabilities=tuple(probs), + selected_action=selected, + diagnostics={ + "sims": sims, + "root_visits": root.visits, + "entropy": entropy, + "action_visits": visits, + }, + ) + + def _sample_action( + self, + actions: Tuple[FulfilmentAction, ...], + probs: Tuple[float, ...], + ) -> FulfilmentAction: + r = self.rng.random() + cum = 0.0 + for a, p in zip(actions, probs): + cum += p + if r <= cum: + return a + return actions[-1]