Abhishek-CS221006's picture
Update env.py
f9875f7 verified
Raw
History Blame Contribute Delete
17.9 kB
"""Core RL environment for clinical trial patient screening."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Dict, List, Optional
from uuid import uuid4
from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import EnvironmentMetadata
from pydantic import BaseModel, Field
try:
from .models import (
ClinicalTrialAction,
ClinicalTrialObservation,
ClinicalTrialReward,
ClinicalTrialState,
)
except ImportError:
from models import (
ClinicalTrialAction,
ClinicalTrialObservation,
ClinicalTrialReward,
ClinicalTrialState,
)
INCREMENTAL_REWARD = 0.20
FINAL_REWARD = 1.00
HALLUCINATION_PENALTY = -0.50
TASK_SEQUENCE = ["easy", "medium", "hard"]
MIN_STRICT_SCORE = 0.01
MAX_STRICT_SCORE = 0.99
DEFAULT_GRADER_SCORE = 0.5
def _normalize(value: Optional[str]) -> str:
return " ".join((value or "").strip().lower().replace("_", " ").split())
class GroundTruth(BaseModel):
"""Deterministic grader targets for a scenario."""
extracted_fields: Dict[str, str] = Field(default_factory=dict)
ranking: List[str] = Field(default_factory=list)
final_decision: str
class ScenarioSpec(BaseModel):
"""Scenario loaded from patient_data.json."""
task_id: str
difficulty: str
title: str
instructions: str
context: Dict[str, Any]
ground_truth: GroundTruth
hidden_exclusions: List[str] = Field(default_factory=list)
max_steps: int = 6
grader_name: str = "deterministic_json_grader"
class ClinicalTrialEnvironment(
Environment[ClinicalTrialAction, ClinicalTrialObservation, ClinicalTrialState]
):
"""Clinical trial screening environment backed by externalized JSON scenarios."""
SUPPORTS_CONCURRENT_SESSIONS = True
def __init__(self) -> None:
super().__init__()
data_path = Path(__file__).resolve().with_name("patient_data.json")
payload = json.loads(data_path.read_text(encoding="utf-8"))
task_payload = payload.get("tasks", {})
self._scenarios: Dict[str, ScenarioSpec] = {
task_id: ScenarioSpec.model_validate({**scenario, "task_id": task_id})
for task_id, scenario in task_payload.items()
}
self._task_cursor = -1
self._current_scenario: Optional[ScenarioSpec] = None
self._submitted_ranking: List[str] = []
self._state = ClinicalTrialState(episode_id=str(uuid4()), step_count=0)
def reset(
self,
seed: Optional[int] = None,
episode_id: Optional[str] = None,
task_id: Optional[str] = None,
**kwargs: Any,
) -> ClinicalTrialObservation:
del seed, kwargs
selected_task_id = task_id or self._next_task_id()
if selected_task_id not in self._scenarios:
available = ", ".join(sorted(self._scenarios))
raise ValueError(f"Unknown task_id '{selected_task_id}'. Expected one of: {available}")
self._current_scenario = self._scenarios[selected_task_id]
self._submitted_ranking = []
self._state = ClinicalTrialState(
episode_id=episode_id or str(uuid4()),
step_count=0,
current_task_id=self._current_scenario.task_id,
difficulty=self._current_scenario.difficulty,
title=self._current_scenario.title,
extracted_fields={},
identified_deviations=[],
final_decision=None,
grading_score=DEFAULT_GRADER_SCORE,
)
return self._build_observation(
reward_details=ClinicalTrialReward(
notes=["Episode reset."], grader_score=DEFAULT_GRADER_SCORE
),
done=False,
)
def step(
self,
action: ClinicalTrialAction,
timeout_s: Optional[float] = None,
**kwargs: Any,
) -> ClinicalTrialObservation:
del timeout_s, kwargs
if self._current_scenario is None:
return self.reset()
self._state.step_count += 1
reward = ClinicalTrialReward()
done = False
terminal_reason: Optional[str] = None
if action.action_type == "extract_data":
self._handle_extraction(action, reward)
elif action.action_type == "flag_deviation":
self._handle_deviation_flag(action, reward)
elif action.action_type == "rank_patients":
ranking_accepted = self._handle_ranking(action, reward)
if ranking_accepted:
done = True
terminal_reason = "ranking_submitted"
elif action.action_type == "submit_decision":
self._state.final_decision = action.final_decision
done = True
terminal_reason = "final_decision_submitted"
elif action.action_type == "delete_evidence":
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append("Destructive action: deleting evidence is not allowed.")
else:
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append(f"Unsupported action type: {action.action_type}")
if self._state.step_count >= self._current_scenario.max_steps and not done:
done = True
terminal_reason = "max_steps_reached"
reward.grader_score = self._grade_for_current_task()
self._state.grading_score = reward.grader_score
if done:
if self._is_final_submission_correct():
reward.final_reward = FINAL_REWARD
reward.notes.append("Correct final screening decision.")
reward.missing_items = self._missing_items()
reward.total_reward = round(
reward.incremental_reward + reward.final_reward + reward.penalty, 4
)
return self._build_observation(
reward_details=reward,
done=done,
terminal_reason=terminal_reason,
)
@property
def state(self) -> ClinicalTrialState:
return self._state
def get_metadata(self) -> EnvironmentMetadata:
return EnvironmentMetadata(
name="clinical_trial_env",
description=(
"Clinical trial patient screening environment with 3 deterministic tasks "
"and explicit graders for easy, medium, and hard difficulty."
),
version="1.0.0",
author="Abhishek-CS221006",
)
def grader(self) -> float:
"""Deterministically compare agent outputs against the current scenario ground truth."""
if self._current_scenario is None:
return DEFAULT_GRADER_SCORE
return self._score_scenario(
self._current_scenario,
self._state,
self._submitted_ranking,
)
def _score_scenario(
self,
scenario: ScenarioSpec,
state: ClinicalTrialState,
submitted_ranking: List[str],
) -> float:
"""Deterministically compare agent outputs against the provided scenario state."""
components: List[float] = []
truth = scenario.ground_truth
if truth.extracted_fields:
field_hits = sum(
1
for field_name, expected in truth.extracted_fields.items()
if _normalize(state.extracted_fields.get(field_name)) == _normalize(expected)
)
score = field_hits / len(truth.extracted_fields)
score = min(max(score, MIN_STRICT_SCORE), MAX_STRICT_SCORE)
components.append(score)
if scenario.hidden_exclusions:
exclusion_hits = sum(
1
for exclusion in scenario.hidden_exclusions
if exclusion in state.identified_deviations
)
score = exclusion_hits / len(scenario.hidden_exclusions)
score = min(max(score, MIN_STRICT_SCORE), MAX_STRICT_SCORE)
components.append(score)
if truth.ranking:
ranking = submitted_ranking
if ranking and len(ranking) == len(truth.ranking):
positional_hits = sum(
1 for actual, expected in zip(ranking, truth.ranking) if actual == expected
) / len(truth.ranking)
pairwise_hits = 0
total_pairs = 0
for index, higher in enumerate(truth.ranking):
for lower in truth.ranking[index + 1 :]:
total_pairs += 1
if ranking.index(higher) < ranking.index(lower):
pairwise_hits += 1
pairwise_score = pairwise_hits / max(total_pairs, 1)
score = (0.6 * positional_hits) + (0.4 * pairwise_score)
else:
score = DEFAULT_GRADER_SCORE
score = min(max(score, MIN_STRICT_SCORE), MAX_STRICT_SCORE)
components.append(score)
if state.final_decision is None:
score = DEFAULT_GRADER_SCORE
else:
final_match = _normalize(state.final_decision) == _normalize(truth.final_decision)
score = MAX_STRICT_SCORE if final_match else MIN_STRICT_SCORE
components.append(score)
if not components:
return MIN_STRICT_SCORE
raw_score = sum(components) / len(components)
strict_score = min(max(raw_score, MIN_STRICT_SCORE), MAX_STRICT_SCORE)
return round(strict_score, 4)
def grade_easy_screening(self) -> float:
"""Task-specific grader for the easy screening task."""
return self._grade_task_by_id("easy")
def grade_medium_ranking(self) -> float:
"""Task-specific grader for the medium ranking task."""
return self._grade_task_by_id("medium")
def grade_hard_exclusions(self) -> float:
"""Task-specific grader for the hard exclusions task."""
return self._grade_task_by_id("hard")
def _grade_task_by_id(self, task_id: str) -> float:
"""Return a task grader score without mutating the environment state."""
scenario = self._scenarios[task_id]
if self._current_scenario is not None and self._current_scenario.task_id == task_id:
state = self._state
ranking = self._submitted_ranking
else:
state = ClinicalTrialState(
current_task_id=scenario.task_id,
difficulty=scenario.difficulty,
title=scenario.title,
extracted_fields={},
identified_deviations=[],
final_decision=None,
grading_score=DEFAULT_GRADER_SCORE,
)
ranking = []
return self._score_scenario(scenario, state, ranking)
def _grade_for_current_task(self) -> float:
"""Resolve and run the grader declared by the current scenario."""
assert self._current_scenario is not None
grader_name = (self._current_scenario.grader_name or "").strip()
grader_fn = getattr(self, grader_name, None)
if callable(grader_fn):
score = float(grader_fn())
else:
score = float(self.grader())
return round(min(max(score, MIN_STRICT_SCORE), MAX_STRICT_SCORE), 4)
def _next_task_id(self) -> str:
self._task_cursor = (self._task_cursor + 1) % len(TASK_SEQUENCE)
return TASK_SEQUENCE[self._task_cursor]
def _handle_extraction(self, action: ClinicalTrialAction, reward: ClinicalTrialReward) -> None:
assert self._current_scenario is not None
if not action.field_name or action.value is None:
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append("extract_data requires field_name and value.")
return
expected_value = self._current_scenario.ground_truth.extracted_fields.get(action.field_name)
if expected_value is None:
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append(f"Hallucinated field: {action.field_name}")
return
if _normalize(action.value) == _normalize(expected_value):
if action.field_name not in self._state.extracted_fields:
reward.incremental_reward += INCREMENTAL_REWARD
reward.matched_items.append(action.field_name)
reward.notes.append(f"Validated extraction for {action.field_name}.")
self._state.extracted_fields[action.field_name] = action.value
else:
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append(f"Incorrect value for {action.field_name}.")
def _handle_deviation_flag(self, action: ClinicalTrialAction, reward: ClinicalTrialReward) -> None:
assert self._current_scenario is not None
if self._current_scenario.task_id != "hard":
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append("Deviation flagging is only valid for the hard task.")
return
submitted = [_normalize(item) for item in action.deviations]
if not submitted:
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append("flag_deviation requires at least one deviation.")
return
for deviation in submitted:
if deviation in self._current_scenario.hidden_exclusions:
if deviation not in self._state.identified_deviations:
self._state.identified_deviations.append(deviation)
reward.incremental_reward += INCREMENTAL_REWARD
reward.matched_items.append(deviation)
reward.notes.append(f"Validated deviation: {deviation}.")
else:
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append(f"Unsupported deviation claim: {deviation}.")
def _handle_ranking(self, action: ClinicalTrialAction, reward: ClinicalTrialReward) -> bool:
assert self._current_scenario is not None
if self._current_scenario.task_id != "medium":
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append("Ranking is only valid for the medium task.")
return False
ranking = action.ranking
valid_patients = [
patient["patient_id"]
for patient in self._current_scenario.context.get("patients", [])
]
if sorted(ranking) != sorted(valid_patients):
reward.penalty += HALLUCINATION_PENALTY
reward.notes.append("Ranking must include each patient exactly once.")
return False
self._submitted_ranking = ranking
self._state.final_decision = "ranking_submitted"
return True
def _is_final_submission_correct(self) -> bool:
assert self._current_scenario is not None
truth = self._current_scenario.ground_truth
if truth.ranking:
return self._submitted_ranking == truth.ranking
return _normalize(self._state.final_decision) == _normalize(truth.final_decision)
def _missing_items(self) -> List[str]:
assert self._current_scenario is not None
truth = self._current_scenario.ground_truth
missing_fields = [
field_name
for field_name, expected in truth.extracted_fields.items()
if _normalize(self._state.extracted_fields.get(field_name)) != _normalize(expected)
]
missing_fields.extend(
exclusion
for exclusion in self._current_scenario.hidden_exclusions
if exclusion not in self._state.identified_deviations
)
if truth.ranking and self._submitted_ranking != truth.ranking:
missing_fields.append("ranking")
if _normalize(self._state.final_decision) != _normalize(truth.final_decision):
missing_fields.append("final_decision")
return missing_fields
def _build_observation(
self,
reward_details: ClinicalTrialReward,
done: bool,
terminal_reason: Optional[str] = None,
) -> ClinicalTrialObservation:
assert self._current_scenario is not None
attempts_remaining = max(self._current_scenario.max_steps - self._state.step_count, 0)
return ClinicalTrialObservation(
task_id=self._current_scenario.task_id,
difficulty=self._current_scenario.difficulty, # type: ignore[arg-type]
title=self._current_scenario.title,
instructions=self._current_scenario.instructions,
context=self._current_scenario.context,
expected_fields=list(self._current_scenario.ground_truth.extracted_fields.keys()),
extracted_fields=dict(self._state.extracted_fields),
identified_deviations=list(self._state.identified_deviations),
attempts_remaining=attempts_remaining,
grader_name=self._current_scenario.grader_name,
reward_details=reward_details,
reward=reward_details.total_reward,
done=done,
metadata={"grading_score": self._state.grading_score},
terminal_reason=terminal_reason,
)
class ClinicalTrialEnv(ClinicalTrialEnvironment):
"""Compatibility alias for manifest entry points expecting env:ClinicalTrialEnv."""
pass
def grade_easy_screening() -> float:
"""Module-level easy task grader for validator discovery."""
return ClinicalTrialEnv().grade_easy_screening()
def grade_medium_ranking() -> float:
"""Module-level medium task grader for validator discovery."""
return ClinicalTrialEnv().grade_medium_ranking()
def grade_hard_exclusions() -> float:
"""Module-level hard task grader for validator discovery."""
return ClinicalTrialEnv().grade_hard_exclusions()