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/highlights.py ADDED
@@ -0,0 +1,947 @@
1
+ """Automatic highlights: the moments where a rollout's physics peaks.
2
+
3
+ A highlight is a marker on the timeline. Two kinds are built in, both
4
+ computed from signals every simulator provides, so they mean the same for any
5
+ robot or task:
6
+
7
+ ``contact``
8
+ The **net contact force**: the norm of the sum of all the force vectors
9
+ of the ``contacts`` stream. A rollout
10
+ without a ``contacts`` stream gets none.
11
+ ``acceleration``
12
+ The **centre-of-mass acceleration**: the second difference of the centre
13
+ of mass, in m/s^2. The centre of mass is the mean of the body positions
14
+ weighted by ``Body.mass``; a scene without masses uses the plain mean of
15
+ the bodies other than the world. A robot standing still reads about 0,
16
+ free fall reads about 9.8, and an impact or a push-off is a peak.
17
+
18
+ A signal is scored per env by a robust z-score ``(x - median) / max(1.4826 *
19
+ MAD, floor)``. The floor keeps flat and intermittent signals from exploding
20
+ (see ``SPREAD_FLOOR``). A highlight is a local maximum within +-0.25 s whose
21
+ score is above 6. The top 10 per env and kind are kept, and the top 50 per
22
+ kind across envs; markers of one kind within 0.15 s merge. ``ratio`` says
23
+ how many times the run's typical peak the value is.
24
+
25
+ Nothing else is built in: there are no task words such as "jump" or "fall"
26
+ here. To mark something that is specific to your robot, register your own
27
+ kind with :func:`register`, or add explicit events with
28
+ ``Annotations.add_event`` (see "Custom markers" in the getting-started
29
+ guide).
30
+
31
+ Everything is vectorized over ``[T, E]`` and reads the run one window at a
32
+ time, decoding env chunks on threads, so 1,000 frames of 4,096 envs take a
33
+ few seconds and a single 400-frame run takes milliseconds. Teleports (env
34
+ resets) are masked out.
35
+
36
+ Example:
37
+ >>> from simscope import highlights
38
+ >>> doc = highlights.load_or_compute(rollout, cache_dir)
39
+ >>> [(h["t"], h["kind"]) for h in doc["highlights"]]
40
+ """
41
+
42
+ import concurrent.futures
43
+ import dataclasses
44
+ import json
45
+ import logging
46
+ import math
47
+ import os
48
+ import pathlib
49
+ import re
50
+ from collections.abc import Callable, Sequence
51
+ from typing import Any
52
+
53
+ import numpy as np
54
+ import numpy.typing as npt
55
+
56
+ from simscope import core, library
57
+ from simscope.io import blockfile, cas, codecs
58
+
59
+ logger = logging.getLogger(__name__)
60
+
61
+ DETECTOR_VERSION = "simscope/3"
62
+ """Version of the built-in detectors; a change invalidates cached results."""
63
+
64
+ FORMAT = "simscope-highlights/2"
65
+ FILE_NAME = "highlights.json"
66
+ KEY_NAME = "highlights.key"
67
+
68
+ CONTACT = "contact"
69
+ ACCELERATION = "acceleration"
70
+
71
+ _LABELS = {CONTACT: "Contact force", ACCELERATION: "Acceleration"}
72
+
73
+ CONTACTS_STREAM = "contacts"
74
+ ROOT_NAMES = ("torso", "base", "trunk", "pelvis", "chassis")
75
+
76
+ GRAVITY = 9.81
77
+ """Standard gravity in m/s^2, for the ``g`` of a detail."""
78
+
79
+ MERGE_S = 0.15
80
+ """Markers of one kind closer than this merge into one."""
81
+ PER_ENV = 10
82
+ """Most highlights per env and kind."""
83
+ PER_KIND = 50
84
+ """Most highlights per kind across envs."""
85
+
86
+ Z_THRESHOLD = 6.0
87
+ """Least robust z-score of a highlight."""
88
+ WINDOW_S = 0.25
89
+ """A highlight is the largest score within this many seconds."""
90
+ MAD_SIGMA = 1.4826
91
+ """Scales the median absolute deviation to a standard deviation."""
92
+ REL_FLOOR = 0.05
93
+ """Scale floor as a fraction of the median, so a steady load is not flat."""
94
+ SPREAD_FLOOR = 0.25
95
+ """Scale floor as a fraction of the signal's usual peak excursion.
96
+
97
+ A gait's contact force is zero half the time, so its MAD is near zero and
98
+ every stride would score as an outlier. The usual peak is the ``k``-th
99
+ largest value, ``k = max(TOP_MIN, TOP_FRACTION * T)``; a highlight has to
100
+ beat it by 50 %."""
101
+ TOP_FRACTION = 0.05
102
+ TOP_MIN = 3
103
+ CONTACT_FLOOR = 1.0
104
+ """Scale floor of the net contact force, in newtons."""
105
+ ACCEL_FLOOR = 0.5
106
+ """Scale floor of the acceleration, in m/s^2 (q16d noise is about 0.04)."""
107
+ TELEPORT_SPEED = 20.0
108
+ """A centre-of-mass step faster than this (m/s) is a reset, not motion."""
109
+ MIN_FRAMES = 8
110
+ """Fewest frames a signal needs for a typical level to mean anything."""
111
+
112
+ KEY_PATTERN = re.compile(r"^[a-z][a-z0-9_]{0,31}$")
113
+ """What a custom kind's key looks like: lower case, digits and ``_``."""
114
+ COLOR_PATTERN = re.compile(r"^#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{6})$")
115
+ """What a custom kind's colour looks like: ``#rgb`` or ``#rrggbb``."""
116
+ MAX_LABEL = 40
117
+ """Longest label of a custom kind, in characters."""
118
+
119
+ _CANDIDATE_CHUNK = 1 << 18
120
+ _MAX_WORKERS = 3
121
+ _WIDE = 64 # floats per env-frame above which threads pay off
122
+ _ENV_CHUNK = 256
123
+
124
+
125
+ @dataclasses.dataclass(frozen=True)
126
+ class Highlight:
127
+ """One marker on the timeline (viewer contracts 9.3).
128
+
129
+ Attributes:
130
+ t: Time in seconds (the start, for a span).
131
+ frame: Frame index.
132
+ env: Env index.
133
+ kind: Key of the kind of marker, such as ``"contact"``.
134
+ score: Number that ranks the marker within its kind (higher is more
135
+ notable). The built-in kinds use the robust z-score.
136
+ value: The headline measurement, in the unit of the kind: newtons of
137
+ ``contact``, m/s^2 of ``acceleration``.
138
+ body: Body index the marker belongs to, or ``None``.
139
+ label: Name shown for this marker; the kind's label if left empty.
140
+ detail: One short phrase with the numbers behind the marker, with
141
+ units, such as ``"412 N, 5.1 times typical"``.
142
+ t1: End of a span in seconds, or ``None`` for a moment.
143
+ frame1: End frame of a span, or ``None``.
144
+ ratio: How many times the typical level the value is, or ``None``.
145
+ also: Keys of other kinds that were merged into this marker.
146
+ """
147
+
148
+ t: float
149
+ frame: int
150
+ env: int
151
+ kind: str
152
+ score: float
153
+ value: float
154
+ body: int | None = None
155
+ label: str = ""
156
+ detail: str = ""
157
+ t1: float | None = None
158
+ frame1: int | None = None
159
+ ratio: float | None = None
160
+ also: tuple[str, ...] = ()
161
+
162
+
163
+ Detector = Callable[[library.Rollout], list[Highlight]]
164
+ """A custom detector: returns highlights whose ``kind`` is its key."""
165
+
166
+
167
+ @dataclasses.dataclass(frozen=True)
168
+ class _Kind:
169
+ """A registered kind of highlight."""
170
+
171
+ key: str
172
+ label: str
173
+ color: str | None
174
+ detector: Detector | None # None: built in, produced by a pass
175
+
176
+
177
+ _REGISTRY: dict[str, _Kind] = {}
178
+
179
+
180
+ def register(
181
+ key: str,
182
+ detector: Detector,
183
+ *,
184
+ label: str,
185
+ color: str | None = None,
186
+ ) -> None:
187
+ """Adds your own kind of marker to the ones :func:`detect` runs.
188
+
189
+ The detector runs once per run, in the background, and its markers are
190
+ cached with the built-in ones. Adding or changing a detector refreshes
191
+ the cache. A detector that raises, or returns something invalid, is
192
+ logged and skipped; the other kinds still work.
193
+
194
+ Args:
195
+ key: Name of the kind: a lower case letter, then lower case letters,
196
+ digits and ``_``, at most 32 characters. It must be new (see
197
+ :func:`unregister`).
198
+ detector: ``detector(rollout)`` returns a list of :class:`Highlight`
199
+ whose ``kind`` equals ``key``. Each has a ``t`` in seconds, a
200
+ ``frame`` inside the run, and an ``env`` of the run. A span also
201
+ sets ``t1`` and ``frame1``. It is called once for the whole run;
202
+ ``detect`` keeps the envs it was asked for.
203
+ label: Display name of the kind (at most 40 characters), used for
204
+ highlights that have no label of their own.
205
+ color: A CSS hex colour for the marker, such as ``"#d9480f"``, or
206
+ ``None`` for the viewer's default.
207
+
208
+ Raises:
209
+ ValueError: If the key, label or colour is invalid, or the key is
210
+ already registered.
211
+ """
212
+ if not isinstance(key, str) or not KEY_PATTERN.match(key):
213
+ raise ValueError(
214
+ f"kind key {key!r} must be lower case letters, digits and _, "
215
+ "start with a letter, and have at most 32 characters"
216
+ )
217
+ if key in _REGISTRY:
218
+ raise ValueError(f"kind {key!r} is already registered")
219
+ if not 1 <= len(label) <= MAX_LABEL:
220
+ raise ValueError(f"label must have 1 to {MAX_LABEL} characters")
221
+ if color is not None and not (
222
+ isinstance(color, str) and COLOR_PATTERN.match(color)
223
+ ):
224
+ raise ValueError(f"color {color!r} must be a hex colour like #d9480f")
225
+ _REGISTRY[key] = _Kind(key, label, color, detector)
226
+
227
+
228
+ def unregister(key: str) -> None:
229
+ """Removes a kind, built-in or registered.
230
+
231
+ Args:
232
+ key: Kind key.
233
+
234
+ Raises:
235
+ KeyError: If no such kind is registered.
236
+ """
237
+ del _REGISTRY[key]
238
+
239
+
240
+ def kinds() -> list[dict[str, str]]:
241
+ """Lists the registered kinds as ``{key, label}`` objects.
242
+
243
+ A custom kind that has a colour also carries ``color``.
244
+ """
245
+ out = []
246
+ for kind in _REGISTRY.values():
247
+ entry = {"key": kind.key, "label": kind.label}
248
+ if kind.detector is not None and kind.color is not None:
249
+ entry["color"] = kind.color
250
+ out.append(entry)
251
+ return out
252
+
253
+
254
+ def root_body(names: Sequence[str] | core.Scene) -> int:
255
+ """Picks the followed (root) body of a scene.
256
+
257
+ That is the first of :data:`ROOT_NAMES` (exact names first, then
258
+ prefixes), else the first body that is not the world. The derived
259
+ root-pose stream and the player's follow camera use this same rule
260
+ (``FOLLOW_ALIASES`` in ``web/src/core/player.js``).
261
+
262
+ Args:
263
+ names: Body names in scene order, or the scene itself.
264
+
265
+ Returns:
266
+ A body index.
267
+ """
268
+ if isinstance(names, core.Scene):
269
+ names = [b.name for b in names.bodies]
270
+ names = [n.lower() for n in names]
271
+ for test in (str.__eq__, str.startswith):
272
+ for wanted in ROOT_NAMES:
273
+ for i, name in enumerate(names):
274
+ if test(name, wanted):
275
+ return i
276
+ for i, name in enumerate(names):
277
+ if name != "world":
278
+ return i
279
+ return 0
280
+
281
+
282
+ # -- reading ---------------------------------------------------------------
283
+
284
+ Reducer = Callable[[npt.NDArray[np.float32], int], npt.NDArray[np.float32]]
285
+ """Maps a window ``[t, e, c]`` of envs starting at ``lo`` to ``[t, e, ...]``."""
286
+
287
+
288
+ def _scan(
289
+ run: library.Rollout,
290
+ name: str,
291
+ envs: npt.NDArray[np.intp],
292
+ comps: Sequence[int],
293
+ reduce: Reducer,
294
+ ) -> npt.NDArray[np.float32]:
295
+ """Reduces a stream window by window, decoding env chunks on threads.
296
+
297
+ Args:
298
+ run: A finished run.
299
+ name: Stream name.
300
+ envs: Env indices to read.
301
+ comps: Components (flattened item indices) to decode.
302
+ reduce: Maps a window ``[t, e, len(comps)]`` and the index of its
303
+ first env within ``envs`` to ``[t, e, *rest]``. It runs on
304
+ worker threads and must not keep the window.
305
+
306
+ Returns:
307
+ ``[n_frames, len(envs), *rest]``.
308
+ """
309
+ reader = run.stream(name)
310
+ k = math.prod(reader.item_shape)
311
+ cols = np.asarray(comps, np.intp)
312
+ n_frames, bf, n_envs = reader.n_frames, reader.block_frames, reader.n_envs
313
+ n_windows = -(-n_frames // bf)
314
+ table = reader.directory
315
+ offsets = table["offset"].astype(np.int64) + blockfile.BLOCK_HEADER.size
316
+ lengths = table["clen"].astype(np.int64)
317
+ codec_ids = table["codec"].astype(np.int64)
318
+ chunks = [envs[i : i + _ENV_CHUNK] for i in range(0, len(envs), _ENV_CHUNK)]
319
+ out: npt.NDArray[np.float32] | None = None
320
+ fd = cas.open_read(run.path / run.manifest.streams[name].file)
321
+
322
+ def job(at: tuple[int, int]) -> None:
323
+ nonlocal out
324
+ w, c = at
325
+ t0 = w * bf
326
+ n = min(bf, n_frames - t0)
327
+ window = np.empty((n, len(chunks[c]), len(cols)), np.float32)
328
+ for j, env in enumerate(chunks[c].tolist()):
329
+ i = w * n_envs + env
330
+ payload = cas.read_at(fd, int(lengths[i]), int(offsets[i]))
331
+ window[:, j] = codecs.decode_components(
332
+ payload, int(codec_ids[i]), n, k, cols
333
+ )
334
+ lo = c * _ENV_CHUNK
335
+ part = reduce(window, lo)
336
+ if out is None: # the first job runs alone, before the threads
337
+ out = np.empty((n_frames, len(envs), *part.shape[2:]), np.float32)
338
+ out[t0 : t0 + n, lo : lo + part.shape[1]] = part
339
+
340
+ jobs = [(w, c) for w in range(n_windows) for c in range(len(chunks))]
341
+ try:
342
+ if not jobs:
343
+ return np.empty((0, len(envs)), np.float32)
344
+ job(jobs[0])
345
+ rest = jobs[1:]
346
+ # Small items decode in microseconds and threads then only contend
347
+ # for the GIL, so only wide ones (poses) get workers.
348
+ workers = _MAX_WORKERS if k >= _WIDE else 1
349
+ if workers == 1 or len(rest) < 4:
350
+ for j in rest:
351
+ job(j)
352
+ else:
353
+ with concurrent.futures.ThreadPoolExecutor(workers) as pool:
354
+ for f in [pool.submit(job, j) for j in rest]:
355
+ f.result()
356
+ finally:
357
+ os.close(fd)
358
+ assert out is not None
359
+ return out
360
+
361
+
362
+ def _net_force(
363
+ window: npt.NDArray[np.float32], _lo: int
364
+ ) -> npt.NDArray[np.float32]:
365
+ """Norm of the sum of the K force vectors, ``[t, e, 3K] -> [t, e]``."""
366
+ total = window.reshape(*window.shape[:2], -1, 3).sum(axis=2)
367
+ return np.sqrt(np.einsum("...i,...i->...", total, total))
368
+
369
+
370
+ def _weighted_mean(
371
+ window: npt.NDArray[np.float32],
372
+ _lo: int,
373
+ weights: npt.NDArray[np.float32],
374
+ ) -> npt.NDArray[np.float32]:
375
+ """Weighted mean of B positions, ``[t, e, 3B] -> [t, e, 3]``."""
376
+ xyz = window.reshape(*window.shape[:2], len(weights), 3)
377
+ return np.einsum("tebi,b->tei", xyz, weights)
378
+
379
+
380
+ class _Context:
381
+ """Signals of one run, each read once and shared by the detectors."""
382
+
383
+ def __init__(
384
+ self, rollout: library.Rollout, envs: Sequence[int] | None
385
+ ) -> None:
386
+ self.rollout = rollout
387
+ n = rollout.n_envs
388
+ ids = np.arange(n) if envs is None else np.unique(np.asarray(envs, int))
389
+ if ids.size and (ids[0] < 0 or ids[-1] >= n):
390
+ raise IndexError(f"env index outside [0, {n})")
391
+ self.envs = ids.astype(np.intp)
392
+ self.dt = rollout.dt
393
+ self.n_frames = rollout.n_frames
394
+
395
+ def bodies(self) -> list[tuple[str, float]]:
396
+ """Reads ``(name, mass)`` of every body of the scene."""
397
+ run = self.rollout
398
+ doc = json.loads(
399
+ cas.ContentStore(run.path.parents[1]).get(
400
+ run.manifest.scene, "scene"
401
+ )
402
+ ) # body names only: the full scene would load every mesh
403
+ return [(b["name"], float(b.get("mass", 0.0))) for b in doc["bodies"]]
404
+
405
+ def centre_of_mass(self) -> npt.NDArray[np.float32] | None:
406
+ """Centre of mass of each env and frame, ``[T, len(envs), 3]``.
407
+
408
+ Weighted by body mass if the scene has masses, else the plain mean
409
+ of the bodies other than the world. ``None`` if there is no body.
410
+ """
411
+ bodies = self.bodies()
412
+ masses = [m for _, m in bodies]
413
+ if any(m > 0 for m in masses):
414
+ used = [i for i, m in enumerate(masses) if m > 0]
415
+ weights = np.asarray([masses[i] for i in used], np.float64)
416
+ else:
417
+ used = [
418
+ i for i, (n, _) in enumerate(bodies) if n.lower() != "world"
419
+ ]
420
+ weights = np.ones(len(used))
421
+ if not used:
422
+ return None
423
+ weights = (weights / weights.sum()).astype(np.float32)
424
+ comps = [7 * i + c for i in used for c in range(3)]
425
+
426
+ def reduce(
427
+ window: npt.NDArray[np.float32], lo: int
428
+ ) -> npt.NDArray[np.float32]:
429
+ return _weighted_mean(window, lo, weights)
430
+
431
+ return _scan(self.rollout, "body_pose", self.envs, comps, reduce)
432
+
433
+ def net_contact_force(self) -> npt.NDArray[np.float32] | None:
434
+ """Net contact force per env and frame, or ``None`` if none."""
435
+ info = self.rollout.manifest.streams.get(CONTACTS_STREAM)
436
+ if info is None or info.kind != "arrows":
437
+ return None
438
+ forces = [
439
+ 6 * i + j for i in range(info.item_shape[0]) for j in (3, 4, 5)
440
+ ]
441
+ return _scan(
442
+ self.rollout, CONTACTS_STREAM, self.envs, forces, _net_force
443
+ )
444
+
445
+
446
+ # -- scoring ---------------------------------------------------------------
447
+
448
+
449
+ def _robust_z(
450
+ x: npt.NDArray[np.float32],
451
+ bad: npt.NDArray[np.bool_],
452
+ floor: float,
453
+ ) -> tuple[npt.NDArray[np.float32], npt.NDArray[np.float32]]:
454
+ """Robust z-scores of each env's series (columns of ``x``).
455
+
456
+ Args:
457
+ x: Signal, ``[T, E]``.
458
+ bad: Frames to ignore (not finite, or teleports), ``[T, E]``.
459
+ floor: Least scale, in the signal's unit.
460
+
461
+ Returns:
462
+ ``(z, typical)``. ``z`` is ``(x - median) / scale`` with ``-inf``
463
+ where ``bad``; the scale is ``max(1.4826 * MAD, floor, 5 % of
464
+ |median|, 25 % of the usual peak's excursion)``. ``typical`` is the
465
+ usual peak (see ``SPREAD_FLOOR``) of each env, at least ``floor``:
466
+ what a peak is measured against.
467
+ """
468
+ if bad.any():
469
+ first = np.median(np.where(bad, 0.0, x), axis=0)
470
+ x = np.where(bad, first, x)
471
+ med = np.median(x, axis=0)
472
+ mad = np.median(np.abs(x - med), axis=0)
473
+ n = x.shape[0]
474
+ k = min(n, max(TOP_MIN, math.ceil(TOP_FRACTION * n)))
475
+ usual = np.partition(x, n - k, axis=0)[n - k] # the k-th largest
476
+ scale = np.maximum(
477
+ MAD_SIGMA * mad,
478
+ np.maximum(
479
+ floor,
480
+ np.maximum(REL_FLOOR * np.abs(med), SPREAD_FLOOR * (usual - med)),
481
+ ),
482
+ )
483
+ z = (x - med) / scale.astype(np.float32)
484
+ z[bad] = -np.inf
485
+ typical = np.maximum(usual, floor).astype(np.float32)
486
+ return z.astype(np.float32, copy=False), typical
487
+
488
+
489
+ def _local_maxima(
490
+ z: npt.NDArray[np.float32], threshold: float, half: int
491
+ ) -> tuple[npt.NDArray[np.intp], npt.NDArray[np.intp]]:
492
+ """Finds frames above ``threshold`` that top their +-``half`` window.
493
+
494
+ Among equal neighbours the earliest wins. Cost is proportional to the
495
+ number of frames above the threshold, not to the size of ``z``.
496
+
497
+ Returns:
498
+ Frame and column indices of the peaks.
499
+ """
500
+ cand = np.argwhere(z > threshold)
501
+ n = z.shape[0]
502
+ offsets = np.concatenate([np.arange(-half, 0), np.arange(1, half + 1)])
503
+ frames: list[npt.NDArray[np.intp]] = []
504
+ cols: list[npt.NDArray[np.intp]] = []
505
+ for i in range(0, len(cand), _CANDIDATE_CHUNK):
506
+ t, e = cand[i : i + _CANDIDATE_CHUNK].T
507
+ at = t[:, None] + offsets[None, :]
508
+ valid = (at >= 0) & (at < n)
509
+ around = np.where(valid, z[np.clip(at, 0, n - 1), e[:, None]], -np.inf)
510
+ centre = z[t, e][:, None]
511
+ left = offsets[None, :] < 0
512
+ beaten = np.where(left, around >= centre, around > centre)
513
+ keep = ~beaten.any(axis=1)
514
+ frames.append(t[keep])
515
+ cols.append(e[keep])
516
+ if not frames:
517
+ return np.empty(0, np.intp), np.empty(0, np.intp)
518
+ return np.concatenate(frames), np.concatenate(cols)
519
+
520
+
521
+ def _peaks(
522
+ x: npt.NDArray[np.float32],
523
+ dt: float,
524
+ *,
525
+ floor: float,
526
+ bad: npt.NDArray[np.bool_] | None = None,
527
+ ) -> tuple[
528
+ npt.NDArray[np.intp],
529
+ npt.NDArray[np.intp],
530
+ npt.NDArray[np.float32],
531
+ npt.NDArray[np.float32],
532
+ ]:
533
+ """Finds the peaks of a signal ``[T, E]``.
534
+
535
+ Args:
536
+ x: Signal.
537
+ dt: Seconds per frame.
538
+ floor: Least scale of the robust z-score, in the signal's unit.
539
+ bad: Frames to ignore (those that are not finite are added).
540
+
541
+ Returns:
542
+ ``(frames, columns, scores, typical)`` of the kept peaks: the top
543
+ ``PER_ENV`` of each column, best first within a column, at most
544
+ ``PER_KIND`` overall (as a set), and the typical level per column.
545
+ """
546
+ empty = np.empty(0, np.intp)
547
+ typical = np.ones(x.shape[1], np.float32)
548
+ if x.shape[0] < MIN_FRAMES or x.shape[1] == 0:
549
+ return empty, empty, np.empty(0, np.float32), typical
550
+ invalid = ~np.isfinite(x)
551
+ if bad is not None:
552
+ invalid |= bad
553
+ z, typical = _robust_z(x, invalid, floor)
554
+ half = max(1, round(WINDOW_S / dt))
555
+ t, col = _local_maxima(z, Z_THRESHOLD, min(half, x.shape[0]))
556
+ if not len(t):
557
+ return empty, empty, np.empty(0, np.float32), typical
558
+ score = z[t, col]
559
+ order = np.lexsort((-score, col)) # by env, best first
560
+ t, col, score = t[order], col[order], score[order]
561
+ starts = np.flatnonzero(np.r_[True, col[1:] != col[:-1]])
562
+ rank = np.arange(len(col)) - np.repeat(
563
+ starts, np.diff(np.r_[starts, len(col)])
564
+ )
565
+ keep = rank < PER_ENV
566
+ t, col, score = t[keep], col[keep], score[keep]
567
+ best = np.argsort(-score, kind="stable")[:PER_KIND]
568
+ return t[best], col[best], score[best], typical
569
+
570
+
571
+ _TIMES = chr(0xD7) # the multiplication sign
572
+
573
+
574
+ def _times(ratio: float) -> str:
575
+ """Formats a ratio with a multiplication sign, such as 5.1 times."""
576
+ return f"{ratio:.1f}{_TIMES}"
577
+
578
+
579
+ # -- built-in kinds --------------------------------------------------------
580
+
581
+
582
+ def _contact(ctx: _Context) -> list[Highlight]:
583
+ """Peaks of the net contact force."""
584
+ force = ctx.net_contact_force()
585
+ if force is None:
586
+ return []
587
+ t, col, score, typical = _peaks(force, ctx.dt, floor=CONTACT_FLOOR)
588
+ found = []
589
+ for i in range(len(t)):
590
+ value = float(force[t[i], col[i]])
591
+ ratio = value / float(typical[col[i]])
592
+ found.append(
593
+ Highlight(
594
+ t=float(t[i] * ctx.dt),
595
+ frame=int(t[i]),
596
+ env=int(ctx.envs[col[i]]),
597
+ kind=CONTACT,
598
+ score=float(score[i]),
599
+ value=value,
600
+ label=_LABELS[CONTACT],
601
+ detail=f"{value:.0f} N, {_times(ratio)} typical",
602
+ ratio=ratio,
603
+ )
604
+ )
605
+ return found
606
+
607
+
608
+ def _acceleration_magnitude(
609
+ com: npt.NDArray[np.float32], dt: float
610
+ ) -> tuple[npt.NDArray[np.float32], npt.NDArray[np.bool_]]:
611
+ """Gets the centre-of-mass acceleration from its positions.
612
+
613
+ The 3-point stencil ``(x[t+1] - 2 x[t] + x[t-1]) / dt^2``.
614
+
615
+ Args:
616
+ com: Centre of mass, ``[T, E, 3]``.
617
+ dt: Seconds per frame.
618
+
619
+ Returns:
620
+ ``(magnitude, bad)``, both ``[T, E]``. ``bad`` marks the two end
621
+ frames, where there is no second difference, and the frames on
622
+ either side of a step faster than a teleport (an env reset).
623
+ """
624
+ acc = np.zeros(com.shape[:2], np.float32)
625
+ bad = np.zeros(com.shape[:2], bool)
626
+ bad[0] = bad[-1] = True
627
+ second = com[2:] - 2 * com[1:-1] + com[:-2]
628
+ acc[1:-1] = np.sqrt(
629
+ np.einsum("...i,...i->...", second, second)
630
+ ) / np.float32(dt**2)
631
+ step = com[1:] - com[:-1]
632
+ jump = np.einsum("...i,...i->...", step, step) > (TELEPORT_SPEED * dt) ** 2
633
+ bad[:-1] |= jump
634
+ bad[1:] |= jump
635
+ return acc, bad
636
+
637
+
638
+ def _acceleration(ctx: _Context) -> list[Highlight]:
639
+ """Peaks of the centre-of-mass acceleration."""
640
+ if ctx.n_frames < 3:
641
+ return []
642
+ com = ctx.centre_of_mass()
643
+ if com is None or com.shape[0] < 3:
644
+ return []
645
+ acc, bad = _acceleration_magnitude(com, ctx.dt)
646
+ t, col, score, typical = _peaks(acc, ctx.dt, floor=ACCEL_FLOOR, bad=bad)
647
+ found = []
648
+ for i in range(len(t)):
649
+ value = float(acc[t[i], col[i]])
650
+ found.append(
651
+ Highlight(
652
+ t=float(t[i] * ctx.dt),
653
+ frame=int(t[i]),
654
+ env=int(ctx.envs[col[i]]),
655
+ kind=ACCELERATION,
656
+ score=float(score[i]),
657
+ value=value,
658
+ label=_LABELS[ACCELERATION],
659
+ detail=f"{value:.0f} m/s², {value / GRAVITY:.1f} g",
660
+ ratio=value / float(typical[col[i]]),
661
+ )
662
+ )
663
+ return found
664
+
665
+
666
+ _PASSES: tuple[tuple[str, Callable[[_Context], list[Highlight]]], ...] = (
667
+ (CONTACT, _contact),
668
+ (ACCELERATION, _acceleration),
669
+ )
670
+
671
+
672
+ def _install_builtins() -> None:
673
+ """Registers the built-in kinds."""
674
+ for key, label in _LABELS.items():
675
+ _REGISTRY[key] = _Kind(key, label, None, None)
676
+
677
+
678
+ _install_builtins()
679
+
680
+
681
+ # -- caps, merging, custom results -----------------------------------------
682
+
683
+
684
+ def _cap(found: list[Highlight]) -> list[Highlight]:
685
+ """Keeps the best ``PER_ENV`` of each env and kind, ``PER_KIND`` per kind.
686
+
687
+ Args:
688
+ found: Highlights of any kinds.
689
+
690
+ Returns:
691
+ The kept highlights, in no particular order.
692
+ """
693
+ best: dict[tuple[str, int], list[Highlight]] = {}
694
+ for h in sorted(found, key=lambda h: -h.score):
695
+ group = best.setdefault((h.kind, h.env), [])
696
+ if len(group) < PER_ENV:
697
+ group.append(h)
698
+ per_kind: dict[str, list[Highlight]] = {}
699
+ for (kind, _), group in best.items():
700
+ per_kind.setdefault(kind, []).extend(group)
701
+ out: list[Highlight] = []
702
+ for group in per_kind.values():
703
+ group.sort(key=lambda h: -h.score)
704
+ out += group[:PER_KIND]
705
+ return out
706
+
707
+
708
+ def _merge(found: list[Highlight]) -> list[Highlight]:
709
+ """Merges moments of one env and kind that are within ``MERGE_S``.
710
+
711
+ The higher score keeps the moment. Spans (markers with ``t1``) never
712
+ merge.
713
+
714
+ Args:
715
+ found: Highlights of any kinds.
716
+
717
+ Returns:
718
+ The merged highlights.
719
+ """
720
+ spans = [h for h in found if h.t1 is not None]
721
+ points = sorted(
722
+ (h for h in found if h.t1 is None), key=lambda h: (-h.score, h.t)
723
+ )
724
+ kept: dict[tuple[str, int], list[Highlight]] = {}
725
+ for h in points:
726
+ near = kept.setdefault((h.kind, h.env), [])
727
+ if all(abs(k.t - h.t) > MERGE_S + 1e-9 for k in near):
728
+ near.append(h)
729
+ return [h for group in kept.values() for h in group] + spans
730
+
731
+
732
+ def _check(
733
+ key: str, found: object, n_frames: int, n_envs: int
734
+ ) -> list[Highlight]:
735
+ """Validates what a custom detector returned.
736
+
737
+ Args:
738
+ key: The detector's kind key.
739
+ found: What it returned.
740
+ n_frames: Frames of the run.
741
+ n_envs: Envs of the run.
742
+
743
+ Returns:
744
+ The highlights, as a list.
745
+
746
+ Raises:
747
+ ValueError: With a message that names the detector, if the result
748
+ is not a list of highlights of its kind with finite times and a
749
+ frame and env inside the run.
750
+ """
751
+ who = f"detector {key!r}"
752
+ if not isinstance(found, list | tuple):
753
+ raise ValueError(
754
+ f"{who} must return a list, got {type(found).__name__}"
755
+ )
756
+ for i, h in enumerate(found):
757
+ if not isinstance(h, Highlight):
758
+ raise ValueError(
759
+ f"{who} returned a {type(h).__name__}, not Highlight"
760
+ )
761
+ where = f"{who}, highlight {i}"
762
+ if h.kind != key:
763
+ raise ValueError(f"{where}: kind is {h.kind!r}, not {key!r}")
764
+ times = [h.t] if h.t1 is None else [h.t, h.t1]
765
+ if not all(math.isfinite(x) for x in times) or h.t < 0:
766
+ raise ValueError(f"{where}: t and t1 must be finite and >= 0")
767
+ if h.t1 is not None and h.t1 < h.t:
768
+ raise ValueError(f"{where}: t1 {h.t1} is before t {h.t}")
769
+ frames = [h.frame] if h.frame1 is None else [h.frame, h.frame1]
770
+ if not all(0 <= f < n_frames for f in frames):
771
+ raise ValueError(
772
+ f"{where}: frame must be in [0, {n_frames}), got {frames}"
773
+ )
774
+ if (h.t1 is None) != (h.frame1 is None):
775
+ raise ValueError(f"{where}: set both t1 and frame1, or neither")
776
+ if not 0 <= h.env < n_envs:
777
+ raise ValueError(f"{where}: env must be in [0, {n_envs})")
778
+ if not (math.isfinite(h.score) and math.isfinite(h.value)):
779
+ raise ValueError(f"{where}: score and value must be finite")
780
+ return list(found)
781
+
782
+
783
+ # -- public API ------------------------------------------------------------
784
+
785
+
786
+ def detect(
787
+ rollout: library.Rollout, *, envs: Sequence[int] | None = None
788
+ ) -> list[Highlight]:
789
+ """Runs every registered detector on a finished run.
790
+
791
+ A custom detector that raises, or whose result is invalid (see
792
+ :func:`register`), is logged and skipped.
793
+
794
+ Args:
795
+ rollout: The run. It must not be recording.
796
+ envs: Only look at these envs (built-in detectors read only them;
797
+ custom detectors see the whole run and their result is
798
+ filtered). ``None`` means all envs.
799
+
800
+ Returns:
801
+ The highlights, sorted by time (then env and kind).
802
+
803
+ Raises:
804
+ IndexError: If an env index is out of range.
805
+ errors.FormatError: If the run is unfinished.
806
+ """
807
+ ctx = _Context(rollout, envs)
808
+ wanted = None if envs is None else set(ctx.envs.tolist())
809
+ found: list[Highlight] = []
810
+ for key, run_pass in _PASSES:
811
+ if key in _REGISTRY:
812
+ found += _cap(run_pass(ctx))
813
+ for kind in list(_REGISTRY.values()):
814
+ if kind.detector is None:
815
+ continue
816
+ try:
817
+ mine = _check(
818
+ kind.key,
819
+ kind.detector(rollout),
820
+ rollout.n_frames,
821
+ rollout.n_envs,
822
+ )
823
+ except Exception as exc: # isolation point: one bad detector only
824
+ logger.warning(
825
+ "skipping detector %r: %s",
826
+ kind.key,
827
+ exc,
828
+ exc_info=not isinstance(exc, ValueError),
829
+ )
830
+ continue
831
+ found += [
832
+ dataclasses.replace(h, label=h.label or kind.label)
833
+ for h in mine
834
+ if wanted is None or h.env in wanted
835
+ ]
836
+ merged = _merge(found)
837
+ merged.sort(key=lambda h: (h.t, h.env, h.kind))
838
+ return merged
839
+
840
+
841
+ def to_json(
842
+ rollout: library.Rollout, highlights: Sequence[Highlight]
843
+ ) -> dict[str, Any]:
844
+ """Builds the ``highlights.json`` object (viewer contracts 9.3).
845
+
846
+ Args:
847
+ rollout: The run the highlights belong to.
848
+ highlights: Highlights from :func:`detect`.
849
+
850
+ Returns:
851
+ The JSON-ready object. ``kinds`` lists the kinds that have at least
852
+ one highlight (as its kind or among its ``also``), in registration
853
+ order; a custom kind with a colour carries it.
854
+ """
855
+ ordered = sorted(highlights, key=lambda h: (h.t, h.env, h.kind))
856
+ present = {h.kind for h in ordered} | {k for h in ordered for k in h.also}
857
+ labels = {k["key"]: k["label"] for k in kinds()}
858
+
859
+ def rounded(x: float | None, digits: int) -> float | None:
860
+ return None if x is None else round(x, digits)
861
+
862
+ return {
863
+ "format": FORMAT,
864
+ "detector": DETECTOR_VERSION,
865
+ "run_id": rollout.manifest.id,
866
+ "kinds": [k for k in kinds() if k["key"] in present],
867
+ "highlights": [
868
+ {
869
+ "t": round(h.t, 6),
870
+ "frame": h.frame,
871
+ "t1": rounded(h.t1, 6),
872
+ "frame1": h.frame1,
873
+ "env": h.env,
874
+ "kind": h.kind,
875
+ "label": h.label or labels.get(h.kind, h.kind),
876
+ "detail": h.detail,
877
+ "score": round(h.score, 3),
878
+ "ratio": rounded(h.ratio, 3),
879
+ "value": round(h.value, 4),
880
+ "body": h.body,
881
+ "also": list(h.also),
882
+ }
883
+ for h in ordered
884
+ ],
885
+ }
886
+
887
+
888
+ def cache_key(rollout: library.Rollout) -> str:
889
+ """Identifies the inputs of a cached result.
890
+
891
+ Args:
892
+ rollout: The run.
893
+
894
+ Returns:
895
+ A string of run id, detector version, the registered custom kinds
896
+ (key, label and colour) and the manifest's modification time.
897
+ """
898
+ man = rollout.path / (
899
+ "rollout.json"
900
+ if (rollout.path / "rollout.json").exists()
901
+ else "rollout.json.partial"
902
+ )
903
+ extra = ",".join(
904
+ sorted(
905
+ f"{k.key}:{k.label}:{k.color or ''}"
906
+ for k in _REGISTRY.values()
907
+ if k.detector is not None
908
+ )
909
+ )
910
+ return (
911
+ f"{rollout.manifest.id} {DETECTOR_VERSION} [{extra}] "
912
+ f"{man.stat().st_mtime_ns}"
913
+ )
914
+
915
+
916
+ def load_or_compute(
917
+ rollout: library.Rollout, cache_dir: os.PathLike[str] | str
918
+ ) -> dict[str, Any]:
919
+ """Returns the highlights document, cached in ``cache_dir``.
920
+
921
+ The cache is ``highlights.json`` plus a ``highlights.key`` file beside
922
+ it; both are replaced atomically, the key last. A different run id,
923
+ detector version, set of registered custom kinds or manifest mtime makes
924
+ the next call recompute.
925
+
926
+ Args:
927
+ rollout: A finished run.
928
+ cache_dir: Directory for the cache (created if missing), for example
929
+ ``.simscope/derived/<run id>``.
930
+
931
+ Returns:
932
+ The ``highlights.json`` object.
933
+ """
934
+ cache = pathlib.Path(cache_dir)
935
+ key = cache_key(rollout)
936
+ doc_path, key_path = cache / FILE_NAME, cache / KEY_NAME
937
+ try:
938
+ if key_path.read_text("utf-8") == key:
939
+ return json.loads(doc_path.read_bytes())
940
+ except (OSError, ValueError):
941
+ pass
942
+ doc = to_json(rollout, detect(rollout))
943
+ cache.mkdir(parents=True, exist_ok=True)
944
+ cas.atomic_write(doc_path, json.dumps(doc, separators=(",", ":")).encode())
945
+ cas.atomic_write(key_path, key.encode("utf-8"))
946
+ logger.info("highlights for %s: %d", rollout.name, len(doc["highlights"]))
947
+ return doc