""" Ray-based parallel episode evaluation — industrial multi-core execution. Uses Ray for conflict-free, bit-identical parallel execution across cores. Each worker gets its own CWM + planner — zero shared state, zero races. Bit-identity guarantee: same seed + same params + same scenarios = identical results, regardless of worker count. Each worker processes one scenario independently. Usage: from malkhut.training.ray_eval import RayEpisodeRunner runner = RayEpisodeRunner(workers=4) results = runner.run_episodes(params, scenarios, seed=42) """ from __future__ import annotations import os from typing import List, Optional, Sequence import ray from malkhut.cwm.core import MinimalCryptoLOBCWM from malkhut.state import FulfilmentPolicyParams from malkhut.training.cma_trainer import EpisodeResult, Scenario @ray.remote def _run_episode_remote(params_bytes: bytes, scenario_bytes: bytes, rng_seed: int, planner_type: str) -> EpisodeResult: """Ray remote function. Each worker creates its own CWM + planner.""" import pickle from malkhut.training.cma_trainer import PolicyEvaluator from malkhut.cwm.core import MinimalCryptoLOBCWM params = pickle.loads(params_bytes) scenario = pickle.loads(scenario_bytes) evaluator = PolicyEvaluator(cwm_factory=MinimalCryptoLOBCWM) return evaluator._run_episode(params, scenario, rng_seed, planner_type) class RayEpisodeRunner: """Ray-based parallel episode evaluation. Conflict-free: each worker is an independent process with its own CWM. Bit-identical: same seed + same params = same results, regardless of workers. Industrial: uses Ray for multi-core scheduling, shared object store. """ def __init__(self, workers: int = 0) -> None: if workers <= 0: workers = min(os.cpu_count() or 4, 8) self.workers = workers self._initialized = False def _ensure_ray(self) -> None: if not self._initialized: if not ray.is_initialized(): ray.init( num_cpus=self.workers, ignore_reinit_error=True, log_to_driver=False, ) self._initialized = True def run_episodes( self, params: FulfilmentPolicyParams, scenarios: Sequence[Scenario], rng_seed: int = 0, planner_type: str = "sm_mcts", ) -> List[EpisodeResult]: """Run episodes in parallel via Ray. Returns results in original order.""" if len(scenarios) <= 1: return self._run_sequential(params, scenarios, rng_seed, planner_type) import pickle self._ensure_ray() # Put shared data in Ray's object store (zero pickle per worker) params_ref = ray.put(pickle.dumps(params)) scenario_refs = [ray.put(pickle.dumps(s)) for s in scenarios] # Launch all episodes as independent remote tasks futures = [ _run_episode_remote.remote( params_ref, scenario_refs[i], rng_seed + i, planner_type ) for i in range(len(scenarios)) ] # Gather results in order results = ray.get(futures) return list(results) def _run_sequential( self, params: FulfilmentPolicyParams, scenarios: Sequence[Scenario], rng_seed: int, planner_type: str, ) -> List[EpisodeResult]: """Sequential fallback for trivial cases.""" 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) ] def shutdown(self) -> None: """Shutdown Ray if we initialized it.""" if self._initialized and ray.is_initialized(): ray.shutdown() self._initialized = False