""" Parallel episode evaluator — multiprocessing-accelerated scenario evaluation. Each scenario is completely independent (own CWM, planner, counterparties, state), so we can evaluate them in parallel with ZERO fidelity loss. Usage: from malkhut.training.parallel_eval import ParallelEpisodeRunner runner = ParallelEpisodeRunner(workers=4) results = runner.run_episodes(params, scenarios, seed=42) """ from __future__ import annotations import os from concurrent.futures import ProcessPoolExecutor, as_completed from dataclasses import dataclass from typing import Any, Callable, List, Sequence, Tuple from malkhut.cwm.core import MinimalCryptoLOBCWM from malkhut.state import FulfilmentPolicyParams from malkhut.training.cma_trainer import EpisodeResult, Scenario def _run_single_episode(args: Tuple) -> EpisodeResult: """Top-level function for multiprocessing. Must be picklable.""" params_pkl, scenario_pkl, rng_seed, planner_type = args import pickle from malkhut.training.cma_trainer import PolicyEvaluator from malkhut.cwm.core import MinimalCryptoLOBCWM params = pickle.loads(params_pkl) scenario = pickle.loads(scenario_pkl) evaluator = PolicyEvaluator(cwm_factory=MinimalCryptoLOBCWM) return evaluator._run_episode(params, scenario, rng_seed, planner_type) class ParallelEpisodeRunner: """Evaluate multiple scenarios in parallel using process pool. Each worker gets its own CWM + planner instance — zero shared state. Deterministic: same seed → same result regardless of worker count. """ def __init__(self, workers: int = 0) -> None: if workers <= 0: workers = min(os.cpu_count() or 4, 8) self.workers = workers def run_episodes( self, params: FulfilmentPolicyParams, scenarios: Sequence[Scenario], rng_seed: int = 0, planner_type: str = "sm_mcts", ) -> List[EpisodeResult]: """Run episodes in parallel. Returns results in original order.""" if self.workers <= 1 or len(scenarios) <= 1: return self._run_sequential(params, scenarios, rng_seed, planner_type) import pickle params_pkl = pickle.dumps(params) tasks = [ (params_pkl, pickle.dumps(s), rng_seed + i, planner_type) for i, s in enumerate(scenarios) ] results: List[EpisodeResult] = [None] * len(scenarios) # type: ignore[list-item] with ProcessPoolExecutor(max_workers=self.workers) as pool: future_to_idx = { pool.submit(_run_single_episode, task): i for i, task in enumerate(tasks) } for future in as_completed(future_to_idx): idx = future_to_idx[future] results[idx] = future.result() return results def _run_sequential( self, params: FulfilmentPolicyParams, scenarios: Sequence[Scenario], rng_seed: int, planner_type: str, ) -> List[EpisodeResult]: """Sequential fallback (single worker or trivial case).""" from malkhut.training.cma_trainer import PolicyEvaluator evaluator = PolicyEvaluator(cwm_factory=MinimalCryptoLOBCWM) return [ evaluator._run_episode(params, s, rng_seed + i, planner_type) for i, s in enumerate(scenarios) ]