115 lines
3.9 KiB
Python
115 lines
3.9 KiB
Python
|
|
"""
|
||
|
|
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
|