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:
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