Abhishek-CS221006 commited on
Commit
f2d8053
·
verified ·
1 Parent(s): 28cd431

Upload 2 files

Browse files
Files changed (2) hide show
  1. env.py +6 -5
  2. models.py +12 -2
env.py CHANGED
@@ -100,10 +100,10 @@ class ClinicalTrialEnvironment(
100
  extracted_fields={},
101
  identified_deviations=[],
102
  final_decision=None,
103
- grading_score=0.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(default=0.0, ge=0.0, le=1.0, description="Deterministic task score.")
 
 
 
 
 
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(default=0.0, ge=0.0, le=1.0, description="Latest grader output.")
 
 
 
 
 
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."""