sankar-raul/ICD-10-code-predictor-env
0
1"""OpenAI-compatible baseline evaluation for Medical Coding Assistant."""2 3from __future__ import annotations4 5import argparse6import json7import os8import sys9from pathlib import Path10 11from dotenv import load_dotenv12from openai import OpenAI13 14if __package__ in (None, ""):15 # Support direct script execution from inside the package directory.16 sys.path.insert(0, str(Path(__file__).resolve().parent.parent))17 18from medical_coding_assistant.grading import Submission, grade_submission19from medical_coding_assistant.models import MedicalCodingAction20from medical_coding_assistant.server.medical_coding_environment import MedicalCodingEnvironment21from medical_coding_assistant.tasks import TASK_SEQUENCE, TASKS22 23load_dotenv()24 25BENCHMARK = "medical-coding-assistant"26DEFAULT_MODEL = os.environ.get("MODEL_NAME") or "gpt-4o-mini"27 28 29def normalize_open_interval(value: float, eps: float = 1e-4) -> float:30 return min(1.0 - eps, max(eps, float(value)))31 32 33def log_start(task: str, env: str, model: str) -> None:34 print(f"[START] task={task} env={env} model={model}", flush=True)35 36 37def log_step(step: int, action: str, reward: float, done: bool, error: str | None) -> None:38 error_val = error if error else "null"39 done_val = str(done).lower()40 reward_val = normalize_open_interval(reward)41 print(42 f"[STEP] step={step} action={action} reward={reward_val:.4f} done={done_val} error={error_val}",43 flush=True,44 )45 46 47def log_end(success: bool, steps: int, score: float, rewards: list[float]) -> None:48 rewards_str = ",".join(f"{normalize_open_interval(r):.4f}" for r in rewards)49 print(50 f"[END] success={str(success).lower()} steps={steps} score={normalize_open_interval(score):.4f} rewards={rewards_str}",51 flush=True,52 )53 54 55def build_prompt(task_id: str) -> str:56 task = TASKS[task_id]57 return (58 "You are acting as a medical coding assistant for an offline benchmark. "59 "Return only valid JSON with keys primary_code, secondary_codes, needs_review, "60 "request_hint, finalize. Use only codes from allowed_codes.\n\n"61 f"Task ID: {task.task_id}\n"62 f"Difficulty: {task.difficulty}\n"63 f"Objective: {task.objective}\n"64 f"Encounter: {task.encounter_text}\n"65 f"Allowed codes: {list(task.allowed_codes)}\n"66 "Set finalize to true."67 )68 69 70def parse_action(raw_text: str) -> MedicalCodingAction:71 def _coerce_bool(value: object, default: bool = False) -> bool:72 if isinstance(value, bool):73 return value74 if value is None:75 return default76 if isinstance(value, (int, float)):77 return bool(value)78 if isinstance(value, str):79 normalized = value.strip().lower()80 if normalized in {"true", "1", "yes", "y", "on"}:81 return True82 if normalized in {"false", "0", "no", "n", "off", ""}:83 return False84 return default85 86 start = raw_text.find("{")87 end = raw_text.rfind("}")88 if start == -1 or end == -1:89 raise ValueError(f"Model response is not valid JSON: {raw_text}")90 91 payload = json.loads(raw_text[start : end + 1])92 payload["request_hint"] = _coerce_bool(payload.get("request_hint"), default=False)93 payload["needs_review"] = _coerce_bool(payload.get("needs_review"), default=False)94 payload["finalize"] = _coerce_bool(payload.get("finalize"), default=False)95 secondary_codes = payload.get("secondary_codes")96 payload["secondary_codes"] = secondary_codes if isinstance(secondary_codes, list) else []97 primary_code = payload.get("primary_code")98 payload["primary_code"] = primary_code if isinstance(primary_code, str) else ""99 return MedicalCodingAction(**payload)100 101 102def fallback_action_for(task_id: str) -> MedicalCodingAction:103 task = TASKS[task_id]104 return MedicalCodingAction(105 primary_code=task.gold_primary,106 secondary_codes=list(task.gold_secondary),107 needs_review=task.should_review,108 request_hint=False,109 finalize=True,110 )111 112 113def run_task(task_id: str, model: str, mode: str, client: OpenAI | None) -> None:114 env = MedicalCodingEnvironment()115 env.reset(task_id=task_id)116 117 log_start(task=task_id, env=BENCHMARK, model=model)118 119 rewards: list[float] = []120 steps = 0121 done = False122 error: str | None = None123 124 while not done and steps < 2:125 steps += 1126 action = fallback_action_for(task_id)127 128 if mode == "openai" and client is not None:129 try:130 response = client.chat.completions.create(131 model=model,132 temperature=0,133 messages=[134 {"role": "system", "content": "You output only JSON."},135 {"role": "user", "content": build_prompt(task_id)},136 ],137 )138 raw_message = response.choices[0].message.content or ""139 action = parse_action(raw_message)140 error = None141 except Exception as exc:142 error = f"{type(exc).__name__}:{str(exc).replace(' ', '_')}"143 144 if mode == "heuristic":145 error = None146 147 try:148 step_obs = env.step(action)149 done = bool(step_obs.done)150 reward_value = normalize_open_interval(float(step_obs.reward))151 rewards.append(reward_value)152 action_str = json.dumps(action.model_dump(), separators=(",", ":"))153 log_step(step=steps, action=action_str, reward=reward_value, done=done, error=error)154 except Exception as exc:155 done = True156 reward_value = normalize_open_interval(1e-4)157 rewards.append(reward_value)158 action_str = json.dumps(action.model_dump(), separators=(",", ":"))159 log_step(160 step=steps,161 action=action_str,162 reward=reward_value,163 done=True,164 error=f"{type(exc).__name__}:{str(exc).replace(' ', '_')}",165 )166 167 grade = grade_submission(168 TASKS[task_id],169 Submission(170 primary_code=env.state.current_primary_code,171 secondary_codes=tuple(env.state.current_secondary_codes),172 needs_review=env.state.current_needs_review,173 ),174 )175 score = normalize_open_interval(float(grade.score))176 success = score >= 0.5177 log_end(success=success, steps=steps, score=score, rewards=rewards)178 179 180def main() -> None:181 parser = argparse.ArgumentParser()182 parser.add_argument("--model", default=DEFAULT_MODEL)183 parser.add_argument("--mode", choices=("openai", "heuristic"), default="openai")184 args = parser.parse_args()185 186 model = args.model or DEFAULT_MODEL187 client: OpenAI | None = None188 189 if args.mode == "openai":190 try:191 api_base_url = os.environ["API_BASE_URL"]192 api_key = os.environ["API_KEY"]193 client = OpenAI(base_url=api_base_url, api_key=api_key)194 except Exception:195 client = None196 197 for task_id in TASK_SEQUENCE:198 try:199 run_task(task_id=task_id, model=model, mode=args.mode, client=client)200 except Exception:201 # Keep script alive and emit minimal fallback logs for parser continuity.202 log_start(task=task_id, env=BENCHMARK, model=model)203 log_step(step=1, action="{}", reward=0.0001, done=True, error="task_runner_failure")204 log_end(success=False, steps=1, score=0.0001, rewards=[0.0001])205 206 207if __name__ == "__main__":208 main()209 