CoolFace
Apppublic

garvitsachdeva/SpindleFlow-RL

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
curriculum.py123 linesDownload Raw Back to training
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