RhutuTuvoc commited on
Commit
b847937
·
1 Parent(s): 68ac963

Use deterministic task graders for all three tasks

Browse files
Files changed (5) hide show
  1. graders.py +266 -24
  2. it_mental_health_environment.py +119 -251
  3. openenv.yaml +5 -2
  4. server/app.py +4 -0
  5. tasks.py +6 -2
graders.py CHANGED
@@ -1,40 +1,273 @@
1
  """
2
- Static grader registry for OpenEnv task discoverability.
3
 
4
- The environment computes the real rubric score in `it_mental_health_environment.py`.
5
- These helpers expose one named grader per task so submission validators can
6
- statically discover at least three task-grader pairs.
 
 
 
 
7
  """
8
 
9
- from typing import Any, Dict
 
 
 
 
 
 
10
 
11
 
12
  def _normalize_reward(reward: float) -> float:
13
- return min(max(float(reward), 0.0), 1.0)
14
 
15
 
16
- def _task_matches(state: Dict[str, Any], expected_task_id: str) -> bool:
17
- if not isinstance(state, dict):
18
- return False
19
 
20
- candidates = [
21
- state.get("task_id"),
22
- state.get("current_task"),
23
- (state.get("metadata") or {}).get("task_id") if isinstance(state.get("metadata"), dict) else None,
24
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
25
  return expected_task_id in candidates
26
 
27
 
28
- def grade_burnout_detection(state: Dict[str, Any], reward: float) -> float:
29
- return _normalize_reward(reward if _task_matches(state, "burnout_detection") else 0.0)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30
 
 
 
 
31
 
32
- def grade_stress_triage(state: Dict[str, Any], reward: float) -> float:
33
- return _normalize_reward(reward if _task_matches(state, "stress_triage") else 0.0)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
34
 
35
 
36
- def grade_intervention_plan(state: Dict[str, Any], reward: float) -> float:
37
- return _normalize_reward(reward if _task_matches(state, "intervention_plan") else 0.0)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
38
 
39
 
40
  GRADERS = {
@@ -44,17 +277,26 @@ GRADERS = {
44
  }
45
 
46
 
47
- TASK_GRADER_PAIRS = [
48
  ("burnout_detection", "burnout_detection_grader"),
49
  ("stress_triage", "stress_triage_grader"),
50
  ("intervention_plan", "intervention_plan_grader"),
51
  ]
52
 
53
 
 
 
 
 
 
 
 
54
  __all__ = [
55
- "grade_burnout_detection",
56
- "grade_stress_triage",
57
- "grade_intervention_plan",
58
  "GRADERS",
 
59
  "TASK_GRADER_PAIRS",
 
 
 
 
60
  ]
 
1
  """
2
+ Deterministic agent graders for the IT Mental Health OpenEnv tasks.
3
 
4
+ Each task has:
5
+ - a concrete objective
6
+ - a deterministic programmatic grader
7
+ - a normalized score in the range [0.0, 1.0]
8
+
9
+ These graders are intentionally rule-based so submission validators can
10
+ statically discover them and evaluation remains reproducible.
11
  """
12
 
13
+ from __future__ import annotations
14
+
15
+ from typing import Any, Dict, Iterable, List, Tuple
16
+
17
+
18
+ TaskState = Dict[str, Any]
19
+ GradeResult = Tuple[float, Dict[str, float], str]
20
 
21
 
22
  def _normalize_reward(reward: float) -> float:
23
+ return round(min(max(float(reward), 0.0), 1.0), 3)
24
 
25
 
26
+ def _safe_lower(value: Any) -> str:
27
+ return str(value or "").lower()
 
28
 
29
+
30
+ def _contains_any(text: str, terms: Iterable[str]) -> bool:
31
+ return any(term in text for term in terms)
32
+
33
+
34
+ def _count_contains(text: str, terms: Iterable[str]) -> int:
35
+ return sum(1 for term in terms if term in text)
36
+
37
+
38
+ def _metadata_task_id(state: TaskState) -> str:
39
+ metadata = state.get("metadata")
40
+ if isinstance(metadata, dict):
41
+ return str(metadata.get("task_id", ""))
42
+ return ""
43
+
44
+
45
+ def _task_matches(state: TaskState, expected_task_id: str) -> bool:
46
+ candidates = {
47
+ str(state.get("task_id", "")),
48
+ str(state.get("current_task", "")),
49
+ _metadata_task_id(state),
50
+ }
51
  return expected_task_id in candidates
52
 
53
 
54
+ def _breakdown_to_reward(
55
+ breakdown: Dict[str, float], weights: Dict[str, float]
56
+ ) -> float:
57
+ total_weight = sum(weights.values()) or 1.0
58
+ weighted = sum(breakdown.get(key, 0.0) * weight for key, weight in weights.items())
59
+ return _normalize_reward(weighted / total_weight)
60
+
61
+
62
+ def _burnout_grade(response: str, ground_truth: Dict[str, Any]) -> GradeResult:
63
+ text = _safe_lower(response)
64
+ dims = ground_truth.get("active_dimensions", [])
65
+ dim_tokens = {
66
+ "Exhaustion": ["exhaustion", "exhausted", "fatigue", "tiredness"],
67
+ "Depersonalization": ["depersonalization", "depersonalisation", "cynicism", "cynical", "detached"],
68
+ "Reduced Personal Accomplishment": [
69
+ "reduced personal accomplishment",
70
+ "reduced accomplishment",
71
+ "inefficacy",
72
+ "ineffective",
73
+ "personal accomplishment",
74
+ ],
75
+ }
76
+
77
+ dim_hits = 0
78
+ for dimension in dims:
79
+ if _contains_any(text, dim_tokens.get(dimension, [dimension.lower()])):
80
+ dim_hits += 1
81
+ dimensions_identified = dim_hits / max(len(dims), 1)
82
+
83
+ severity_order = ["low", "moderate", "high", "critical"]
84
+ expected_severity = _safe_lower(ground_truth.get("severity"))
85
+ mentioned = [level for level in severity_order if level in text]
86
+ if expected_severity in mentioned:
87
+ severity_accuracy = 1.0
88
+ elif mentioned and expected_severity in severity_order:
89
+ gt_idx = severity_order.index(expected_severity)
90
+ min_distance = min(abs(severity_order.index(level) - gt_idx) for level in mentioned)
91
+ severity_accuracy = 0.5 if min_distance == 1 else 0.0
92
+ else:
93
+ severity_accuracy = 0.0
94
+
95
+ red_flag_terms = [
96
+ str(ground_truth.get("hours", "")),
97
+ str(ground_truth.get("vacation_months", "")),
98
+ "week",
99
+ "hours",
100
+ "vacation",
101
+ "leave",
102
+ "headache",
103
+ "migraine",
104
+ "sleep",
105
+ "desk",
106
+ "detached",
107
+ "cynical",
108
+ "incomplete",
109
+ "productivity",
110
+ ]
111
+ red_flags_quality = min(1.0, _count_contains(text, red_flag_terms) / 3.0)
112
+
113
+ escalation_expected = bool(ground_truth.get("escalation_needed"))
114
+ escalation_positive = _contains_any(text, ["yes", "escalat", "immediate hr", "urgent hr"])
115
+ escalation_reasoning = 1.0 if escalation_positive == escalation_expected else 0.0
116
+
117
+ structure_clarity = 1.0 if _contains_any(response, ["1.", "2.", "3.", "4.", "##", "**"]) else 0.4
118
+
119
+ breakdown = {
120
+ "dimensions_identified": round(dimensions_identified, 3),
121
+ "severity_accuracy": round(severity_accuracy, 3),
122
+ "red_flags_quality": round(red_flags_quality, 3),
123
+ "escalation_reasoning": round(escalation_reasoning, 3),
124
+ "structure_clarity": round(structure_clarity, 3),
125
+ }
126
+ weights = {
127
+ "dimensions_identified": 0.30,
128
+ "severity_accuracy": 0.20,
129
+ "red_flags_quality": 0.20,
130
+ "escalation_reasoning": 0.20,
131
+ "structure_clarity": 0.10,
132
+ }
133
+ reward = _breakdown_to_reward(breakdown, weights)
134
+ feedback = (
135
+ f"[Deterministic Grader] burnout_detection score={reward:.2f}. "
136
+ f"Expected severity={ground_truth.get('severity')}; "
137
+ f"dimension hits={dim_hits}/{max(len(dims), 1)}."
138
+ )
139
+ return reward, breakdown, feedback
140
+
141
+
142
+ def _stress_triage_grade(response: str, ground_truth: Dict[str, Any]) -> GradeResult:
143
+ text = _safe_lower(response)
144
+ tiers = ground_truth.get("correct_tiers", {})
145
+ names = ground_truth.get("names", [])
146
+ priority_order = [str(name).lower() for name in ground_truth.get("priority_order", [])]
147
+
148
+ tier_hits = 0
149
+ for name, tier in tiers.items():
150
+ name_lower = name.lower()
151
+ tier_lower = tier.lower()
152
+ if name_lower in text and tier_lower in text:
153
+ tier_hits += 1
154
+ tier_accuracy = tier_hits / max(len(tiers), 1)
155
+
156
+ top_two_correct = sum(1 for name in priority_order[:2] if name and name in text[:800])
157
+ priority_ranking = 1.0 if top_two_correct == 2 else 0.5 if top_two_correct == 1 else 0.0
158
 
159
+ immediate_actions = 1.0 if _contains_any(text, ["24 hour", "24-hour", "within 24", "immediate", "today"]) else 0.3
160
+ medium_term_support = 1.0 if _contains_any(text, ["2 week", "2-week", "within 2 weeks", "fortnight", "two weeks"]) else 0.3
161
+ completeness = 1.0 if all(name.lower() in text for name in names) else 0.0
162
 
163
+ breakdown = {
164
+ "tier_accuracy": round(tier_accuracy, 3),
165
+ "priority_ranking": round(priority_ranking, 3),
166
+ "immediate_actions": round(immediate_actions, 3),
167
+ "medium_term_support": round(medium_term_support, 3),
168
+ "completeness": round(completeness, 3),
169
+ }
170
+ weights = {
171
+ "tier_accuracy": 0.35,
172
+ "priority_ranking": 0.20,
173
+ "immediate_actions": 0.15,
174
+ "medium_term_support": 0.15,
175
+ "completeness": 0.15,
176
+ }
177
+ reward = _breakdown_to_reward(breakdown, weights)
178
+ feedback = (
179
+ f"[Deterministic Grader] stress_triage score={reward:.2f}. "
180
+ f"Correct tier assignments={tier_hits}/{max(len(tiers), 1)}."
181
+ )
182
+ return reward, breakdown, feedback
183
 
184
 
185
+ def _intervention_plan_grade(response: str, ground_truth: Dict[str, Any]) -> GradeResult:
186
+ text = _safe_lower(response)
187
+
188
+ week_hits = sum(1 for week in ["week 1", "week 2", "week 3", "week 4"] if week in text)
189
+ four_week_structure = week_hits / 4.0
190
+
191
+ action_terms = [
192
+ "on-call",
193
+ "overtime",
194
+ "survey",
195
+ "1:1",
196
+ "1-1",
197
+ "check-in",
198
+ "training",
199
+ "leave",
200
+ "workload",
201
+ "rotation",
202
+ "policy",
203
+ ]
204
+ action_concreteness = min(1.0, _count_contains(text, action_terms) / 6.0)
205
+ responsibility = 1.0 if _contains_any(text, ["hr", "manager", "eap"]) else 0.0
206
+
207
+ kpi_terms = ["kpi", "metric", "measure", "indicator", "tracking", "90-day"]
208
+ kpis_quality = 1.0 if _count_contains(text, kpi_terms) >= 2 else 0.5 if _count_contains(text, kpi_terms) == 1 else 0.0
209
+
210
+ risk_and_budget = 1.0 if ("risk" in text and _contains_any(text, ["budget", "low", "medium", "high", "$"])) else 0.0
211
+
212
+ proportionality_terms = [
213
+ str(ground_truth.get("team_size", "")),
214
+ str(ground_truth.get("affected", "")),
215
+ str(ground_truth.get("overtime", "")),
216
+ str(ground_truth.get("oncall_days", "")),
217
+ str(ground_truth.get("hr_complaints", "")),
218
+ ]
219
+ proportionality = 1.0 if _count_contains(text, [term for term in proportionality_terms if term]) >= 2 else 0.5
220
+
221
+ breakdown = {
222
+ "four_week_structure": round(four_week_structure, 3),
223
+ "action_concreteness": round(action_concreteness, 3),
224
+ "responsibility": round(responsibility, 3),
225
+ "kpis_quality": round(kpis_quality, 3),
226
+ "risk_and_budget": round(risk_and_budget, 3),
227
+ "proportionality": round(proportionality, 3),
228
+ }
229
+ weights = {
230
+ "four_week_structure": 0.30,
231
+ "action_concreteness": 0.20,
232
+ "responsibility": 0.10,
233
+ "kpis_quality": 0.15,
234
+ "risk_and_budget": 0.10,
235
+ "proportionality": 0.15,
236
+ }
237
+ reward = _breakdown_to_reward(breakdown, weights)
238
+ feedback = (
239
+ f"[Deterministic Grader] intervention_plan score={reward:.2f}. "
240
+ f"Week coverage={week_hits}/4."
241
+ )
242
+ return reward, breakdown, feedback
243
+
244
+
245
+ def grade_response(task_id: str, response: str, ground_truth: Dict[str, Any]) -> GradeResult:
246
+ if task_id == "burnout_detection":
247
+ return _burnout_grade(response, ground_truth)
248
+ if task_id == "stress_triage":
249
+ return _stress_triage_grade(response, ground_truth)
250
+ if task_id == "intervention_plan":
251
+ return _intervention_plan_grade(response, ground_truth)
252
+ return 0.0, {}, f"[Deterministic Grader] unknown task_id={task_id}"
253
+
254
+
255
+ def grade_burnout_detection(state: TaskState, reward: float) -> float:
256
+ if not _task_matches(state, "burnout_detection"):
257
+ return 0.0
258
+ return _normalize_reward(reward)
259
+
260
+
261
+ def grade_stress_triage(state: TaskState, reward: float) -> float:
262
+ if not _task_matches(state, "stress_triage"):
263
+ return 0.0
264
+ return _normalize_reward(reward)
265
+
266
+
267
+ def grade_intervention_plan(state: TaskState, reward: float) -> float:
268
+ if not _task_matches(state, "intervention_plan"):
269
+ return 0.0
270
+ return _normalize_reward(reward)
271
 
272
 
273
  GRADERS = {
 
277
  }
278
 
279
 
280
+ TASK_GRADER_PAIRS: List[Tuple[str, str]] = [
281
  ("burnout_detection", "burnout_detection_grader"),
282
  ("stress_triage", "stress_triage_grader"),
283
  ("intervention_plan", "intervention_plan_grader"),
284
  ]
285
 
286
 
287
+ TASK_GRADER_OBJECTIVES = {
288
+ "burnout_detection": "Identify burnout dimensions, severity, red flags, and escalation need from an employee profile.",
289
+ "stress_triage": "Triage three employees by urgency and recommend immediate plus medium-term support.",
290
+ "intervention_plan": "Produce a four-week intervention plan with owners, KPIs, risk, and budget.",
291
+ }
292
+
293
+
294
  __all__ = [
 
 
 
295
  "GRADERS",
296
+ "TASK_GRADER_OBJECTIVES",
297
  "TASK_GRADER_PAIRS",
298
+ "grade_burnout_detection",
299
+ "grade_intervention_plan",
300
+ "grade_response",
301
+ "grade_stress_triage",
302
  ]
it_mental_health_environment.py CHANGED
@@ -1,26 +1,29 @@
1
  """
2
- IT Mental Health OpenEnv - Environment Logic (v2 - improved)
3
 
4
- FIX 1: LLM-as-judge grader replaces keyword matching.
5
- FIX 2: Randomised episode generation — every reset() is unique.
 
 
6
  """
7
 
8
- import os
9
- import uuid
10
  import random
11
- import json
12
- from typing import Optional
13
- from openai import OpenAI
14
 
15
  try:
16
- from models import MentalHealthAction, MentalHealthObservation, MentalHealthReward, MentalHealthState
17
  except ImportError:
18
  try:
19
- from server.models import MentalHealthAction, MentalHealthObservation, MentalHealthReward, MentalHealthState
20
  except ImportError:
 
21
  import sys
 
22
  sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
23
- from models import MentalHealthAction, MentalHealthObservation, MentalHealthReward, MentalHealthState
 
24
 
25
  TASK_ORDER = ["burnout_detection", "stress_triage", "intervention_plan"]
26
  TASK_DIFFICULTY = {
@@ -29,27 +32,19 @@ TASK_DIFFICULTY = {
29
  "intervention_plan": "hard",
30
  }
31
 
32
- # ── LLM Judge client ──────────────────────────────────────────────────────────
33
- _llm_client: Optional[OpenAI] = None
34
-
35
- def _get_llm_client() -> Optional[OpenAI]:
36
- global _llm_client
37
- use_llm_judge = os.environ.get("USE_LLM_JUDGE", "").strip().lower() in {"1", "true", "yes"}
38
- if not use_llm_judge:
39
- return None
40
- if _llm_client is None:
41
- api_key = os.environ.get("HF_TOKEN", "") or os.environ.get("OPENAI_API_KEY", "")
42
- base_url = os.environ.get("API_BASE_URL", "https://api-inference.huggingface.co/v1")
43
- if api_key:
44
- _llm_client = OpenAI(api_key=api_key, base_url=base_url)
45
- return _llm_client
46
 
47
-
48
- # ── Randomised scenario data pools ───────────────────────────────────────────
49
- NAMES = ["Alex","Jordan","Sam","Riley","Morgan","Casey","Taylor","Drew","Jamie","Avery"]
50
- ROLES = ["Software Engineer","DevOps Engineer","QA Engineer","Data Scientist",
51
- "Backend Developer","Frontend Developer","ML Engineer","Platform Engineer"]
52
- YEARS_EXP = [1,2,3,4,5,7,8,10]
 
 
 
 
 
 
53
 
54
  EXHAUSTION_SYMPTOMS = [
55
  "sleeping only 4-5 hours/night",
@@ -78,21 +73,21 @@ REDUCED_ACCOMPLISHMENT_SYMPTOMS = [
78
  ]
79
 
80
  SEVERITY_MAP = {
81
- "Low": {"hours_range":(42,50), "vacation_months":(1,4), "dimensions":1},
82
- "Moderate": {"hours_range":(50,60), "vacation_months":(4,8), "dimensions":2},
83
- "High": {"hours_range":(60,68), "vacation_months":(8,14), "dimensions":2},
84
- "Critical": {"hours_range":(68,80), "vacation_months":(14,24),"dimensions":3},
85
  }
86
 
87
  STRESS_TIER_TEMPLATES = {
88
- "GREEN": ["Feeling slightly overwhelmed with the sprint but I think I'll manage."],
89
- "AMBER": ["Haven't slept great — around 6 hours most nights. Starting to affect focus."],
90
- "RED": ["I snapped at a teammate yesterday. I'm scared of how I'm feeling lately."],
91
  "CRITICAL": ["I've been having chest tightness every time I get a pager alert. Three weeks. Haven't told anyone."],
92
  }
93
 
94
 
95
- def generate_burnout_scenario(rng):
96
  name = rng.choice(NAMES)
97
  role = rng.choice(ROLES)
98
  exp = rng.choice(YEARS_EXP)
@@ -101,7 +96,7 @@ def generate_burnout_scenario(rng):
101
  vacation = rng.randint(*cfg["vacation_months"])
102
  n_dims = cfg["dimensions"]
103
 
104
- all_dims = ["exhaustion","depersonalization","reduced_accomplishment"]
105
  active_dims = rng.sample(all_dims, n_dims)
106
  symptoms = []
107
  ground_truth_dims = []
@@ -115,9 +110,9 @@ def generate_burnout_scenario(rng):
115
  symptoms += rng.sample(REDUCED_ACCOMPLISHMENT_SYMPTOMS, 2)
116
  ground_truth_dims.append("Reduced Personal Accomplishment")
117
  rng.shuffle(symptoms)
118
- sym_text = "\n".join(f"- {s}" for s in symptoms)
119
 
120
- scenario = f"""[TASK: Burnout Detection — EASY]
121
 
122
  You are an occupational psychologist reviewing an IT employee profile.
123
 
@@ -128,74 +123,81 @@ Employee Profile:
128
  - Observed symptoms:
129
  {sym_text}
130
 
131
- Your task:
132
- 1. Identify which MBI dimensions are present: Exhaustion / Depersonalization / Reduced Personal Accomplishment
133
  2. Rate severity: Low / Moderate / High / Critical
134
- 3. List TOP 3 red-flag signals from the profile
135
- 4. State whether immediate HR escalation is needed (Yes/No) and why
136
 
137
  Respond with clear headings."""
138
 
139
- gt = {"active_dimensions":ground_truth_dims, "severity":severity_label,
140
- "escalation_needed":severity_label in ("High","Critical"),
141
- "name":name, "hours":hours, "vacation_months":vacation}
142
- return scenario, gt
 
 
 
 
 
143
 
144
 
145
- def generate_stress_triage_scenario(rng):
146
- tiers = ["CRITICAL","RED",rng.choice(["GREEN","AMBER"])]
147
  rng.shuffle(tiers)
148
  selected_names = rng.sample(NAMES, len(tiers))
149
  employees = []
150
  for idx, tier in enumerate(tiers):
151
- employees.append({
152
- "name": selected_names[idx],
153
- "role": rng.choice(ROLES),
154
- "tier": tier,
155
- "quote": rng.choice(STRESS_TIER_TEMPLATES[tier]),
156
- })
157
- tier_order = {"CRITICAL":0,"RED":1,"AMBER":2,"GREEN":3}
158
- priority_order = sorted(range(3), key=lambda i: tier_order[employees[i]["tier"]])
159
 
 
 
160
  cases = "".join(
161
- f'{i+1}. {e["name"]} ({e["role"]}): "{e["quote"]}"\n'
162
- for i,e in enumerate(employees)
163
  )
164
- scenario = f"""[TASK: Stress Triage — MEDIUM]
165
 
166
  You are reviewing urgent mental health flags from an IT pulse survey.
167
  Triage these 3 employees and prioritise them.
168
 
169
  Cases:
170
  {cases}
171
- For each case:
172
  a) Assign: GREEN / AMBER / RED / CRITICAL
173
- b) Primary stressor: Workload / Physiological / Relationship / Cognitive
174
- c) ONE immediate action (within 24 hours)
175
- d) ONE medium-term support (within 2 weeks)
176
 
177
  Rank the 3 cases by intervention priority (1 = most urgent)."""
178
 
179
- gt = {
180
  "employees": employees,
181
- "correct_tiers": {e["name"]:e["tier"] for e in employees},
182
- "priority_order": [employees[i]["name"] for i in priority_order],
183
- "names": [e["name"] for e in employees],
184
  }
185
- return scenario, gt
186
 
187
 
188
- def generate_intervention_plan_scenario(rng):
189
- team_size = rng.choice([8,10,12,15,18,20])
190
- affected = int(team_size * rng.choice([0.4,0.5,0.6,0.7]))
191
- overtime = rng.choice([10,15,18,20,25])
192
- oncall = rng.choice([2,3,4,5])
193
- no_tb = rng.choice([6,9,12,15,18,24])
194
- no_1on1 = rng.choice([3,4,5,6,8])
195
- on_leave = rng.randint(1,3)
196
- hr_c = rng.randint(1,4)
197
 
198
- scenario = f"""[TASK: Intervention Plan — HARD]
199
 
200
  You are a workplace mental health consultant. Design a 4-week intervention for a
201
  software team with systemic burnout ({affected} of {team_size} members affected).
@@ -205,167 +207,32 @@ Team Context:
205
  - On-call: 1 person every {oncall} days (24/7)
206
  - No team-building in {no_tb} months
207
  - No manager 1:1s in {no_1on1} months
208
- - {hr_c} formal HR complaints about workload
209
  - {on_leave} members on anxiety-related medical leave
210
 
211
- Design a 4-week plan:
212
- Week 1: Immediate stabilisation
213
- Week 2: Assessment & listening
214
- Week 3: Process reforms (on-call, overtime, workload)
215
- Week 4: Sustainable systems
216
-
217
- For each week: 2+ concrete actions, responsible party (HR/Manager/EAP), measurable outcome.
218
-
219
- Also include:
220
- - 3 measurable KPIs for 90-day tracking
221
- - 1 key risk if plan is NOT executed
222
- - Budget: Low (<$500) / Medium ($500-$5000) / High (>$5000)"""
223
-
224
- gt = {"team_size":team_size,"affected":affected,"overtime":overtime,
225
- "oncall_days":oncall,"on_leave":on_leave,"hr_complaints":hr_c}
226
- return scenario, gt
227
-
228
-
229
- # ── Clinical rubrics ──────────────────────────────────────────────────────────
230
- RUBRICS = {
231
- "burnout_detection": {
232
- "dimensions_identified": {"max":3.0,"desc":"1pt per correct MBI dimension identified (0-3)"},
233
- "severity_accuracy": {"max":2.0,"desc":"2=exact match, 1=adjacent, 0=wrong"},
234
- "red_flags_quality": {"max":2.0,"desc":"Flags are specific and drawn from the profile (0-2)"},
235
- "escalation_reasoning": {"max":2.0,"desc":"Correct Yes/No + sound clinical reasoning (0-2)"},
236
- "structure_clarity": {"max":1.0,"desc":"Clear headings, professional format (0-1)"},
237
- },
238
- "stress_triage": {
239
- "tier_accuracy": {"max":3.0,"desc":"1pt per correct tier assigned (0-3)"},
240
- "priority_ranking": {"max":2.0,"desc":"2=correct, 1=partially correct, 0=wrong"},
241
- "immediate_actions": {"max":2.0,"desc":"Specific 24h actions proportionate to tier (0-2)"},
242
- "medium_term_support": {"max":2.0,"desc":"Realistic 2-week support for each case (0-2)"},
243
- "completeness": {"max":1.0,"desc":"All 3 employees addressed with all 4 elements (0-1)"},
244
- },
245
- "intervention_plan": {
246
- "four_week_structure": {"max":3.0,"desc":"All 4 weeks with distinct focus (0-3)"},
247
- "action_concreteness": {"max":2.0,"desc":"Actions specific, not vague (0-2)"},
248
- "responsibility": {"max":1.0,"desc":"Responsible party per action (0-1)"},
249
- "kpis_quality": {"max":2.0,"desc":"3 measurable, relevant KPIs (0-2)"},
250
- "risk_and_budget": {"max":1.0,"desc":"Risk realistic + budget category present (0-1)"},
251
- "proportionality": {"max":1.0,"desc":"Plan addresses this team's specific numbers (0-1)"},
252
- },
253
- }
254
-
255
- JUDGE_SYSTEM = """You are a clinical rubric evaluator for an AI mental health RL environment.
256
- Score the agent's response against the rubric. Award marks based on reasoning quality and accuracy — NOT keyword presence alone.
257
- Return ONLY valid JSON, no markdown fences:
258
- {"scores": {"dim_name": float, ...}, "feedback": "2-3 sentence critique"}"""
259
-
260
-
261
- def _llm_judge_grade(task_id, scenario, response, ground_truth):
262
- rubric = RUBRICS[task_id]
263
- total_max = sum(v["max"] for v in rubric.values())
264
- rubric_text = "\n".join(f' "{k}": max={v["max"]} — {v["desc"]}' for k,v in rubric.items())
265
-
266
- prompt = f"""Task: {task_id}
267
- Scenario: {scenario}
268
- Ground truth: {json.dumps(ground_truth)}
269
- Agent response: {response}
270
- Rubric:
271
- {rubric_text}
272
- Score now. JSON only."""
273
 
274
- client = _get_llm_client()
275
- if client is None:
276
- return _heuristic_grade(task_id, response, ground_truth)
277
 
278
- model = os.environ.get("MODEL_NAME","meta-llama/Llama-3.1-8B-Instruct")
279
- try:
280
- resp = client.chat.completions.create(
281
- model=model,
282
- messages=[{"role":"system","content":JUDGE_SYSTEM},
283
- {"role":"user","content":prompt}],
284
- max_tokens=400, temperature=0.0,
285
- )
286
- raw = resp.choices[0].message.content.strip().replace("```json","").replace("```","").strip()
287
- parsed = json.loads(raw)
288
- scores = parsed.get("scores",{})
289
- feedback = parsed.get("feedback","")
290
- breakdown, total = {}, 0.0
291
- for dim, cfg in rubric.items():
292
- s = max(0.0, min(cfg["max"], float(scores.get(dim, 0.0))))
293
- breakdown[dim] = round(s/cfg["max"], 3)
294
- total += s
295
- reward = round(max(0.0, min(1.0, total/total_max)), 3)
296
- return reward, breakdown, f"[LLM Judge] {feedback}"
297
- except Exception as e:
298
- r,b,f = _heuristic_grade(task_id, response, ground_truth)
299
- return r, b, f"[Fallback — judge error: {e}] {f}"
300
-
301
-
302
- def _heuristic_grade(task_id, response, ground_truth):
303
- """Ground-truth-aware heuristic fallback (better than pure keyword matching)."""
304
- r = response.lower()
305
- breakdown, score = {}, 0.0
306
-
307
- if task_id == "burnout_detection":
308
- dims = ground_truth.get("active_dimensions",[])
309
- # Match on a distinctive token per dimension, not just the first word.
310
- dim_keys = {
311
- "Exhaustion": ["exhaustion","exhausted"],
312
- "Depersonalization": ["depersonalization","depersonalisation","cynicism","cynical"],
313
- "Reduced Personal Accomplishment": ["reduced personal accomplishment","reduced accomplishment","inefficacy","personal accomplishment"],
314
- }
315
- dim_hits = sum(1 for d in dims if any(k in r for k in dim_keys.get(d, [d.lower()])))
316
- breakdown["dimensions_identified"] = dim_hits/max(len(dims),1)
317
- # Severity: exact match = 1.0, adjacent tier = 0.5, otherwise 0.
318
- sev_order = ["low","moderate","high","critical"]
319
- sev = ground_truth.get("severity","High").lower()
320
- mentioned = [s for s in sev_order if s in r]
321
- if sev in mentioned:
322
- breakdown["severity_accuracy"] = 1.0
323
- elif mentioned and sev in sev_order:
324
- gt_idx = sev_order.index(sev)
325
- best = min(abs(sev_order.index(m) - gt_idx) for m in mentioned)
326
- breakdown["severity_accuracy"] = 0.5 if best == 1 else 0.0
327
- else:
328
- breakdown["severity_accuracy"] = 0.0
329
- breakdown["red_flags_quality"] = min(1.0, len(response)/300)
330
- esc = ground_truth.get("escalation_needed",True)
331
- esc_hit = "yes" in r or "escalat" in r or "immediate" in r
332
- breakdown["escalation_reasoning"] = 1.0 if esc==esc_hit else 0.3
333
- breakdown["structure_clarity"] = 1.0 if any(h in response for h in ["##","**","1.","1:"]) else 0.3
334
-
335
- elif task_id == "stress_triage":
336
- tiers = ground_truth.get("correct_tiers",{})
337
- tier_hits = sum(1 for nm,t in tiers.items() if nm.lower() in r and t.lower() in r)
338
- breakdown["tier_accuracy"] = tier_hits/max(len(tiers),1)
339
- priority = ground_truth.get("priority_order",[])
340
- breakdown["priority_ranking"] = 1.0 if (priority and priority[0].lower() in r[:500]) else 0.4
341
- breakdown["immediate_actions"] = min(1.0, r.count("24")*0.3 + r.count("immediat")*0.4)
342
- breakdown["medium_term_support"] = 1.0 if ("week" in r or "fortnight" in r) else 0.3
343
- names = ground_truth.get("names",[])
344
- breakdown["completeness"] = 1.0 if all(n.lower() in r for n in names) else 0.4
345
-
346
- elif task_id == "intervention_plan":
347
- weeks = sum(1 for w in ["week 1","week 2","week 3","week 4"] if w in r)
348
- breakdown["four_week_structure"] = weeks/4.0
349
- breakdown["action_concreteness"] = min(1.0, len(response)/600)
350
- breakdown["responsibility"] = 1.0 if any(x in r for x in ["hr","manager","eap"]) else 0.0
351
- kpi_hits = sum(1 for k in ["kpi","metric","measure","indicator"] if k in r)
352
- breakdown["kpis_quality"] = min(1.0, kpi_hits*0.4)
353
- breakdown["risk_and_budget"] = 1.0 if ("risk" in r and any(b in r for b in ["low","medium","high","$"])) else 0.3
354
- gt_size = str(ground_truth.get("team_size",""))
355
- breakdown["proportionality"] = 1.0 if gt_size in response else 0.5
356
-
357
- score = sum(breakdown.values())/max(len(breakdown),1)
358
- reward = round(max(0.0, min(1.0, score)), 3)
359
- return reward, breakdown, f"[Heuristic] {task_id} = {reward:.2f}. Breakdown: {breakdown}"
360
-
361
-
362
- # ── Environment class ─────────────────────────────────────────────────────────
363
  class ITMentalHealthEnvironment:
364
- """
365
- OpenEnv RL environment — IT sector mental health.
366
- Every reset() is unique (randomised profiles/scenarios).
367
- Graded by LLM-as-judge against clinical rubric.
368
- """
369
  def __init__(self):
370
  self._state = MentalHealthState()
371
  self._task_index = 0
@@ -375,7 +242,7 @@ class ITMentalHealthEnvironment:
375
  def _generate_all(self):
376
  self._current_scenarios = {
377
  "burnout_detection": generate_burnout_scenario(self._rng),
378
- "stress_triage": generate_stress_triage_scenario(self._rng),
379
  "intervention_plan": generate_intervention_plan_scenario(self._rng),
380
  }
381
 
@@ -383,9 +250,7 @@ class ITMentalHealthEnvironment:
383
  difficulty = TASK_DIFFICULTY[task]
384
  task_step_count = self._state.task_step_counts.get(task, 0) + 1
385
  task_cumulative_score = round(self._state.task_scores.get(task, 0.0) + reward, 3)
386
- difficulty_cumulative_score = round(
387
- self._state.difficulty_scores.get(difficulty, 0.0) + reward, 3
388
- )
389
  return {
390
  "difficulty": difficulty,
391
  "step_score": round(reward, 3),
@@ -393,6 +258,7 @@ class ITMentalHealthEnvironment:
393
  "task_cumulative_score": task_cumulative_score,
394
  "difficulty_cumulative_score": difficulty_cumulative_score,
395
  "overall_cumulative_score": round(self._state.cumulative_reward + reward, 3),
 
396
  }
397
 
398
  def reset(self, seed=None):
@@ -401,8 +267,11 @@ class ITMentalHealthEnvironment:
401
  self._task_index = 0
402
  self._generate_all()
403
  self._state = MentalHealthState(
404
- episode_id=str(uuid.uuid4()), step_count=0,
405
- current_task=TASK_ORDER[0], cumulative_reward=0.0, tasks_completed=[],
 
 
 
406
  task_scores={task_id: 0.0 for task_id in TASK_ORDER},
407
  task_step_counts={task_id: 0 for task_id in TASK_ORDER},
408
  difficulty_scores={"easy": 0.0, "medium": 0.0, "hard": 0.0},
@@ -410,7 +279,8 @@ class ITMentalHealthEnvironment:
410
  task = TASK_ORDER[0]
411
  scenario_text, _ = self._current_scenarios[task]
412
  return MentalHealthObservation(
413
- scenario=scenario_text, feedback=f"New episode (seed={actual_seed}).",
 
414
  task_id=task,
415
  metadata={
416
  "seed": actual_seed,
@@ -420,6 +290,7 @@ class ITMentalHealthEnvironment:
420
  "task_cumulative_score": 0.0,
421
  "difficulty_cumulative_score": 0.0,
422
  "overall_cumulative_score": 0.0,
 
423
  },
424
  )
425
 
@@ -427,7 +298,6 @@ class ITMentalHealthEnvironment:
427
  if not self._state.episode_id or not self._current_scenarios:
428
  raise ValueError("Episode not initialized. Call /reset before /step.")
429
 
430
- # Guard: refuse to grade once the episode is finished.
431
  if self._task_index >= len(TASK_ORDER):
432
  observation = MentalHealthObservation(
433
  scenario="Episode already finished. Call /reset to start a new one.",
@@ -441,12 +311,10 @@ class ITMentalHealthEnvironment:
441
  task = self._state.current_task
442
  if action.task_id != task:
443
  raise ValueError(f"task_id mismatch: expected '{task}', got '{action.task_id}'.")
444
- scenario_text, ground_truth = self._current_scenarios.get(task, ("", {}))
445
 
446
- reward, breakdown, feedback = _llm_judge_grade(
447
- task_id=task, scenario=scenario_text,
448
- response=action.response, ground_truth=ground_truth,
449
- )
450
  score_metadata = self._score_metadata(task, reward)
451
  self._state.cumulative_reward += reward
452
  self._state.task_scores[task] = score_metadata["task_cumulative_score"]
 
1
  """
2
+ IT Mental Health OpenEnv - Environment Logic.
3
 
4
+ This environment defines three tasks with deterministic programmatic graders:
5
+ - burnout_detection (easy)
6
+ - stress_triage (medium)
7
+ - intervention_plan (hard)
8
  """
9
 
 
 
10
  import random
11
+ import uuid
12
+
13
+ from graders import grade_response
14
 
15
  try:
16
+ from models import MentalHealthObservation, MentalHealthReward, MentalHealthState
17
  except ImportError:
18
  try:
19
+ from server.models import MentalHealthObservation, MentalHealthReward, MentalHealthState
20
  except ImportError:
21
+ import os
22
  import sys
23
+
24
  sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
25
+ from models import MentalHealthObservation, MentalHealthReward, MentalHealthState
26
+
27
 
28
  TASK_ORDER = ["burnout_detection", "stress_triage", "intervention_plan"]
29
  TASK_DIFFICULTY = {
 
32
  "intervention_plan": "hard",
33
  }
34
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
35
 
36
+ NAMES = ["Alex", "Jordan", "Sam", "Riley", "Morgan", "Casey", "Taylor", "Drew", "Jamie", "Avery"]
37
+ ROLES = [
38
+ "Software Engineer",
39
+ "DevOps Engineer",
40
+ "QA Engineer",
41
+ "Data Scientist",
42
+ "Backend Developer",
43
+ "Frontend Developer",
44
+ "ML Engineer",
45
+ "Platform Engineer",
46
+ ]
47
+ YEARS_EXP = [1, 2, 3, 4, 5, 7, 8, 10]
48
 
49
  EXHAUSTION_SYMPTOMS = [
50
  "sleeping only 4-5 hours/night",
 
73
  ]
74
 
75
  SEVERITY_MAP = {
76
+ "Low": {"hours_range": (42, 50), "vacation_months": (1, 4), "dimensions": 1},
77
+ "Moderate": {"hours_range": (50, 60), "vacation_months": (4, 8), "dimensions": 2},
78
+ "High": {"hours_range": (60, 68), "vacation_months": (8, 14), "dimensions": 2},
79
+ "Critical": {"hours_range": (68, 80), "vacation_months": (14, 24), "dimensions": 3},
80
  }
81
 
82
  STRESS_TIER_TEMPLATES = {
83
+ "GREEN": ["Feeling slightly overwhelmed with the sprint but I think I'll manage."],
84
+ "AMBER": ["Haven't slept great - around 6 hours most nights. Starting to affect focus."],
85
+ "RED": ["I snapped at a teammate yesterday. I'm scared of how I'm feeling lately."],
86
  "CRITICAL": ["I've been having chest tightness every time I get a pager alert. Three weeks. Haven't told anyone."],
87
  }
88
 
89
 
90
+ def generate_burnout_scenario(rng: random.Random):
91
  name = rng.choice(NAMES)
92
  role = rng.choice(ROLES)
93
  exp = rng.choice(YEARS_EXP)
 
96
  vacation = rng.randint(*cfg["vacation_months"])
97
  n_dims = cfg["dimensions"]
98
 
99
+ all_dims = ["exhaustion", "depersonalization", "reduced_accomplishment"]
100
  active_dims = rng.sample(all_dims, n_dims)
101
  symptoms = []
102
  ground_truth_dims = []
 
110
  symptoms += rng.sample(REDUCED_ACCOMPLISHMENT_SYMPTOMS, 2)
111
  ground_truth_dims.append("Reduced Personal Accomplishment")
112
  rng.shuffle(symptoms)
113
+ sym_text = "\n".join(f"- {item}" for item in symptoms)
114
 
115
+ scenario = f"""[TASK: Burnout Detection - EASY]
116
 
117
  You are an occupational psychologist reviewing an IT employee profile.
118
 
 
123
  - Observed symptoms:
124
  {sym_text}
125
 
126
+ Objective:
127
+ 1. Identify which MBI dimensions are present
128
  2. Rate severity: Low / Moderate / High / Critical
129
+ 3. List the top 3 red-flag signals
130
+ 4. State whether immediate HR escalation is needed and why
131
 
132
  Respond with clear headings."""
133
 
134
+ ground_truth = {
135
+ "active_dimensions": ground_truth_dims,
136
+ "severity": severity_label,
137
+ "escalation_needed": severity_label in ("High", "Critical"),
138
+ "name": name,
139
+ "hours": hours,
140
+ "vacation_months": vacation,
141
+ }
142
+ return scenario, ground_truth
143
 
144
 
145
+ def generate_stress_triage_scenario(rng: random.Random):
146
+ tiers = ["CRITICAL", "RED", rng.choice(["GREEN", "AMBER"])]
147
  rng.shuffle(tiers)
148
  selected_names = rng.sample(NAMES, len(tiers))
149
  employees = []
150
  for idx, tier in enumerate(tiers):
151
+ employees.append(
152
+ {
153
+ "name": selected_names[idx],
154
+ "role": rng.choice(ROLES),
155
+ "tier": tier,
156
+ "quote": rng.choice(STRESS_TIER_TEMPLATES[tier]),
157
+ }
158
+ )
159
 
160
+ tier_order = {"CRITICAL": 0, "RED": 1, "AMBER": 2, "GREEN": 3}
161
+ priority_order = sorted(range(3), key=lambda index: tier_order[employees[index]["tier"]])
162
  cases = "".join(
163
+ f'{index + 1}. {employee["name"]} ({employee["role"]}): "{employee["quote"]}"\n'
164
+ for index, employee in enumerate(employees)
165
  )
166
+ scenario = f"""[TASK: Stress Triage - MEDIUM]
167
 
168
  You are reviewing urgent mental health flags from an IT pulse survey.
169
  Triage these 3 employees and prioritise them.
170
 
171
  Cases:
172
  {cases}
173
+ Objective:
174
  a) Assign: GREEN / AMBER / RED / CRITICAL
175
+ b) Name the primary stressor
176
+ c) Give one immediate action within 24 hours
177
+ d) Give one medium-term support action within 2 weeks
178
 
179
  Rank the 3 cases by intervention priority (1 = most urgent)."""
180
 
181
+ ground_truth = {
182
  "employees": employees,
183
+ "correct_tiers": {employee["name"]: employee["tier"] for employee in employees},
184
+ "priority_order": [employees[index]["name"] for index in priority_order],
185
+ "names": [employee["name"] for employee in employees],
186
  }
187
+ return scenario, ground_truth
188
 
189
 
190
+ def generate_intervention_plan_scenario(rng: random.Random):
191
+ team_size = rng.choice([8, 10, 12, 15, 18, 20])
192
+ affected = int(team_size * rng.choice([0.4, 0.5, 0.6, 0.7]))
193
+ overtime = rng.choice([10, 15, 18, 20, 25])
194
+ oncall = rng.choice([2, 3, 4, 5])
195
+ no_tb = rng.choice([6, 9, 12, 15, 18, 24])
196
+ no_1on1 = rng.choice([3, 4, 5, 6, 8])
197
+ on_leave = rng.randint(1, 3)
198
+ hr_complaints = rng.randint(1, 4)
199
 
200
+ scenario = f"""[TASK: Intervention Plan - HARD]
201
 
202
  You are a workplace mental health consultant. Design a 4-week intervention for a
203
  software team with systemic burnout ({affected} of {team_size} members affected).
 
207
  - On-call: 1 person every {oncall} days (24/7)
208
  - No team-building in {no_tb} months
209
  - No manager 1:1s in {no_1on1} months
210
+ - {hr_complaints} formal HR complaints about workload
211
  - {on_leave} members on anxiety-related medical leave
212
 
213
+ Objective:
214
+ - Week 1: Immediate stabilisation
215
+ - Week 2: Assessment and listening
216
+ - Week 3: Process reforms
217
+ - Week 4: Sustainable systems
218
+
219
+ For each week include 2+ actions, responsible owner, measurable outcome.
220
+ Also include 3 KPIs, 1 key risk, and a budget category."""
221
+
222
+ ground_truth = {
223
+ "team_size": team_size,
224
+ "affected": affected,
225
+ "overtime": overtime,
226
+ "oncall_days": oncall,
227
+ "on_leave": on_leave,
228
+ "hr_complaints": hr_complaints,
229
+ }
230
+ return scenario, ground_truth
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
231
 
 
 
 
232
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
233
  class ITMentalHealthEnvironment:
234
+ """OpenEnv-compatible environment with deterministic task graders."""
235
+
 
 
 
236
  def __init__(self):
237
  self._state = MentalHealthState()
238
  self._task_index = 0
 
242
  def _generate_all(self):
243
  self._current_scenarios = {
244
  "burnout_detection": generate_burnout_scenario(self._rng),
245
+ "stress_triage": generate_stress_triage_scenario(self._rng),
246
  "intervention_plan": generate_intervention_plan_scenario(self._rng),
247
  }
248
 
 
250
  difficulty = TASK_DIFFICULTY[task]
251
  task_step_count = self._state.task_step_counts.get(task, 0) + 1
252
  task_cumulative_score = round(self._state.task_scores.get(task, 0.0) + reward, 3)
253
+ difficulty_cumulative_score = round(self._state.difficulty_scores.get(difficulty, 0.0) + reward, 3)
 
 
254
  return {
255
  "difficulty": difficulty,
256
  "step_score": round(reward, 3),
 
258
  "task_cumulative_score": task_cumulative_score,
259
  "difficulty_cumulative_score": difficulty_cumulative_score,
260
  "overall_cumulative_score": round(self._state.cumulative_reward + reward, 3),
261
+ "task_id": task,
262
  }
263
 
264
  def reset(self, seed=None):
 
267
  self._task_index = 0
268
  self._generate_all()
269
  self._state = MentalHealthState(
270
+ episode_id=str(uuid.uuid4()),
271
+ step_count=0,
272
+ current_task=TASK_ORDER[0],
273
+ cumulative_reward=0.0,
274
+ tasks_completed=[],
275
  task_scores={task_id: 0.0 for task_id in TASK_ORDER},
276
  task_step_counts={task_id: 0 for task_id in TASK_ORDER},
277
  difficulty_scores={"easy": 0.0, "medium": 0.0, "hard": 0.0},
 
279
  task = TASK_ORDER[0]
280
  scenario_text, _ = self._current_scenarios[task]
281
  return MentalHealthObservation(
282
+ scenario=scenario_text,
283
+ feedback=f"New episode (seed={actual_seed}).",
284
  task_id=task,
285
  metadata={
286
  "seed": actual_seed,
 
290
  "task_cumulative_score": 0.0,
291
  "difficulty_cumulative_score": 0.0,
292
  "overall_cumulative_score": 0.0,
293
+ "task_id": task,
294
  },
295
  )
296
 
 
298
  if not self._state.episode_id or not self._current_scenarios:
299
  raise ValueError("Episode not initialized. Call /reset before /step.")
300
 
 
301
  if self._task_index >= len(TASK_ORDER):
302
  observation = MentalHealthObservation(
303
  scenario="Episode already finished. Call /reset to start a new one.",
 
311
  task = self._state.current_task
312
  if action.task_id != task:
313
  raise ValueError(f"task_id mismatch: expected '{task}', got '{action.task_id}'.")
 
314
 
315
+ _, ground_truth = self._current_scenarios.get(task, ("", {}))
316
+ reward, breakdown, feedback = grade_response(task, action.response, ground_truth)
317
+
 
318
  score_metadata = self._score_metadata(task, reward)
319
  self._state.cumulative_reward += reward
320
  self._state.task_scores[task] = score_metadata["task_cumulative_score"]
openenv.yaml CHANGED
@@ -6,8 +6,8 @@ description: >
6
  A reinforcement learning environment for mental health assessment and
7
  intervention planning in IT and software engineering teams. Agents solve
8
  three progressively harder tasks: burnout detection, stress triage, and
9
- team-level intervention planning. Each task returns a normalized reward in
10
- the 0.0-1.0 range based on rubric coverage, structure, and clinical accuracy.
11
 
12
  type: typed
13
 
@@ -25,6 +25,7 @@ tasks:
25
  - id: burnout_detection
26
  difficulty: easy
27
  description: Identify Maslach burnout dimensions and severity from an employee profile
 
28
  reward_range: [0.0, 1.0]
29
  grader: burnout_detection_grader
30
  graders:
@@ -33,6 +34,7 @@ tasks:
33
  - id: stress_triage
34
  difficulty: medium
35
  description: Triage 3 IT employees by stress tier and recommend immediate actions
 
36
  reward_range: [0.0, 1.0]
37
  grader: stress_triage_grader
38
  graders:
@@ -41,6 +43,7 @@ tasks:
41
  - id: intervention_plan
42
  difficulty: hard
43
  description: Design a 4-week mental health intervention plan for a burned-out IT team
 
44
  reward_range: [0.0, 1.0]
45
  grader: intervention_plan_grader
46
  graders:
 
6
  A reinforcement learning environment for mental health assessment and
7
  intervention planning in IT and software engineering teams. Agents solve
8
  three progressively harder tasks: burnout detection, stress triage, and
9
+ team-level intervention planning. Each task has a deterministic programmatic
10
+ grader and returns a normalized reward in the 0.0-1.0 range.
11
 
12
  type: typed
13
 
 
25
  - id: burnout_detection
26
  difficulty: easy
27
  description: Identify Maslach burnout dimensions and severity from an employee profile
28
+ objective: Identify burnout dimensions, severity, top red flags, and HR escalation need from a single employee profile
29
  reward_range: [0.0, 1.0]
30
  grader: burnout_detection_grader
31
  graders:
 
34
  - id: stress_triage
35
  difficulty: medium
36
  description: Triage 3 IT employees by stress tier and recommend immediate actions
37
+ objective: Classify three employee cases by urgency, rank intervention priority, and recommend immediate plus 2-week support
38
  reward_range: [0.0, 1.0]
39
  grader: stress_triage_grader
40
  graders:
 
43
  - id: intervention_plan
44
  difficulty: hard
45
  description: Design a 4-week mental health intervention plan for a burned-out IT team
46
+ objective: Create a four-week intervention plan with owners, KPIs, risk, and budget for a burned-out software team
47
  reward_range: [0.0, 1.0]
48
  grader: intervention_plan_grader
49
  graders:
server/app.py CHANGED
@@ -66,16 +66,19 @@ TASK_METADATA = {
66
  "burnout_detection": {
67
  "difficulty": "easy",
68
  "description": "Identify Maslach burnout dimensions, severity, red flags, and whether HR escalation is needed.",
 
69
  "graders": ["burnout_detection_grader"],
70
  },
71
  "stress_triage": {
72
  "difficulty": "medium",
73
  "description": "Triage three IT employees by stress tier and recommend immediate and medium-term support.",
 
74
  "graders": ["stress_triage_grader"],
75
  },
76
  "intervention_plan": {
77
  "difficulty": "hard",
78
  "description": "Design a four-week intervention plan for a software team facing systemic burnout.",
 
79
  "graders": ["intervention_plan_grader"],
80
  },
81
  }
@@ -258,6 +261,7 @@ def list_tasks():
258
  "id": tid,
259
  "difficulty": TASK_METADATA[tid]["difficulty"],
260
  "description": TASK_METADATA[tid]["description"],
 
261
  "grader": TASK_METADATA[tid]["graders"][0],
262
  "graders": TASK_METADATA[tid]["graders"],
263
  "scoring": {
 
66
  "burnout_detection": {
67
  "difficulty": "easy",
68
  "description": "Identify Maslach burnout dimensions, severity, red flags, and whether HR escalation is needed.",
69
+ "objective": "Assess one employee profile and determine burnout dimensions, severity, red flags, and escalation need.",
70
  "graders": ["burnout_detection_grader"],
71
  },
72
  "stress_triage": {
73
  "difficulty": "medium",
74
  "description": "Triage three IT employees by stress tier and recommend immediate and medium-term support.",
75
+ "objective": "Classify three employees by urgency, rank intervention priority, and recommend immediate plus 2-week support.",
76
  "graders": ["stress_triage_grader"],
77
  },
78
  "intervention_plan": {
79
  "difficulty": "hard",
80
  "description": "Design a four-week intervention plan for a software team facing systemic burnout.",
81
+ "objective": "Produce a four-week intervention plan with owners, measurable outcomes, KPIs, risk, and budget.",
82
  "graders": ["intervention_plan_grader"],
83
  },
84
  }
 
261
  "id": tid,
262
  "difficulty": TASK_METADATA[tid]["difficulty"],
263
  "description": TASK_METADATA[tid]["description"],
264
+ "objective": TASK_METADATA[tid]["objective"],
265
  "grader": TASK_METADATA[tid]["graders"][0],
266
  "graders": TASK_METADATA[tid]["graders"],
267
  "scoring": {
tasks.py CHANGED
@@ -5,23 +5,27 @@ Static task registry for OpenEnv task discoverability.
5
  TASKS = [
6
  {
7
  "id": "burnout_detection",
 
 
8
  "grader": "burnout_detection_grader",
9
  "graders": ["burnout_detection_grader"],
10
  },
11
  {
12
  "id": "stress_triage",
 
 
13
  "grader": "stress_triage_grader",
14
  "graders": ["stress_triage_grader"],
15
  },
16
  {
17
  "id": "intervention_plan",
 
 
18
  "grader": "intervention_plan_grader",
19
  "graders": ["intervention_plan_grader"],
20
  },
21
  ]
22
 
23
-
24
  TASK_IDS = [task["id"] for task in TASKS]
25
 
26
-
27
  __all__ = ["TASKS", "TASK_IDS"]
 
5
  TASKS = [
6
  {
7
  "id": "burnout_detection",
8
+ "difficulty": "easy",
9
+ "objective": "Identify burnout dimensions, severity, top red flags, and whether immediate HR escalation is needed.",
10
  "grader": "burnout_detection_grader",
11
  "graders": ["burnout_detection_grader"],
12
  },
13
  {
14
  "id": "stress_triage",
15
+ "difficulty": "medium",
16
+ "objective": "Classify three employee cases by stress tier, rank urgency, and recommend immediate plus 2-week support.",
17
  "grader": "stress_triage_grader",
18
  "graders": ["stress_triage_grader"],
19
  },
20
  {
21
  "id": "intervention_plan",
22
+ "difficulty": "hard",
23
+ "objective": "Create a four-week intervention plan with owners, KPIs, risk, and budget for a burned-out team.",
24
  "grader": "intervention_plan_grader",
25
  "graders": ["intervention_plan_grader"],
26
  },
27
  ]
28
 
 
29
  TASK_IDS = [task["id"] for task in TASKS]
30
 
 
31
  __all__ = ["TASKS", "TASK_IDS"]