""" 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] # Vectorized UCB computation import numpy as np from malkhut.cwm.numba_core import ucb_select_vectorized visits_arr = np.array(self.visits, dtype=np.float64) values_arr = np.array(self.total_value, dtype=np.float64) idx = ucb_select_vectorized(visits_arr, values_arr, parent_visits, c, rng.randint(0, 2**31)) return int(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]