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/derived.py ADDED
@@ -0,0 +1,697 @@
1
+ """Derived data: files computed once per run and cached (contracts 5).
2
+
3
+ Everything here is "just more files in the existing formats", so the browser
4
+ reads it through the same path-addressed ``Source`` as the recording itself:
5
+
6
+ * ``root_pose.blk``: the followed body of every env, ``[E, 1, 7]`` ``q16d``
7
+ blocks in the normal block format, for runs with more than
8
+ :data:`CROWD_ENVS` envs. It is what the crowd tier draws.
9
+ * ``summaries.json``: one number per env for each of a few columns, which
10
+ sorts the env picker.
11
+ * ``envelopes/<stream>.json``: the p5, p50 and p95 of a scalar or vector
12
+ stream across envs, per frame.
13
+ * ``highlights.json``: the contact-force and acceleration peaks of
14
+ :mod:`simscope.highlights`; ``n_highlights`` counts them per env.
15
+
16
+ Files live in ``.simscope/derived/<run id>/`` next to a ``stamp.json`` that
17
+ holds the digest of the run's manifest: when the manifest changes, every
18
+ derived file of the run is discarded. Nothing is computed for a run that is
19
+ still recording.
20
+
21
+ Every pass reads one window at a time, and only the components it needs: a
22
+ block stores its components as separate byte planes, so the root body is
23
+ inflated from a 140-float item and unshuffled as 7 floats. A 4,096-env,
24
+ 1,000-frame, 20-body run never holds more than one window of one body.
25
+ """
26
+
27
+ import concurrent.futures
28
+ import hashlib
29
+ import json
30
+ import logging
31
+ import math
32
+ import os
33
+ import pathlib
34
+ import shutil
35
+ import tempfile
36
+ import threading
37
+ from collections.abc import Callable
38
+ from typing import Any
39
+
40
+ import numpy as np
41
+ import numpy.typing as npt
42
+
43
+ from simscope import core, highlights, library
44
+ from simscope.io import blockfile, cas, codecs, manifest
45
+
46
+ logger = logging.getLogger(__name__)
47
+
48
+ VERSION = 4
49
+ """Bump to discard every cached derived file.
50
+
51
+ 2: highlights are ``simscope-highlights/2`` (viewer v3.1).
52
+ 3: detector simscope/2.1, a fall is a body that stayed down.
53
+ 4: detector simscope/3, only contact and acceleration kinds."""
54
+ CROWD_ENVS = 64
55
+ """Runs with more envs than this get a ``root_pose.blk``."""
56
+ ROOT_POSE = "root_pose.blk"
57
+ SUMMARIES = "summaries.json"
58
+ HIGHLIGHTS = "highlights.json"
59
+ ENVELOPES = "envelopes/"
60
+ STAMP = "stamp.json"
61
+ _ROOT_STATS = "_root_stats.json"
62
+ _ENVELOPE_PERCENTILES = (5, 50, 95)
63
+ _WINDOW_BYTES = 64 << 20
64
+ """Most decoded bytes of one envelope window read at a time."""
65
+ _CONTACT_BYTES = 32 << 20
66
+ """Most decoded bytes of one contacts chunk read at a time."""
67
+ _SIG_DIGITS = 7
68
+
69
+ _locks: dict[pathlib.Path, threading.Lock] = {}
70
+ _locks_guard = threading.Lock()
71
+
72
+
73
+ # -- locations and validity --
74
+
75
+
76
+ def cache_root(lib_root: pathlib.Path) -> pathlib.Path:
77
+ """Returns the folder that holds every run's derived cache.
78
+
79
+ That is ``<library>/.simscope/derived``, or a per-library folder under
80
+ the system temp directory when the library is read-only.
81
+
82
+ Args:
83
+ lib_root: The library folder.
84
+
85
+ Returns:
86
+ The cache folder (not necessarily created yet).
87
+ """
88
+ home = pathlib.Path(lib_root) / ".simscope"
89
+ parent = home if home.exists() else pathlib.Path(lib_root)
90
+ if os.access(parent, os.W_OK):
91
+ return home / "derived"
92
+ digest = hashlib.sha1(
93
+ os.path.realpath(lib_root).encode(), usedforsecurity=False
94
+ ).hexdigest()[:12]
95
+ return pathlib.Path(tempfile.gettempdir()) / "simscope-derived" / digest
96
+
97
+
98
+ def cache_dir(lib_root: pathlib.Path, rollout: library.Rollout) -> pathlib.Path:
99
+ """Returns the cache folder of one run, keyed by the run's ULID.
100
+
101
+ Args:
102
+ lib_root: The library folder.
103
+ rollout: The run.
104
+
105
+ Returns:
106
+ ``.simscope/derived/<run id>`` (created on first write).
107
+ """
108
+ return cache_root(lib_root) / rollout.manifest.id
109
+
110
+
111
+ def manifest_digest(rollout: library.Rollout) -> str:
112
+ """Hashes the run's manifest file, which is what invalidates the cache.
113
+
114
+ Args:
115
+ rollout: The run.
116
+
117
+ Returns:
118
+ A hex digest, or ``"missing"`` if the manifest cannot be read.
119
+ """
120
+ for name in (manifest.MANIFEST_NAME, manifest.PARTIAL_NAME):
121
+ try:
122
+ data = (rollout.path / name).read_bytes()
123
+ except OSError:
124
+ continue
125
+ return hashlib.sha1(data, usedforsecurity=False).hexdigest()
126
+ return "missing"
127
+
128
+
129
+ def _lock_for(cache: pathlib.Path) -> threading.Lock:
130
+ """Returns the lock that serializes work on one cache folder."""
131
+ with _locks_guard:
132
+ return _locks.setdefault(cache, threading.Lock())
133
+
134
+
135
+ def _stamp_ok(rollout: library.Rollout, cache: pathlib.Path) -> bool:
136
+ """Tells whether the cache folder was made for this manifest."""
137
+ try:
138
+ stamp = json.loads((cache / STAMP).read_bytes())
139
+ except (OSError, ValueError):
140
+ return False
141
+ return stamp == {"version": VERSION, "manifest": manifest_digest(rollout)}
142
+
143
+
144
+ def fingerprint(rollout: library.Rollout) -> tuple[str, str]:
145
+ """Returns what the cache of a run is valid for, to carry it over.
146
+
147
+ Args:
148
+ rollout: The run.
149
+
150
+ Returns:
151
+ The manifest digest and the highlights cache key.
152
+ """
153
+ return manifest_digest(rollout), highlights.cache_key(rollout)
154
+
155
+
156
+ def carry_over(
157
+ lib_root: pathlib.Path,
158
+ rollout: library.Rollout,
159
+ before: tuple[str, str],
160
+ ) -> None:
161
+ """Keeps a run's cache valid after its manifest was rewritten.
162
+
163
+ A rename rewrites ``rollout.json`` (only its ``name`` changes), which
164
+ would make the stamp and the highlights key look stale and throw the
165
+ cache away. If the cache was valid for ``before`` (:func:`fingerprint`
166
+ taken before the rewrite), it is stamped for the run as it is now.
167
+
168
+ Args:
169
+ lib_root: The library folder.
170
+ rollout: The run, opened after the rewrite.
171
+ before: The fingerprint from before the rewrite.
172
+ """
173
+ digest, key = before
174
+ cache = cache_dir(lib_root, rollout)
175
+ with _lock_for(cache):
176
+ try:
177
+ stamp = json.loads((cache / STAMP).read_bytes())
178
+ except (OSError, ValueError):
179
+ stamp = None
180
+ if stamp == {"version": VERSION, "manifest": digest}:
181
+ now = {"version": VERSION, "manifest": manifest_digest(rollout)}
182
+ cas.atomic_write(cache / STAMP, _dumps(now))
183
+ key_path = cache / highlights.KEY_NAME
184
+ try:
185
+ same = key_path.read_text("utf-8") == key
186
+ except (OSError, ValueError):
187
+ same = False
188
+ if same:
189
+ cas.atomic_write(
190
+ key_path, highlights.cache_key(rollout).encode("utf-8")
191
+ )
192
+
193
+
194
+ def _reset(rollout: library.Rollout, cache: pathlib.Path) -> None:
195
+ """Discards the cache folder's files and stamps it for this manifest."""
196
+ if cache.exists():
197
+ for child in cache.iterdir():
198
+ if child.is_dir():
199
+ shutil.rmtree(child, ignore_errors=True)
200
+ else:
201
+ child.unlink(missing_ok=True)
202
+ stamp = {"version": VERSION, "manifest": manifest_digest(rollout)}
203
+ cas.atomic_write(cache / STAMP, _dumps(stamp))
204
+
205
+
206
+ def _dumps(obj: Any) -> bytes:
207
+ """Serializes compact JSON."""
208
+ return json.dumps(obj, separators=(",", ":"), allow_nan=False).encode()
209
+
210
+
211
+ def envelope_name(stream: str) -> str:
212
+ """Returns the derived file name of a stream's envelope."""
213
+ return f"{ENVELOPES}{stream}.json"
214
+
215
+
216
+ def applicable(rollout: library.Rollout, what: str) -> bool:
217
+ """Tells whether a derived file exists for this run at all.
218
+
219
+ Args:
220
+ rollout: The run.
221
+ what: A path relative to ``derived/<run>/``: ``root_pose.blk``,
222
+ ``summaries.json``, ``highlights.json`` or
223
+ ``envelopes/<stream>.json``.
224
+
225
+ Returns:
226
+ False for a run that is still recording, for a root stream when
227
+ there are at most :data:`CROWD_ENVS` envs, for an envelope of a
228
+ stream that is not scalar or vector (or of a single env), for
229
+ highlights when :mod:`simscope.highlights` is not available, and for
230
+ anything else that is not derived data.
231
+ """
232
+ m = rollout.manifest
233
+ if m.status != "complete":
234
+ return False
235
+ if what == ROOT_POSE:
236
+ return m.n_envs > CROWD_ENVS
237
+ if what == SUMMARIES:
238
+ return True
239
+ if what == HIGHLIGHTS:
240
+ return True
241
+ if what.startswith(ENVELOPES) and what.endswith(".json"):
242
+ info = m.streams.get(what[len(ENVELOPES) : -len(".json")])
243
+ return (
244
+ info is not None
245
+ and info.kind in ("scalar", "vector")
246
+ and m.n_envs > 1
247
+ )
248
+ return False
249
+
250
+
251
+ def fresh(
252
+ rollout: library.Rollout, cache: pathlib.Path, what: str
253
+ ) -> pathlib.Path | None:
254
+ """Returns the cached file if it is up to date, without computing.
255
+
256
+ Args:
257
+ rollout: The run.
258
+ cache: The run's cache folder.
259
+ what: See :func:`applicable`.
260
+
261
+ Returns:
262
+ The file path, or ``None`` if it must be computed first.
263
+ """
264
+ path = cache / what
265
+ if path.is_file() and _stamp_ok(rollout, cache):
266
+ return path
267
+ return None
268
+
269
+
270
+ def highlight_count(cache: pathlib.Path) -> int | None:
271
+ """Counts the highlights in a run's cached file, if it has one.
272
+
273
+ Args:
274
+ cache: The run's cache folder.
275
+
276
+ Returns:
277
+ The count, or ``None`` if nothing is cached.
278
+ """
279
+ try:
280
+ doc = json.loads((cache / HIGHLIGHTS).read_bytes())
281
+ return len(doc["highlights"])
282
+ except (OSError, ValueError, KeyError, TypeError):
283
+ return None
284
+
285
+
286
+ def ensure(
287
+ rollout: library.Rollout, cache: pathlib.Path, what: str
288
+ ) -> pathlib.Path | None:
289
+ """Computes a derived file if it is missing or stale, and returns it.
290
+
291
+ Blocks until the file exists; the server runs it in a worker thread.
292
+ Two callers asking for the same run wait for each other, so a pass runs
293
+ once. The file is written atomically.
294
+
295
+ Args:
296
+ rollout: The run (complete).
297
+ cache: The run's cache folder, from :func:`cache_dir`.
298
+ what: See :func:`applicable`.
299
+
300
+ Returns:
301
+ The file path, or ``None`` if :func:`applicable` is false.
302
+
303
+ Raises:
304
+ errors.FormatError: If the run's data is corrupt.
305
+ OSError: If the cache cannot be written.
306
+ """
307
+ if not applicable(rollout, what):
308
+ return None
309
+ with _lock_for(cache):
310
+ if not _stamp_ok(rollout, cache):
311
+ _reset(rollout, cache)
312
+ path = cache / what
313
+ if path.is_file():
314
+ return path
315
+ if what == ROOT_POSE:
316
+ _root_pass(rollout, cache)
317
+ elif what == SUMMARIES:
318
+ _write_json(path, _summaries(rollout, cache))
319
+ elif what == HIGHLIGHTS:
320
+ return _highlights(rollout, cache)
321
+ else:
322
+ stream = what[len(ENVELOPES) : -len(".json")]
323
+ _write_json(path, _envelope(rollout, stream))
324
+ return path
325
+
326
+
327
+ def _write_json(path: pathlib.Path, obj: Any) -> None:
328
+ """Writes a JSON file atomically, creating its folder."""
329
+ cas.atomic_write(path, _dumps(obj))
330
+
331
+
332
+ # -- reading a stream by component --
333
+
334
+
335
+ def _threads() -> int:
336
+ """Threads one pass may use for inflating and encoding blocks."""
337
+ return max(1, min(4, os.cpu_count() or 1))
338
+
339
+
340
+ def _renormalize(x: npt.NDArray[np.float32]) -> None:
341
+ """Renormalizes the quaternion of each ``[..., 7]`` pose in place."""
342
+ q = x[..., 3:]
343
+ norm = np.sqrt(np.sum(q * q, axis=-1, keepdims=True))
344
+ np.divide(q, norm, out=q, where=norm > 0)
345
+
346
+
347
+ class _Source:
348
+ """A complete stream, read window by window and component by component."""
349
+
350
+ def __init__(self, rollout: library.Rollout, name: str) -> None:
351
+ info = rollout.manifest.streams[name]
352
+ self.path = rollout.path / info.file
353
+ self.kind = info.kind
354
+ self.reader = rollout.stream(name)
355
+ self.n_envs = self.reader.n_envs
356
+ self.n_frames = self.reader.n_frames
357
+ self.block_frames = self.reader.block_frames
358
+ self.k = math.prod(self.reader.item_shape)
359
+ self.n_windows = -(-self.n_frames // self.block_frames)
360
+ self._dir = self.reader.directory
361
+
362
+ def window(
363
+ self,
364
+ w: int,
365
+ c0: int,
366
+ c1: int,
367
+ pool: concurrent.futures.Executor,
368
+ envs: tuple[int, int] | None = None,
369
+ ) -> npt.NDArray[np.float32]:
370
+ """Reads components ``c0 <= c < c1`` of one window.
371
+
372
+ Args:
373
+ w: Window index.
374
+ c0: First component.
375
+ c1: One past the last component.
376
+ pool: Executor that inflates blocks in parallel.
377
+ envs: ``(first, stop)`` env range; all envs by default.
378
+
379
+ Returns:
380
+ A float32 array ``[n, envs, c1 - c0]``.
381
+ """
382
+ e0, e1 = envs or (0, self.n_envs)
383
+ t0 = w * self.block_frames
384
+ n = min(self.block_frames, self.n_frames - t0)
385
+ out = np.empty((n, e1 - e0, c1 - c0), np.float32)
386
+ d = self._dir
387
+ pose = self.kind == "pose" and c0 % 7 == 0 and c1 % 7 == 0
388
+ fd = cas.open_read(self.path)
389
+ try:
390
+
391
+ def run(lo: int, hi: int) -> None:
392
+ for env in range(lo, hi):
393
+ ent = d[w * self.n_envs + env]
394
+ payload = cas.read_at(
395
+ fd, int(ent["clen"]), int(ent["offset"]) + 32
396
+ )
397
+ codec = int(ent["codec"])
398
+ x = codecs.decode_components(
399
+ payload, codec, n, self.k, slice(c0, c1)
400
+ )
401
+ if pose and codec == codecs.CODEC_Q16D:
402
+ _renormalize(x.reshape(n, -1, 7))
403
+ out[:, env - e0] = x
404
+
405
+ _fan_out(pool, e0, e1, run)
406
+ finally:
407
+ os.close(fd)
408
+ return out
409
+
410
+
411
+ def _fan_out(
412
+ pool: concurrent.futures.Executor,
413
+ lo: int,
414
+ hi: int,
415
+ fn: Callable[[int, int], None],
416
+ ) -> None:
417
+ """Runs ``fn(a, b)`` over chunks of ``range(lo, hi)`` on the pool."""
418
+ step = max(1, -(-(hi - lo) // (4 * _threads())))
419
+ futures = [
420
+ pool.submit(fn, a, min(a + step, hi)) for a in range(lo, hi, step)
421
+ ]
422
+ for f in futures:
423
+ f.result()
424
+
425
+
426
+ # -- number formatting --
427
+
428
+
429
+ def _sig_list(a: npt.ArrayLike) -> list[Any]:
430
+ """Rounds floats to 7 significant digits and lists them for JSON.
431
+
432
+ Short literals keep derived files small. Non-finite values become
433
+ ``None`` (``null``).
434
+
435
+ Args:
436
+ a: Any-shaped array.
437
+
438
+ Returns:
439
+ Nested lists of the same shape.
440
+ """
441
+ x = np.asarray(a, np.float64)
442
+ finite = np.isfinite(x)
443
+ out = np.zeros_like(x)
444
+ v = x[finite]
445
+ with np.errstate(divide="ignore"):
446
+ mag = np.floor(np.log10(np.abs(v)))
447
+ mag = np.where(np.isfinite(mag), mag, 0.0)
448
+ p = mag - (_SIG_DIGITS - 1) # decimal exponent of the last kept digit
449
+ small = p < 0
450
+ scale = 10.0 ** np.clip(np.abs(p), 0, 300)
451
+ rounded = np.where(
452
+ small, np.rint(v * scale) / scale, np.rint(v / scale) * scale
453
+ )
454
+ out[finite] = np.where(np.abs(p) > 22, v, rounded)
455
+ if finite.all():
456
+ return out.tolist()
457
+ obj = out.astype(object)
458
+ obj[~finite] = None
459
+ return obj.tolist()
460
+
461
+
462
+ # -- root pose and its statistics --
463
+
464
+
465
+ def root_body(scene: core.Scene) -> int:
466
+ """Picks the followed body (see :func:`simscope.highlights.root_body`).
467
+
468
+ Args:
469
+ scene: The run's scene.
470
+
471
+ Returns:
472
+ A body index (never the world body when another exists).
473
+ """
474
+ return highlights.root_body(scene)
475
+
476
+
477
+ def _root_pass(rollout: library.Rollout, cache: pathlib.Path) -> None:
478
+ """Reads ``body_pose`` once for the root body's stream and statistics.
479
+
480
+ Writes ``root_pose.blk`` when the run has more than :data:`CROWD_ENVS`
481
+ envs, and ``_root_stats.json`` (minimum height and peak speed per env)
482
+ always, so ``summaries.json`` does not read the poses again.
483
+ """
484
+ src = _Source(rollout, manifest.BODY_POSE)
485
+ body = root_body(rollout.scene)
486
+ n_envs, bf, dt = src.n_envs, src.block_frames, rollout.dt
487
+ write = n_envs > CROWD_ENVS
488
+ min_z = np.full(n_envs, np.inf)
489
+ peak = np.zeros(n_envs)
490
+ prev: npt.NDArray[np.float32] | None = None
491
+ tmp = cache / (ROOT_POSE + ".part")
492
+ cache.mkdir(parents=True, exist_ok=True)
493
+ writer = (
494
+ blockfile.BlockWriter(
495
+ tmp,
496
+ item_shape=(1, core.POSE_DIM),
497
+ n_envs=n_envs,
498
+ kind="pose",
499
+ codec="q16d",
500
+ block_frames=bf,
501
+ )
502
+ if write
503
+ else None
504
+ )
505
+ try:
506
+ with concurrent.futures.ThreadPoolExecutor(_threads()) as pool:
507
+ for w in range(src.n_windows):
508
+ r = src.window(w, 7 * body, 7 * body + 7, pool) # [n, E, 7]
509
+ pos = r[..., :3]
510
+ np.minimum(min_z, pos[..., 2].min(axis=0), out=min_z)
511
+ seq = pos if prev is None else np.concatenate([prev, pos])
512
+ if len(seq) > 1:
513
+ speed = np.linalg.norm(np.diff(seq, axis=0), axis=-1) / dt
514
+ np.maximum(peak, speed.max(axis=0), out=peak)
515
+ prev = pos[-1:]
516
+ if writer is not None:
517
+ writer.write_encoded(w * bf, _encode_window(r, pool))
518
+ if writer is not None:
519
+ writer.finalize()
520
+ cas.replace_file(tmp, cache / ROOT_POSE)
521
+ except BaseException:
522
+ if writer is not None:
523
+ writer.close()
524
+ tmp.unlink(missing_ok=True)
525
+ raise
526
+ has_frames = src.n_frames > 0
527
+ _write_json(
528
+ cache / _ROOT_STATS,
529
+ {
530
+ "min_height": _sig_list(min_z) if has_frames else None,
531
+ "peak_speed": _sig_list(peak) if has_frames else None,
532
+ },
533
+ )
534
+
535
+
536
+ def _encode_window(
537
+ r: npt.NDArray[np.float32], pool: concurrent.futures.Executor
538
+ ) -> list[blockfile.EncodedBlock]:
539
+ """Encodes a ``[n, E, 7]`` root window as one ``q16d`` block per env."""
540
+ n, n_envs, k = r.shape
541
+ per_env = np.ascontiguousarray(r.transpose(1, 0, 2))
542
+ blocks: list[blockfile.EncodedBlock | None] = [None] * n_envs
543
+
544
+ def run(lo: int, hi: int) -> None:
545
+ for env in range(lo, hi):
546
+ cid, payload = codecs.encode_block_auto(per_env[env], "q16d")
547
+ blocks[env] = blockfile.EncodedBlock(
548
+ env, n, cid, codecs.payload_ulen(cid, n, k), payload
549
+ )
550
+
551
+ _fan_out(pool, 0, n_envs, run)
552
+ return [b for b in blocks if b is not None]
553
+
554
+
555
+ def _root_stats(
556
+ rollout: library.Rollout, cache: pathlib.Path
557
+ ) -> dict[str, list[float] | None]:
558
+ """Returns the cached root statistics, computing them on first use."""
559
+ path = cache / _ROOT_STATS
560
+ if not path.is_file():
561
+ _root_pass(rollout, cache)
562
+ return json.loads(path.read_bytes())
563
+
564
+
565
+ # -- summaries --
566
+
567
+
568
+ def _summaries(rollout: library.Rollout, cache: pathlib.Path) -> dict[str, Any]:
569
+ """Builds ``summaries.json``: per-env columns for the env picker."""
570
+ m = rollout.manifest
571
+ stats = _root_stats(rollout, cache)
572
+ columns: list[dict[str, str]] = []
573
+ values: dict[str, Any] = {}
574
+
575
+ def add(key: str, label: str, unit: str, better: str, v: Any) -> None:
576
+ columns.append(
577
+ {"key": key, "label": label, "unit": unit, "better": better}
578
+ )
579
+ values[key] = v
580
+
581
+ reward = m.streams.get("reward")
582
+ if reward is not None and reward.kind == "scalar":
583
+ add("return", "Return", "", "high", _sig_list(_total(rollout)))
584
+ if stats["min_height"] is not None:
585
+ add("min_height", "Min height", "m", "high", stats["min_height"])
586
+ add("peak_speed", "Peak speed", "m/s", "high", stats["peak_speed"])
587
+ contacts = m.streams.get("contacts")
588
+ if contacts is not None and contacts.kind == "arrows":
589
+ add(
590
+ "peak_contact_force",
591
+ "Peak contact force",
592
+ "N",
593
+ "low",
594
+ _sig_list(_peak_contact(rollout)),
595
+ )
596
+ counts = _highlight_counts(rollout, cache)
597
+ if counts is not None:
598
+ add("n_highlights", "Highlights", "", "low", counts)
599
+ return {"columns": columns, "values": values}
600
+
601
+
602
+ def _total(rollout: library.Rollout) -> npt.NDArray[np.float64]:
603
+ """Sums the ``reward`` stream over time, per env."""
604
+ src = _Source(rollout, "reward")
605
+ total = np.zeros(src.n_envs)
606
+ with concurrent.futures.ThreadPoolExecutor(_threads()) as pool:
607
+ for w in range(src.n_windows):
608
+ total += src.window(w, 0, 1, pool).sum(axis=0, dtype=np.float64)[
609
+ :, 0
610
+ ]
611
+ return total
612
+
613
+
614
+ def _peak_contact(rollout: library.Rollout) -> npt.NDArray[np.float64]:
615
+ """Finds the largest contact force magnitude per env, over time."""
616
+ src = _Source(rollout, "contacts")
617
+ peak = np.zeros(src.n_envs)
618
+ chunk = max(1, _CONTACT_BYTES // (src.block_frames * src.k * 4))
619
+ with concurrent.futures.ThreadPoolExecutor(_threads()) as pool:
620
+ for w in range(src.n_windows):
621
+ for e0 in range(0, src.n_envs, chunk):
622
+ e1 = min(e0 + chunk, src.n_envs)
623
+ a = src.window(w, 0, src.k, pool, (e0, e1))
624
+ force = a.reshape(len(a), e1 - e0, -1, 6)[..., 3:]
625
+ mag = np.sqrt(np.sum(force * force, axis=-1, dtype=np.float64))
626
+ np.maximum(peak[e0:e1], mag.max(axis=(0, 2)), out=peak[e0:e1])
627
+ return peak
628
+
629
+
630
+ def _highlight_counts(
631
+ rollout: library.Rollout, cache: pathlib.Path
632
+ ) -> list[int] | None:
633
+ """Counts highlights per env, or ``None`` if there are none to count."""
634
+ if not applicable(rollout, HIGHLIGHTS):
635
+ return None
636
+ try:
637
+ doc = _highlights_doc(rollout, cache)
638
+ except Exception:
639
+ logger.exception("highlights for %s failed", rollout.name)
640
+ return None
641
+ envs = np.array([h["env"] for h in doc["highlights"]], np.int64)
642
+ n_envs = rollout.manifest.n_envs
643
+ envs = envs[(envs >= 0) & (envs < n_envs)]
644
+ return np.bincount(envs, minlength=n_envs).tolist()
645
+
646
+
647
+ # -- envelopes --
648
+
649
+
650
+ def _envelope(rollout: library.Rollout, stream: str) -> dict[str, Any]:
651
+ """Computes p5, p50 and p95 across envs for every frame and component."""
652
+ src = _Source(rollout, stream)
653
+ t_total, n_envs, k = src.n_frames, src.n_envs, src.k
654
+ bf = src.block_frames
655
+ group = max(1, min(k, _WINDOW_BYTES // (bf * n_envs * 4)))
656
+ out = np.empty((3, t_total, k), np.float32)
657
+ with concurrent.futures.ThreadPoolExecutor(_threads()) as pool:
658
+ for c0 in range(0, k, group):
659
+ c1 = min(c0 + group, k)
660
+ for w in range(src.n_windows):
661
+ win = src.window(w, c0, c1, pool) # [n, E, c]
662
+ t0 = w * bf
663
+ # Envs last so each percentile reads a contiguous row.
664
+ flat = np.ascontiguousarray(win.transpose(0, 2, 1))
665
+ pct = np.percentile(flat, _ENVELOPE_PERCENTILES, axis=-1)
666
+ out[:, t0 : t0 + len(win), c0:c1] = pct
667
+ by_component = out.transpose(0, 2, 1) # [3, K, T]
668
+ return {
669
+ "dt": rollout.dt,
670
+ "t0": 0,
671
+ "components": k,
672
+ "p5": _sig_list(by_component[0]),
673
+ "p50": _sig_list(by_component[1]),
674
+ "p95": _sig_list(by_component[2]),
675
+ }
676
+
677
+
678
+ # -- highlights --
679
+
680
+
681
+ def _highlights_doc(
682
+ rollout: library.Rollout, cache: pathlib.Path
683
+ ) -> dict[str, Any]:
684
+ """Loads or computes the highlights document of a run."""
685
+ doc = highlights.load_or_compute(rollout, cache)
686
+ path = cache / HIGHLIGHTS
687
+ if not path.is_file():
688
+ _write_json(path, doc)
689
+ return doc
690
+
691
+
692
+ def _highlights(
693
+ rollout: library.Rollout, cache: pathlib.Path
694
+ ) -> pathlib.Path | None:
695
+ """Produces ``highlights.json`` in the cache and returns its path."""
696
+ _highlights_doc(rollout, cache)
697
+ return cache / HIGHLIGHTS