Spaces:
Sleeping
Sleeping
Upload 2 files
Browse files
env.py
CHANGED
|
@@ -100,10 +100,10 @@ class ClinicalTrialEnvironment(
|
|
| 100 |
extracted_fields={},
|
| 101 |
identified_deviations=[],
|
| 102 |
final_decision=None,
|
| 103 |
-
grading_score=0.
|
| 104 |
)
|
| 105 |
return self._build_observation(
|
| 106 |
-
reward_details=ClinicalTrialReward(notes=["Episode reset."]),
|
| 107 |
done=False,
|
| 108 |
)
|
| 109 |
|
|
@@ -146,13 +146,14 @@ class ClinicalTrialEnvironment(
|
|
| 146 |
done = True
|
| 147 |
terminal_reason = "max_steps_reached"
|
| 148 |
|
|
|
|
|
|
|
|
|
|
| 149 |
if done:
|
| 150 |
-
reward.grader_score = self.grader()
|
| 151 |
if self._is_final_submission_correct():
|
| 152 |
reward.final_reward = FINAL_REWARD
|
| 153 |
reward.notes.append("Correct final screening decision.")
|
| 154 |
reward.missing_items = self._missing_items()
|
| 155 |
-
self._state.grading_score = reward.grader_score
|
| 156 |
|
| 157 |
reward.total_reward = round(
|
| 158 |
reward.incremental_reward + reward.final_reward + reward.penalty, 4
|
|
@@ -344,4 +345,4 @@ class ClinicalTrialEnvironment(
|
|
| 344 |
class ClinicalTrialEnv(ClinicalTrialEnvironment):
|
| 345 |
"""Compatibility alias for manifest entry points expecting env:ClinicalTrialEnv."""
|
| 346 |
|
| 347 |
-
pass
|
|
|
|
| 100 |
extracted_fields={},
|
| 101 |
identified_deviations=[],
|
| 102 |
final_decision=None,
|
| 103 |
+
grading_score=0.5,
|
| 104 |
)
|
| 105 |
return self._build_observation(
|
| 106 |
+
reward_details=ClinicalTrialReward(notes=["Episode reset."], grader_score=0.5),
|
| 107 |
done=False,
|
| 108 |
)
|
| 109 |
|
|
|
|
| 146 |
done = True
|
| 147 |
terminal_reason = "max_steps_reached"
|
| 148 |
|
| 149 |
+
reward.grader_score = self.grader()
|
| 150 |
+
self._state.grading_score = reward.grader_score
|
| 151 |
+
|
| 152 |
if done:
|
|
|
|
| 153 |
if self._is_final_submission_correct():
|
| 154 |
reward.final_reward = FINAL_REWARD
|
| 155 |
reward.notes.append("Correct final screening decision.")
|
| 156 |
reward.missing_items = self._missing_items()
|
|
|
|
| 157 |
|
| 158 |
reward.total_reward = round(
|
| 159 |
reward.incremental_reward + reward.final_reward + reward.penalty, 4
|
|
|
|
| 345 |
class ClinicalTrialEnv(ClinicalTrialEnvironment):
|
| 346 |
"""Compatibility alias for manifest entry points expecting env:ClinicalTrialEnv."""
|
| 347 |
|
| 348 |
+
pass
|
models.py
CHANGED
|
@@ -21,7 +21,12 @@ class ClinicalTrialReward(BaseModel):
|
|
| 21 |
final_reward: float = Field(default=0.0, description="Terminal reward for a correct final decision.")
|
| 22 |
penalty: float = Field(default=0.0, description="Penalty for hallucinations or destructive actions.")
|
| 23 |
total_reward: float = Field(default=0.0, description="Net reward for the step.")
|
| 24 |
-
grader_score: float = Field(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 25 |
matched_items: List[str] = Field(default_factory=list, description="Correctly matched items this step.")
|
| 26 |
missing_items: List[str] = Field(default_factory=list, description="Expected items still missing at grading time.")
|
| 27 |
notes: List[str] = Field(default_factory=list, description="Human-readable reward rationale.")
|
|
@@ -78,7 +83,12 @@ class ClinicalTrialState(State):
|
|
| 78 |
extracted_fields: Dict[str, str] = Field(default_factory=dict, description="Accepted extracted fields.")
|
| 79 |
identified_deviations: List[str] = Field(default_factory=list, description="Accepted deviations.")
|
| 80 |
final_decision: Optional[str] = Field(default=None, description="Submitted terminal decision.")
|
| 81 |
-
grading_score: float = Field(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 82 |
|
| 83 |
def __call__(self) -> "ClinicalTrialState":
|
| 84 |
"""Support env.state() as a compatibility alias for the OpenEnv state property."""
|
|
|
|
| 21 |
final_reward: float = Field(default=0.0, description="Terminal reward for a correct final decision.")
|
| 22 |
penalty: float = Field(default=0.0, description="Penalty for hallucinations or destructive actions.")
|
| 23 |
total_reward: float = Field(default=0.0, description="Net reward for the step.")
|
| 24 |
+
grader_score: float = Field(
|
| 25 |
+
default=0.5,
|
| 26 |
+
gt=0.0,
|
| 27 |
+
lt=1.0,
|
| 28 |
+
description="Deterministic task score in the strict open interval (0, 1).",
|
| 29 |
+
)
|
| 30 |
matched_items: List[str] = Field(default_factory=list, description="Correctly matched items this step.")
|
| 31 |
missing_items: List[str] = Field(default_factory=list, description="Expected items still missing at grading time.")
|
| 32 |
notes: List[str] = Field(default_factory=list, description="Human-readable reward rationale.")
|
|
|
|
| 83 |
extracted_fields: Dict[str, str] = Field(default_factory=dict, description="Accepted extracted fields.")
|
| 84 |
identified_deviations: List[str] = Field(default_factory=list, description="Accepted deviations.")
|
| 85 |
final_decision: Optional[str] = Field(default=None, description="Submitted terminal decision.")
|
| 86 |
+
grading_score: float = Field(
|
| 87 |
+
default=0.5,
|
| 88 |
+
gt=0.0,
|
| 89 |
+
lt=1.0,
|
| 90 |
+
description="Latest grader output in the strict open interval (0, 1).",
|
| 91 |
+
)
|
| 92 |
|
| 93 |
def __call__(self) -> "ClinicalTrialState":
|
| 94 |
"""Support env.state() as a compatibility alias for the OpenEnv state property."""
|