sankar-raul/ICD-10-code-predictor-env
0
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 