CoolFace
Apppublic

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

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
medical_coding_environment.py243 linesDownload Raw Back to server
1"""Core OpenEnv environment for Medical Coding Assistant."""2 3from __future__ import annotations4 5from typing import Any6from uuid import uuid47 8from openenv.core.env_server.interfaces import Environment9from openenv.core.env_server.types import EnvironmentMetadata10 11try:12    from medical_coding_assistant.grading import Submission, grade_submission13    from medical_coding_assistant.models import (14        MedicalCodingAction,15        MedicalCodingObservation,16        RewardBreakdown,17        MedicalCodingState,18    )19    from medical_coding_assistant.tasks import TASK_SEQUENCE, TASKS, TaskCase20except ModuleNotFoundError:21    from grading import Submission, grade_submission22    from models import MedicalCodingAction, MedicalCodingObservation, MedicalCodingState, RewardBreakdown23    from tasks import TASK_SEQUENCE, TASKS, TaskCase24 25 26class MedicalCodingEnvironment(27    Environment[MedicalCodingAction, MedicalCodingObservation, MedicalCodingState]28):29    """A deterministic clinical coding workflow benchmark."""30 31    SUPPORTS_CONCURRENT_SESSIONS: bool = True32    MAX_STEPS: int = 633 34    def __init__(self) -> None:35        super().__init__()36        self._task_index = 037        self._task: TaskCase = TASKS[TASK_SEQUENCE[0]]38        self._last_signature = ""39        self._state = MedicalCodingState(episode_id=str(uuid4()), step_count=0)40 41    def reset(42        self,43        seed: int | None = None,44        episode_id: str | None = None,45        **kwargs: Any,46    ) -> MedicalCodingObservation:47        task_id = kwargs.get("task_id")48        if task_id is None:49            task_id = TASK_SEQUENCE[self._task_index % len(TASK_SEQUENCE)]50            self._task_index += 151        if task_id not in TASKS:52            raise ValueError(f"Unknown task_id '{task_id}'. Available tasks: {list(TASKS)}")53 54        self._task = TASKS[task_id]55        self._last_signature = ""56        self._state = MedicalCodingState(57            episode_id=episode_id or str(uuid4()),58            step_count=0,59            current_task_id=self._task.task_id,60            difficulty=self._task.difficulty,  # type: ignore[arg-type]61            current_primary_code="",62            current_secondary_codes=[],63            current_needs_review=False,64            best_score=0.0,65            hints_used=0,66            repeated_actions=0,67            completed=False,68        )69        return self._build_observation(70            reward=0.0,71            done=False,72            feedback=["Environment reset. Start by drafting codes or requesting a hint."],73            reward_breakdown=RewardBreakdown(total=0.0),74        )75 76    def step(77        self,78        action: MedicalCodingAction,79        timeout_s: float | None = None,80        **kwargs: Any,81    ) -> MedicalCodingObservation:82        if self._state.completed:83            return self._build_observation(84                reward=0.0,85                done=True,86                feedback=["Task is already complete. Call reset() for a new task."],87            )88 89        self._state.step_count += 190        feedback: list[str] = []91        score_delta = 0.092        hint_penalty = 0.093        invalid_code_penalty = 0.094        loop_penalty = 0.095        timeout_penalty = 0.096 97        if action.request_hint:98            if self._state.hints_used < len(self._task.hints):99                self._state.hints_used += 1100                hint_penalty -= 0.02101                feedback.append("Hint revealed with a small efficiency penalty.")102            else:103                hint_penalty -= 0.05104                feedback.append("No hints remain.")105 106        proposed_primary = action.primary_code.strip().upper() or self._state.current_primary_code107        proposed_secondary = [108            code.strip().upper()109            for code in (action.secondary_codes or self._state.current_secondary_codes)110            if code.strip()111        ]112        proposed_secondary = list(dict.fromkeys(proposed_secondary))113        proposed_review = action.needs_review114 115        invalid_codes = [116            code117            for code in [proposed_primary, *proposed_secondary]118            if code and code not in self._task.allowed_codes119        ]120        if invalid_codes:121            invalid_code_penalty -= min(0.2, 0.1 * len(invalid_codes))122            feedback.append(f"Unsupported code(s): {invalid_codes}.")123 124        signature = "|".join(125            [126                proposed_primary,127                ",".join(proposed_secondary),128                str(proposed_review),129                str(action.request_hint),130                str(action.finalize),131            ]132        )133        if signature == self._last_signature:134            self._state.repeated_actions += 1135            loop_penalty -= 0.05136            feedback.append("Repeated the same draft; small loop penalty applied.")137        self._last_signature = signature138 139        self._state.current_primary_code = proposed_primary140        self._state.current_secondary_codes = proposed_secondary141        self._state.current_needs_review = proposed_review142 143        grade = grade_submission(144            self._task,145            Submission(146                primary_code=proposed_primary,147                secondary_codes=tuple(proposed_secondary),148                needs_review=proposed_review,149            ),150        )151        delta = round(grade.score - self._state.best_score, 4)152        if delta > 0:153            score_delta = delta154            feedback.append(f"Draft improved by {delta:.2f}.")155        self._state.best_score = max(self._state.best_score, grade.score)156        feedback.extend(list(grade.feedback))157 158        done = False159        if action.finalize or grade.score >= 1.0:160            done = True161            self._state.completed = True162            feedback.append("Task finalized.")163        elif self._state.step_count >= self.MAX_STEPS:164            done = True165            self._state.completed = True166            timeout_penalty -= 0.1167            feedback.append("Step budget exhausted.")168 169        total_reward = round(170            score_delta171            + hint_penalty172            + invalid_code_penalty173            + loop_penalty174            + timeout_penalty,175            4,176        )177        reward_breakdown = RewardBreakdown(178            score_delta=score_delta,179            hint_penalty=hint_penalty,180            invalid_code_penalty=invalid_code_penalty,181            loop_penalty=loop_penalty,182            timeout_penalty=timeout_penalty,183            total=total_reward,184        )185 186        return self._build_observation(187            reward=total_reward,188            done=done,189            feedback=feedback,190            grader_score=grade.score,191            reward_breakdown=reward_breakdown,192        )193 194    def _build_observation(195        self,196        reward: float,197        done: bool,198        feedback: list[str],199        grader_score: float | None = None,200        reward_breakdown: RewardBreakdown | None = None,201    ) -> MedicalCodingObservation:202        score = self._state.best_score if grader_score is None else grader_score203        breakdown = reward_breakdown or RewardBreakdown(total=reward)204        info = {205            "task_id": self._task.task_id,206            "difficulty": self._task.difficulty,207            "hints_remaining": max(0, len(self._task.hints) - self._state.hints_used),208            "grader_score": score,209            "reward_breakdown": breakdown.model_dump(),210        }211        return MedicalCodingObservation(212            task_id=self._task.task_id,213            difficulty=self._task.difficulty,  # type: ignore[arg-type]214            objective=self._task.objective,215            encounter_text=self._task.encounter_text,216            allowed_codes=list(self._task.allowed_codes),217            revealed_hints=list(self._task.hints[: self._state.hints_used]),218            current_primary_code=self._state.current_primary_code,219            current_secondary_codes=list(self._state.current_secondary_codes),220            current_needs_review=self._state.current_needs_review,221            attempts_remaining=max(0, self.MAX_STEPS - self._state.step_count),222            progress_score=self._state.best_score,223            grader_feedback=feedback,224            reward_breakdown=breakdown,225            done=done,226            reward=reward,227            metadata={228                "info": info,229                "grader": {"score": score},230            },231        )232 233    @property234    def state(self) -> MedicalCodingState:235        return self._state236 237    def get_metadata(self) -> EnvironmentMetadata:238        return EnvironmentMetadata(239            name="MedicalCodingEnvironment",240            description="Deterministic ICD-10 coding workflow benchmark with graded real-world tasks.",241            version="0.1.0",242        )243