prashant-9457/my-openenv-task
0
1# -*- coding: utf-8 -*-2# inference.py - ICU Resource Allocation OpenEnv3 4import os5import sys6import time7import requests8from openai import OpenAI9 10API_BASE_URL = os.getenv("API_BASE_URL") # Must be set by evaluator — no fallback11MODEL_NAME = os.getenv("MODEL_NAME", "meta-llama/Llama-3.2-3B-Instruct")12API_KEY = os.getenv("API_KEY")13 14ENV_BASE_URL = "http://localhost:7860"15BENCHMARK = "icu-resource-allocation"16MAX_STEPS = 4817SUCCESS_THRESHOLD = 0.4018REQUEST_TIMEOUT = 3019 20ACTION_NAMES = {21 0: "HOLD", 1: "ADMIT_CRITICAL", 2: "ADMIT_FIFO",22 3: "TRANSFER_OUT", 4: "CALL_EXTRA_NURSE",23 5: "SPECIALIST_CONSULT", 6: "EXPEDITE_BED",24}25 26SYSTEM_PROMPT = (27 "You are an ICU charge coordinator. Reply with ONE digit 0-6 only. "28 "0=HOLD 1=ADMIT_CRITICAL 2=ADMIT_FIFO 3=TRANSFER_OUT "29 "4=CALL_EXTRA_NURSE 5=SPECIALIST_CONSULT 6=EXPEDITE_BED."30)31 32 33def log_start(task, env_name, model):34 print(f"[START] task={task} env={env_name} model={model}", flush=True)35 36def log_step(step, action, reward, done, error):37 error_val = error if error else "null"38 print(f"[STEP] step={step} action={action} reward={reward:.2f} done={str(done).lower()} error={error_val}", flush=True)39 40def log_end(success, steps, score, rewards):41 rewards_str = ",".join(f"{r:.2f}" for r in rewards)42 print(f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", flush=True)43 44 45def _wait_for_server(max_wait=60):46 for _ in range(max_wait):47 try:48 r = requests.get(ENV_BASE_URL + "/", timeout=5)49 if r.status_code == 200:50 return True51 except Exception:52 pass53 time.sleep(1)54 return False55 56def _reset_env(seed=42):57 r = requests.post(ENV_BASE_URL + "/reset", json={"seed": seed}, timeout=REQUEST_TIMEOUT)58 r.raise_for_status()59 return r.json()60 61def _step_env(action):62 r = requests.post(ENV_BASE_URL + "/step", json={"action": action}, timeout=REQUEST_TIMEOUT)63 r.raise_for_status()64 data = r.json()65 return data["observation"], float(data["reward"]), bool(data["done"]), data.get("info", {})66 67def _obs_to_prompt(obs):68 return (69 f"Beds {obs.get('beds_occupied',0)}/20 free={obs.get('beds_available',0)} "70 f"Queue CRITICAL={obs.get('queue_critical',0)} "71 f"ratio={obs.get('nurse_patient_ratio',1.2)} "72 f"Step {obs.get('step',0)}/48. Reply ONE digit 0-6."73 )74 75def _fallback(obs):76 if obs.get("queue_critical", 0) > 0 and obs.get("beds_available", 0) > 0:77 return 178 if obs.get("nurse_patient_ratio", 1.0) > 2.2:79 return 480 if obs.get("queue_total", 0) > 0 and obs.get("beds_available", 0) == 0:81 return 382 if obs.get("queue_total", 0) > 0 and obs.get("beds_available", 0) > 0:83 return 284 return 085 86def _get_action(client, obs):87 try:88 resp = client.chat.completions.create(89 model=MODEL_NAME,90 messages=[91 {"role": "system", "content": SYSTEM_PROMPT},92 {"role": "user", "content": _obs_to_prompt(obs)},93 ],94 max_tokens=5,95 temperature=0.0,96 )97 raw = (resp.choices[0].message.content or "").strip()98 print(f"[DEBUG] LLM raw={raw!r}", flush=True)99 if raw and raw[0].isdigit():100 a = int(raw[0])101 if 0 <= a <= 6:102 return a103 except Exception as e:104 print(f"[DEBUG] LLM call failed: {e} — using fallback", flush=True)105 return _fallback(obs)106 107def _score(task_id, m):108 try:109 if task_id == "task_easy":110 raw = 0.60*(1.0-min(0.999,m["deaths"]*0.25)) + 0.40*(1.0-m["ratio_breach_frac"])111 elif task_id == "task_medium":112 raw = (0.40*max(0.0,1.0-m["deaths"]*0.30) +113 0.30*(1.0-m["ratio_breach_frac"]) +114 0.30*max(0.0,1.0-m["wait_violations"]*0.15))115 elif task_id == "task_hard":116 bu = m["budget_used_pct"]117 bs = 1.0 if bu<=0.85 else max(0.0,1.0-(bu-0.85)*4)118 ss = 1.0 if m["sofa_trend"]<=0 else max(0.0,1.0-m["sofa_trend"]/5.0)119 raw = (0.30*max(0.0,1.0-m["deaths"]*0.35)+0.20*max(0.0,1.0-m["adverse"]*0.10)+120 0.20*max(0.0,1.0-m["wait_violations"]*0.12)+0.15*bs+0.15*ss)121 else:122 raw = 0.0123 return round(min(0.999, max(0.001, raw*0.92+0.04)), 3)124 except Exception:125 return 0.001126 127def run_task(task_id, client):128 rewards, ratio_breaches, sofa_traj = [], [], []129 steps_taken, score, success, obs = 0, 0.001, False, {}130 log_start(task=task_id, env_name=BENCHMARK, model=MODEL_NAME)131 try:132 obs = _reset_env(seed=42)133 done = False134 for step_n in range(1, MAX_STEPS + 1):135 if done:136 break137 action_int = _get_action(client, obs)138 action_str = ACTION_NAMES.get(action_int, str(action_int))139 obs, reward, done, _ = _step_env(action_int)140 rewards.append(reward)141 ratio_breaches.append(float(obs.get("nurse_patient_ratio", 1.0)) > 2.0)142 sofa_traj.append(float(obs.get("avg_icu_sofa", 0.0)))143 steps_taken = step_n144 log_step(step_n, action_str, reward, done, None)145 if steps_taken > 0:146 n = max(1, len(ratio_breaches))147 metrics = {148 "deaths": int(obs.get("deaths_in_queue", 0)),149 "adverse": int(obs.get("adverse_events", 0)),150 "wait_violations": int(obs.get("wait_violations", 0)),151 "ratio_breach_frac": sum(ratio_breaches) / n,152 "budget_used_pct": float(obs.get("budget_utilisation_pct", 0)) / 100.0,153 "sofa_trend": (sofa_traj[-1]-sofa_traj[0]) if len(sofa_traj)>=2 else 0.0,154 }155 score = _score(task_id, metrics)156 success = score >= SUCCESS_THRESHOLD157 except Exception as e:158 print(f"[DEBUG] CRASH in {task_id}: {e}", flush=True)159 import traceback; traceback.print_exc()160 finally:161 log_end(success=success, steps=steps_taken, score=score, rewards=rewards)162 163 164def main():165 # Print ALL environment variables related to API so we can diagnose166 print(f"[DEBUG] API_BASE_URL={API_BASE_URL}", flush=True)167 print(f"[DEBUG] MODEL_NAME={MODEL_NAME}", flush=True)168 print(f"[DEBUG] API_KEY set={bool(API_KEY)}", flush=True)169 print(f"[DEBUG] All env keys with API: {[k for k in os.environ if 'API' in k.upper()]}", flush=True)170 print(f"[DEBUG] All env keys with TOKEN: {[k for k in os.environ if 'TOKEN' in k.upper()]}", flush=True)171 print(f"[DEBUG] All env keys with KEY: {[k for k in os.environ if 'KEY' in k.upper()]}", flush=True)172 173 if not API_KEY:174 print("ERROR: API_KEY not set", flush=True)175 sys.exit(1)176 177 if not API_BASE_URL:178 print("ERROR: API_BASE_URL not set — cannot proceed without the LiteLLM proxy URL", flush=True)179 sys.exit(1)180 181 print("[DEBUG] Waiting for env server...", flush=True)182 if not _wait_for_server(max_wait=60):183 print("[DEBUG] Server not ready, continuing anyway", flush=True)184 185 client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)186 print(f"[DEBUG] client created base_url={client.base_url}", flush=True)187 188 for task_id in ["task_easy", "task_medium", "task_hard"]:189 run_task(task_id, client)190 191 192if __name__ == "__main__":193 try:194 main()195 except Exception as e:196 print(f"[DEBUG] FATAL: {e}", flush=True)197 import traceback; traceback.print_exc()198 for task_id in ["task_easy", "task_medium", "task_hard"]:199 log_start(task_id, BENCHMARK, MODEL_NAME)200 log_end(success=False, steps=0, score=0.001, rewards=[])201 finally:202 sys.exit(0)203 