garvitsachdeva/SpindleFlow-RL
0
1"""2CurriculumManager — performance-gated phase advancement.3 4Phases advance when rolling_mean_reward >= phase_advance_threshold,5not after a fixed episode count. Thresholds and window size come from config.6"""7 8from __future__ import annotations9from collections import deque10from dataclasses import dataclass11import yaml12 13 14@dataclass15class CurriculumPhase:16 phase: int17 name: str18 episode_budget: int19 task_types: list[str]20 enable_tier2: bool21 enable_tier3: bool22 23 24class CurriculumManager:25 """26 Tracks curriculum progress and transitions between phases.27 Advances when the rolling mean reward over the last N episodes28 exceeds a configurable threshold — not after a fixed episode count.29 """30 31 _PHASE_NAMES = {32 1: "Simple Delegation",33 2: "Moderate Tasks + Conflict",34 3: "Complex + Enterprise",35 }36 _TIER2_PHASES = {2, 3}37 _TIER3_PHASES = {3}38 39 def __init__(self, config_path: str = "configs/training_config.yaml"):40 with open(config_path) as f:41 cfg = yaml.safe_load(f)["curriculum"]42 43 # Performance-gated advancement parameters44 self._window_size = cfg.get("phase_advance_window", 50)45 self._thresholds = {46 1: cfg.get("phase1_advance_threshold", 0.30),47 2: cfg.get("phase2_advance_threshold", 0.50),48 }49 self._min_episodes = cfg.get("phase_min_episodes", 100)50 51 # Task types still read from config (used by TaskBank)52 self._phase_task_types = {53 1: cfg.get("phase1_task_types", ["atomic", "simple"]),54 2: cfg.get("phase2_task_types", ["moderate"]),55 3: cfg.get("phase3_task_types", ["complex", "enterprise"]),56 }57 # Legacy budget fields — kept for get_current_phase() / progress_str()58 self._phase_budgets = {59 1: cfg.get("phase1_episodes", 200),60 2: cfg.get("phase2_episodes", 400),61 3: cfg.get("phase3_episodes", 600),62 }63 64 self.current_phase = 165 self.episodes_in_phase = 066 self.total_episodes = 067 self._reward_window: deque[float] = deque(maxlen=self._window_size)68 69 def on_episode_end(self, episode_reward: float = 0.0) -> bool:70 """71 Called after each episode with the terminal reward.72 Returns True if the phase advanced.73 """74 self.total_episodes += 175 self.episodes_in_phase += 176 self._reward_window.append(episode_reward)77 78 if (79 self.current_phase < 380 and self.episodes_in_phase >= self._min_episodes81 and len(self._reward_window) >= self._window_size82 ):83 rolling_mean = sum(self._reward_window) / len(self._reward_window)84 threshold = self._thresholds.get(self.current_phase, float("inf"))85 if rolling_mean >= threshold:86 self.current_phase += 187 self.episodes_in_phase = 088 self._reward_window.clear()89 print(90 f"\n[Curriculum] >> Advanced to Phase {self.current_phase} "91 f"(rolling mean {rolling_mean:.3f} >= {threshold:.3f})"92 )93 return True94 return False95 96 @property97 def phase(self) -> int:98 return self.current_phase99 100 def rolling_mean(self) -> float:101 if not self._reward_window:102 return 0.0103 return sum(self._reward_window) / len(self._reward_window)104 105 def get_current_phase(self) -> CurriculumPhase:106 p = self.current_phase107 return CurriculumPhase(108 phase=p,109 name=self._PHASE_NAMES[p],110 episode_budget=self._phase_budgets[p],111 task_types=self._phase_task_types[p],112 enable_tier2=p in self._TIER2_PHASES,113 enable_tier3=p in self._TIER3_PHASES,114 )115 116 def progress_str(self) -> str:117 threshold = self._thresholds.get(self.current_phase, "—")118 return (119 f"Phase {self.current_phase}/3 | "120 f"Rolling mean: {self.rolling_mean():.3f} / {threshold} | "121 f"Episodes in phase: {self.episodes_in_phase}"122 )123 