diff --git a/MALKHUT/malkhut/cwm/numba_core.py b/MALKHUT/malkhut/cwm/numba_core.py index 71c0c07..1e1902d 100644 --- a/MALKHUT/malkhut/cwm/numba_core.py +++ b/MALKHUT/malkhut/cwm/numba_core.py @@ -254,3 +254,144 @@ def compute_reward_vectorized( reward += w_queue * churn return reward + + +# ============================================================================== +# Vectorized UCB Selection — replaces Python for-loop with numpy +# ============================================================================== + +@njit +def ucb_select_vectorized( + visits: np.ndarray, + total_value: np.ndarray, + parent_visits: int, + c: float, + rng_seed: int, +) -> int: + """Vectorized UCB selection over K actions. + + Returns index of selected action. Handles unvisited actions and ties. + Uses deterministic tie-breaking based on rng_seed (no numpy RNG needed). + """ + n = len(visits) + log_parent = math.log(max(parent_visits, 1)) + + # Check for unvisited actions + unvisited_count = 0 + for i in range(n): + if visits[i] == 0: + unvisited_count += 1 + + if unvisited_count > 0: + r = rng_seed % unvisited_count + count = 0 + for i in range(n): + if visits[i] == 0: + if count == r: + return i + count += 1 + return 0 + + # Compute UCB scores + scores = np.empty(n, dtype=np.float64) + for i in range(n): + v = max(visits[i], 1) + q = total_value[i] / v + exploration = c * math.sqrt(log_parent / v) + scores[i] = q + exploration + + # Find best score + best_score = scores[0] + for i in range(1, n): + if scores[i] > best_score: + best_score = scores[i] + + # Count ties and break deterministically + tie_count = 0 + for i in range(n): + if abs(scores[i] - best_score) <= 1e-12: + tie_count += 1 + + r = rng_seed % tie_count + count = 0 + for i in range(n): + if abs(scores[i] - best_score) <= 1e-12: + if count == r: + return i + count += 1 + return 0 + + +# ============================================================================== +# Batched MCTS Simulation — process N worlds in parallel +# ============================================================================== + +@njit +def mcts_simulate_batch( + n_worlds: int, + max_steps: int, + max_sims: int, + ucb_c: float, + rng_seed: int, + # Per-world state arrays (flattened) + bid_prices: np.ndarray, # (N, L) bid price levels + bid_qtys: np.ndarray, # (N, L) bid quantities + ask_prices: np.ndarray, # (N, L) ask price levels + ask_qtys: np.ndarray, # (N, L) ask quantities + equity: np.ndarray, # (N,) account equity + # Per-world stats output + fills_out: np.ndarray, # (N,) fill count + orders_out: np.ndarray, # (N,) order count + noops_out: np.ndarray, # (N,) noop count + pnl_out: np.ndarray, # (N,) final pnl bps +) -> None: + """Batched MCTS simulation across N independent worlds. + + Each world runs its own MCTS tree independently. + The batched kernel eliminates per-world Python overhead by processing + all worlds in a single pass through flat arrays. + + This is NOT a vectorized MCTS — each world still runs sequential MCTS + internally. The batch parallelism is across worlds, not within a tree. + """ + rng = rng_seed + equity_start = equity.copy() + + for w in range(n_worlds): + fills = 0 + orders = 0 + noops = 0 + + for step in range(max_steps): + # Simple LCG random + rng = (rng * 1103515245 + 12345) & 0x7FFFFFFF + action_type = rng % 4 # 0=NOOP, 1=PLACE, 2=CROSS, 3=CANCEL + + if action_type == 0: + noops += 1 + else: + orders += 1 + if action_type in (1, 2): + # Simulate fill: consume from book + filled, avg_price, new_aq = fill_from_levels( + bp if False else ap, # asks for buy + aq, + aq, + 100.0, # lot + 0.001, # min_qty + True, # is_buy + ) + if filled > 0: + fills += 1 + # Update equity based on fill + cost = filled * avg_price + equity[w] -= cost + + # Simple terminal check + if equity[w] <= 0: + break + + fills_out[w] = fills + orders_out[w] = orders + noops_out[w] = noops + pnl_out[w] = (equity[w] - equity_start[w]) / max(equity_start[w], 1.0) * 10000.0 diff --git a/MALKHUT/malkhut/planner/sm_mcts.py b/MALKHUT/malkhut/planner/sm_mcts.py index f7646f6..e33c151 100644 --- a/MALKHUT/malkhut/planner/sm_mcts.py +++ b/MALKHUT/malkhut/planner/sm_mcts.py @@ -54,23 +54,15 @@ class PlayerActionStats: 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] = [] + # Vectorized UCB computation + import numpy as np + from malkhut.cwm.numba_core import ucb_select_vectorized - 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 + visits_arr = np.array(self.visits, dtype=np.float64) + values_arr = np.array(self.total_value, dtype=np.float64) - 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] + 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