from __future__ import annotations import json import math import random import time from dataclasses import asdict, dataclass, field from pathlib import Path from typing import Any @dataclass(frozen=True) class TaskCase: task: str seed: int @property def case_id(self) -> str: return f"{self.task}:{self.seed}" @classmethod def from_json(cls, row: dict[str, Any]) -> "TaskCase": return cls(task=str(row["task"]), seed=int(row["seed"])) def to_json(self) -> dict[str, Any]: return {"task": self.task, "seed": self.seed} @dataclass class CaseStats: attempts: int = 0 successes: int = 0 failures: int = 0 consecutive_failures: int = 0 cooldown_until_epoch: int = 0 best_steps: int | None = None last_success: bool | None = None last_seen_epoch: int = 0 last_updated: float = 0.0 @property def success_rate(self) -> float: return self.successes / max(self.attempts, 1) def to_json(self) -> dict[str, Any]: return asdict(self) @classmethod def from_json(cls, row: dict[str, Any]) -> "CaseStats": return cls(**row) @dataclass class CurriculumConfig: # Failure Curriculum Filtering: repeatedly hopeless cases leave the active pool. min_attempts_before_filter: int = 3 fail_streak_to_cooldown: int = 3 cooldown_epochs: int = 5 # Same idea at task-family level. Exact seed cooldown is too sparse when a # whole AndroidWorld task family is currently unrewardable. task_min_attempts_before_filter: int = 8 task_fail_streak_to_cooldown: int = 6 task_cooldown_epochs: int = 5 task_low_success_rate: float = 0.15 # Too-easy cases are still sampled, just less often. easy_success_rate: float = 0.90 easy_after_attempts: int = 3 easy_weight: float = 0.25 # Difficulty-Adaptive Positive Replay: rare successes on harder cases matter more. positive_replay_bonus: float = 2.0 @dataclass class CurriculumState: epoch: int = 0 cases: dict[str, CaseStats] = field(default_factory=dict) tasks: dict[str, CaseStats] = field(default_factory=dict) successful_replay: list[dict[str, Any]] = field(default_factory=list) @classmethod def load(cls, path: str | Path) -> "CurriculumState": p = Path(path) if not p.exists(): return cls() raw = json.loads(p.read_text(encoding="utf-8")) return cls( epoch=int(raw.get("epoch", 0)), cases={k: CaseStats.from_json(v) for k, v in raw.get("cases", {}).items()}, tasks={k: CaseStats.from_json(v) for k, v in raw.get("tasks", {}).items()}, successful_replay=list(raw.get("successful_replay", [])), ) def save(self, path: str | Path) -> None: p = Path(path) p.parent.mkdir(parents=True, exist_ok=True) row = { "epoch": self.epoch, "cases": {k: v.to_json() for k, v in sorted(self.cases.items())}, "tasks": {k: v.to_json() for k, v in sorted(self.tasks.items())}, "successful_replay": self.successful_replay[-1000:], } p.write_text(json.dumps(row, indent=2, ensure_ascii=False), encoding="utf-8") def stats_for(self, case: TaskCase) -> CaseStats: return self.cases.setdefault(case.case_id, CaseStats()) def task_stats_for(self, task: str) -> CaseStats: return self.tasks.setdefault(task, CaseStats()) def load_split(path: str | Path) -> dict[str, list[TaskCase]]: raw = json.loads(Path(path).read_text(encoding="utf-8")) return { "train": [TaskCase.from_json(r) for r in raw.get("train", [])], "dev_eval": [TaskCase.from_json(r) for r in raw.get("dev_eval", [])], "final_eval": [TaskCase.from_json(r) for r in raw.get("final_eval", [])], } def is_in_cooldown(stats: CaseStats, epoch: int) -> bool: return stats.cooldown_until_epoch > epoch def maybe_filter_case(stats: CaseStats, epoch: int, cfg: CurriculumConfig) -> None: if stats.attempts < cfg.min_attempts_before_filter: return if stats.consecutive_failures >= cfg.fail_streak_to_cooldown: stats.cooldown_until_epoch = max( stats.cooldown_until_epoch, epoch + cfg.cooldown_epochs, ) def maybe_filter_task(stats: CaseStats, epoch: int, cfg: CurriculumConfig) -> None: if stats.attempts < cfg.task_min_attempts_before_filter: return if stats.consecutive_failures >= cfg.task_fail_streak_to_cooldown or stats.success_rate <= cfg.task_low_success_rate: stats.cooldown_until_epoch = max( stats.cooldown_until_epoch, epoch + cfg.task_cooldown_epochs, ) def sampling_weight(stats: CaseStats, epoch: int, cfg: CurriculumConfig, task_stats: CaseStats | None = None) -> float: if is_in_cooldown(stats, epoch): return 0.0 if task_stats is not None and is_in_cooldown(task_stats, epoch): return 0.0 if stats.attempts >= cfg.easy_after_attempts and stats.success_rate >= cfg.easy_success_rate: return cfg.easy_weight if stats.successes > 0: # Keep "sometimes succeeds" tasks hot; they provide contrast. difficulty = 1.0 - stats.success_rate return 1.0 + difficulty return 1.0 def choose_case( cases: list[TaskCase], state: CurriculumState, rng: random.Random, cfg: CurriculumConfig, ) -> TaskCase: if not cases: raise ValueError("No task cases available") weights = [sampling_weight(state.stats_for(c), state.epoch, cfg, state.task_stats_for(c.task)) for c in cases] if not any(w > 0 for w in weights): # Everything is cooling down. Advance the curriculum clock enough to retry. cooldowns = [] for case in cases: cooldowns.append(state.stats_for(case).cooldown_until_epoch) cooldowns.append(state.task_stats_for(case.task).cooldown_until_epoch) state.epoch = min(c for c in cooldowns if c > state.epoch) weights = [sampling_weight(state.stats_for(c), state.epoch, cfg, state.task_stats_for(c.task)) for c in cases] return rng.choices(cases, weights=weights, k=1)[0] def record_rollout( state: CurriculumState, case: TaskCase, success: bool, steps: int, cfg: CurriculumConfig, trajectory_path: str | None = None, ) -> CaseStats: stats = state.stats_for(case) task_stats = state.task_stats_for(case.task) stats.attempts += 1 task_stats.attempts += 1 stats.last_seen_epoch = state.epoch task_stats.last_seen_epoch = state.epoch stats.last_updated = time.time() task_stats.last_updated = stats.last_updated stats.last_success = bool(success) task_stats.last_success = bool(success) if success: stats.successes += 1 task_stats.successes += 1 stats.consecutive_failures = 0 task_stats.consecutive_failures = 0 stats.best_steps = steps if stats.best_steps is None else min(stats.best_steps, steps) task_stats.best_steps = steps if task_stats.best_steps is None else min(task_stats.best_steps, steps) hardness = 1.0 - stats.success_rate state.successful_replay.append( { "case_id": case.case_id, "task": case.task, "seed": case.seed, "steps": steps, "weight": 1.0 + cfg.positive_replay_bonus * max(hardness, 0.05), "trajectory_path": trajectory_path, "epoch": state.epoch, "ts": stats.last_updated, } ) else: stats.failures += 1 task_stats.failures += 1 stats.consecutive_failures += 1 task_stats.consecutive_failures += 1 maybe_filter_case(stats, state.epoch, cfg) maybe_filter_task(task_stats, state.epoch, cfg) return stats def summarize_state(cases: list[TaskCase], state: CurriculumState, cfg: CurriculumConfig) -> dict[str, Any]: active = cooldown = easy = attempted = successes = 0 rates = [] for case in cases: stats = state.stats_for(case) if stats.attempts: attempted += 1 successes += stats.successes rates.append(stats.success_rate) if is_in_cooldown(stats, state.epoch): cooldown += 1 else: active += 1 if stats.attempts >= cfg.easy_after_attempts and stats.success_rate >= cfg.easy_success_rate: easy += 1 task_cooldown = sum(1 for stats in state.tasks.values() if is_in_cooldown(stats, state.epoch)) task_attempted = sum(1 for stats in state.tasks.values() if stats.attempts) return { "epoch": state.epoch, "cases": len(cases), "active": active, "cooldown": cooldown, "task_cooldown": task_cooldown, "task_attempted": task_attempted, "easy": easy, "attempted": attempted, "successes": successes, "mean_case_success_rate": sum(rates) / max(len(rates), 1), "replay_items": len(state.successful_replay), }