File size: 3,281 Bytes
796da7c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional

from .common.serialization import to_jsonable
from .world.state import WorldState


@dataclass
class PredictionRecord:
    step: int
    action_id: str
    predicted_r_level: Optional[int]
    predicted_confidence: Optional[float]
    actual_r_level: int
    parameters: Dict[str, Any] = field(default_factory=dict)


@dataclass
class EpisodeResult:
    task_id: str
    task_name: str
    scenario_id: str
    terminated_by: str
    step_count: int
    max_steps: int
    success: bool
    prediction_records: List[PredictionRecord]
    final_world_state_summary: Dict[str, Any]
    final_locked_actions: Dict[str, str]
    final_critical_options: Dict[str, bool]
    available_actions: List[str]
    preservation_targets: List[str]

    def to_dict(self) -> Dict[str, Any]:
        return to_jsonable(self)


@dataclass
class EpisodeTracker:
    task_id: str = ""
    scenario_id: str = ""
    max_steps: int = 0
    step_count: int = 0
    prediction_records: List[PredictionRecord] = field(default_factory=list)
    _preservation_targets: List[str] = field(default_factory=list)

    def reset(self, task_id: str, scenario_id: str, max_steps: int, preservation_targets: List[str]) -> None:
        self.task_id = task_id
        self.scenario_id = scenario_id
        self.max_steps = max_steps
        self.step_count = 0
        self.prediction_records = []
        self._preservation_targets = list(preservation_targets)

    def increment_step(self) -> int:
        self.step_count += 1
        return self.step_count

    def record_prediction(
        self,
        action_id: str,
        predicted_r_level: Optional[int],
        predicted_confidence: Optional[float],
        actual_r_level: int,
        parameters: Optional[Dict[str, Any]] = None,
    ) -> None:
        self.prediction_records.append(
            PredictionRecord(
                step=self.step_count,
                action_id=action_id,
                predicted_r_level=predicted_r_level,
                predicted_confidence=predicted_confidence,
                actual_r_level=actual_r_level,
                parameters=dict(parameters or {}),
            )
        )

    def finalize(self, final_world_state: WorldState, task_spec: Any, terminated_by: str) -> EpisodeResult:
        return EpisodeResult(
            task_id=getattr(task_spec, "task_id", self.task_id),
            task_name=getattr(task_spec, "name", self.task_id),
            scenario_id=final_world_state.scenario_id,
            terminated_by=terminated_by,
            step_count=self.step_count,
            max_steps=self.max_steps,
            success=bool(getattr(task_spec, "success_fn", lambda ws, task: False)(final_world_state, task_spec)),
            prediction_records=list(self.prediction_records),
            final_world_state_summary=final_world_state.to_summary_dict(),
            final_locked_actions=dict(final_world_state.locked_actions),
            final_critical_options=dict(final_world_state.critical_options),
            available_actions=list(getattr(task_spec, "available_actions", [])),
            preservation_targets=list(self._preservation_targets),
        )