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

92 lines
3.3 KiB
Python
Raw Normal View History

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