Fix duplicate WebSocket log output; update benchmark numbers to latest run (avg 75.9%, ~10min)
bfbbcd1 | import os | |
| import sys | |
| import time | |
| from typing import List, Optional | |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) | |
| from openai import OpenAI | |
| import requests as http_requests | |
| from client import CascadeContainmentEnv | |
| from models import ContainmentAction | |
| from baseline.policy import get_client, build_prompt, call_llm, parse_action, build_prompt_with_memory | |
| from core.trajectory import EpisodicMemory | |
| from core.policy_update import compute_advantage, update_memory | |
| from core.reward import normalise_score | |
| API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1") | |
| MODEL_NAME = os.getenv("MODEL_NAME", "meta-llama/Llama-3.1-8B-Instruct") | |
| ENV_BASE_URL = os.getenv("ENV_BASE_URL", "http://localhost:7860") | |
| BENCHMARK = "cascade-containment" | |
| N_ROLLOUTS = { | |
| "easy": 2, | |
| "medium": 3, | |
| "hard": 3, | |
| } | |
| # If a rollout already hits this score, skip remaining rollouts for the task. | |
| # Keeps runtime predictable when judges evaluate with slower models. | |
| EARLY_STOP_THRESHOLD = { | |
| "easy": 0.85, | |
| "medium": 0.72, | |
| "hard": 0.65, | |
| } | |
| # ── Mandatory structured log format ────────────────────────────────────────── | |
| def log_start(task: str, env: str, model: str) -> None: | |
| print(f"[START] task={task} env={env} model={model}", flush=True) | |
| def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None: | |
| print( | |
| f"[STEP] step={step} action={action} reward={reward:.2f} " | |
| f"done={str(done).lower()} error={error if error else 'null'}", | |
| flush=True, | |
| ) | |
| def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None: | |
| print( | |
| f"[END] success={str(success).lower()} steps={steps} score={score:.3f} " | |
| f"rewards={','.join(f'{r:.2f}' for r in rewards)}", | |
| flush=True, | |
| ) | |
| def run_rollout( | |
| env, | |
| task_name: str, | |
| client: OpenAI, | |
| memory: EpisodicMemory, | |
| rollout_idx: int, | |
| ) -> tuple: | |
| result = env.reset(task_name=task_name) | |
| obs = result.observation | |
| done = result.done | |
| total_reward = 0.0 | |
| step = 0 | |
| trajectory = [] | |
| rewards = [] | |
| log_start(task=f"{task_name}-r{rollout_idx}", env=BENCHMARK, model=MODEL_NAME) | |
| seen_steps: set = set() # deduplicate WebSocket replay artefacts | |
| end_logged: bool = False | |
| try: | |
| while not done: | |
| prompt = build_prompt_with_memory(obs, memory) | |
| response = call_llm(prompt, client) | |
| action = parse_action(response, len(obs.districts)) | |
| action_str = f"{action.action_type}(district={action.district_id})" | |
| try: | |
| result = env.step(action) | |
| except Exception as e: | |
| log_step(step=step + 1, action=action_str, reward=0.0, done=True, error=str(e)[:80]) | |
| if not end_logged: | |
| end_logged = True | |
| log_end(success=False, steps=step, score=0.0, rewards=rewards) | |
| return total_reward, step, trajectory, 0.0 | |
| next_obs = result.observation | |
| reward = result.reward or 0.0 | |
| done = result.done | |
| total_reward += reward | |
| step += 1 | |
| rewards.append(reward) | |
| trajectory.append({"obs": obs, "action": action, "reward": reward}) | |
| # Only log each step number once — WebSocket can replay buffered responses | |
| if step not in seen_steps: | |
| seen_steps.add(step) | |
| log_step(step=step, action=action_str, reward=reward, done=done, error=None) | |
| obs = next_obs | |
| if done: | |
| break | |
| score = 0.0 | |
| try: | |
| grade_resp = http_requests.get(ENV_BASE_URL.rstrip('/') + '/grade', timeout=10) | |
| if grade_resp.status_code == 200: | |
| score = grade_resp.json().get("final_score", 0.0) | |
| except Exception: | |
| num_districts = {"easy": 2, "medium": 4, "hard": 6}.get(task_name, 2) | |
| score = normalise_score(total_reward, step, num_districts) | |
| success = score >= 0.40 | |
| except Exception: | |
| if not end_logged: | |
| end_logged = True | |
| log_end(success=False, steps=step, score=0.0, rewards=rewards) | |
| return total_reward, step, trajectory, 0.0 | |
| if not end_logged: | |
| end_logged = True | |
| log_end(success=success, steps=step, score=score, rewards=rewards) | |
| return total_reward, step, trajectory, score | |
| def run_task(env, task_name: str, client: OpenAI) -> float: | |
| n_rollouts = N_ROLLOUTS[task_name] | |
| threshold = EARLY_STOP_THRESHOLD[task_name] | |
| memory = EpisodicMemory(max_size=20) | |
| rollouts = [] | |
| for i in range(1, n_rollouts + 1): | |
| total_reward, steps, trajectory, score = run_rollout( | |
| env, task_name, client, memory, rollout_idx=i | |
| ) | |
| rollouts.append((total_reward, steps, score)) | |
| completed_rewards = [r[0] for r in rollouts] | |
| advantage = compute_advantage(total_reward, completed_rewards[:-1]) | |
| update_memory(memory, trajectory, advantage) | |
| if score >= threshold: | |
| break | |
| return max(r[2] for r in rollouts) | |
| def main() -> dict: | |
| client = get_client() | |
| scores = {} | |
| start = time.time() | |
| with CascadeContainmentEnv(base_url=ENV_BASE_URL).sync() as env: | |
| for task_name in ["easy", "medium", "hard"]: | |
| try: | |
| scores[task_name] = run_task(env, task_name, client) | |
| except Exception as e: | |
| scores[task_name] = 0.0 | |
| print(f"[DEBUG] Task {task_name} failed: {e}", flush=True) | |
| scores["average"] = round( | |
| sum(v for k, v in scores.items() if k != "average") / 3, 4 | |
| ) | |
| elapsed = round(time.time() - start, 1) | |
| print( | |
| f"\n# SCORES easy={scores.get('easy', 0):.4f} " | |
| f"medium={scores.get('medium', 0):.4f} " | |
| f"hard={scores.get('hard', 0):.4f} " | |
| f"average={scores.get('average', 0):.4f} " | |
| f"elapsed={elapsed}s", | |
| flush=True, | |
| ) | |
| return scores | |
| if __name__ == "__main__": | |
| scores = main() | |
| if scores.get("average", 0.0) == 0.0: | |
| print("\n[DEBUG] All scores zero — check environment variables:", flush=True) | |
| print(f" ENV_BASE_URL = {ENV_BASE_URL}", flush=True) | |
| print(f" API_BASE_URL = {API_BASE_URL}", flush=True) | |
| print(f" MODEL_NAME = {MODEL_NAME}", flush=True) | |
| sys.exit(1) | |