xlab-api 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- xlab_api/__init__.py +66 -0
- xlab_api/data.py +90 -0
- xlab_api/py.typed +1 -0
- xlab_api/spec/__init__.py +118 -0
- xlab_api/spec/_names.py +15 -0
- xlab_api/spec/articulation.py +34 -0
- xlab_api/spec/capability.py +198 -0
- xlab_api/spec/entity.py +156 -0
- xlab_api/spec/freeze.py +92 -0
- xlab_api/spec/lifecycle.py +428 -0
- xlab_api/spec/mdp.py +384 -0
- xlab_api/spec/motion_reference.py +323 -0
- xlab_api/spec/portability.py +106 -0
- xlab_api/spec/relations.py +29 -0
- xlab_api/spec/rigid_object.py +103 -0
- xlab_api/spec/robot.py +300 -0
- xlab_api/spec/sensor.py +481 -0
- xlab_api/spec/task.py +524 -0
- xlab_api/spec/validation.py +364 -0
- xlab_api/spec/volume.py +150 -0
- xlab_api-0.1.0.dist-info/METADATA +21 -0
- xlab_api-0.1.0.dist-info/RECORD +24 -0
- xlab_api-0.1.0.dist-info/WHEEL +5 -0
- xlab_api-0.1.0.dist-info/top_level.txt +1 -0
xlab_api/spec/freeze.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
"""Create the immutable declaration snapshot consumed by compilers and contracts."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
from copy import copy
|
|
7
|
+
from dataclasses import fields, is_dataclass
|
|
8
|
+
from types import MappingProxyType
|
|
9
|
+
from typing import TYPE_CHECKING, Any
|
|
10
|
+
|
|
11
|
+
if TYPE_CHECKING:
|
|
12
|
+
from .task import TaskSpec
|
|
13
|
+
|
|
14
|
+
_IN_PROGRESS = object()
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _freeze(value: Any, memo: dict[int, Any]) -> Any:
|
|
18
|
+
"""Copy one declaration value while replacing mutable containers."""
|
|
19
|
+
if value is None or isinstance(value, (bool, int, float, complex, str, bytes)):
|
|
20
|
+
return value
|
|
21
|
+
|
|
22
|
+
value_id = id(value)
|
|
23
|
+
previous = memo.get(value_id)
|
|
24
|
+
if previous is _IN_PROGRESS:
|
|
25
|
+
raise ValueError("TaskSpec declarations must not contain container cycles")
|
|
26
|
+
if previous is not None:
|
|
27
|
+
return previous
|
|
28
|
+
|
|
29
|
+
if isinstance(value, Mapping):
|
|
30
|
+
memo[value_id] = _IN_PROGRESS
|
|
31
|
+
backing: dict[Any, Any] = {}
|
|
32
|
+
for key, item in value.items():
|
|
33
|
+
backing[_freeze(key, memo)] = _freeze(item, memo)
|
|
34
|
+
frozen = MappingProxyType(backing)
|
|
35
|
+
memo[value_id] = frozen
|
|
36
|
+
return frozen
|
|
37
|
+
|
|
38
|
+
if is_dataclass(value) and not isinstance(value, type):
|
|
39
|
+
parameters = getattr(type(value), "__dataclass_params__", None)
|
|
40
|
+
if parameters is None or not parameters.frozen:
|
|
41
|
+
raise TypeError(
|
|
42
|
+
"TaskSpec declarations may contain only frozen dataclasses; "
|
|
43
|
+
f"got {type(value).__module__}.{type(value).__qualname__}"
|
|
44
|
+
)
|
|
45
|
+
memo[value_id] = _IN_PROGRESS
|
|
46
|
+
frozen = copy(value)
|
|
47
|
+
for declaration_field in fields(value):
|
|
48
|
+
object.__setattr__(
|
|
49
|
+
frozen,
|
|
50
|
+
declaration_field.name,
|
|
51
|
+
_freeze(getattr(value, declaration_field.name), memo),
|
|
52
|
+
)
|
|
53
|
+
memo[value_id] = frozen
|
|
54
|
+
return frozen
|
|
55
|
+
|
|
56
|
+
if isinstance(value, (tuple, list)):
|
|
57
|
+
memo[value_id] = _IN_PROGRESS
|
|
58
|
+
frozen = tuple(_freeze(item, memo) for item in value)
|
|
59
|
+
memo[value_id] = frozen
|
|
60
|
+
return frozen
|
|
61
|
+
|
|
62
|
+
if isinstance(value, (set, frozenset)):
|
|
63
|
+
memo[value_id] = _IN_PROGRESS
|
|
64
|
+
frozen = frozenset(_freeze(item, memo) for item in value)
|
|
65
|
+
memo[value_id] = frozen
|
|
66
|
+
return frozen
|
|
67
|
+
|
|
68
|
+
# Functions, enums, paths, scalar library values, and other already-immutable
|
|
69
|
+
# leaves are retained. Contract serialization remains the authority that
|
|
70
|
+
# rejects values which cannot be represented reproducibly.
|
|
71
|
+
return value
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def freeze_task_spec(spec: TaskSpec) -> TaskSpec:
|
|
75
|
+
"""Return a deeply immutable copy of a validated task declaration.
|
|
76
|
+
|
|
77
|
+
Task configuration objects remain mutable while a family assembles them.
|
|
78
|
+
The application registry calls this once, after initial validation, and
|
|
79
|
+
validates the returned snapshot again. Hashing and compilation therefore
|
|
80
|
+
consume exactly the same object and cannot observe later source mutations.
|
|
81
|
+
"""
|
|
82
|
+
from .task import TaskSpec
|
|
83
|
+
|
|
84
|
+
if not isinstance(spec, TaskSpec):
|
|
85
|
+
raise TypeError(f"freeze_task_spec expects TaskSpec, got {type(spec).__name__}")
|
|
86
|
+
frozen = _freeze(spec, {})
|
|
87
|
+
if not isinstance(frozen, TaskSpec): # pragma: no cover - guarded above
|
|
88
|
+
raise TypeError("TaskSpec freezing produced an invalid root object")
|
|
89
|
+
return frozen
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
__all__ = ["freeze_task_spec"]
|
|
@@ -0,0 +1,428 @@
|
|
|
1
|
+
"""Engine-neutral timing, reset, and recoverable-component contracts.
|
|
2
|
+
|
|
3
|
+
Lifecycle declarations use integer clock ratios. Floating point time is only a
|
|
4
|
+
presentation value; scheduling is derived from physics ticks so a long run does
|
|
5
|
+
not drift differently on two engines.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import math
|
|
11
|
+
from collections.abc import Mapping
|
|
12
|
+
from dataclasses import dataclass, field
|
|
13
|
+
from fractions import Fraction
|
|
14
|
+
from typing import TYPE_CHECKING, Literal
|
|
15
|
+
|
|
16
|
+
if TYPE_CHECKING:
|
|
17
|
+
from .task import SimSpec, TaskSpec
|
|
18
|
+
|
|
19
|
+
ClockReset = Literal["never", "episode"]
|
|
20
|
+
ComponentPhase = Literal[
|
|
21
|
+
"startup",
|
|
22
|
+
"pre_step",
|
|
23
|
+
"pre_physics",
|
|
24
|
+
"post_physics",
|
|
25
|
+
"post_step",
|
|
26
|
+
"on_reset",
|
|
27
|
+
"on_demand",
|
|
28
|
+
]
|
|
29
|
+
ComponentReset = Literal["stateless", "full", "partial"]
|
|
30
|
+
ComponentState = Literal["stateless", "snapshot"]
|
|
31
|
+
|
|
32
|
+
_BUILTIN_CLOCKS = frozenset({"physics", "policy", "episode"})
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _is_dotted_name(value: str) -> bool:
|
|
36
|
+
return bool(value) and all(part.isidentifier() for part in value.split("."))
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass(frozen=True, slots=True)
|
|
40
|
+
class ClockDomainSpec:
|
|
41
|
+
"""A named integer subdivision of another clock.
|
|
42
|
+
|
|
43
|
+
``tick_divider=4`` means one tick here for every four ticks of ``parent``.
|
|
44
|
+
``phase`` is expressed in parent ticks and must be smaller than the divider.
|
|
45
|
+
The built-in ``physics``, ``policy``, and ``episode`` clocks are derived from
|
|
46
|
+
:class:`SimSpec`; declarations add clocks below those roots.
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
name: str
|
|
50
|
+
parent: str = "physics"
|
|
51
|
+
tick_divider: int = 1
|
|
52
|
+
phase: int = 0
|
|
53
|
+
reset: ClockReset = "never"
|
|
54
|
+
|
|
55
|
+
def __post_init__(self) -> None:
|
|
56
|
+
if not _is_dotted_name(self.name):
|
|
57
|
+
raise ValueError(
|
|
58
|
+
f"Clock domain names must be dotted identifiers, got {self.name!r}."
|
|
59
|
+
)
|
|
60
|
+
if self.name in _BUILTIN_CLOCKS:
|
|
61
|
+
raise ValueError(
|
|
62
|
+
f"Clock domain {self.name!r} is built in and cannot be redeclared."
|
|
63
|
+
)
|
|
64
|
+
if not _is_dotted_name(self.parent):
|
|
65
|
+
raise ValueError(
|
|
66
|
+
f"Clock parent names must be dotted identifiers, got {self.parent!r}."
|
|
67
|
+
)
|
|
68
|
+
if isinstance(self.tick_divider, bool) or self.tick_divider < 1:
|
|
69
|
+
raise ValueError("Clock tick_divider must be a positive integer.")
|
|
70
|
+
if isinstance(self.phase, bool) or not 0 <= self.phase < self.tick_divider:
|
|
71
|
+
raise ValueError(
|
|
72
|
+
f"Clock phase must satisfy 0 <= phase < tick_divider, got "
|
|
73
|
+
f"phase={self.phase}, divider={self.tick_divider}."
|
|
74
|
+
)
|
|
75
|
+
if self.reset not in {"never", "episode"}:
|
|
76
|
+
raise ValueError(f"Unknown clock reset semantics {self.reset!r}.")
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
@dataclass(frozen=True, slots=True)
|
|
80
|
+
class ResolvedClockDomain:
|
|
81
|
+
"""A clock reduced to exact physics-step units."""
|
|
82
|
+
|
|
83
|
+
name: str
|
|
84
|
+
period_numerator: int
|
|
85
|
+
period_denominator: int
|
|
86
|
+
phase_numerator: int
|
|
87
|
+
phase_denominator: int
|
|
88
|
+
reset: ClockReset
|
|
89
|
+
|
|
90
|
+
def __post_init__(self) -> None:
|
|
91
|
+
period = self.period_physics_steps
|
|
92
|
+
phase = self.phase_physics_steps
|
|
93
|
+
if period <= 0:
|
|
94
|
+
raise ValueError("Resolved clock periods must be positive.")
|
|
95
|
+
if not 0 <= phase < period:
|
|
96
|
+
raise ValueError("Resolved clock phase is outside its period.")
|
|
97
|
+
|
|
98
|
+
@classmethod
|
|
99
|
+
def from_fractions(
|
|
100
|
+
cls,
|
|
101
|
+
name: str,
|
|
102
|
+
period: Fraction,
|
|
103
|
+
phase: Fraction,
|
|
104
|
+
reset: ClockReset,
|
|
105
|
+
) -> ResolvedClockDomain:
|
|
106
|
+
return cls(
|
|
107
|
+
name=name,
|
|
108
|
+
period_numerator=period.numerator,
|
|
109
|
+
period_denominator=period.denominator,
|
|
110
|
+
phase_numerator=phase.numerator,
|
|
111
|
+
phase_denominator=phase.denominator,
|
|
112
|
+
reset=reset,
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
@property
|
|
116
|
+
def period_physics_steps(self) -> Fraction:
|
|
117
|
+
return Fraction(self.period_numerator, self.period_denominator)
|
|
118
|
+
|
|
119
|
+
@property
|
|
120
|
+
def phase_physics_steps(self) -> Fraction:
|
|
121
|
+
return Fraction(self.phase_numerator, self.phase_denominator)
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
@dataclass(frozen=True, slots=True)
|
|
125
|
+
class ComponentLifecycleSpec:
|
|
126
|
+
"""When a component runs, how it resets, and whether state is recoverable."""
|
|
127
|
+
|
|
128
|
+
clock: str
|
|
129
|
+
phase: ComponentPhase
|
|
130
|
+
reset: ComponentReset
|
|
131
|
+
state: ComponentState
|
|
132
|
+
latency_ticks: int = 0
|
|
133
|
+
|
|
134
|
+
def __post_init__(self) -> None:
|
|
135
|
+
if not _is_dotted_name(self.clock):
|
|
136
|
+
raise ValueError(
|
|
137
|
+
f"Component clock must be a dotted identifier, got {self.clock!r}."
|
|
138
|
+
)
|
|
139
|
+
if self.phase not in {
|
|
140
|
+
"startup",
|
|
141
|
+
"pre_step",
|
|
142
|
+
"pre_physics",
|
|
143
|
+
"post_physics",
|
|
144
|
+
"post_step",
|
|
145
|
+
"on_reset",
|
|
146
|
+
"on_demand",
|
|
147
|
+
}:
|
|
148
|
+
raise ValueError(f"Unknown component phase {self.phase!r}.")
|
|
149
|
+
if self.reset not in {"stateless", "full", "partial"}:
|
|
150
|
+
raise ValueError(f"Unknown component reset semantics {self.reset!r}.")
|
|
151
|
+
if self.state not in {"stateless", "snapshot"}:
|
|
152
|
+
raise ValueError(f"Unknown component state semantics {self.state!r}.")
|
|
153
|
+
if isinstance(self.latency_ticks, bool) or self.latency_ticks < 0:
|
|
154
|
+
raise ValueError("Component latency_ticks must be a non-negative integer.")
|
|
155
|
+
if self.state == "snapshot" and self.reset == "stateless":
|
|
156
|
+
raise ValueError(
|
|
157
|
+
"A recoverable stateful component must declare full or partial reset semantics."
|
|
158
|
+
)
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
@dataclass(frozen=True)
|
|
162
|
+
class LifecycleSpec:
|
|
163
|
+
"""The stable lifecycle extension point carried by every :class:`TaskSpec`.
|
|
164
|
+
|
|
165
|
+
Built-in component contracts are derived from scene and MDP declarations.
|
|
166
|
+
``components`` is a complete, explicit replacement for a derived contract
|
|
167
|
+
when a component has stronger semantics than its family default. Keys below
|
|
168
|
+
``controller/`` declare application controllers in addition to derived MDP
|
|
169
|
+
and sensor components.
|
|
170
|
+
"""
|
|
171
|
+
|
|
172
|
+
clocks: tuple[ClockDomainSpec, ...] = ()
|
|
173
|
+
components: Mapping[str, ComponentLifecycleSpec] = field(default_factory=dict)
|
|
174
|
+
trace_schema_version: int = 1
|
|
175
|
+
snapshot_schema_version: int = 1
|
|
176
|
+
|
|
177
|
+
def __post_init__(self) -> None:
|
|
178
|
+
object.__setattr__(self, "clocks", tuple(self.clocks))
|
|
179
|
+
object.__setattr__(self, "components", dict(self.components))
|
|
180
|
+
names = tuple(clock.name for clock in self.clocks)
|
|
181
|
+
duplicates = sorted({name for name in names if names.count(name) > 1})
|
|
182
|
+
if duplicates:
|
|
183
|
+
raise ValueError(
|
|
184
|
+
f"Lifecycle clock names must be unique; repeated: {duplicates}."
|
|
185
|
+
)
|
|
186
|
+
empty_components = sorted(name for name in self.components if not name)
|
|
187
|
+
if empty_components:
|
|
188
|
+
raise ValueError("Lifecycle component keys must be non-empty.")
|
|
189
|
+
if self.trace_schema_version != 1 or self.snapshot_schema_version != 1:
|
|
190
|
+
raise ValueError(
|
|
191
|
+
"This release supports lifecycle trace and snapshot schema version 1 only."
|
|
192
|
+
)
|
|
193
|
+
|
|
194
|
+
def resolved_clocks(self, sim: SimSpec) -> dict[str, ResolvedClockDomain]:
|
|
195
|
+
"""Resolve every clock to integer physics steps and reject cycles."""
|
|
196
|
+
resolved = {
|
|
197
|
+
"physics": ResolvedClockDomain.from_fractions(
|
|
198
|
+
"physics", Fraction(1), Fraction(0), "never"
|
|
199
|
+
),
|
|
200
|
+
"policy": ResolvedClockDomain.from_fractions(
|
|
201
|
+
"policy", Fraction(sim.decimation), Fraction(0), "never"
|
|
202
|
+
),
|
|
203
|
+
"episode": ResolvedClockDomain.from_fractions(
|
|
204
|
+
"episode", Fraction(sim.decimation), Fraction(0), "episode"
|
|
205
|
+
),
|
|
206
|
+
}
|
|
207
|
+
pending = {clock.name: clock for clock in self.clocks}
|
|
208
|
+
while pending:
|
|
209
|
+
progressed = False
|
|
210
|
+
for name, clock in tuple(pending.items()):
|
|
211
|
+
parent = resolved.get(clock.parent)
|
|
212
|
+
if parent is None:
|
|
213
|
+
continue
|
|
214
|
+
period = parent.period_physics_steps * clock.tick_divider
|
|
215
|
+
phase = (
|
|
216
|
+
parent.phase_physics_steps
|
|
217
|
+
+ parent.period_physics_steps * clock.phase
|
|
218
|
+
) % period
|
|
219
|
+
reset: ClockReset = (
|
|
220
|
+
"episode"
|
|
221
|
+
if parent.reset == "episode" or clock.reset == "episode"
|
|
222
|
+
else "never"
|
|
223
|
+
)
|
|
224
|
+
resolved[name] = ResolvedClockDomain.from_fractions(
|
|
225
|
+
name, period, phase, reset
|
|
226
|
+
)
|
|
227
|
+
pending.pop(name)
|
|
228
|
+
progressed = True
|
|
229
|
+
if progressed:
|
|
230
|
+
continue
|
|
231
|
+
unresolved = {name: clock.parent for name, clock in sorted(pending.items())}
|
|
232
|
+
raise ValueError(
|
|
233
|
+
f"Lifecycle clocks contain an unknown parent or cycle: {unresolved}."
|
|
234
|
+
)
|
|
235
|
+
return resolved
|
|
236
|
+
|
|
237
|
+
|
|
238
|
+
def _period_clock(
|
|
239
|
+
*,
|
|
240
|
+
component_key: str,
|
|
241
|
+
update_period: float | None,
|
|
242
|
+
sim: SimSpec,
|
|
243
|
+
clocks: dict[str, ResolvedClockDomain],
|
|
244
|
+
) -> str:
|
|
245
|
+
if update_period is None:
|
|
246
|
+
return "physics"
|
|
247
|
+
raw_ratio = update_period / sim.physics_dt
|
|
248
|
+
physics_steps = Fraction(str(raw_ratio)).limit_denominator(1_000_000)
|
|
249
|
+
if physics_steps <= 0 or not math.isclose(
|
|
250
|
+
raw_ratio, float(physics_steps), rel_tol=0.0, abs_tol=1.0e-12
|
|
251
|
+
):
|
|
252
|
+
raise ValueError(
|
|
253
|
+
f"Lifecycle component {component_key!r} update period {update_period} s "
|
|
254
|
+
f"cannot be represented as a stable rational multiple of "
|
|
255
|
+
f"physics_dt={sim.physics_dt} s."
|
|
256
|
+
)
|
|
257
|
+
matching = [
|
|
258
|
+
name
|
|
259
|
+
for name, clock in clocks.items()
|
|
260
|
+
if clock.period_physics_steps == physics_steps
|
|
261
|
+
and clock.phase_physics_steps == 0
|
|
262
|
+
and clock.reset == "never"
|
|
263
|
+
]
|
|
264
|
+
if matching:
|
|
265
|
+
return min(matching, key=lambda name: (name not in _BUILTIN_CLOCKS, name))
|
|
266
|
+
generated_name = component_key.replace("/", ".")
|
|
267
|
+
clocks[generated_name] = ResolvedClockDomain.from_fractions(
|
|
268
|
+
generated_name, physics_steps, Fraction(0), "never"
|
|
269
|
+
)
|
|
270
|
+
return generated_name
|
|
271
|
+
|
|
272
|
+
|
|
273
|
+
def _derived_components(
|
|
274
|
+
task: TaskSpec,
|
|
275
|
+
clocks: dict[str, ResolvedClockDomain],
|
|
276
|
+
) -> dict[str, ComponentLifecycleSpec]:
|
|
277
|
+
components: dict[str, ComponentLifecycleSpec] = {}
|
|
278
|
+
|
|
279
|
+
def sensor_contract(
|
|
280
|
+
name: str,
|
|
281
|
+
update_period: float | None,
|
|
282
|
+
*,
|
|
283
|
+
stateful: bool,
|
|
284
|
+
latency_s: float = 0.0,
|
|
285
|
+
) -> None:
|
|
286
|
+
key = f"sensor/{name}"
|
|
287
|
+
clock_name = _period_clock(
|
|
288
|
+
component_key=key,
|
|
289
|
+
update_period=update_period,
|
|
290
|
+
sim=task.sim,
|
|
291
|
+
clocks=clocks,
|
|
292
|
+
)
|
|
293
|
+
clock = clocks[clock_name]
|
|
294
|
+
latency_ratio = latency_s / (clock.period_physics_steps * task.sim.physics_dt)
|
|
295
|
+
latency_ticks = round(latency_ratio)
|
|
296
|
+
if not math.isclose(latency_ratio, latency_ticks, rel_tol=0.0, abs_tol=1.0e-9):
|
|
297
|
+
raise ValueError(
|
|
298
|
+
f"Lifecycle component {key!r} latency {latency_s} s is not an "
|
|
299
|
+
f"integer number of {clock_name!r} ticks."
|
|
300
|
+
)
|
|
301
|
+
components[key] = ComponentLifecycleSpec(
|
|
302
|
+
clock=clock_name,
|
|
303
|
+
phase="post_physics",
|
|
304
|
+
reset="partial",
|
|
305
|
+
state="snapshot" if stateful else "stateless",
|
|
306
|
+
latency_ticks=latency_ticks,
|
|
307
|
+
)
|
|
308
|
+
|
|
309
|
+
for sensor in task.scene.contact_sensors:
|
|
310
|
+
sensor_contract(
|
|
311
|
+
sensor.name,
|
|
312
|
+
task.sim.physics_dt,
|
|
313
|
+
stateful=sensor.track_air_time or sensor.history_length > 0,
|
|
314
|
+
)
|
|
315
|
+
for sensor in task.scene.ray_casters:
|
|
316
|
+
sensor_contract(
|
|
317
|
+
sensor.name,
|
|
318
|
+
sensor.update_period,
|
|
319
|
+
stateful=True,
|
|
320
|
+
)
|
|
321
|
+
for sensor in task.scene.motion_references:
|
|
322
|
+
sensor_contract(sensor.name, sensor.update_period, stateful=True)
|
|
323
|
+
for sensor in task.scene.volume_points:
|
|
324
|
+
sensor_contract(sensor.name, sensor.update_period, stateful=True)
|
|
325
|
+
for sensor in task.scene.native_sensors:
|
|
326
|
+
sensor_contract(
|
|
327
|
+
sensor.name,
|
|
328
|
+
sensor.update_period,
|
|
329
|
+
stateful=True,
|
|
330
|
+
latency_s=sensor.latency,
|
|
331
|
+
)
|
|
332
|
+
|
|
333
|
+
family_defaults = {
|
|
334
|
+
"action": ComponentLifecycleSpec("policy", "pre_step", "partial", "snapshot"),
|
|
335
|
+
"command": ComponentLifecycleSpec("policy", "pre_step", "partial", "snapshot"),
|
|
336
|
+
"observation": ComponentLifecycleSpec(
|
|
337
|
+
"policy", "post_physics", "stateless", "stateless"
|
|
338
|
+
),
|
|
339
|
+
"reward": ComponentLifecycleSpec(
|
|
340
|
+
"policy", "post_physics", "stateless", "stateless"
|
|
341
|
+
),
|
|
342
|
+
"termination": ComponentLifecycleSpec(
|
|
343
|
+
"policy", "post_physics", "stateless", "stateless"
|
|
344
|
+
),
|
|
345
|
+
"curriculum": ComponentLifecycleSpec(
|
|
346
|
+
"episode", "on_reset", "partial", "snapshot"
|
|
347
|
+
),
|
|
348
|
+
}
|
|
349
|
+
for key, term in task.mdp.terms().items():
|
|
350
|
+
family = key.partition("/")[0]
|
|
351
|
+
if family == "event":
|
|
352
|
+
if term.mode == "startup": # type: ignore[attr-defined]
|
|
353
|
+
contract = ComponentLifecycleSpec(
|
|
354
|
+
"physics", "startup", "full", "snapshot"
|
|
355
|
+
)
|
|
356
|
+
elif term.mode == "reset": # type: ignore[attr-defined]
|
|
357
|
+
contract = ComponentLifecycleSpec(
|
|
358
|
+
"episode", "on_reset", "partial", "snapshot"
|
|
359
|
+
)
|
|
360
|
+
else:
|
|
361
|
+
contract = ComponentLifecycleSpec(
|
|
362
|
+
"physics", "post_physics", "partial", "snapshot"
|
|
363
|
+
)
|
|
364
|
+
else:
|
|
365
|
+
contract = family_defaults[family]
|
|
366
|
+
if family == "observation" and getattr(term, "history_length", 0) > 0:
|
|
367
|
+
contract = ComponentLifecycleSpec(
|
|
368
|
+
contract.clock,
|
|
369
|
+
contract.phase,
|
|
370
|
+
"partial",
|
|
371
|
+
"snapshot",
|
|
372
|
+
)
|
|
373
|
+
components[key] = contract
|
|
374
|
+
return components
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
def resolve_lifecycle_contract(
|
|
378
|
+
task: TaskSpec,
|
|
379
|
+
) -> tuple[dict[str, ResolvedClockDomain], dict[str, ComponentLifecycleSpec]]:
|
|
380
|
+
"""Return the complete clock/component contract for a materialized task."""
|
|
381
|
+
clocks = task.lifecycle.resolved_clocks(task.sim)
|
|
382
|
+
components = _derived_components(task, clocks)
|
|
383
|
+
unknown = sorted(set(task.lifecycle.components) - set(components))
|
|
384
|
+
unknown = [name for name in unknown if not name.startswith("controller/")]
|
|
385
|
+
if unknown:
|
|
386
|
+
raise ValueError(
|
|
387
|
+
f"Lifecycle overrides name undeclared components: {unknown}. "
|
|
388
|
+
f"Declared components: {sorted(components)}."
|
|
389
|
+
)
|
|
390
|
+
components.update(task.lifecycle.components)
|
|
391
|
+
unknown_clocks = sorted(
|
|
392
|
+
{component.clock for component in components.values()} - set(clocks)
|
|
393
|
+
)
|
|
394
|
+
if unknown_clocks:
|
|
395
|
+
raise ValueError(
|
|
396
|
+
f"Lifecycle components refer to unknown clocks: {unknown_clocks}."
|
|
397
|
+
)
|
|
398
|
+
for name, component in components.items():
|
|
399
|
+
if not name.startswith("controller/"):
|
|
400
|
+
continue
|
|
401
|
+
if name == "controller/":
|
|
402
|
+
raise ValueError("Lifecycle controller keys must include a name.")
|
|
403
|
+
if component.state != "snapshot" or component.reset not in {
|
|
404
|
+
"full",
|
|
405
|
+
"partial",
|
|
406
|
+
}:
|
|
407
|
+
raise ValueError(
|
|
408
|
+
f"Lifecycle controller {name!r} must declare recoverable state "
|
|
409
|
+
"and full or partial reset semantics."
|
|
410
|
+
)
|
|
411
|
+
if component.phase not in {"pre_step", "pre_physics"}:
|
|
412
|
+
raise ValueError(
|
|
413
|
+
f"Lifecycle controller {name!r} must run in pre_step or pre_physics."
|
|
414
|
+
)
|
|
415
|
+
if clocks[component.clock].reset != "never":
|
|
416
|
+
raise ValueError(
|
|
417
|
+
f"Lifecycle controller {name!r} must use a non-resetting clock."
|
|
418
|
+
)
|
|
419
|
+
return clocks, components
|
|
420
|
+
|
|
421
|
+
|
|
422
|
+
__all__ = [
|
|
423
|
+
"ClockDomainSpec",
|
|
424
|
+
"ComponentLifecycleSpec",
|
|
425
|
+
"LifecycleSpec",
|
|
426
|
+
"ResolvedClockDomain",
|
|
427
|
+
"resolve_lifecycle_contract",
|
|
428
|
+
]
|