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/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
|