malkhut(T3): planner + adversarial counterparties
Planner: Decoupled UCB/UCT simultaneous-move MCTS (sm_mcts.py), compact action space (action_menu.py), planner alternatives (alternatives.py). Counterparties: 4+ adversarial agent ecology — ToxicTaker, LatencyArb, MarketMaker, NoiseTrader + extended: LiquidationFlow, WhaleOrder, MomentumFollower, SpoofDetector, QueueChaser.
This commit is contained in:
126
MALKHUT/malkhut/counterparties.py
Normal file
126
MALKHUT/malkhut/counterparties.py
Normal file
@@ -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(),
|
||||
)
|
||||
142
MALKHUT/malkhut/counterparties_extended.py
Normal file
142
MALKHUT/malkhut/counterparties_extended.py
Normal file
@@ -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(),
|
||||
)
|
||||
2
MALKHUT/malkhut/planner/__init__.py
Normal file
2
MALKHUT/malkhut/planner/__init__.py
Normal file
@@ -0,0 +1,2 @@
|
||||
from malkhut.planner.sm_mcts import DecoupledUCBPlanner
|
||||
from malkhut.planner.action_menu import build_our_actions
|
||||
129
MALKHUT/malkhut/planner/action_menu.py
Normal file
129
MALKHUT/malkhut/planner/action_menu.py
Normal file
@@ -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)
|
||||
662
MALKHUT/malkhut/planner/alternatives.py
Normal file
662
MALKHUT/malkhut/planner/alternatives.py
Normal file
@@ -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}")
|
||||
275
MALKHUT/malkhut/planner/sm_mcts.py
Normal file
275
MALKHUT/malkhut/planner/sm_mcts.py
Normal file
@@ -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]
|
||||
Reference in New Issue
Block a user