simscope 0.1.1__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.
- simscope/__init__.py +6 -0
- simscope/__main__.py +8 -0
- simscope/_assets/simscope-app.css +2 -0
- simscope/_assets/simscope-app.js +4311 -0
- simscope/_assets/simscope-player.js +4325 -0
- simscope/_assets/simscope-web.LICENSES.txt +407 -0
- simscope/_icon.py +22 -0
- simscope/_mjviser.py +203 -0
- simscope/annotations.py +1132 -0
- simscope/cli.py +482 -0
- simscope/core.py +257 -0
- simscope/derived.py +697 -0
- simscope/export.py +799 -0
- simscope/highlights.py +947 -0
- simscope/importers.py +874 -0
- simscope/index.py +579 -0
- simscope/io/__init__.py +45 -0
- simscope/io/blockfile.py +938 -0
- simscope/io/cas.py +294 -0
- simscope/io/codecs.py +566 -0
- simscope/io/errors.py +9 -0
- simscope/io/manifest.py +358 -0
- simscope/io/pack.py +563 -0
- simscope/io/scene.py +239 -0
- simscope/isaaclab.py +1460 -0
- simscope/library.py +705 -0
- simscope/mujoco.py +578 -0
- simscope/py.typed +0 -0
- simscope/recorder.py +784 -0
- simscope/server/__init__.py +9 -0
- simscope/server/app.py +149 -0
- simscope/server/blocks.py +191 -0
- simscope/server/jobs.py +166 -0
- simscope/server/routes.py +707 -0
- simscope/server/security.py +218 -0
- simscope/server/state.py +751 -0
- simscope/server/static.py +84 -0
- simscope/transforms.py +147 -0
- simscope-0.1.1.dist-info/METADATA +132 -0
- simscope-0.1.1.dist-info/RECORD +45 -0
- simscope-0.1.1.dist-info/WHEEL +4 -0
- simscope-0.1.1.dist-info/entry_points.txt +3 -0
- simscope-0.1.1.dist-info/licenses/LICENSE.md +201 -0
- simscope-0.1.1.dist-info/licenses/THIRD_PARTY_NOTICES.md +267 -0
- simscope-0.1.1.dist-info/licenses/src/simscope/_assets/simscope-web.LICENSES.txt +407 -0
simscope/recorder.py
ADDED
|
@@ -0,0 +1,784 @@
|
|
|
1
|
+
"""Crash-safe, low-overhead rollout recorder.
|
|
2
|
+
|
|
3
|
+
``Recorder.log`` only validates cheaply and copies one frame into a
|
|
4
|
+
preallocated window buffer ``[block_frames, E, K]`` per stream. When a
|
|
5
|
+
window fills it is handed to one background thread that applies pose sign
|
|
6
|
+
continuity, encodes (deflate releases the GIL, so this overlaps with the
|
|
7
|
+
simulator) and appends the blocks to the ``.blk`` files. Every finished
|
|
8
|
+
window is flushed to the OS, so a crash loses at most the window in flight;
|
|
9
|
+
``Library.recover`` restores the rest.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import concurrent.futures
|
|
13
|
+
import json
|
|
14
|
+
import logging
|
|
15
|
+
import math
|
|
16
|
+
import os
|
|
17
|
+
import pathlib
|
|
18
|
+
import queue
|
|
19
|
+
import shutil
|
|
20
|
+
import threading
|
|
21
|
+
import time
|
|
22
|
+
from collections.abc import Mapping, Sequence
|
|
23
|
+
from types import TracebackType
|
|
24
|
+
from typing import Any
|
|
25
|
+
|
|
26
|
+
import numpy as np
|
|
27
|
+
import numpy.typing as npt
|
|
28
|
+
|
|
29
|
+
from simscope import core, transforms
|
|
30
|
+
from simscope.io import blockfile, cas, codecs, manifest
|
|
31
|
+
from simscope.io import scene as scene_io
|
|
32
|
+
|
|
33
|
+
logger = logging.getLogger(__name__)
|
|
34
|
+
|
|
35
|
+
_SIGN_CHUNK = 1 << 21
|
|
36
|
+
"""Max elements per sign-continuity pass, to bound temporary memory."""
|
|
37
|
+
_POLL_S = 0.05
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class RecorderError(RuntimeError):
|
|
41
|
+
"""Raised for lifecycle errors, such as logging to a closed recorder."""
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class _Stream:
|
|
45
|
+
"""Static description of one recorded stream."""
|
|
46
|
+
|
|
47
|
+
def __init__(
|
|
48
|
+
self,
|
|
49
|
+
index: int,
|
|
50
|
+
name: str,
|
|
51
|
+
kind: core.StreamKind,
|
|
52
|
+
item_shape: tuple[int, ...],
|
|
53
|
+
n_envs: int,
|
|
54
|
+
codec: str,
|
|
55
|
+
labels: tuple[str, ...] | None = None,
|
|
56
|
+
units: str | None = None,
|
|
57
|
+
) -> None:
|
|
58
|
+
self.index = index
|
|
59
|
+
self.name = name
|
|
60
|
+
self.kind = kind
|
|
61
|
+
self.item_shape = item_shape
|
|
62
|
+
self.k = int(np.prod(item_shape, dtype=np.int64))
|
|
63
|
+
self.frame_shape = (n_envs, *item_shape)
|
|
64
|
+
self.codec = codec
|
|
65
|
+
self.labels = labels
|
|
66
|
+
self.units = units
|
|
67
|
+
self.info = manifest.StreamInfo(
|
|
68
|
+
f"{name}.blk", kind, item_shape, labels, units
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
class _Window:
|
|
73
|
+
"""One block-sized buffer per stream, plus where it belongs in time."""
|
|
74
|
+
|
|
75
|
+
__slots__ = ("arrays", "frames", "n", "t0")
|
|
76
|
+
|
|
77
|
+
def __init__(
|
|
78
|
+
self, streams: Sequence[_Stream], block_frames: int, n_envs: int
|
|
79
|
+
) -> None:
|
|
80
|
+
self.arrays = [
|
|
81
|
+
np.empty((block_frames, n_envs, s.k), np.float32) for s in streams
|
|
82
|
+
]
|
|
83
|
+
# Views shaped [block_frames, E, *item]; writes land in `arrays`.
|
|
84
|
+
self.frames = [
|
|
85
|
+
a.reshape(block_frames, n_envs, *s.item_shape)
|
|
86
|
+
for a, s in zip(self.arrays, streams, strict=True)
|
|
87
|
+
]
|
|
88
|
+
self.n = 0
|
|
89
|
+
self.t0 = 0
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
class _StreamWriter(blockfile.BlockWriter):
|
|
93
|
+
"""A ``BlockWriter`` fed with pre-encoded windows from the worker."""
|
|
94
|
+
|
|
95
|
+
def write_window(
|
|
96
|
+
self, t0: int, blocks: Sequence[blockfile.EncodedBlock], n: int
|
|
97
|
+
) -> None:
|
|
98
|
+
"""Appends one encoded window and flushes it to the OS.
|
|
99
|
+
|
|
100
|
+
Args:
|
|
101
|
+
t0: First frame index of the window.
|
|
102
|
+
blocks: One encoded block per env.
|
|
103
|
+
n: Frames in the window.
|
|
104
|
+
"""
|
|
105
|
+
del n # write_encoded advances the frame count from the blocks.
|
|
106
|
+
self.write_encoded(t0, blocks)
|
|
107
|
+
if self._file is not None:
|
|
108
|
+
self._file.flush()
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def _frame_error(name: str, want: tuple[int, ...], got: tuple[int, ...]):
|
|
112
|
+
"""Builds the error for a frame of the wrong shape."""
|
|
113
|
+
return ValueError(f"stream {name!r}: expected shape {want}, got {got}")
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class Recorder:
|
|
117
|
+
"""Records one rollout into a library.
|
|
118
|
+
|
|
119
|
+
Create it with ``Library.record``. Use it as a context manager: leaving
|
|
120
|
+
the ``with`` block normally calls :meth:`close` (status ``"complete"``);
|
|
121
|
+
leaving it on an exception calls :meth:`abort`, which keeps the run
|
|
122
|
+
recoverable (``rollout.json.partial``). Not thread-safe: call ``log``
|
|
123
|
+
from one thread.
|
|
124
|
+
|
|
125
|
+
Attributes:
|
|
126
|
+
name: Run name.
|
|
127
|
+
path: Run directory.
|
|
128
|
+
dt: Seconds per frame.
|
|
129
|
+
n_envs: Envs per frame.
|
|
130
|
+
block_frames: Frames per block.
|
|
131
|
+
"""
|
|
132
|
+
|
|
133
|
+
def __init__(
|
|
134
|
+
self,
|
|
135
|
+
root: pathlib.Path,
|
|
136
|
+
name: str,
|
|
137
|
+
*,
|
|
138
|
+
scene: core.Scene,
|
|
139
|
+
dt: float,
|
|
140
|
+
n_envs: int = 1,
|
|
141
|
+
env_origins: npt.ArrayLike | None = None,
|
|
142
|
+
source: Mapping[str, Any] | None = None,
|
|
143
|
+
tags: Sequence[str] = (),
|
|
144
|
+
meta: Mapping[str, Any] | None = None,
|
|
145
|
+
codec: str = "f32s",
|
|
146
|
+
block_frames: int = blockfile.DEFAULT_BLOCK_FRAMES,
|
|
147
|
+
overwrite: bool = False,
|
|
148
|
+
max_pending: int = 2,
|
|
149
|
+
encode_threads: int | None = None,
|
|
150
|
+
created: str | None = None,
|
|
151
|
+
) -> None:
|
|
152
|
+
"""Validates arguments. Nothing is written until the first ``log``.
|
|
153
|
+
|
|
154
|
+
Args:
|
|
155
|
+
root: The library root.
|
|
156
|
+
name: Run name, ``[A-Za-z0-9][A-Za-z0-9._-]{0,127}``.
|
|
157
|
+
scene: The static scene. Its body count fixes the pose shape.
|
|
158
|
+
dt: Seconds between frames (positive).
|
|
159
|
+
n_envs: Number of parallel envs.
|
|
160
|
+
env_origins: World offset of each env, shape ``[n_envs, 3]``, or
|
|
161
|
+
``None`` for zeros.
|
|
162
|
+
source: Provenance, such as ``{"simulator": "mujoco", ...}``.
|
|
163
|
+
tags: Record-time tags.
|
|
164
|
+
meta: Free-form JSON-serializable metadata.
|
|
165
|
+
codec: Codec of the pose stream: ``"f32s"`` (lossless) or
|
|
166
|
+
``"q16d"``. Other streams use their own codec (default
|
|
167
|
+
``"f32s"``).
|
|
168
|
+
block_frames: Frames per block.
|
|
169
|
+
overwrite: Replace an existing run of this name. The whole run
|
|
170
|
+
directory (including its annotations) is deleted when
|
|
171
|
+
recording starts.
|
|
172
|
+
created: When the run was recorded, ``YYYY-MM-DDTHH:MM:SSZ``
|
|
173
|
+
(UTC). Importers pass the time of the original recording;
|
|
174
|
+
``None`` means now.
|
|
175
|
+
max_pending: Full windows that may wait for the encoder before
|
|
176
|
+
``log`` blocks.
|
|
177
|
+
encode_threads: Threads the background writer fans a window's
|
|
178
|
+
per-env blocks out to (deflate releases the GIL, so this
|
|
179
|
+
scales). ``None`` picks 1 for fewer than 4 envs, else up
|
|
180
|
+
to 4.
|
|
181
|
+
|
|
182
|
+
Raises:
|
|
183
|
+
ValueError: If an argument is invalid.
|
|
184
|
+
FileExistsError: If the run exists and ``overwrite`` is false.
|
|
185
|
+
"""
|
|
186
|
+
self.name = manifest.validate_run_name(name)
|
|
187
|
+
self._root = pathlib.Path(root)
|
|
188
|
+
self.path = self._root / "runs" / name
|
|
189
|
+
if not dt > 0 or not math.isfinite(dt):
|
|
190
|
+
raise ValueError(f"dt must be positive and finite, got {dt!r}")
|
|
191
|
+
if n_envs < 1 or block_frames < 1 or max_pending < 1:
|
|
192
|
+
raise ValueError("n_envs, block_frames, max_pending must be >= 1")
|
|
193
|
+
codecs.codec_id(codec)
|
|
194
|
+
self.dt = float(dt)
|
|
195
|
+
self.n_envs = int(n_envs)
|
|
196
|
+
self.block_frames = int(block_frames)
|
|
197
|
+
self._scene = scene
|
|
198
|
+
self._codec = codec
|
|
199
|
+
self._overwrite = overwrite
|
|
200
|
+
self._max_pending = int(max_pending)
|
|
201
|
+
if encode_threads is None:
|
|
202
|
+
encode_threads = 1 if n_envs < 4 else min(4, os.cpu_count() or 1)
|
|
203
|
+
if encode_threads < 1:
|
|
204
|
+
raise ValueError("encode_threads must be at least 1")
|
|
205
|
+
self._encode_threads = int(encode_threads)
|
|
206
|
+
self._pool: concurrent.futures.ThreadPoolExecutor | None = None
|
|
207
|
+
self._source = dict(source or {})
|
|
208
|
+
self._tags = tuple(str(t) for t in tags)
|
|
209
|
+
self._meta = dict(meta or {})
|
|
210
|
+
json.dumps([self._source, self._meta], allow_nan=False) # early check
|
|
211
|
+
self._origins = self._parse_origins(env_origins)
|
|
212
|
+
self._check_free()
|
|
213
|
+
self._id = manifest.new_ulid()
|
|
214
|
+
self._created = created or manifest.utc_now()
|
|
215
|
+
self._specs: list[_Stream] = []
|
|
216
|
+
self._by_name: dict[str, _Stream] = {}
|
|
217
|
+
self._started = False
|
|
218
|
+
self._closed = False
|
|
219
|
+
self._error: BaseException | None = None
|
|
220
|
+
self._stats = {"stalls": 0, "stall_seconds": 0.0, "close_seconds": 0.0}
|
|
221
|
+
|
|
222
|
+
# -- setup --
|
|
223
|
+
|
|
224
|
+
def _parse_origins(
|
|
225
|
+
self, origins: npt.ArrayLike | None
|
|
226
|
+
) -> tuple[core.Vec3, ...] | None:
|
|
227
|
+
"""Validates env origins into a tuple of triples."""
|
|
228
|
+
if origins is None:
|
|
229
|
+
return None
|
|
230
|
+
arr = np.asarray(origins, dtype=np.float64)
|
|
231
|
+
if arr.shape != (self.n_envs, 3) or not np.isfinite(arr).all():
|
|
232
|
+
raise ValueError(
|
|
233
|
+
f"env_origins must be finite with shape ({self.n_envs}, 3)"
|
|
234
|
+
)
|
|
235
|
+
return tuple((float(x), float(y), float(z)) for x, y, z in arr)
|
|
236
|
+
|
|
237
|
+
def _check_free(self) -> None:
|
|
238
|
+
"""Raises if a finished or in-progress run of this name exists."""
|
|
239
|
+
if self._overwrite:
|
|
240
|
+
return
|
|
241
|
+
for fname in (manifest.MANIFEST_NAME, manifest.PARTIAL_NAME):
|
|
242
|
+
if (self.path / fname).exists():
|
|
243
|
+
raise FileExistsError(
|
|
244
|
+
f"run {self.name!r} exists in {self._root}; "
|
|
245
|
+
"pass overwrite=True to replace it"
|
|
246
|
+
)
|
|
247
|
+
|
|
248
|
+
def add_stream(
|
|
249
|
+
self,
|
|
250
|
+
name: str,
|
|
251
|
+
kind: core.StreamKind,
|
|
252
|
+
item_shape: Sequence[int] = (),
|
|
253
|
+
*,
|
|
254
|
+
labels: Sequence[str] | None = None,
|
|
255
|
+
units: str | None = None,
|
|
256
|
+
codec: str = "f32s",
|
|
257
|
+
scale: float | None = None,
|
|
258
|
+
) -> None:
|
|
259
|
+
"""Declares an extra stream. Call before the first ``log``.
|
|
260
|
+
|
|
261
|
+
Args:
|
|
262
|
+
name: Stream name; also the keyword given to ``log``.
|
|
263
|
+
kind: How viewers draw it (``"scalar"``, ``"vector"``,
|
|
264
|
+
``"arrows"``, ``"points"`` or ``"polyline"``).
|
|
265
|
+
item_shape: Shape of one env-frame item (``()`` for a scalar,
|
|
266
|
+
``(K,)`` for a vector, ``(K, 6)`` for arrows, ...).
|
|
267
|
+
labels: Optional component names.
|
|
268
|
+
units: Optional unit string.
|
|
269
|
+
codec: ``"f32s"`` (lossless, default) or ``"q16d"``.
|
|
270
|
+
scale: For ``arrows`` streams, metres drawn per unit of vector.
|
|
271
|
+
``None`` means viewers use 1.
|
|
272
|
+
|
|
273
|
+
Raises:
|
|
274
|
+
RecorderError: If recording already started.
|
|
275
|
+
ValueError: If the name is taken or invalid, or the kind and
|
|
276
|
+
shape do not fit.
|
|
277
|
+
"""
|
|
278
|
+
if self._started or self._closed:
|
|
279
|
+
raise RecorderError("streams must be declared before the first log")
|
|
280
|
+
try:
|
|
281
|
+
manifest.validate_run_name(name)
|
|
282
|
+
except ValueError as exc:
|
|
283
|
+
raise ValueError(f"invalid stream name {name!r}") from exc
|
|
284
|
+
if name == manifest.BODY_POSE or name in self._by_name:
|
|
285
|
+
raise ValueError(f"stream {name!r} is already declared")
|
|
286
|
+
if kind not in core.STREAM_KINDS or kind == "pose":
|
|
287
|
+
raise ValueError(f"extra stream kind must be non-pose: {kind!r}")
|
|
288
|
+
codecs.codec_id(codec)
|
|
289
|
+
shape = tuple(int(d) for d in item_shape)
|
|
290
|
+
if kind == "scalar" and shape != ():
|
|
291
|
+
raise ValueError("scalar streams have item_shape ()")
|
|
292
|
+
if kind == "vector" and len(shape) != 1:
|
|
293
|
+
raise ValueError("vector streams need item_shape (K,)")
|
|
294
|
+
if kind == "arrows" and (len(shape) != 2 or shape[1] != 6):
|
|
295
|
+
raise ValueError("arrows streams need item_shape (K, 6)")
|
|
296
|
+
if kind in ("points", "polyline") and (
|
|
297
|
+
len(shape) != 2 or shape[1] != 3
|
|
298
|
+
):
|
|
299
|
+
raise ValueError(f"{kind} streams need item_shape (K, 3)")
|
|
300
|
+
if 0 in shape or len(shape) > 4:
|
|
301
|
+
raise ValueError(f"invalid item_shape {shape}")
|
|
302
|
+
if labels is not None and len(labels) != (shape[-1] if shape else 1):
|
|
303
|
+
raise ValueError("labels must have one entry per component")
|
|
304
|
+
spec = _Stream(
|
|
305
|
+
len(self._specs) + 1,
|
|
306
|
+
name,
|
|
307
|
+
kind,
|
|
308
|
+
shape,
|
|
309
|
+
self.n_envs,
|
|
310
|
+
codec,
|
|
311
|
+
None if labels is None else tuple(labels),
|
|
312
|
+
units,
|
|
313
|
+
)
|
|
314
|
+
if scale is not None:
|
|
315
|
+
if kind != "arrows":
|
|
316
|
+
raise ValueError("scale applies to arrows streams only")
|
|
317
|
+
spec.info.scale = float(scale)
|
|
318
|
+
self._register(spec)
|
|
319
|
+
|
|
320
|
+
def _register(self, spec: _Stream) -> None:
|
|
321
|
+
"""Adds a stream description."""
|
|
322
|
+
self._specs.append(spec)
|
|
323
|
+
self._by_name[spec.name] = spec
|
|
324
|
+
|
|
325
|
+
def _start(
|
|
326
|
+
self,
|
|
327
|
+
pose_shape: tuple[int, ...],
|
|
328
|
+
extra_shapes: Mapping[str, tuple[int, ...]],
|
|
329
|
+
) -> None:
|
|
330
|
+
"""Validates the first frame, then creates files and the worker.
|
|
331
|
+
|
|
332
|
+
Args:
|
|
333
|
+
pose_shape: Shape of one frame of poses, ``[E, B, 7]`` or
|
|
334
|
+
``[B, 7]`` when ``n_envs == 1``.
|
|
335
|
+
extra_shapes: Per-frame shape of every stream given to the
|
|
336
|
+
first ``log``. Undeclared scalars are auto-declared.
|
|
337
|
+
|
|
338
|
+
Raises:
|
|
339
|
+
ValueError: If shapes disagree with the scene or declarations.
|
|
340
|
+
"""
|
|
341
|
+
n_bodies = self._scene.n_bodies
|
|
342
|
+
want = (self.n_envs, n_bodies, core.POSE_DIM)
|
|
343
|
+
if pose_shape not in (want, want[1:]) or (
|
|
344
|
+
pose_shape == want[1:] and self.n_envs != 1
|
|
345
|
+
):
|
|
346
|
+
raise _frame_error("poses", want, pose_shape)
|
|
347
|
+
pose = _Stream(
|
|
348
|
+
0,
|
|
349
|
+
manifest.BODY_POSE,
|
|
350
|
+
"pose",
|
|
351
|
+
(n_bodies, core.POSE_DIM),
|
|
352
|
+
self.n_envs,
|
|
353
|
+
self._codec,
|
|
354
|
+
)
|
|
355
|
+
specs = [pose, *self._specs]
|
|
356
|
+
for name, shape in extra_shapes.items():
|
|
357
|
+
if name in self._by_name:
|
|
358
|
+
continue
|
|
359
|
+
if shape not in ((), (self.n_envs,)):
|
|
360
|
+
raise ValueError(
|
|
361
|
+
f"stream {name!r} was not declared; call add_stream() "
|
|
362
|
+
"before the first log (only scalars can be inferred)"
|
|
363
|
+
)
|
|
364
|
+
spec = _Stream(len(specs), name, "scalar", (), self.n_envs, "f32s")
|
|
365
|
+
specs.append(spec)
|
|
366
|
+
self._by_name[name] = spec
|
|
367
|
+
self._specs = specs
|
|
368
|
+
self._open_files()
|
|
369
|
+
|
|
370
|
+
def _open_files(self) -> None:
|
|
371
|
+
"""Writes the scene and partial manifest, opens writers and worker."""
|
|
372
|
+
self._check_free()
|
|
373
|
+
if self._overwrite and self.path.exists():
|
|
374
|
+
shutil.rmtree(self.path)
|
|
375
|
+
self.path.mkdir(parents=True, exist_ok=True)
|
|
376
|
+
ref = scene_io.put_scene(cas.ContentStore(self._root), self._scene)
|
|
377
|
+
self._manifest = manifest.RolloutManifest(
|
|
378
|
+
id=self._id,
|
|
379
|
+
name=self.name,
|
|
380
|
+
created=self._created,
|
|
381
|
+
dt=self.dt,
|
|
382
|
+
n_frames=0,
|
|
383
|
+
n_envs=self.n_envs,
|
|
384
|
+
n_bodies=self._scene.n_bodies,
|
|
385
|
+
scene=ref,
|
|
386
|
+
streams={s.name: s.info for s in self._specs},
|
|
387
|
+
status="recording",
|
|
388
|
+
env_origins=self._origins,
|
|
389
|
+
source=self._source,
|
|
390
|
+
tags=self._tags,
|
|
391
|
+
meta=self._meta,
|
|
392
|
+
)
|
|
393
|
+
self._writers = [
|
|
394
|
+
_StreamWriter(
|
|
395
|
+
self.path / s.info.file,
|
|
396
|
+
item_shape=s.item_shape,
|
|
397
|
+
n_envs=self.n_envs,
|
|
398
|
+
kind=s.kind,
|
|
399
|
+
codec=s.codec,
|
|
400
|
+
block_frames=self.block_frames,
|
|
401
|
+
)
|
|
402
|
+
for s in self._specs
|
|
403
|
+
]
|
|
404
|
+
manifest.write_manifest(self.path, self._manifest, partial=True)
|
|
405
|
+
self._n_extra = len(self._specs) - 1
|
|
406
|
+
self._pose_shape = self._specs[0].frame_shape
|
|
407
|
+
n_windows = self._max_pending + 1
|
|
408
|
+
windows = [
|
|
409
|
+
_Window(self._specs, self.block_frames, self.n_envs)
|
|
410
|
+
for _ in range(n_windows)
|
|
411
|
+
]
|
|
412
|
+
self._win = windows[0]
|
|
413
|
+
self._free_q: queue.Queue[_Window] = queue.Queue()
|
|
414
|
+
for w in windows[1:]:
|
|
415
|
+
self._free_q.put(w)
|
|
416
|
+
self._work_q: queue.Queue[_Window | None] = queue.Queue()
|
|
417
|
+
self._fill = 0
|
|
418
|
+
self._next_t0 = 0
|
|
419
|
+
self._prev_quat: npt.NDArray[np.float32] | None = None
|
|
420
|
+
if self._encode_threads > 1 and self.n_envs > 1:
|
|
421
|
+
self._pool = concurrent.futures.ThreadPoolExecutor(
|
|
422
|
+
self._encode_threads, thread_name_prefix="simscope-encode"
|
|
423
|
+
)
|
|
424
|
+
self._thread = threading.Thread(
|
|
425
|
+
target=self._worker, name="simscope-recorder", daemon=True
|
|
426
|
+
)
|
|
427
|
+
self._thread.start()
|
|
428
|
+
self._started = True
|
|
429
|
+
|
|
430
|
+
# -- logging --
|
|
431
|
+
|
|
432
|
+
@property
|
|
433
|
+
def n_frames(self) -> int:
|
|
434
|
+
"""Frames logged so far."""
|
|
435
|
+
return self._next_t0 + self._fill if self._started else 0
|
|
436
|
+
|
|
437
|
+
@property
|
|
438
|
+
def stats(self) -> dict[str, float]:
|
|
439
|
+
"""Backpressure counters.
|
|
440
|
+
|
|
441
|
+
``stalls`` counts the times ``log`` had to wait for the encoder,
|
|
442
|
+
``stall_seconds`` the total wait, and ``close_seconds`` how long the
|
|
443
|
+
last ``close`` took.
|
|
444
|
+
"""
|
|
445
|
+
return dict(self._stats)
|
|
446
|
+
|
|
447
|
+
def log(self, poses: npt.ArrayLike, **streams: npt.ArrayLike) -> None:
|
|
448
|
+
"""Records one frame.
|
|
449
|
+
|
|
450
|
+
Args:
|
|
451
|
+
poses: Body poses ``[E, B, 7]`` (or ``[B, 7]`` when there is one
|
|
452
|
+
env), xyzw quaternions.
|
|
453
|
+
**streams: One value per declared stream, each ``[E, *item]``.
|
|
454
|
+
Scalar streams also accept a Python float (used for every
|
|
455
|
+
env) and, when there is one env, no leading axis. A name
|
|
456
|
+
that was never declared is declared as a scalar on the
|
|
457
|
+
first frame if its value is a float or has shape ``[E]``.
|
|
458
|
+
|
|
459
|
+
Raises:
|
|
460
|
+
RecorderError: If the recorder is closed.
|
|
461
|
+
ValueError: If a shape is wrong or a stream is missing or
|
|
462
|
+
undeclared. The frame is not recorded.
|
|
463
|
+
Exception: A failure of the background writer, re-raised.
|
|
464
|
+
"""
|
|
465
|
+
if self._error is not None or self._closed:
|
|
466
|
+
self._raise_pending()
|
|
467
|
+
if not self._started:
|
|
468
|
+
self._start(
|
|
469
|
+
np.shape(poses), {k: np.shape(v) for k, v in streams.items()}
|
|
470
|
+
)
|
|
471
|
+
elif len(streams) != self._n_extra:
|
|
472
|
+
self._check_names(streams)
|
|
473
|
+
i = self._fill
|
|
474
|
+
frames = self._win.frames
|
|
475
|
+
by_name = self._by_name
|
|
476
|
+
shape = getattr(poses, "shape", None)
|
|
477
|
+
if shape == self._pose_shape:
|
|
478
|
+
np.copyto(frames[0][i], poses)
|
|
479
|
+
else:
|
|
480
|
+
frames[0][i] = self._coerce(self._specs[0], poses)
|
|
481
|
+
for name, value in streams.items():
|
|
482
|
+
spec = by_name.get(name)
|
|
483
|
+
if spec is None:
|
|
484
|
+
self._check_names(streams)
|
|
485
|
+
raise AssertionError("unreachable") # pragma: no cover
|
|
486
|
+
if getattr(value, "shape", None) == spec.frame_shape:
|
|
487
|
+
np.copyto(frames[spec.index][i], value)
|
|
488
|
+
else:
|
|
489
|
+
frames[spec.index][i] = self._coerce(spec, value)
|
|
490
|
+
self._fill = i + 1
|
|
491
|
+
if self._fill == self.block_frames:
|
|
492
|
+
self._submit()
|
|
493
|
+
|
|
494
|
+
def _check_names(self, streams: Mapping[str, Any]) -> None:
|
|
495
|
+
"""Raises a clear error for missing or undeclared streams."""
|
|
496
|
+
extra = sorted(set(streams) - self._by_name.keys())
|
|
497
|
+
missing = sorted(self._by_name.keys() - set(streams))
|
|
498
|
+
if extra:
|
|
499
|
+
raise ValueError(
|
|
500
|
+
f"undeclared streams {extra}: call add_stream() before the "
|
|
501
|
+
"first log"
|
|
502
|
+
)
|
|
503
|
+
if missing:
|
|
504
|
+
raise ValueError(f"missing streams {missing} in this frame")
|
|
505
|
+
|
|
506
|
+
def _coerce(self, spec: _Stream, value: Any) -> npt.NDArray[np.float32]:
|
|
507
|
+
"""Converts one frame value to an array that fits ``spec``.
|
|
508
|
+
|
|
509
|
+
Returns:
|
|
510
|
+
An array broadcastable to ``spec.frame_shape``: exactly that
|
|
511
|
+
shape, that shape with the env axis added, or (scalars only) a
|
|
512
|
+
0-d array.
|
|
513
|
+
"""
|
|
514
|
+
a = np.asarray(value, dtype=np.float32)
|
|
515
|
+
if a.shape == spec.frame_shape:
|
|
516
|
+
return a
|
|
517
|
+
if self.n_envs == 1 and a.shape == spec.frame_shape[1:]:
|
|
518
|
+
return a[None]
|
|
519
|
+
if spec.kind == "scalar" and a.ndim == 0:
|
|
520
|
+
return a
|
|
521
|
+
raise _frame_error(spec.name, spec.frame_shape, a.shape)
|
|
522
|
+
|
|
523
|
+
def log_frames(
|
|
524
|
+
self, poses: npt.ArrayLike, **streams: npt.ArrayLike
|
|
525
|
+
) -> None:
|
|
526
|
+
"""Records ``n`` frames at once, without a Python loop over frames.
|
|
527
|
+
|
|
528
|
+
Args:
|
|
529
|
+
poses: Body poses ``[n, E, B, 7]`` (or ``[n, B, 7]`` when there
|
|
530
|
+
is one env), for example ``PoseBuffer.drain()`` output.
|
|
531
|
+
**streams: One array per declared stream, each
|
|
532
|
+
``[n, E, *item]`` (scalars may be ``[n]`` when there is one
|
|
533
|
+
env).
|
|
534
|
+
|
|
535
|
+
Raises:
|
|
536
|
+
RecorderError: If the recorder is closed.
|
|
537
|
+
ValueError: If a shape is wrong or a stream is missing or
|
|
538
|
+
undeclared. Nothing is recorded.
|
|
539
|
+
Exception: A failure of the background writer, re-raised.
|
|
540
|
+
"""
|
|
541
|
+
if self._error is not None or self._closed:
|
|
542
|
+
self._raise_pending()
|
|
543
|
+
arrays = {
|
|
544
|
+
k: np.asarray(v, dtype=np.float32) for k, v in streams.items()
|
|
545
|
+
}
|
|
546
|
+
pose_arr = np.asarray(poses, dtype=np.float32)
|
|
547
|
+
if pose_arr.ndim < 1:
|
|
548
|
+
raise ValueError("poses must have a leading frame axis")
|
|
549
|
+
n = pose_arr.shape[0]
|
|
550
|
+
if not self._started:
|
|
551
|
+
self._start(
|
|
552
|
+
pose_arr.shape[1:], {k: a.shape[1:] for k, a in arrays.items()}
|
|
553
|
+
)
|
|
554
|
+
elif len(arrays) != self._n_extra:
|
|
555
|
+
self._check_names(arrays)
|
|
556
|
+
batch = [self._coerce_batch(self._specs[0], pose_arr, n)]
|
|
557
|
+
for name, a in arrays.items():
|
|
558
|
+
if name not in self._by_name:
|
|
559
|
+
self._check_names(arrays)
|
|
560
|
+
spec = self._by_name[name]
|
|
561
|
+
batch.append((spec, self._coerce_batch(spec, a, n)[1]))
|
|
562
|
+
pos = 0
|
|
563
|
+
while pos < n:
|
|
564
|
+
take = min(n - pos, self.block_frames - self._fill)
|
|
565
|
+
lo, hi = self._fill, self._fill + take
|
|
566
|
+
frames = self._win.frames
|
|
567
|
+
for spec, a in batch:
|
|
568
|
+
np.copyto(frames[spec.index][lo:hi], a[pos : pos + take])
|
|
569
|
+
self._fill = hi
|
|
570
|
+
pos += take
|
|
571
|
+
if self._fill == self.block_frames:
|
|
572
|
+
self._submit()
|
|
573
|
+
|
|
574
|
+
def _coerce_batch(
|
|
575
|
+
self, spec: _Stream, a: npt.NDArray[np.float32], n: int
|
|
576
|
+
) -> tuple[_Stream, npt.NDArray[np.float32]]:
|
|
577
|
+
"""Checks a ``[n, E, *item]`` array and adds the env axis if needed."""
|
|
578
|
+
want = (n, *spec.frame_shape)
|
|
579
|
+
if a.shape == want:
|
|
580
|
+
return spec, a
|
|
581
|
+
if self.n_envs == 1 and a.shape == (n, *spec.frame_shape[1:]):
|
|
582
|
+
return spec, a[:, None]
|
|
583
|
+
raise _frame_error(spec.name, want, a.shape)
|
|
584
|
+
|
|
585
|
+
# -- window hand-off --
|
|
586
|
+
|
|
587
|
+
def _submit(self) -> None:
|
|
588
|
+
"""Queues the full current window and switches to a free one."""
|
|
589
|
+
win = self._win
|
|
590
|
+
win.n = self._fill
|
|
591
|
+
win.t0 = self._next_t0
|
|
592
|
+
self._next_t0 += self._fill
|
|
593
|
+
self._work_q.put(win)
|
|
594
|
+
try:
|
|
595
|
+
self._win = self._free_q.get_nowait()
|
|
596
|
+
except queue.Empty:
|
|
597
|
+
self._win = self._wait_for_window()
|
|
598
|
+
self._fill = 0
|
|
599
|
+
|
|
600
|
+
def _wait_for_window(self) -> _Window:
|
|
601
|
+
"""Blocks until the encoder frees a window (backpressure)."""
|
|
602
|
+
start = time.perf_counter()
|
|
603
|
+
while True:
|
|
604
|
+
try:
|
|
605
|
+
win = self._free_q.get(timeout=_POLL_S)
|
|
606
|
+
break
|
|
607
|
+
except queue.Empty:
|
|
608
|
+
if not self._thread.is_alive():
|
|
609
|
+
raise RecorderError("recorder thread died") from None
|
|
610
|
+
self._stats["stalls"] += 1
|
|
611
|
+
self._stats["stall_seconds"] += time.perf_counter() - start
|
|
612
|
+
return win
|
|
613
|
+
|
|
614
|
+
def _worker(self) -> None:
|
|
615
|
+
"""Encodes and writes queued windows until it gets ``None``."""
|
|
616
|
+
while True:
|
|
617
|
+
win = self._work_q.get()
|
|
618
|
+
try:
|
|
619
|
+
if win is None:
|
|
620
|
+
return
|
|
621
|
+
if self._error is None:
|
|
622
|
+
self._write_window(win)
|
|
623
|
+
except BaseException as exc:
|
|
624
|
+
logger.error("recording worker failed: %r", exc)
|
|
625
|
+
self._error = exc
|
|
626
|
+
finally:
|
|
627
|
+
if win is not None:
|
|
628
|
+
self._free_q.put(win)
|
|
629
|
+
self._work_q.task_done()
|
|
630
|
+
|
|
631
|
+
def _write_window(self, win: _Window) -> None:
|
|
632
|
+
"""Encodes a window for every stream and appends it."""
|
|
633
|
+
n = win.n
|
|
634
|
+
for spec, arr, writer in zip(
|
|
635
|
+
self._specs, win.arrays, self._writers, strict=True
|
|
636
|
+
):
|
|
637
|
+
data = arr[:n]
|
|
638
|
+
if spec.kind == "pose":
|
|
639
|
+
self._fix_signs(data)
|
|
640
|
+
blocks = self._encode(data, spec.codec)
|
|
641
|
+
writer.write_window(win.t0, blocks, n)
|
|
642
|
+
# Keep the partial manifest's frame count current, so the index (and
|
|
643
|
+
# any `simscope serve` watching the folder) sees live progress.
|
|
644
|
+
self._manifest.n_frames = win.t0 + n
|
|
645
|
+
manifest.write_manifest(self.path, self._manifest, partial=True)
|
|
646
|
+
|
|
647
|
+
def _encode(
|
|
648
|
+
self, data: npt.NDArray[np.float32], codec: str
|
|
649
|
+
) -> Sequence[blockfile.EncodedBlock]:
|
|
650
|
+
"""Encodes a window, one block per env, using the thread pool."""
|
|
651
|
+
if self._pool is None:
|
|
652
|
+
return blockfile.encode_window(data, codec)
|
|
653
|
+
n, n_envs, k = data.shape
|
|
654
|
+
|
|
655
|
+
def encode_env(env: int) -> blockfile.EncodedBlock:
|
|
656
|
+
cid, payload = codecs.encode_block_auto(data[:, env, :], codec)
|
|
657
|
+
return blockfile.EncodedBlock(
|
|
658
|
+
env, n, cid, codecs.payload_ulen(cid, n, k), payload
|
|
659
|
+
)
|
|
660
|
+
|
|
661
|
+
return list(self._pool.map(encode_env, range(n_envs)))
|
|
662
|
+
|
|
663
|
+
def _fix_signs(self, data: npt.NDArray[np.float32]) -> None:
|
|
664
|
+
"""Applies quaternion sign continuity to a pose window in place."""
|
|
665
|
+
n = data.shape[0]
|
|
666
|
+
n_bodies = self._scene.n_bodies
|
|
667
|
+
q = data.reshape(n, self.n_envs, n_bodies, core.POSE_DIM)[..., 3:]
|
|
668
|
+
prev = self._prev_quat
|
|
669
|
+
if prev is None:
|
|
670
|
+
prev_full = None
|
|
671
|
+
self._prev_quat = prev = np.empty(q.shape[1:], np.float32)
|
|
672
|
+
else:
|
|
673
|
+
prev_full = prev
|
|
674
|
+
step = max(1, _SIGN_CHUNK // (n * n_bodies * 4))
|
|
675
|
+
for lo in range(0, self.n_envs, step):
|
|
676
|
+
sl = slice(lo, lo + step)
|
|
677
|
+
fixed = transforms.enforce_sign_continuity(
|
|
678
|
+
q[:, sl], None if prev_full is None else prev_full[sl]
|
|
679
|
+
)
|
|
680
|
+
q[:, sl] = fixed
|
|
681
|
+
prev[sl] = fixed[-1]
|
|
682
|
+
|
|
683
|
+
# -- lifecycle --
|
|
684
|
+
|
|
685
|
+
def _raise_pending(self) -> None:
|
|
686
|
+
"""Raises the worker's error, or a closed-recorder error."""
|
|
687
|
+
if self._error is not None:
|
|
688
|
+
raise self._error
|
|
689
|
+
raise RecorderError(f"recorder for {self.name!r} is closed")
|
|
690
|
+
|
|
691
|
+
def flush(self) -> None:
|
|
692
|
+
"""Waits until every full window is on disk.
|
|
693
|
+
|
|
694
|
+
Raises:
|
|
695
|
+
Exception: A failure of the background writer, re-raised.
|
|
696
|
+
"""
|
|
697
|
+
if not self._started or self._closed:
|
|
698
|
+
return
|
|
699
|
+
self._work_q.join()
|
|
700
|
+
if self._error is not None:
|
|
701
|
+
raise self._error
|
|
702
|
+
|
|
703
|
+
def _drain(self) -> None:
|
|
704
|
+
"""Queues the partial window, then stops and joins the worker."""
|
|
705
|
+
if self._fill:
|
|
706
|
+
win = self._win
|
|
707
|
+
win.n = self._fill
|
|
708
|
+
win.t0 = self._next_t0
|
|
709
|
+
self._next_t0 += self._fill
|
|
710
|
+
self._fill = 0
|
|
711
|
+
self._work_q.put(win)
|
|
712
|
+
self._work_q.put(None)
|
|
713
|
+
self._thread.join()
|
|
714
|
+
if self._pool is not None:
|
|
715
|
+
self._pool.shutdown()
|
|
716
|
+
|
|
717
|
+
def close(self) -> None:
|
|
718
|
+
"""Finishes the run: flushes, writes directories and the manifest.
|
|
719
|
+
|
|
720
|
+
Does nothing if no frame was logged (no run is created). Safe to
|
|
721
|
+
call twice.
|
|
722
|
+
|
|
723
|
+
Raises:
|
|
724
|
+
Exception: A failure of the background writer. The files stay
|
|
725
|
+
recoverable (``rollout.json.partial`` remains).
|
|
726
|
+
"""
|
|
727
|
+
if self._closed:
|
|
728
|
+
return
|
|
729
|
+
self._closed = True
|
|
730
|
+
if not self._started:
|
|
731
|
+
return
|
|
732
|
+
start = time.perf_counter()
|
|
733
|
+
try:
|
|
734
|
+
self._drain()
|
|
735
|
+
if self._error is not None:
|
|
736
|
+
raise self._error
|
|
737
|
+
for writer in self._writers:
|
|
738
|
+
writer.finalize()
|
|
739
|
+
self._manifest.n_frames = self._next_t0
|
|
740
|
+
self._manifest.status = "complete"
|
|
741
|
+
manifest.write_manifest(self.path, self._manifest, partial=False)
|
|
742
|
+
finally:
|
|
743
|
+
for writer in self._writers:
|
|
744
|
+
writer.close()
|
|
745
|
+
self._stats["close_seconds"] = time.perf_counter() - start
|
|
746
|
+
|
|
747
|
+
def abort(self) -> None:
|
|
748
|
+
"""Stops without finishing the run, keeping it recoverable.
|
|
749
|
+
|
|
750
|
+
Flushes what it can and leaves ``rollout.json.partial`` in place.
|
|
751
|
+
Never raises; use ``Library.recover`` to complete the run.
|
|
752
|
+
"""
|
|
753
|
+
if self._closed:
|
|
754
|
+
return
|
|
755
|
+
self._closed = True
|
|
756
|
+
if not self._started:
|
|
757
|
+
return
|
|
758
|
+
try:
|
|
759
|
+
self._drain()
|
|
760
|
+
finally:
|
|
761
|
+
for writer in self._writers:
|
|
762
|
+
writer.close()
|
|
763
|
+
if self._error is not None:
|
|
764
|
+
logger.error(
|
|
765
|
+
"run %r aborted after a writer error: %r",
|
|
766
|
+
self.name,
|
|
767
|
+
self._error,
|
|
768
|
+
)
|
|
769
|
+
|
|
770
|
+
def __enter__(self) -> "Recorder":
|
|
771
|
+
"""Returns the recorder."""
|
|
772
|
+
return self
|
|
773
|
+
|
|
774
|
+
def __exit__(
|
|
775
|
+
self,
|
|
776
|
+
exc_type: type[BaseException] | None,
|
|
777
|
+
exc: BaseException | None,
|
|
778
|
+
tb: TracebackType | None,
|
|
779
|
+
) -> None:
|
|
780
|
+
"""Closes on success; aborts (keeping the run recoverable) on error."""
|
|
781
|
+
if exc_type is None:
|
|
782
|
+
self.close()
|
|
783
|
+
else:
|
|
784
|
+
self.abort()
|