Spaces:
Sleeping
Sleeping
File size: 2,122 Bytes
ba693b7 52cb935 ba693b7 52cb935 ba693b7 52cb935 a523ae1 52cb935 a523ae1 52cb935 a523ae1 ba693b7 52cb935 a523ae1 52cb935 ba693b7 52cb935 a523ae1 52cb935 ba693b7 52cb935 ba693b7 52cb935 ba693b7 52cb935 ba693b7 52cb935 ba693b7 52cb935 ba693b7 a523ae1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | """
task2_medium.py
βββββββββββββββ
Task 2 β Anomaly Timing Detection (Medium)
"""
from __future__ import annotations
from typing import Callable, Dict, Any, List, Optional
from cognitive_env import CogTraceEnv
from patient_simulator import PatientConfig
def _score_timing(true_anomaly_day: int, first_alert_day: Optional[int], episode_length: int) -> float:
if first_alert_day is None:
return 0.001
delta = abs(true_anomaly_day - first_alert_day)
if delta == 0:
return 0.999
elif delta == 1:
return 0.75
elif delta == 2:
return 0.50
elif delta == 3:
return 0.25
else:
return 0.001
def grade(results: List[Dict[str, Any]]) -> float:
if not results:
return 0.001
score = sum(r["score"] for r in results) / len(results)
return max(0.001, min(0.999, score))
def build_env(seed: int = 0) -> CogTraceEnv:
stage = (seed % 3) + 1
cfg = PatientConfig(
true_stage=stage,
episode_length=7,
decline_rate=0.0,
noise_level=1.0,
anomaly_day=None,
anomaly_duration=2,
seed=seed,
patient_id=f"task2_seed{seed}",
)
return CogTraceEnv(config=cfg)
def run_episode(
env: CogTraceEnv,
agent_fn: Callable[[Dict[str, Any], int], int],
) -> Dict[str, Any]:
obs = env.reset()
true_anomaly_day = env._sim.anomaly_day
episode_length = env.config.episode_length
first_alert_day: Optional[int] = None
step = 0
while True:
obs_dict = obs.model_dump()
action = agent_fn(obs_dict, step)
action = max(0, min(3, int(action)))
if action > 0 and first_alert_day is None:
first_alert_day = step
obs, reward, done, info = env.step(action)
step += 1
if done:
break
score = _score_timing(true_anomaly_day, first_alert_day, episode_length)
return {
"anomaly_day": true_anomaly_day,
"first_alert_day": first_alert_day if first_alert_day is not None else -1,
"score": score,
} |