File size: 1,645 Bytes
ba693b7
 
 
52cb935
ba693b7
 
 
 
52cb935
ba693b7
52cb935
 
 
ba693b7
 
52cb935
 
 
a523ae1
52cb935
 
 
 
 
a523ae1
ba693b7
 
52cb935
 
a523ae1
 
 
ba693b7
 
52cb935
a523ae1
52cb935
 
 
 
 
 
 
 
 
 
 
ba693b7
 
52cb935
 
 
 
 
 
 
 
 
 
 
 
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
"""
task1_easy.py
─────────────
Task 1 β€” Cognitive Stage Classification (Easy)
"""

from __future__ import annotations

from typing import Callable, Dict, Any, List

from cognitive_env import CogTraceEnv
from patient_simulator import PatientConfig, PatientSimulator
from models import Observation


def _score_prediction(true_stage: int, predicted_stage: int) -> float:
    delta = abs(true_stage - predicted_stage)
    if delta == 0:
        return 0.999
    elif delta == 1:
        return 0.50
    elif delta == 2:
        return 0.20
    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:
    true_stage = seed % 5
    cfg = PatientConfig(
        true_stage=true_stage,
        episode_length=1,
        decline_rate=0.0,
        noise_level=0.5,
        anomaly_day=None,
        anomaly_duration=0,
        seed=seed,
        patient_id=f"task1_seed{seed}",
    )
    return CogTraceEnv(config=cfg)


def run_episode(env: CogTraceEnv, agent_fn: Callable[[Dict[str, Any]], int]) -> Dict[str, Any]:
    obs = env.reset()
    obs_dict = obs.model_dump()
    true_stage = env._sim.true_stage(0)
    predicted = agent_fn(obs_dict)
    predicted = max(0, min(4, int(predicted)))
    score = _score_prediction(true_stage, predicted)
    return {
        "true_stage":      true_stage,
        "predicted_stage": predicted,
        "score":           score,
        "observation":     obs_dict,
    }