RhutuTuvoc commited on
Commit
566e5c7
·
1 Parent(s): a0b0950

add 3 tasks

Browse files
Files changed (4) hide show
  1. inference.py +21 -8
  2. it_mental_health_environment.py +59 -12
  3. models.py +17 -45
  4. server/app.py +34 -10
inference.py CHANGED
@@ -23,7 +23,7 @@ from openai import OpenAI
23
 
24
  API_BASE_URL = os.getenv("API_BASE_URL", "https://api-inference.huggingface.co/v1")
25
  MODEL_NAME = os.getenv("MODEL_NAME", "meta-llama/Llama-3.1-8B-Instruct")
26
- HF_TOKEN = os.getenv("HF_TOKEN")
27
  LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
28
  ENV_BASE_URL = os.getenv("ENV_BASE_URL", "http://localhost:7860")
29
 
@@ -70,12 +70,14 @@ def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> No
70
  )
71
 
72
 
73
- def call_env(endpoint: str, payload: Optional[dict] = None, method: str = "POST") -> dict:
 
 
74
  url = f"{ENV_BASE_URL}/{endpoint}"
75
  if method == "GET":
76
- response = requests.get(url, timeout=30)
77
  else:
78
- response = requests.post(url, json=payload or {}, timeout=30)
79
  response.raise_for_status()
80
  return response.json()
81
 
@@ -105,11 +107,13 @@ def main() -> None:
105
 
106
  try:
107
  if not HF_TOKEN:
108
- raise RuntimeError("HF_TOKEN must be set for inference.")
109
 
110
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
111
- observation = call_env("reset", {})
 
112
  current_task = observation.get("task_id", "unknown")
 
113
  log_start(task=current_task, env=BENCHMARK, model=MODEL_NAME)
114
  start_logged = True
115
 
@@ -118,22 +122,31 @@ def main() -> None:
118
  break
119
 
120
  current_task = observation.get("task_id", "unknown")
121
- action_text = get_model_response(client, observation["scenario"])
 
 
 
 
122
 
123
  error = None
124
  done = False
125
  reward = 0.0
126
 
127
  try:
 
 
 
128
  observation = call_env(
 
129
  "step",
130
  {
131
  "response": action_text,
132
  "task_id": current_task,
133
  "confidence": 0.85,
134
- "metadata": {},
135
  },
136
  )
 
137
  reward = float(observation.get("reward", 0.0) or 0.0)
138
  done = bool(observation.get("done", False))
139
  except Exception as exc:
 
23
 
24
  API_BASE_URL = os.getenv("API_BASE_URL", "https://api-inference.huggingface.co/v1")
25
  MODEL_NAME = os.getenv("MODEL_NAME", "meta-llama/Llama-3.1-8B-Instruct")
26
+ HF_TOKEN = os.getenv("HF_TOKEN") or os.getenv("OPENAI_API_KEY")
27
  LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
28
  ENV_BASE_URL = os.getenv("ENV_BASE_URL", "http://localhost:7860")
29
 
 
70
  )
71
 
72
 
73
+ def call_env(
74
+ http: requests.Session, endpoint: str, payload: Optional[dict] = None, method: str = "POST"
75
+ ) -> dict:
76
  url = f"{ENV_BASE_URL}/{endpoint}"
77
  if method == "GET":
78
+ response = http.get(url, timeout=30)
79
  else:
80
+ response = http.post(url, json=payload or {}, timeout=30)
81
  response.raise_for_status()
82
  return response.json()
83
 
 
107
 
108
  try:
109
  if not HF_TOKEN:
110
+ raise RuntimeError("HF_TOKEN or OPENAI_API_KEY must be set for inference.")
111
 
112
  client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)
113
+ http = requests.Session()
114
+ observation = call_env(http, "reset", {})
115
  current_task = observation.get("task_id", "unknown")
116
+ session_id = observation.get("metadata", {}).get("session_id")
117
  log_start(task=current_task, env=BENCHMARK, model=MODEL_NAME)
118
  start_logged = True
119
 
 
122
  break
123
 
124
  current_task = observation.get("task_id", "unknown")
125
+ scenario = observation.get("scenario")
126
+ if not scenario:
127
+ raise RuntimeError("Environment response missing 'scenario'.")
128
+
129
+ action_text = get_model_response(client, scenario)
130
 
131
  error = None
132
  done = False
133
  reward = 0.0
134
 
135
  try:
136
+ metadata = {}
137
+ if session_id:
138
+ metadata["session_id"] = session_id
139
  observation = call_env(
140
+ http,
141
  "step",
142
  {
143
  "response": action_text,
144
  "task_id": current_task,
145
  "confidence": 0.85,
146
+ "metadata": metadata,
147
  },
148
  )
149
+ session_id = observation.get("metadata", {}).get("session_id", session_id)
150
  reward = float(observation.get("reward", 0.0) or 0.0)
151
  done = bool(observation.get("done", False))
152
  except Exception as exc:
it_mental_health_environment.py CHANGED
@@ -13,24 +13,32 @@ from typing import Optional
13
  from openai import OpenAI
14
 
15
  try:
16
- from models import MentalHealthAction, MentalHealthObservation, MentalHealthState
17
  except ImportError:
18
  try:
19
- from server.models import MentalHealthAction, MentalHealthObservation, 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, MentalHealthState
24
 
25
  TASK_ORDER = ["burnout_detection", "stress_triage", "intervention_plan"]
 
 
 
 
 
26
 
27
  # ── LLM Judge client ──────────────────────────────────────────────────────────
28
  _llm_client: Optional[OpenAI] = None
29
 
30
  def _get_llm_client() -> Optional[OpenAI]:
31
  global _llm_client
 
 
 
32
  if _llm_client is None:
33
- api_key = os.environ.get("HF_TOKEN", "")
34
  base_url = os.environ.get("API_BASE_URL", "https://api-inference.huggingface.co/v1")
35
  if api_key:
36
  _llm_client = OpenAI(api_key=api_key, base_url=base_url)
@@ -371,6 +379,22 @@ class ITMentalHealthEnvironment:
371
  "intervention_plan": generate_intervention_plan_scenario(self._rng),
372
  }
373
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
374
  def reset(self, seed=None):
375
  actual_seed = seed if seed is not None else random.randint(0, 2**31)
376
  self._rng = random.Random(actual_seed)
@@ -379,13 +403,24 @@ class ITMentalHealthEnvironment:
379
  self._state = MentalHealthState(
380
  episode_id=str(uuid.uuid4()), step_count=0,
381
  current_task=TASK_ORDER[0], cumulative_reward=0.0, tasks_completed=[],
 
 
 
382
  )
383
  task = TASK_ORDER[0]
384
  scenario_text, _ = self._current_scenarios[task]
385
  return MentalHealthObservation(
386
  scenario=scenario_text, feedback=f"New episode (seed={actual_seed}).",
387
- reward=0.0, done=False, score_breakdown={}, task_id=task,
388
- metadata={"seed": actual_seed},
 
 
 
 
 
 
 
 
389
  )
390
 
391
  def step(self, action):
@@ -394,12 +429,13 @@ class ITMentalHealthEnvironment:
394
 
395
  # Guard: refuse to grade once the episode is finished.
396
  if self._task_index >= len(TASK_ORDER):
397
- return MentalHealthObservation(
398
  scenario="Episode already finished. Call /reset to start a new one.",
399
  feedback="No-op: episode is done.",
400
- reward=0.0, done=True, score_breakdown={},
401
  task_id=self._state.current_task,
402
  )
 
 
403
 
404
  self._state.step_count += 1
405
  task = self._state.current_task
@@ -411,7 +447,11 @@ class ITMentalHealthEnvironment:
411
  task_id=task, scenario=scenario_text,
412
  response=action.response, ground_truth=ground_truth,
413
  )
 
414
  self._state.cumulative_reward += reward
 
 
 
415
  self._state.tasks_completed.append(task)
416
 
417
  self._task_index += 1
@@ -425,11 +465,18 @@ class ITMentalHealthEnvironment:
425
  self._state.current_task = next_task
426
  next_scenario, _ = self._current_scenarios[next_task]
427
 
428
- return MentalHealthObservation(
429
- scenario=next_scenario, feedback=feedback,
430
- reward=reward, done=done, score_breakdown=breakdown, task_id=next_task,
 
 
 
 
 
 
 
431
  )
 
432
 
433
- @property
434
  def state(self):
435
  return self._state
 
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 = {
27
+ "burnout_detection": "easy",
28
+ "stress_triage": "medium",
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)
 
379
  "intervention_plan": generate_intervention_plan_scenario(self._rng),
380
  }
381
 
382
+ def _score_metadata(self, task: str, reward: float):
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),
392
+ "task_step_number": task_step_count,
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):
399
  actual_seed = seed if seed is not None else random.randint(0, 2**31)
400
  self._rng = random.Random(actual_seed)
 
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},
409
  )
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,
417
+ "difficulty": TASK_DIFFICULTY[task],
418
+ "step_score": 0.0,
419
+ "task_step_number": 0,
420
+ "task_cumulative_score": 0.0,
421
+ "difficulty_cumulative_score": 0.0,
422
+ "overall_cumulative_score": 0.0,
423
+ },
424
  )
425
 
426
  def step(self, action):
 
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.",
434
  feedback="No-op: episode is done.",
 
435
  task_id=self._state.current_task,
436
  )
437
+ reward = MentalHealthReward(value=0.0, score_breakdown={}, feedback="No-op: episode is done.")
438
+ return observation, reward, True, {}
439
 
440
  self._state.step_count += 1
441
  task = self._state.current_task
 
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"]
453
+ self._state.task_step_counts[task] = score_metadata["task_step_number"]
454
+ self._state.difficulty_scores[TASK_DIFFICULTY[task]] = score_metadata["difficulty_cumulative_score"]
455
  self._state.tasks_completed.append(task)
456
 
457
  self._task_index += 1
 
465
  self._state.current_task = next_task
466
  next_scenario, _ = self._current_scenarios[next_task]
467
 
468
+ observation = MentalHealthObservation(
469
+ scenario=next_scenario,
470
+ feedback=feedback,
471
+ task_id=next_task,
472
+ metadata=score_metadata,
473
+ )
474
+ reward_model = MentalHealthReward(
475
+ value=reward,
476
+ score_breakdown=breakdown,
477
+ feedback=feedback,
478
  )
479
+ return observation, reward_model, done, score_metadata
480
 
 
481
  def state(self):
482
  return self._state
models.py CHANGED
@@ -1,66 +1,38 @@
1
  """
2
- IT Mental Health OpenEnv - Typed Models
3
- Action / Observation / State for the IT Burnout & Mental Wellness environment.
4
  """
5
 
6
- from dataclasses import dataclass, field
7
- from typing import Optional, Dict, Any, List
8
 
 
9
 
10
- @dataclass
11
- class MentalHealthAction:
12
- """
13
- Action submitted by the agent for each mental health scenario step.
14
 
15
- Fields:
16
- response (str): The agent's textual response / intervention recommendation.
17
- task_id (str): Which task is being attempted: 'burnout_detection',
18
- 'stress_triage', or 'intervention_plan'.
19
- confidence (float): Agent's self-reported confidence 0.0–1.0.
20
- metadata (dict): Optional extra fields (reasoning chain, flags, etc.)
21
- """
22
  response: str
23
  task_id: str = "burnout_detection"
24
  confidence: float = 1.0
25
- metadata: Dict[str, Any] = field(default_factory=dict)
26
 
27
 
28
- @dataclass
29
- class MentalHealthObservation:
30
- """
31
- Observation returned after each step.
32
 
33
- Fields:
34
- scenario (str): The current scenario description shown to the agent.
35
- feedback (str): Evaluator feedback on the last action.
36
- reward (float): Reward for the last step (0.0 – 1.0).
37
- done (bool): Whether the episode is finished.
38
- score_breakdown (dict): Partial scores by rubric dimension.
39
- task_id (str): Current task identifier.
40
- """
41
  scenario: str
42
  feedback: str
43
- reward: float
44
- done: bool
45
- score_breakdown: Dict[str, float] = field(default_factory=dict)
46
  task_id: str = "burnout_detection"
47
- metadata: Dict[str, Any] = field(default_factory=dict)
48
-
49
 
50
- @dataclass
51
- class MentalHealthState:
52
- """
53
- Episode state / metadata.
54
 
55
- Fields:
56
- episode_id (str): UUID for the current episode.
57
- step_count (int): How many steps have elapsed.
58
- current_task (str): Active task identifier.
59
- cumulative_reward (float): Total reward accumulated so far.
60
- tasks_completed (list): List of task IDs already completed.
61
- """
62
  episode_id: Optional[str] = None
63
  step_count: int = 0
64
  current_task: str = "burnout_detection"
65
  cumulative_reward: float = 0.0
66
- tasks_completed: List[str] = field(default_factory=list)
 
 
 
 
1
  """
2
+ IT Mental Health OpenEnv - Typed Pydantic models.
 
3
  """
4
 
5
+ from typing import Any, Dict, List, Optional
 
6
 
7
+ from pydantic import BaseModel, Field
8
 
 
 
 
 
9
 
10
+ class MentalHealthAction(BaseModel):
 
 
 
 
 
 
11
  response: str
12
  task_id: str = "burnout_detection"
13
  confidence: float = 1.0
14
+ metadata: Dict[str, Any] = Field(default_factory=dict)
15
 
16
 
17
+ class MentalHealthReward(BaseModel):
18
+ value: float
19
+ score_breakdown: Dict[str, float] = Field(default_factory=dict)
20
+ feedback: str = ""
21
 
22
+
23
+ class MentalHealthObservation(BaseModel):
 
 
 
 
 
 
24
  scenario: str
25
  feedback: str
 
 
 
26
  task_id: str = "burnout_detection"
27
+ metadata: Dict[str, Any] = Field(default_factory=dict)
 
28
 
 
 
 
 
29
 
30
+ class MentalHealthState(BaseModel):
 
 
 
 
 
 
31
  episode_id: Optional[str] = None
32
  step_count: int = 0
33
  current_task: str = "burnout_detection"
34
  cumulative_reward: float = 0.0
35
+ tasks_completed: List[str] = Field(default_factory=list)
36
+ task_scores: Dict[str, float] = Field(default_factory=dict)
37
+ task_step_counts: Dict[str, int] = Field(default_factory=dict)
38
+ difficulty_scores: Dict[str, float] = Field(default_factory=dict)
server/app.py CHANGED
@@ -13,7 +13,7 @@ from fastapi import FastAPI, Header, HTTPException, Query, Request, Response
13
  from fastapi.middleware.cors import CORSMiddleware
14
  from pydantic import BaseModel, Field
15
 
16
- from it_mental_health_environment import ITMentalHealthEnvironment, TASK_ORDER
17
  from models import MentalHealthAction
18
 
19
  app = FastAPI(
@@ -106,6 +106,9 @@ class StateResponse(BaseModel):
106
  current_task: str
107
  cumulative_reward: float
108
  tasks_completed: list
 
 
 
109
  session_id: str
110
 
111
 
@@ -175,9 +178,9 @@ def reset(
175
  return ObservationResponse(
176
  scenario=obs.scenario,
177
  feedback=obs.feedback,
178
- reward=obs.reward,
179
- done=obs.done,
180
- score_breakdown=obs.score_breakdown,
181
  task_id=obs.task_id,
182
  metadata=_response_metadata(obs.metadata, session_id),
183
  )
@@ -202,17 +205,17 @@ def step(
202
  metadata=req.metadata or {},
203
  )
204
  try:
205
- obs = env.step(action)
206
  except ValueError as exc:
207
  raise HTTPException(status_code=400, detail=str(exc)) from exc
208
  return ObservationResponse(
209
  scenario=obs.scenario,
210
  feedback=obs.feedback,
211
- reward=obs.reward,
212
- done=obs.done,
213
- score_breakdown=obs.score_breakdown,
214
  task_id=obs.task_id,
215
- metadata=_response_metadata(obs.metadata, session_id),
216
  )
217
 
218
 
@@ -228,13 +231,16 @@ def state(
228
  header_session_id=x_session_id,
229
  ) or ANONYMOUS_SESSION_ID
230
  env = session_store.get_or_create(resolved_session_id)
231
- s = env.state
232
  return StateResponse(
233
  episode_id=s.episode_id,
234
  step_count=s.step_count,
235
  current_task=s.current_task,
236
  cumulative_reward=s.cumulative_reward,
237
  tasks_completed=s.tasks_completed,
 
 
 
238
  session_id=resolved_session_id,
239
  )
240
 
@@ -248,6 +254,11 @@ def list_tasks():
248
  "task_id": tid,
249
  "difficulty": TASK_METADATA[tid]["difficulty"],
250
  "description": TASK_METADATA[tid]["description"],
 
 
 
 
 
251
  }
252
  for tid in TASK_ORDER
253
  ]
@@ -270,6 +281,19 @@ def schema():
270
  "done": "bool",
271
  "score_breakdown": "dict - partial scores per rubric dimension",
272
  "task_id": "str",
 
 
 
 
 
 
 
 
 
 
 
 
 
273
  },
274
  }
275
 
 
13
  from fastapi.middleware.cors import CORSMiddleware
14
  from pydantic import BaseModel, Field
15
 
16
+ from it_mental_health_environment import ITMentalHealthEnvironment, TASK_DIFFICULTY, TASK_ORDER
17
  from models import MentalHealthAction
18
 
19
  app = FastAPI(
 
106
  current_task: str
107
  cumulative_reward: float
108
  tasks_completed: list
109
+ task_scores: Dict[str, float]
110
+ task_step_counts: Dict[str, int]
111
+ difficulty_scores: Dict[str, float]
112
  session_id: str
113
 
114
 
 
178
  return ObservationResponse(
179
  scenario=obs.scenario,
180
  feedback=obs.feedback,
181
+ reward=0.0,
182
+ done=False,
183
+ score_breakdown={},
184
  task_id=obs.task_id,
185
  metadata=_response_metadata(obs.metadata, session_id),
186
  )
 
205
  metadata=req.metadata or {},
206
  )
207
  try:
208
+ obs, reward_model, done, info = env.step(action)
209
  except ValueError as exc:
210
  raise HTTPException(status_code=400, detail=str(exc)) from exc
211
  return ObservationResponse(
212
  scenario=obs.scenario,
213
  feedback=obs.feedback,
214
+ reward=reward_model.value,
215
+ done=done,
216
+ score_breakdown=reward_model.score_breakdown,
217
  task_id=obs.task_id,
218
+ metadata=_response_metadata(info or obs.metadata, session_id),
219
  )
220
 
221
 
 
231
  header_session_id=x_session_id,
232
  ) or ANONYMOUS_SESSION_ID
233
  env = session_store.get_or_create(resolved_session_id)
234
+ s = env.state()
235
  return StateResponse(
236
  episode_id=s.episode_id,
237
  step_count=s.step_count,
238
  current_task=s.current_task,
239
  cumulative_reward=s.cumulative_reward,
240
  tasks_completed=s.tasks_completed,
241
+ task_scores=s.task_scores,
242
+ task_step_counts=s.task_step_counts,
243
+ difficulty_scores=s.difficulty_scores,
244
  session_id=resolved_session_id,
245
  )
246
 
 
254
  "task_id": tid,
255
  "difficulty": TASK_METADATA[tid]["difficulty"],
256
  "description": TASK_METADATA[tid]["description"],
257
+ "scoring": {
258
+ "step_score": "score for the current graded step in this task",
259
+ "task_cumulative_score": "sum of all step scores recorded for this task",
260
+ "difficulty_cumulative_score": "sum of all step scores for this difficulty bucket",
261
+ },
262
  }
263
  for tid in TASK_ORDER
264
  ]
 
281
  "done": "bool",
282
  "score_breakdown": "dict - partial scores per rubric dimension",
283
  "task_id": "str",
284
+ "metadata": {
285
+ "difficulty": "str - easy / medium / hard for the graded step",
286
+ "step_score": "float - score for this step",
287
+ "task_step_number": "int - graded step count within this task",
288
+ "task_cumulative_score": "float - cumulative score for this task",
289
+ "difficulty_cumulative_score": "float - cumulative score for this difficulty level",
290
+ "overall_cumulative_score": "float - episode-wide cumulative score",
291
+ },
292
+ },
293
+ "state": {
294
+ "task_scores": "dict - cumulative score per task_id",
295
+ "task_step_counts": "dict - graded step count per task_id",
296
+ "difficulty_scores": "dict - cumulative score per difficulty bucket",
297
  },
298
  }
299