Spaces:
Runtime error
Runtime error
RhutuTuvoc commited on
Commit ·
b847937
1
Parent(s): 68ac963
Use deterministic task graders for all three tasks
Browse files- graders.py +266 -24
- it_mental_health_environment.py +119 -251
- openenv.yaml +5 -2
- server/app.py +4 -0
- tasks.py +6 -2
graders.py
CHANGED
|
@@ -1,40 +1,273 @@
|
|
| 1 |
"""
|
| 2 |
-
|
| 3 |
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
"""
|
| 8 |
|
| 9 |
-
from
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
|
| 11 |
|
| 12 |
def _normalize_reward(reward: float) -> float:
|
| 13 |
-
return min(max(float(reward), 0.0), 1.0)
|
| 14 |
|
| 15 |
|
| 16 |
-
def
|
| 17 |
-
|
| 18 |
-
return False
|
| 19 |
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 25 |
return expected_task_id in candidates
|
| 26 |
|
| 27 |
|
| 28 |
-
def
|
| 29 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
-
|
| 33 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
|
| 35 |
|
| 36 |
-
def
|
| 37 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 3 |
|
| 4 |
-
|
| 5 |
-
|
|
|
|
|
|
|
| 6 |
"""
|
| 7 |
|
| 8 |
-
import os
|
| 9 |
-
import uuid
|
| 10 |
import random
|
| 11 |
-
import
|
| 12 |
-
|
| 13 |
-
from
|
| 14 |
|
| 15 |
try:
|
| 16 |
-
from models import
|
| 17 |
except ImportError:
|
| 18 |
try:
|
| 19 |
-
from server.models import
|
| 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
|
|
|
|
| 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 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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":
|
| 82 |
-
"Moderate": {"hours_range":(50,60), "vacation_months":(4,8),
|
| 83 |
-
"High":
|
| 84 |
-
"Critical": {"hours_range":(68,80), "vacation_months":(14,24),"dimensions":3},
|
| 85 |
}
|
| 86 |
|
| 87 |
STRESS_TIER_TEMPLATES = {
|
| 88 |
-
"GREEN":
|
| 89 |
-
"AMBER":
|
| 90 |
-
"RED":
|
| 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"- {
|
| 119 |
|
| 120 |
-
scenario = f"""[TASK: Burnout Detection
|
| 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 |
-
|
| 132 |
-
1. Identify which MBI dimensions are present
|
| 133 |
2. Rate severity: Low / Moderate / High / Critical
|
| 134 |
-
3. List
|
| 135 |
-
4. State whether immediate HR escalation is needed
|
| 136 |
|
| 137 |
Respond with clear headings."""
|
| 138 |
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
|
| 142 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 157 |
-
|
| 158 |
-
|
| 159 |
|
|
|
|
|
|
|
| 160 |
cases = "".join(
|
| 161 |
-
f'{
|
| 162 |
-
for
|
| 163 |
)
|
| 164 |
-
scenario = f"""[TASK: Stress Triage
|
| 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 |
-
|
| 172 |
a) Assign: GREEN / AMBER / RED / CRITICAL
|
| 173 |
-
b)
|
| 174 |
-
c)
|
| 175 |
-
d)
|
| 176 |
|
| 177 |
Rank the 3 cases by intervention priority (1 = most urgent)."""
|
| 178 |
|
| 179 |
-
|
| 180 |
"employees": employees,
|
| 181 |
-
"correct_tiers": {
|
| 182 |
-
"priority_order": [employees[
|
| 183 |
-
"names": [
|
| 184 |
}
|
| 185 |
-
return scenario,
|
| 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 |
-
|
| 197 |
|
| 198 |
-
scenario = f"""[TASK: Intervention Plan
|
| 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 |
-
- {
|
| 209 |
- {on_leave} members on anxiety-related medical leave
|
| 210 |
|
| 211 |
-
|
| 212 |
-
Week 1: Immediate stabilisation
|
| 213 |
-
Week 2: Assessment
|
| 214 |
-
Week 3: Process reforms
|
| 215 |
-
Week 4: Sustainable systems
|
| 216 |
-
|
| 217 |
-
For each week
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
| 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 |
-
|
| 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":
|
| 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()),
|
| 405 |
-
|
|
|
|
|
|
|
|
|
|
| 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,
|
|
|
|
| 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 |
-
|
| 447 |
-
|
| 448 |
-
|
| 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
|
| 10 |
-
the 0.0-1.0 range
|
| 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"]
|