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.
@@ -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
+ ]