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