Files
sentiment-engine/MALKHUT/malkhut/training/ray_eval.py

115 lines
3.9 KiB
Python
Raw Normal View History

"""
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