CoolFace
Apppublic

sankar-raul/ICD-10-code-predictor-env

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