Spaces:
Sleeping
Sleeping
| #!/usr/bin/env python3 | |
| """ | |
| Sanity-check hard-task strategy separation. | |
| """ | |
| import json | |
| import random | |
| import sys | |
| from pathlib import Path | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| if str(PROJECT_ROOT) not in sys.path: | |
| sys.path.insert(0, str(PROJECT_ROOT)) | |
| from agents.baseline import run_episode as run_baseline_episode | |
| from agents.random_agent import run_episode as run_random_episode | |
| from env.environment import SynapseXEnvironment | |
| from env.grader import TASK_REGISTRY, TASK_SEEDS, grade | |
| from env.models import Action, ActionPayload | |
| def run_policy(name: str, runner) -> dict: | |
| env = SynapseXEnvironment(task_config=TASK_REGISTRY["hard"], seed=TASK_SEEDS["hard"]) | |
| actions = runner(env) | |
| result = grade("hard", actions) | |
| return {"name": name, **result.model_dump()} | |
| def always_execute(env: SynapseXEnvironment) -> list[ActionPayload]: | |
| obs = env.reset() | |
| actions: list[ActionPayload] = [] | |
| for _ in range(env.MAX_STEPS): | |
| if obs.episode_done: | |
| break | |
| active = [task for task in obs.tasks if task.is_active] | |
| if not active: | |
| break | |
| action = Action(action_type="execute", task_id=sorted(active, key=lambda task: task.id)[0].id) | |
| actions.append(action.model_dump()) | |
| obs = env.step(action).observation | |
| return actions | |
| def always_delay(env: SynapseXEnvironment) -> list[ActionPayload]: | |
| obs = env.reset() | |
| actions: list[ActionPayload] = [] | |
| for _ in range(env.MAX_STEPS): | |
| if obs.episode_done: | |
| break | |
| active = [task for task in obs.tasks if task.is_active] | |
| if not active: | |
| break | |
| action = Action(action_type="delay", task_id=sorted(active, key=lambda task: task.id)[0].id) | |
| actions.append(action.model_dump()) | |
| obs = env.step(action).observation | |
| return actions | |
| def main(): | |
| results = [ | |
| run_policy("random", lambda env: run_random_episode(env, rng=random.Random(42))), | |
| run_policy("always_execute", always_execute), | |
| run_policy("always_delay", always_delay), | |
| run_policy("baseline", run_baseline_episode), | |
| ] | |
| print(f"Hard-mode strategy comparison (seed={TASK_SEEDS['hard']})") | |
| for result in results: | |
| print(f"{result['name']:15s} score={result['score']:.4f} reward_score={result['reward_score']:.4f}") | |
| print() | |
| print(json.dumps(results, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |