CoolFace
Apppublic

prashant-9457/my-openenv-task

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py203 linesDownload Raw Back to root
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