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.
Files changed (45) hide show
  1. simscope/__init__.py +6 -0
  2. simscope/__main__.py +8 -0
  3. simscope/_assets/simscope-app.css +2 -0
  4. simscope/_assets/simscope-app.js +4311 -0
  5. simscope/_assets/simscope-player.js +4325 -0
  6. simscope/_assets/simscope-web.LICENSES.txt +407 -0
  7. simscope/_icon.py +22 -0
  8. simscope/_mjviser.py +203 -0
  9. simscope/annotations.py +1132 -0
  10. simscope/cli.py +482 -0
  11. simscope/core.py +257 -0
  12. simscope/derived.py +697 -0
  13. simscope/export.py +799 -0
  14. simscope/highlights.py +947 -0
  15. simscope/importers.py +874 -0
  16. simscope/index.py +579 -0
  17. simscope/io/__init__.py +45 -0
  18. simscope/io/blockfile.py +938 -0
  19. simscope/io/cas.py +294 -0
  20. simscope/io/codecs.py +566 -0
  21. simscope/io/errors.py +9 -0
  22. simscope/io/manifest.py +358 -0
  23. simscope/io/pack.py +563 -0
  24. simscope/io/scene.py +239 -0
  25. simscope/isaaclab.py +1460 -0
  26. simscope/library.py +705 -0
  27. simscope/mujoco.py +578 -0
  28. simscope/py.typed +0 -0
  29. simscope/recorder.py +784 -0
  30. simscope/server/__init__.py +9 -0
  31. simscope/server/app.py +149 -0
  32. simscope/server/blocks.py +191 -0
  33. simscope/server/jobs.py +166 -0
  34. simscope/server/routes.py +707 -0
  35. simscope/server/security.py +218 -0
  36. simscope/server/state.py +751 -0
  37. simscope/server/static.py +84 -0
  38. simscope/transforms.py +147 -0
  39. simscope-0.1.1.dist-info/METADATA +132 -0
  40. simscope-0.1.1.dist-info/RECORD +45 -0
  41. simscope-0.1.1.dist-info/WHEEL +4 -0
  42. simscope-0.1.1.dist-info/entry_points.txt +3 -0
  43. simscope-0.1.1.dist-info/licenses/LICENSE.md +201 -0
  44. simscope-0.1.1.dist-info/licenses/THIRD_PARTY_NOTICES.md +267 -0
  45. 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()