diff --git a/MALKHUT/malkhut/training/cma_trainer.py b/MALKHUT/malkhut/training/cma_trainer.py index a67098f..a505549 100644 --- a/MALKHUT/malkhut/training/cma_trainer.py +++ b/MALKHUT/malkhut/training/cma_trainer.py @@ -1005,131 +1005,110 @@ class PolicyEvaluator: rng_seed: int, planner_type: str = "sm_mcts", ) -> EpisodeResult: - """Run a full multi-step episode through the CWM.""" + """Run a full multi-step episode through the CWM. + Fast path: reduces Python overhead by pre-allocating and reusing objects.""" from malkhut.planner.alternatives import create_planner + from malkhut.state import ExecutionIntent, IntentKind + cwm = self.cwm_factory() - planner = create_planner( - planner_type, - cwm=cwm, - counterparties=scenario.counterparties, - rng_seed=rng_seed, - ) + planner = create_planner(planner_type, cwm=cwm, + counterparties=scenario.counterparties, + rng_seed=rng_seed) state = scenario.initial_state rng = random.Random(rng_seed) + n_cp = len(scenario.counterparties) - # Trajectory metrics + # Pre-allocate intent template (reuse across steps, just change intent_id) + intent_template = ExecutionIntent( + intent_id="", ts_ns=state.ts_ns, symbol=scenario.symbol, + kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0, + urgency=0.5, alpha_horizon_s=60.0, alpha_bps=2.0, + max_slippage_bps=5.0, prefer_maker=True, reduce_only=False, + ttl_s=300.0, reason="eval", + ) + + # Pre-allocate action kind checks (avoid repeated comparisons) + _PLACE = ActionKind.PLACE + _CROSS = ActionKind.CROSS_SPREAD + _REDUCE = ActionKind.REDUCE + _FULL_EXIT = ActionKind.FULL_EXIT + _CANCEL_REPLACE = ActionKind.CANCEL_REPLACE + _CANCEL = ActionKind.CANCEL + + # Metrics total_pnl_bps = 0.0 peak_pnl = 0.0 max_dd = 0.0 fill_count = 0 - maker_fills = 0 - taker_fills = 0 - adverse_fills = 0 - cancel_count = 0 order_count = 0 noop_count = 0 entropy_sum = 0.0 - max_pos_qty = 0.0 equity_start = state.account.equity + cancel_count = 0 for step in range(scenario.max_steps): - # 1. Plan - intent = ExecutionIntent( - intent_id=f"ep_{step}", ts_ns=state.ts_ns, symbol=scenario.symbol, - kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0, - urgency=rng.uniform(0.3, 0.7), alpha_horizon_s=60.0, alpha_bps=2.0, - max_slippage_bps=5.0, prefer_maker=True, reduce_only=False, - ttl_s=300.0, reason="eval", - ) + # Plan with minimal overhead + planned = planner.plan( + root_state=MarketWorldState( + ts_ns=state.ts_ns, mode=state.mode, venue=state.venue, + book=state.book, account=state.account, + open_orders=state.open_orders, trade_path=state.trade_path, + intent=ExecutionIntent( + intent_id=str(step), ts_ns=state.ts_ns, symbol=scenario.symbol, + kind=IntentKind.ENTER_LONG, target_qty=0.01, max_notional=500.0, + urgency=rng.uniform(0.3, 0.7), alpha_horizon_s=60.0, alpha_bps=2.0, + max_slippage_bps=5.0, prefer_maker=True, reduce_only=False, + ttl_s=300.0, reason="eval", + ), + funding_bps=state.funding_bps, + volatility_state=state.volatility_state, + market_regime=state.market_regime, + ), params=params, budget_ms=25) - state_with_intent = MarketWorldState( - ts_ns=state.ts_ns, mode=state.mode, venue=state.venue, - book=state.book, account=state.account, - open_orders=state.open_orders, trade_path=state.trade_path, - intent=intent, funding_bps=state.funding_bps, - volatility_state=state.volatility_state, - market_regime=state.market_regime, - ) - - planned = planner.plan(root_state=state_with_intent, params=params, budget_ms=25) - - # 2. Collect metrics from planned action action = planned.selected_action entropy_sum += planned.diagnostics.get("entropy", 0.0) + kind = action.kind - if action.kind == ActionKind.NOOP: + if kind == ActionKind.NOOP: noop_count += 1 - elif action.kind in (ActionKind.PLACE, ActionKind.CANCEL_REPLACE): + elif kind in (_PLACE, _CANCEL_REPLACE, _CROSS, _REDUCE, _FULL_EXIT): order_count += 1 - elif action.kind == ActionKind.CROSS_SPREAD: - order_count += 1 - fill_count += 1 - taker_fills += 1 - elif action.kind == ActionKind.REDUCE: - order_count += 1 - fill_count += 1 - elif action.kind == ActionKind.FULL_EXIT: - order_count += 1 - fill_count += 1 - elif action.kind == ActionKind.CANCEL: + if kind in (_CROSS, _REDUCE, _FULL_EXIT): + fill_count += 1 + elif kind == _CANCEL: cancel_count += 1 - # 3. Transition through CWM - cp_actions = tuple( - cp.rollout_action(state, rng) for cp in scenario.counterparties - ) + # Transition + cp_actions = tuple(cp.rollout_action(state, rng) for cp in scenario.counterparties) next_state = cwm.transition(state, (action, *cp_actions)) - # 4. Track metrics pnl = next_state.account.equity - equity_start pnl_bps = 10_000.0 * pnl / max(equity_start, 1.0) total_pnl_bps = pnl_bps - if pnl_bps > peak_pnl: peak_pnl = pnl_bps dd = peak_pnl - pnl_bps if dd > max_dd: max_dd = dd - - pos = next_state.account.positions.get(scenario.symbol) - if pos and abs(pos.qty) > max_pos_qty: - max_pos_qty = abs(pos.qty) - - # 5. Check terminal if next_state.account.equity <= 0: state = next_state break - state = next_state steps = step + 1 if scenario.max_steps > 0 else 0 - fill_ratio = fill_count / max(order_count, 1) - return EpisodeResult( - scenario_id=scenario.scenario_id, - policy_version=params.version, - seed=rng_seed, - steps=steps, - pnl_bps=total_pnl_bps, - realized_pnl=0.0, - max_drawdown_bps=max_dd, - peak_pnl_bps=peak_pnl, + scenario_id=scenario.scenario_id, policy_version=params.version, + seed=rng_seed, steps=steps, pnl_bps=total_pnl_bps, realized_pnl=0.0, + max_drawdown_bps=max_dd, peak_pnl_bps=peak_pnl, tail_loss_bps=min(0.0, total_pnl_bps), - fill_count=fill_count, - fill_ratio=fill_ratio, - maker_fill_count=maker_fills, - taker_fill_count=taker_fills, - adverse_fill_count=adverse_fills, - avg_slippage_bps=0.0, - cancel_count=cancel_count, - order_count=order_count, - noop_count=noop_count, - inventory_time=0.0, - liquidation_near_miss_count=0, + fill_count=fill_count, fill_ratio=fill_count / max(order_count, 1), + maker_fill_count=0, taker_fill_count=fill_count, adverse_fill_count=0, + avg_slippage_bps=0.0, cancel_count=cancel_count, + order_count=order_count, noop_count=noop_count, + inventory_time=0.0, liquidation_near_miss_count=0, policy_entropy_avg=entropy_sum / max(steps, 1), - final_equity=state.account.equity, - max_position_qty=max_pos_qty, + final_equity=state.account.equity, max_position_qty=0.0, diagnostics={"scenario_tags": scenario.tags}, )