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/importers.py
ADDED
|
@@ -0,0 +1,874 @@
|
|
|
1
|
+
"""Importers for rollouts that already exist on disk.
|
|
2
|
+
|
|
3
|
+
Two sources are supported, both written through ``Library.record`` so the
|
|
4
|
+
content store deduplicates their meshes across runs (decision D15):
|
|
5
|
+
|
|
6
|
+
* ``.rbundle`` files (artifacts-server / dial-mpc), via
|
|
7
|
+
:func:`import_rbundle`;
|
|
8
|
+
* self-contained Brax HTML viewers (``var system = "..."``), via
|
|
9
|
+
:func:`import_brax_html`.
|
|
10
|
+
|
|
11
|
+
:func:`import_path` sniffs a file (or walks a directory) and picks the
|
|
12
|
+
importer. Conventions of each source are documented in the section
|
|
13
|
+
comments below. Verified against the reference code and the real data:
|
|
14
|
+
|
|
15
|
+
``.rbundle`` (``core/viz3d/bundle.py``)
|
|
16
|
+
``[4B "RBDL"][u64 LE header length][JSON header][tail]``; the tail is
|
|
17
|
+
gzip when ``header.compression == "gzip"``. ``buffers[k].off`` is a byte
|
|
18
|
+
offset into the *uncompressed* tail and ``count`` an element count.
|
|
19
|
+
``body_pos`` is ``[T, B, 3]`` and ``body_quat`` ``[T, B, 4]`` in
|
|
20
|
+
**wxyz** (MuJoCo ``xquat``), world frame, Z-up, metres. Body 0 is the
|
|
21
|
+
MuJoCo world body. Geoms carry MuJoCo ``geom_size`` conventions (the
|
|
22
|
+
same as MuJoCo: box half-extents, capsule/cylinder ``[r, h]``
|
|
23
|
+
with the axis along local z, plane ``[hx, hy, spacing]``) and
|
|
24
|
+
``local_quat`` in wxyz. Mesh geoms hold the vertices already in the
|
|
25
|
+
geom's local frame (``verts_count`` and ``faces_count`` count *elements*,
|
|
26
|
+
three per vertex or face) and their ``size`` is the mesh bounding box,
|
|
27
|
+
which is dropped. ``forces`` is ``[T, K, 2, 3]`` holding
|
|
28
|
+
``[anchor, vector]`` in world frame (Newtons); inactive feet are NaN. It
|
|
29
|
+
is recorded as the ``contacts`` stream (viewer contracts 7), scaled so
|
|
30
|
+
one body weight draws as 1 m.
|
|
31
|
+
``predictions`` is ``[T, H, L, 3]``. The scene has no body hierarchy,
|
|
32
|
+
so every body gets ``parent = -1``.
|
|
33
|
+
|
|
34
|
+
Brax HTML (``brax/io/json.py``, ``brax/visualizer/js/system.js``)
|
|
35
|
+
The page holds ``var system = "<base64 of zlib(JSON)>"`` with keys
|
|
36
|
+
``opt`` (``opt.timestep`` is the playback frame interval, so
|
|
37
|
+
``dt = opt.timestep``), ``link_names``, ``name``, ``geoms`` and
|
|
38
|
+
``states.x``. Z-up, metres, same as MuJoCo. ``states.x[t]`` has
|
|
39
|
+
``pos`` ``[L, 3]`` and ``rot`` ``[L, 4]`` (wxyz) for the ``L`` links, in
|
|
40
|
+
the world frame; the world body is not in the state, so our body 0 is a
|
|
41
|
+
fixed identity ``"world"`` and link ``i`` is body ``i + 1``. ``geoms``
|
|
42
|
+
maps a link name to its geoms; ``link_idx`` (``-1`` for the world, which
|
|
43
|
+
holds the floor and obstacles) gives the link, ``pos`` and ``rot``
|
|
44
|
+
(wxyz) place the geom in the link frame. ``name`` is the type
|
|
45
|
+
(``Plane``, ``Sphere``, ``Capsule``, ``Cylinder``, ``Box``, ``Mesh``;
|
|
46
|
+
``HeightMap`` is not drawn by Brax and is skipped) and ``size`` is
|
|
47
|
+
MuJoCo's ``geom_size``: box half-extents, sphere ``[r]``, capsule and
|
|
48
|
+
cylinder ``[r, half-length]`` with the axis along local z, plane
|
|
49
|
+
``[hx, hy, spacing]`` where zero half-extents mean infinite. Meshes carry
|
|
50
|
+
``vert`` ``[V, 3]`` and ``face`` ``[F, 3]`` in the geom's local frame
|
|
51
|
+
(``size`` is then the bounding box, which is dropped). Numbers are
|
|
52
|
+
rounded to six decimals, so quaternions are renormalized on import.
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
import base64
|
|
56
|
+
import concurrent.futures
|
|
57
|
+
import dataclasses
|
|
58
|
+
import datetime
|
|
59
|
+
import gzip
|
|
60
|
+
import hashlib
|
|
61
|
+
import json
|
|
62
|
+
import logging
|
|
63
|
+
import math
|
|
64
|
+
import os
|
|
65
|
+
import pathlib
|
|
66
|
+
import re
|
|
67
|
+
import struct
|
|
68
|
+
import zlib
|
|
69
|
+
from collections.abc import Collection, Iterator, Mapping, Sequence
|
|
70
|
+
from typing import Any
|
|
71
|
+
|
|
72
|
+
import numpy as np
|
|
73
|
+
import numpy.typing as npt
|
|
74
|
+
|
|
75
|
+
from simscope import core, library, transforms
|
|
76
|
+
|
|
77
|
+
logger = logging.getLogger(__name__)
|
|
78
|
+
|
|
79
|
+
RBUNDLE_MAGIC = b"RBDL"
|
|
80
|
+
_SMALL_JSON_BYTES = 64 * 1024
|
|
81
|
+
_GRAVITY = 9.81
|
|
82
|
+
_NAME_CHARS = re.compile(r"[^A-Za-z0-9._-]+")
|
|
83
|
+
_MAX_NAME = 120
|
|
84
|
+
"""Longest base name, leaving room for a de-duplication suffix."""
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class ImportFormatError(ValueError):
|
|
88
|
+
"""Raised when an input file is not a supported rollout."""
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
# -- shared helpers --
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def unique_run_name(
|
|
95
|
+
library_: library.Library,
|
|
96
|
+
base: str,
|
|
97
|
+
overwrite: bool = False,
|
|
98
|
+
taken: Collection[str] = (),
|
|
99
|
+
) -> str:
|
|
100
|
+
"""Builds a valid, unused run name from free text.
|
|
101
|
+
|
|
102
|
+
Args:
|
|
103
|
+
library_: The target library.
|
|
104
|
+
base: Desired name; invalid characters become ``_``.
|
|
105
|
+
overwrite: Allow the name of a run already in the library.
|
|
106
|
+
taken: Names to avoid in any case, such as those planned for other
|
|
107
|
+
files of the same batch.
|
|
108
|
+
|
|
109
|
+
Returns:
|
|
110
|
+
The sanitized name, with ``-2``, ``-3``, ... appended if taken.
|
|
111
|
+
"""
|
|
112
|
+
name = _NAME_CHARS.sub("_", base).strip("._-")[:_MAX_NAME] or "run"
|
|
113
|
+
candidate, n = name, 1
|
|
114
|
+
while candidate in taken or (
|
|
115
|
+
not overwrite and library_.run_dir(candidate).exists()
|
|
116
|
+
):
|
|
117
|
+
n += 1
|
|
118
|
+
candidate = f"{name}-{n}"
|
|
119
|
+
return candidate
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def default_name(path: os.PathLike[str] | str) -> str:
|
|
123
|
+
"""Returns the unsanitized run name for a file.
|
|
124
|
+
|
|
125
|
+
Args:
|
|
126
|
+
path: An importable file.
|
|
127
|
+
|
|
128
|
+
Returns:
|
|
129
|
+
The parent folder name for a file called ``rollout.rbundle``, else
|
|
130
|
+
the file stem.
|
|
131
|
+
"""
|
|
132
|
+
p = pathlib.Path(path).resolve()
|
|
133
|
+
return p.parent.name if p.name == "rollout.rbundle" else p.stem
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def _json_safe(obj: Any) -> Any:
|
|
137
|
+
"""Replaces non-finite floats by ``None`` so the value is strict JSON."""
|
|
138
|
+
if isinstance(obj, float):
|
|
139
|
+
return obj if math.isfinite(obj) else None
|
|
140
|
+
if isinstance(obj, Mapping):
|
|
141
|
+
return {str(k): _json_safe(v) for k, v in obj.items()}
|
|
142
|
+
if isinstance(obj, list | tuple):
|
|
143
|
+
return [_json_safe(v) for v in obj]
|
|
144
|
+
return obj
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
def _read_small_json(path: pathlib.Path) -> Any | None:
|
|
148
|
+
"""Reads a JSON file if it exists and is small, else returns ``None``."""
|
|
149
|
+
try:
|
|
150
|
+
if path.stat().st_size > _SMALL_JSON_BYTES:
|
|
151
|
+
return None
|
|
152
|
+
return _json_safe(json.loads(path.read_text(encoding="utf-8")))
|
|
153
|
+
except (OSError, ValueError):
|
|
154
|
+
return None
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
def _vec3(values: Sequence[float]) -> core.Vec3:
|
|
158
|
+
"""Converts three numbers to a float triple."""
|
|
159
|
+
x, y, z = values[:3]
|
|
160
|
+
return (float(x), float(y), float(z))
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
class _SceneBuilder:
|
|
164
|
+
"""Accumulates materials and de-duplicated meshes for one scene."""
|
|
165
|
+
|
|
166
|
+
def __init__(self) -> None:
|
|
167
|
+
self.geoms: list[core.Geom] = []
|
|
168
|
+
self.materials: list[core.Material] = []
|
|
169
|
+
self.meshes: list[core.Mesh] = []
|
|
170
|
+
self._material_ids: dict[core.Rgba, int] = {}
|
|
171
|
+
self._mesh_ids: dict[bytes, int] = {}
|
|
172
|
+
|
|
173
|
+
def material(self, rgba: Sequence[float]) -> int:
|
|
174
|
+
"""Returns the index of the material with this color."""
|
|
175
|
+
key = tuple(float(np.float32(c)) for c in rgba)
|
|
176
|
+
if len(key) != 4:
|
|
177
|
+
raise ImportFormatError(f"rgba needs 4 values, got {rgba!r}")
|
|
178
|
+
idx = self._material_ids.get(key)
|
|
179
|
+
if idx is None:
|
|
180
|
+
idx = len(self.materials)
|
|
181
|
+
self._material_ids[key] = idx
|
|
182
|
+
self.materials.append(core.Material(rgba=key))
|
|
183
|
+
return idx
|
|
184
|
+
|
|
185
|
+
def mesh(
|
|
186
|
+
self, vertices: npt.NDArray[np.float32], faces: npt.NDArray[np.uint32]
|
|
187
|
+
) -> int:
|
|
188
|
+
"""Returns the index of a mesh, adding it unless identical bytes exist.
|
|
189
|
+
|
|
190
|
+
Args:
|
|
191
|
+
vertices: ``[V, 3]`` float32.
|
|
192
|
+
faces: ``[F, 3]`` uint32.
|
|
193
|
+
|
|
194
|
+
Returns:
|
|
195
|
+
The mesh index in the scene.
|
|
196
|
+
"""
|
|
197
|
+
digest = hashlib.blake2b(digest_size=16)
|
|
198
|
+
digest.update(struct.pack("<II", len(vertices), len(faces)))
|
|
199
|
+
digest.update(vertices)
|
|
200
|
+
digest.update(faces)
|
|
201
|
+
key = digest.digest()
|
|
202
|
+
idx = self._mesh_ids.get(key)
|
|
203
|
+
if idx is None:
|
|
204
|
+
idx = len(self.meshes)
|
|
205
|
+
self._mesh_ids[key] = idx
|
|
206
|
+
self.meshes.append(core.Mesh(vertices, faces))
|
|
207
|
+
return idx
|
|
208
|
+
|
|
209
|
+
def build(self, bodies: Sequence[core.Body]) -> core.Scene:
|
|
210
|
+
"""Returns the finished scene."""
|
|
211
|
+
return core.Scene(
|
|
212
|
+
bodies=tuple(bodies),
|
|
213
|
+
geoms=tuple(self.geoms),
|
|
214
|
+
materials=tuple(self.materials) or (core.Material(),),
|
|
215
|
+
meshes=tuple(self.meshes),
|
|
216
|
+
)
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
class _Stream:
|
|
220
|
+
"""An extra stream to record next to ``body_pose``."""
|
|
221
|
+
|
|
222
|
+
def __init__(
|
|
223
|
+
self,
|
|
224
|
+
kind: core.StreamKind,
|
|
225
|
+
data: npt.NDArray[np.float32],
|
|
226
|
+
scale: float | None = None,
|
|
227
|
+
units: str | None = None,
|
|
228
|
+
) -> None:
|
|
229
|
+
self.kind = kind
|
|
230
|
+
self.data = data
|
|
231
|
+
self.scale = scale
|
|
232
|
+
self.units = units
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
_STAMP = re.compile(
|
|
236
|
+
r"(?<!\d)(20\d{2})(\d{2})(\d{2})[T_-]?(\d{2})(\d{2})(\d{2})(?:\d{3,6})?(?!\d)"
|
|
237
|
+
)
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def recorded_time(path: os.PathLike[str] | str) -> str:
|
|
241
|
+
"""Estimates when a source file was recorded, as a UTC timestamp.
|
|
242
|
+
|
|
243
|
+
A timestamp embedded in the file or its folder name wins (such as
|
|
244
|
+
``20260911-132437_x.html`` or ``crate-20260910T005846Z``), because
|
|
245
|
+
copying a file changes its modification time. Otherwise the file's
|
|
246
|
+
modification time is used.
|
|
247
|
+
|
|
248
|
+
Args:
|
|
249
|
+
path: The source file.
|
|
250
|
+
|
|
251
|
+
Returns:
|
|
252
|
+
``YYYY-MM-DDTHH:MM:SSZ`` in UTC.
|
|
253
|
+
"""
|
|
254
|
+
path = pathlib.Path(path)
|
|
255
|
+
for text in (path.name, path.parent.name):
|
|
256
|
+
m = _STAMP.search(text)
|
|
257
|
+
if m:
|
|
258
|
+
try:
|
|
259
|
+
year, month, day, hour, minute, second = map(int, m.groups())
|
|
260
|
+
when = datetime.datetime(year, month, day, hour, minute, second)
|
|
261
|
+
except ValueError:
|
|
262
|
+
continue
|
|
263
|
+
return when.strftime("%Y-%m-%dT%H:%M:%SZ")
|
|
264
|
+
mtime = datetime.datetime.fromtimestamp(path.stat().st_mtime, datetime.UTC)
|
|
265
|
+
return mtime.strftime("%Y-%m-%dT%H:%M:%SZ")
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def _write_run(
|
|
269
|
+
library_: library.Library,
|
|
270
|
+
name: str,
|
|
271
|
+
*,
|
|
272
|
+
scene: core.Scene,
|
|
273
|
+
dt: float,
|
|
274
|
+
poses: npt.NDArray[np.float32],
|
|
275
|
+
streams: Mapping[str, _Stream],
|
|
276
|
+
source: Mapping[str, Any],
|
|
277
|
+
tags: Sequence[str],
|
|
278
|
+
meta: Mapping[str, Any],
|
|
279
|
+
overwrite: bool,
|
|
280
|
+
created: str | None = None,
|
|
281
|
+
) -> None:
|
|
282
|
+
"""Records a whole run in one ``log_frames`` call."""
|
|
283
|
+
with library_.record(
|
|
284
|
+
name,
|
|
285
|
+
scene=scene,
|
|
286
|
+
dt=dt,
|
|
287
|
+
source=source,
|
|
288
|
+
tags=tags,
|
|
289
|
+
meta=meta,
|
|
290
|
+
overwrite=overwrite,
|
|
291
|
+
created=created,
|
|
292
|
+
) as rec:
|
|
293
|
+
for sname, stream in streams.items():
|
|
294
|
+
rec.add_stream(
|
|
295
|
+
sname,
|
|
296
|
+
stream.kind,
|
|
297
|
+
stream.data.shape[2:],
|
|
298
|
+
scale=stream.scale,
|
|
299
|
+
units=stream.units,
|
|
300
|
+
)
|
|
301
|
+
rec.log_frames(poses, **{k: s.data for k, s in streams.items()})
|
|
302
|
+
|
|
303
|
+
|
|
304
|
+
def _all_tags(user_tags: Sequence[str], auto: str) -> tuple[str, ...]:
|
|
305
|
+
"""User tags followed by the automatic source tag, without repeats."""
|
|
306
|
+
return tuple(dict.fromkeys([*user_tags, auto]))
|
|
307
|
+
|
|
308
|
+
|
|
309
|
+
def _poses(
|
|
310
|
+
pos: npt.NDArray[np.float32], quat_wxyz: npt.NDArray[np.float32]
|
|
311
|
+
) -> npt.NDArray[np.float32]:
|
|
312
|
+
"""Stacks ``[T, B, 3]`` positions and wxyz quats into xyzw poses."""
|
|
313
|
+
return np.concatenate(
|
|
314
|
+
(pos, transforms.wxyz_to_xyzw(quat_wxyz)), axis=-1, dtype=np.float32
|
|
315
|
+
)
|
|
316
|
+
|
|
317
|
+
|
|
318
|
+
# -- .rbundle --
|
|
319
|
+
|
|
320
|
+
_RB_KINDS = frozenset(
|
|
321
|
+
("box", "sphere", "capsule", "cylinder", "ellipsoid", "plane", "mesh")
|
|
322
|
+
)
|
|
323
|
+
_MAX_PREDICTION_LINKS = 32
|
|
324
|
+
|
|
325
|
+
|
|
326
|
+
def _parse_rbundle(
|
|
327
|
+
data: bytes | memoryview,
|
|
328
|
+
) -> tuple[dict[str, Any], memoryview]:
|
|
329
|
+
"""Splits an ``.rbundle`` into its JSON header and uncompressed tail.
|
|
330
|
+
|
|
331
|
+
Raises:
|
|
332
|
+
ImportFormatError: If the magic or header is invalid.
|
|
333
|
+
"""
|
|
334
|
+
view = memoryview(data)
|
|
335
|
+
if len(view) < 12 or bytes(view[:4]) != RBUNDLE_MAGIC:
|
|
336
|
+
raise ImportFormatError("not an .rbundle (bad magic)")
|
|
337
|
+
(hlen,) = struct.unpack("<Q", view[4:12])
|
|
338
|
+
if 12 + hlen > len(view):
|
|
339
|
+
raise ImportFormatError("truncated .rbundle header")
|
|
340
|
+
try:
|
|
341
|
+
header = json.loads(bytes(view[12 : 12 + hlen]))
|
|
342
|
+
except ValueError as exc:
|
|
343
|
+
raise ImportFormatError("invalid .rbundle header") from exc
|
|
344
|
+
tail = view[12 + hlen :]
|
|
345
|
+
if header.get("compression") == "gzip":
|
|
346
|
+
tail = memoryview(gzip.decompress(tail))
|
|
347
|
+
return header, tail
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
def _buffer(
|
|
351
|
+
tail: memoryview, spec: Mapping[str, Any], dtype: type[np.generic]
|
|
352
|
+
) -> npt.NDArray:
|
|
353
|
+
"""Views a header buffer spec as an array."""
|
|
354
|
+
arr = np.frombuffer(tail, dtype, count=spec["count"], offset=spec["off"])
|
|
355
|
+
return arr.reshape(spec["shape"])
|
|
356
|
+
|
|
357
|
+
|
|
358
|
+
def _rbundle_geom(
|
|
359
|
+
builder: _SceneBuilder,
|
|
360
|
+
gj: Mapping[str, Any],
|
|
361
|
+
tail: memoryview,
|
|
362
|
+
n_bodies: int,
|
|
363
|
+
) -> core.Geom | None:
|
|
364
|
+
"""Converts one header geom, or returns ``None`` if unsupported."""
|
|
365
|
+
kind = gj["type"]
|
|
366
|
+
if kind not in _RB_KINDS:
|
|
367
|
+
logger.warning(
|
|
368
|
+
"skipping geom %r of unsupported kind %r", gj["name"], kind
|
|
369
|
+
)
|
|
370
|
+
return None
|
|
371
|
+
if not 0 <= gj["body"] < n_bodies:
|
|
372
|
+
raise ImportFormatError(
|
|
373
|
+
f"geom {gj['name']!r} has bad body {gj['body']}"
|
|
374
|
+
)
|
|
375
|
+
mesh = None
|
|
376
|
+
size = _vec3(gj["size"])
|
|
377
|
+
if kind == "mesh":
|
|
378
|
+
spec = gj.get("mesh")
|
|
379
|
+
if spec is None:
|
|
380
|
+
return None
|
|
381
|
+
verts = np.frombuffer(
|
|
382
|
+
tail, np.float32, spec["verts_count"], spec["verts_off"]
|
|
383
|
+
).reshape(-1, 3)
|
|
384
|
+
faces = np.frombuffer(
|
|
385
|
+
tail, np.uint32, spec["faces_count"], spec["faces_off"]
|
|
386
|
+
).reshape(-1, 3)
|
|
387
|
+
mesh = builder.mesh(verts, faces)
|
|
388
|
+
size = (0.0, 0.0, 0.0) # the source stores the mesh's bounds here
|
|
389
|
+
quat = transforms.wxyz_to_xyzw(gj["local_quat"])
|
|
390
|
+
return core.Geom(
|
|
391
|
+
body=int(gj["body"]),
|
|
392
|
+
kind=kind,
|
|
393
|
+
size=_vec3(size),
|
|
394
|
+
pos=_vec3(gj["local_pos"]),
|
|
395
|
+
quat=(float(quat[0]), float(quat[1]), float(quat[2]), float(quat[3])),
|
|
396
|
+
material=builder.material(gj["rgba"]),
|
|
397
|
+
mesh=mesh,
|
|
398
|
+
role="collision" if gj.get("is_collision") else "visual",
|
|
399
|
+
name=str(gj.get("name", "")),
|
|
400
|
+
)
|
|
401
|
+
|
|
402
|
+
|
|
403
|
+
def _force_scale(
|
|
404
|
+
arrows: npt.NDArray[np.float32], meta: Mapping[str, Any]
|
|
405
|
+
) -> float:
|
|
406
|
+
"""Metres of arrow per Newton: one body weight draws as 1 m.
|
|
407
|
+
|
|
408
|
+
The mass comes from the bundle's metadata. Without it, the weight is
|
|
409
|
+
estimated as the mean total vertical force, since a supported robot's
|
|
410
|
+
ground reaction averages to its weight.
|
|
411
|
+
"""
|
|
412
|
+
mass = meta.get("robot_mass_kg") or meta.get("total_mass_kg")
|
|
413
|
+
if mass:
|
|
414
|
+
weight = max(float(mass), 1e-3) * _GRAVITY
|
|
415
|
+
else:
|
|
416
|
+
weight = float(arrows[..., 5].sum(axis=1).mean())
|
|
417
|
+
return 1.0 / weight if weight > 1.0 else 1.0 / _GRAVITY
|
|
418
|
+
|
|
419
|
+
|
|
420
|
+
def _rbundle_streams(
|
|
421
|
+
header: Mapping[str, Any],
|
|
422
|
+
tail: memoryview,
|
|
423
|
+
) -> dict[str, _Stream]:
|
|
424
|
+
"""Builds the contact-force and prediction streams in the header."""
|
|
425
|
+
bufs = header["buffers"]
|
|
426
|
+
streams: dict[str, _Stream] = {}
|
|
427
|
+
if "forces" in bufs:
|
|
428
|
+
forces = _buffer(tail, bufs["forces"], np.float32)
|
|
429
|
+
# [T, K, 2, 3] -> [T, K, 6]; inactive feet (NaN) become zero-length.
|
|
430
|
+
arrows = np.nan_to_num(forces.reshape(len(forces), -1, 6), nan=0.0)
|
|
431
|
+
streams["contacts"] = _Stream(
|
|
432
|
+
"arrows", arrows, _force_scale(arrows, header["meta"]), "N"
|
|
433
|
+
)
|
|
434
|
+
if "predictions" in bufs:
|
|
435
|
+
preds = _buffer(tail, bufs["predictions"], np.float32)
|
|
436
|
+
if preds.size and preds.ndim == 4:
|
|
437
|
+
n_links = min(preds.shape[2], _MAX_PREDICTION_LINKS)
|
|
438
|
+
for link in range(n_links):
|
|
439
|
+
streams[f"predictions_{link}"] = _Stream(
|
|
440
|
+
"polyline", np.ascontiguousarray(preds[:, :, link])
|
|
441
|
+
)
|
|
442
|
+
return {
|
|
443
|
+
k: _Stream(
|
|
444
|
+
s.kind,
|
|
445
|
+
np.ascontiguousarray(s.data[:, None], np.float32),
|
|
446
|
+
s.scale,
|
|
447
|
+
s.units,
|
|
448
|
+
)
|
|
449
|
+
for k, s in streams.items()
|
|
450
|
+
}
|
|
451
|
+
|
|
452
|
+
|
|
453
|
+
def import_rbundle(
|
|
454
|
+
lib: library.Library,
|
|
455
|
+
path: os.PathLike[str] | str,
|
|
456
|
+
*,
|
|
457
|
+
name: str | None = None,
|
|
458
|
+
tags: Sequence[str] = (),
|
|
459
|
+
overwrite: bool = False,
|
|
460
|
+
) -> str:
|
|
461
|
+
"""Imports an ``.rbundle`` into a library.
|
|
462
|
+
|
|
463
|
+
Args:
|
|
464
|
+
lib: The target library.
|
|
465
|
+
path: The ``.rbundle`` file. Sibling ``config.json`` and
|
|
466
|
+
``metrics.json`` (each under 64 KB) are copied into ``meta``.
|
|
467
|
+
name: Run name. Default: the parent folder name for a file called
|
|
468
|
+
``rollout.rbundle``, else the file stem. It is sanitized and
|
|
469
|
+
de-duplicated with a ``-N`` suffix unless ``overwrite``.
|
|
470
|
+
tags: Extra tags; ``source:rbundle`` is always added.
|
|
471
|
+
overwrite: Replace an existing run of the same name.
|
|
472
|
+
|
|
473
|
+
Returns:
|
|
474
|
+
The name of the new run.
|
|
475
|
+
|
|
476
|
+
Raises:
|
|
477
|
+
ImportFormatError: If the file is not a valid ``.rbundle``.
|
|
478
|
+
OSError: If the file cannot be read.
|
|
479
|
+
"""
|
|
480
|
+
path = pathlib.Path(path).resolve()
|
|
481
|
+
header, tail = _parse_rbundle(path.read_bytes())
|
|
482
|
+
bufs = header["buffers"]
|
|
483
|
+
pos = _buffer(tail, bufs["body_pos"], np.float32)
|
|
484
|
+
quat = _buffer(tail, bufs["body_quat"], np.float32)
|
|
485
|
+
n_bodies = pos.shape[1]
|
|
486
|
+
builder = _SceneBuilder()
|
|
487
|
+
for gj in header["geoms"]:
|
|
488
|
+
geom = _rbundle_geom(builder, gj, tail, n_bodies)
|
|
489
|
+
if geom is not None:
|
|
490
|
+
builder.geoms.append(geom)
|
|
491
|
+
meta_in = header.get("meta", {})
|
|
492
|
+
names = meta_in.get("body_names") or []
|
|
493
|
+
bodies = [
|
|
494
|
+
core.Body(
|
|
495
|
+
str(names[i])
|
|
496
|
+
if i < len(names)
|
|
497
|
+
else ("world" if i == 0 else f"body_{i}")
|
|
498
|
+
)
|
|
499
|
+
for i in range(n_bodies)
|
|
500
|
+
]
|
|
501
|
+
scene = builder.build(bodies)
|
|
502
|
+
meta: dict[str, Any] = {"rbundle": _json_safe(meta_in)}
|
|
503
|
+
for key in ("config", "metrics"):
|
|
504
|
+
value = _read_small_json(path.with_name(f"{key}.json"))
|
|
505
|
+
if value is not None:
|
|
506
|
+
meta[key] = value
|
|
507
|
+
run = unique_run_name(lib, name or default_name(path), overwrite)
|
|
508
|
+
_write_run(
|
|
509
|
+
lib,
|
|
510
|
+
run,
|
|
511
|
+
scene=scene,
|
|
512
|
+
dt=float(meta_in.get("dt") or 1.0 / float(meta_in.get("fps", 50.0))),
|
|
513
|
+
poses=_poses(pos, quat),
|
|
514
|
+
streams=_rbundle_streams(header, tail),
|
|
515
|
+
source={
|
|
516
|
+
"simulator": "mujoco",
|
|
517
|
+
"importer": "rbundle",
|
|
518
|
+
"original": str(path),
|
|
519
|
+
},
|
|
520
|
+
tags=_all_tags(tags, "source:rbundle"),
|
|
521
|
+
meta=meta,
|
|
522
|
+
overwrite=overwrite,
|
|
523
|
+
created=recorded_time(path),
|
|
524
|
+
)
|
|
525
|
+
return run
|
|
526
|
+
|
|
527
|
+
|
|
528
|
+
# -- Brax HTML --
|
|
529
|
+
|
|
530
|
+
_BRAX_KINDS: dict[str, core.GeomKind] = {
|
|
531
|
+
"Plane": "plane",
|
|
532
|
+
"Sphere": "sphere",
|
|
533
|
+
"Capsule": "capsule",
|
|
534
|
+
"Cylinder": "cylinder",
|
|
535
|
+
"Box": "box",
|
|
536
|
+
"Mesh": "mesh",
|
|
537
|
+
}
|
|
538
|
+
_BRAX_SYSTEM = re.compile(rb'var system = "([^"]*)"')
|
|
539
|
+
_BRAX_TITLE = re.compile(rb"<title>(.*?)</title>", re.DOTALL)
|
|
540
|
+
_SNIFF_BYTES = 64 * 1024
|
|
541
|
+
|
|
542
|
+
|
|
543
|
+
def _load_brax_system(data: bytes) -> tuple[dict[str, Any], str]:
|
|
544
|
+
"""Extracts the ``system`` JSON and page title from a Brax page.
|
|
545
|
+
|
|
546
|
+
Raises:
|
|
547
|
+
ImportFormatError: If the page has no decodable ``var system``.
|
|
548
|
+
"""
|
|
549
|
+
match = _BRAX_SYSTEM.search(data)
|
|
550
|
+
if match is None:
|
|
551
|
+
raise ImportFormatError("no `var system` in the HTML page")
|
|
552
|
+
try:
|
|
553
|
+
raw = zlib.decompress(base64.b64decode(match.group(1)))
|
|
554
|
+
system = json.loads(raw)
|
|
555
|
+
except (ValueError, zlib.error) as exc:
|
|
556
|
+
raise ImportFormatError("cannot decode `var system`") from exc
|
|
557
|
+
title = _BRAX_TITLE.search(data[:_SNIFF_BYTES])
|
|
558
|
+
text = title.group(1).decode("utf-8", "replace").strip() if title else ""
|
|
559
|
+
return system, text
|
|
560
|
+
|
|
561
|
+
|
|
562
|
+
def _unit_quats(values: Any) -> npt.NDArray[np.float32]:
|
|
563
|
+
"""Converts wxyz quaternions to normalized xyzw float32."""
|
|
564
|
+
q = np.asarray(values, dtype=np.float64)
|
|
565
|
+
norm = np.linalg.norm(q, axis=-1, keepdims=True)
|
|
566
|
+
q = q / np.where(norm > 0, norm, 1.0)
|
|
567
|
+
return transforms.wxyz_to_xyzw(q)
|
|
568
|
+
|
|
569
|
+
|
|
570
|
+
def _brax_geom(
|
|
571
|
+
builder: _SceneBuilder, gj: Mapping[str, Any], n_links: int
|
|
572
|
+
) -> core.Geom | None:
|
|
573
|
+
"""Converts one Brax geom, or returns ``None`` if it is not drawable."""
|
|
574
|
+
kind = _BRAX_KINDS.get(gj["name"])
|
|
575
|
+
if kind is None:
|
|
576
|
+
logger.warning("skipping Brax geom of type %r", gj["name"])
|
|
577
|
+
return None
|
|
578
|
+
link = int(gj.get("link_idx", -1))
|
|
579
|
+
if not -1 <= link < n_links:
|
|
580
|
+
raise ImportFormatError(f"geom has bad link_idx {link}")
|
|
581
|
+
mesh = None
|
|
582
|
+
size = _vec3([*gj["size"], 0.0, 0.0, 0.0])
|
|
583
|
+
if kind == "mesh":
|
|
584
|
+
verts = np.array(gj["vert"], dtype=np.float32).reshape(-1, 3)
|
|
585
|
+
faces = np.array(gj["face"], dtype=np.uint32).reshape(-1, 3)
|
|
586
|
+
mesh = builder.mesh(verts, faces)
|
|
587
|
+
size = (0.0, 0.0, 0.0)
|
|
588
|
+
quat = _unit_quats(gj["rot"])
|
|
589
|
+
return core.Geom(
|
|
590
|
+
body=link + 1,
|
|
591
|
+
kind=kind,
|
|
592
|
+
size=size,
|
|
593
|
+
pos=_vec3(gj["pos"]),
|
|
594
|
+
quat=(
|
|
595
|
+
float(quat[0]),
|
|
596
|
+
float(quat[1]),
|
|
597
|
+
float(quat[2]),
|
|
598
|
+
float(quat[3]),
|
|
599
|
+
),
|
|
600
|
+
material=builder.material(gj["rgba"]),
|
|
601
|
+
mesh=mesh,
|
|
602
|
+
name=f"{gj['name'].lower()}_{link}",
|
|
603
|
+
)
|
|
604
|
+
|
|
605
|
+
|
|
606
|
+
def _brax_poses(states: Sequence[Mapping[str, Any]]) -> npt.NDArray[np.float32]:
|
|
607
|
+
"""Stacks ``states.x`` into ``[T, 1 + L, 7]`` poses with a world body."""
|
|
608
|
+
pos = np.asarray([s["pos"] for s in states], dtype=np.float32)
|
|
609
|
+
quat = _unit_quats([s["rot"] for s in states])
|
|
610
|
+
poses = np.zeros((len(states), pos.shape[1] + 1, 7), np.float32)
|
|
611
|
+
poses[:, 0, 6] = 1.0
|
|
612
|
+
poses[:, 1:, :3] = pos
|
|
613
|
+
poses[:, 1:, 3:] = quat
|
|
614
|
+
return poses
|
|
615
|
+
|
|
616
|
+
|
|
617
|
+
def import_brax_html(
|
|
618
|
+
lib: library.Library,
|
|
619
|
+
path: os.PathLike[str] | str,
|
|
620
|
+
*,
|
|
621
|
+
name: str | None = None,
|
|
622
|
+
tags: Sequence[str] = (),
|
|
623
|
+
overwrite: bool = False,
|
|
624
|
+
) -> str:
|
|
625
|
+
"""Imports a self-contained Brax HTML viewer into a library.
|
|
626
|
+
|
|
627
|
+
Args:
|
|
628
|
+
lib: The target library.
|
|
629
|
+
path: The HTML file containing ``var system = "..."``.
|
|
630
|
+
name: Run name; default the file stem. Sanitized and de-duplicated
|
|
631
|
+
with a ``-N`` suffix unless ``overwrite``.
|
|
632
|
+
tags: Extra tags; ``source:brax`` is always added.
|
|
633
|
+
overwrite: Replace an existing run of the same name.
|
|
634
|
+
|
|
635
|
+
Returns:
|
|
636
|
+
The name of the new run.
|
|
637
|
+
|
|
638
|
+
Raises:
|
|
639
|
+
ImportFormatError: If the page has no valid Brax ``system``.
|
|
640
|
+
OSError: If the file cannot be read.
|
|
641
|
+
"""
|
|
642
|
+
path = pathlib.Path(path).resolve()
|
|
643
|
+
system, title = _load_brax_system(path.read_bytes())
|
|
644
|
+
states = system["states"]["x"]
|
|
645
|
+
if not states:
|
|
646
|
+
raise ImportFormatError("the Brax page has no frames")
|
|
647
|
+
poses = _brax_poses(states)
|
|
648
|
+
n_links = poses.shape[1] - 1
|
|
649
|
+
builder = _SceneBuilder()
|
|
650
|
+
for link_geoms in system["geoms"].values():
|
|
651
|
+
for gj in link_geoms:
|
|
652
|
+
geom = _brax_geom(builder, gj, n_links)
|
|
653
|
+
if geom is not None:
|
|
654
|
+
builder.geoms.append(geom)
|
|
655
|
+
names = [str(n) for n in system.get("link_names", [])]
|
|
656
|
+
names += [f"link {i}" for i in range(len(names), n_links)]
|
|
657
|
+
bodies = [core.Body("world"), *(core.Body(n) for n in names[:n_links])]
|
|
658
|
+
meta = {"brax": {"title": title, "name": system.get("name", "")}}
|
|
659
|
+
dt = float(system["opt"]["timestep"])
|
|
660
|
+
run = unique_run_name(lib, name or default_name(path), overwrite)
|
|
661
|
+
_write_run(
|
|
662
|
+
lib,
|
|
663
|
+
run,
|
|
664
|
+
scene=builder.build(bodies),
|
|
665
|
+
dt=dt,
|
|
666
|
+
poses=poses,
|
|
667
|
+
streams={},
|
|
668
|
+
source={
|
|
669
|
+
"simulator": "brax",
|
|
670
|
+
"importer": "brax",
|
|
671
|
+
"original": str(path),
|
|
672
|
+
},
|
|
673
|
+
tags=_all_tags(tags, "source:brax"),
|
|
674
|
+
meta=meta,
|
|
675
|
+
overwrite=overwrite,
|
|
676
|
+
created=recorded_time(path),
|
|
677
|
+
)
|
|
678
|
+
return run
|
|
679
|
+
|
|
680
|
+
|
|
681
|
+
# -- dispatch --
|
|
682
|
+
|
|
683
|
+
|
|
684
|
+
def _sniff(path: pathlib.Path) -> str | None:
|
|
685
|
+
"""Returns ``"rbundle"``, ``"brax"`` or ``None`` for a file."""
|
|
686
|
+
try:
|
|
687
|
+
with path.open("rb") as f:
|
|
688
|
+
head = f.read(_SNIFF_BYTES)
|
|
689
|
+
except OSError:
|
|
690
|
+
return None
|
|
691
|
+
if head.startswith(RBUNDLE_MAGIC):
|
|
692
|
+
return "rbundle"
|
|
693
|
+
if b'var system = "' in head:
|
|
694
|
+
return "brax"
|
|
695
|
+
return None
|
|
696
|
+
|
|
697
|
+
|
|
698
|
+
def find_importable(
|
|
699
|
+
paths: Sequence[os.PathLike[str] | str],
|
|
700
|
+
) -> list[tuple[pathlib.Path, str]]:
|
|
701
|
+
"""Finds importable files in files and directories (recursively).
|
|
702
|
+
|
|
703
|
+
Args:
|
|
704
|
+
paths: Files or directories. Files are sniffed by content; in
|
|
705
|
+
directories only ``*.rbundle`` and ``*.html`` files are tried.
|
|
706
|
+
|
|
707
|
+
Returns:
|
|
708
|
+
``(path, format)`` pairs, sorted by path, where format is
|
|
709
|
+
``"rbundle"`` or ``"brax"``.
|
|
710
|
+
|
|
711
|
+
Raises:
|
|
712
|
+
FileNotFoundError: If a path does not exist.
|
|
713
|
+
"""
|
|
714
|
+
found: dict[pathlib.Path, str] = {}
|
|
715
|
+
for p in map(pathlib.Path, paths):
|
|
716
|
+
if p.is_dir():
|
|
717
|
+
candidates = [
|
|
718
|
+
f
|
|
719
|
+
for f in p.rglob("*")
|
|
720
|
+
if f.suffix in (".rbundle", ".html") and f.is_file()
|
|
721
|
+
]
|
|
722
|
+
elif p.is_file():
|
|
723
|
+
candidates = [p]
|
|
724
|
+
else:
|
|
725
|
+
raise FileNotFoundError(f"no such file or folder: {p}")
|
|
726
|
+
for f in candidates:
|
|
727
|
+
fmt = _sniff(f)
|
|
728
|
+
if fmt is not None:
|
|
729
|
+
found[f.resolve()] = fmt
|
|
730
|
+
elif f == p:
|
|
731
|
+
logger.warning("%s is not a supported rollout file", f)
|
|
732
|
+
return sorted(found.items())
|
|
733
|
+
|
|
734
|
+
|
|
735
|
+
def import_one(
|
|
736
|
+
lib: library.Library,
|
|
737
|
+
path: os.PathLike[str] | str,
|
|
738
|
+
fmt: str | None = None,
|
|
739
|
+
**kwargs: Any,
|
|
740
|
+
) -> str:
|
|
741
|
+
"""Imports one file with the importer that matches its content.
|
|
742
|
+
|
|
743
|
+
Args:
|
|
744
|
+
lib: The target library.
|
|
745
|
+
path: The file.
|
|
746
|
+
fmt: ``"rbundle"`` or ``"brax"``, or ``None`` to sniff the file.
|
|
747
|
+
**kwargs: ``name``, ``tags`` and ``overwrite``, as for
|
|
748
|
+
:func:`import_rbundle`.
|
|
749
|
+
|
|
750
|
+
Returns:
|
|
751
|
+
The run name.
|
|
752
|
+
|
|
753
|
+
Raises:
|
|
754
|
+
ImportFormatError: If the file is not a supported rollout.
|
|
755
|
+
"""
|
|
756
|
+
fmt = fmt or _sniff(pathlib.Path(path))
|
|
757
|
+
if fmt == "rbundle":
|
|
758
|
+
return import_rbundle(lib, path, **kwargs)
|
|
759
|
+
if fmt == "brax":
|
|
760
|
+
return import_brax_html(lib, path, **kwargs)
|
|
761
|
+
raise ImportFormatError(f"{path} is not a supported rollout file")
|
|
762
|
+
|
|
763
|
+
|
|
764
|
+
def import_path(
|
|
765
|
+
lib: library.Library, path: os.PathLike[str] | str, **kwargs: Any
|
|
766
|
+
) -> list[str]:
|
|
767
|
+
"""Imports a file, or every importable file under a directory.
|
|
768
|
+
|
|
769
|
+
Args:
|
|
770
|
+
lib: The target library.
|
|
771
|
+
path: A ``.rbundle`` or Brax HTML file, or a directory searched
|
|
772
|
+
recursively.
|
|
773
|
+
**kwargs: ``tags`` and ``overwrite`` (and ``name`` for a single
|
|
774
|
+
file), as for :func:`import_rbundle`.
|
|
775
|
+
|
|
776
|
+
Returns:
|
|
777
|
+
The names of the new runs, in path order.
|
|
778
|
+
|
|
779
|
+
Raises:
|
|
780
|
+
FileNotFoundError: If ``path`` does not exist.
|
|
781
|
+
ImportFormatError: If ``path`` is a file of an unknown format.
|
|
782
|
+
"""
|
|
783
|
+
p = pathlib.Path(path)
|
|
784
|
+
if p.is_file() and _sniff(p) is None:
|
|
785
|
+
raise ImportFormatError(f"{p} is not a supported rollout file")
|
|
786
|
+
return [
|
|
787
|
+
import_one(lib, f, fmt, **kwargs) for f, fmt in find_importable([p])
|
|
788
|
+
]
|
|
789
|
+
|
|
790
|
+
|
|
791
|
+
# -- batches --
|
|
792
|
+
|
|
793
|
+
_MAX_JOBS = 8
|
|
794
|
+
"""Default worker cap: a Brax page needs a few hundred MB while parsed."""
|
|
795
|
+
|
|
796
|
+
|
|
797
|
+
@dataclasses.dataclass(frozen=True)
|
|
798
|
+
class ImportResult:
|
|
799
|
+
"""The outcome of importing one file.
|
|
800
|
+
|
|
801
|
+
Attributes:
|
|
802
|
+
path: The input file.
|
|
803
|
+
name: The new run, or ``None`` if the import failed.
|
|
804
|
+
n_frames: Frames in the new run (0 on failure).
|
|
805
|
+
error: A one-line message if the import failed, else ``None``.
|
|
806
|
+
"""
|
|
807
|
+
|
|
808
|
+
path: pathlib.Path
|
|
809
|
+
name: str | None
|
|
810
|
+
n_frames: int = 0
|
|
811
|
+
error: str | None = None
|
|
812
|
+
|
|
813
|
+
|
|
814
|
+
def _import_task(
|
|
815
|
+
task: tuple[pathlib.Path, pathlib.Path, str, str, tuple[str, ...], bool],
|
|
816
|
+
) -> ImportResult:
|
|
817
|
+
"""Imports one planned file; runs in a worker process."""
|
|
818
|
+
root, path, fmt, name, tags, overwrite = task
|
|
819
|
+
lib = library.Library(root)
|
|
820
|
+
try:
|
|
821
|
+
run = import_one(
|
|
822
|
+
lib, path, fmt, name=name, tags=tags, overwrite=overwrite
|
|
823
|
+
)
|
|
824
|
+
with lib.open(run) as rollout:
|
|
825
|
+
n_frames = rollout.manifest.n_frames
|
|
826
|
+
except (
|
|
827
|
+
Exception
|
|
828
|
+
) as exc: # isolation point: one bad file must not stop a batch
|
|
829
|
+
logger.debug("import of %s failed", path, exc_info=True)
|
|
830
|
+
return ImportResult(path, None, 0, f"{type(exc).__name__}: {exc}")
|
|
831
|
+
finally:
|
|
832
|
+
lib.close()
|
|
833
|
+
return ImportResult(path, run, n_frames)
|
|
834
|
+
|
|
835
|
+
|
|
836
|
+
def import_files(
|
|
837
|
+
lib: library.Library,
|
|
838
|
+
files: Sequence[tuple[pathlib.Path, str]],
|
|
839
|
+
*,
|
|
840
|
+
tags: Sequence[str] = (),
|
|
841
|
+
overwrite: bool = False,
|
|
842
|
+
jobs: int | None = None,
|
|
843
|
+
) -> Iterator[ImportResult]:
|
|
844
|
+
"""Imports many files, in parallel processes, yielding in input order.
|
|
845
|
+
|
|
846
|
+
Run names are planned up front, so workers never race for a name.
|
|
847
|
+
Writes into the shared content store are atomic, so concurrent workers
|
|
848
|
+
are safe. A file that fails is reported in its result; the others go on.
|
|
849
|
+
|
|
850
|
+
Args:
|
|
851
|
+
lib: The target library.
|
|
852
|
+
files: ``(path, format)`` pairs from :func:`find_importable`.
|
|
853
|
+
tags: Extra tags for every run.
|
|
854
|
+
overwrite: Replace existing runs of the same names.
|
|
855
|
+
jobs: Worker processes; ``None`` picks ``min(files, cpus, 8)``, and
|
|
856
|
+
1 imports in this process.
|
|
857
|
+
|
|
858
|
+
Yields:
|
|
859
|
+
One result per file, in the order given.
|
|
860
|
+
"""
|
|
861
|
+
taken: set[str] = set()
|
|
862
|
+
tasks = []
|
|
863
|
+
for path, fmt in files:
|
|
864
|
+
name = unique_run_name(lib, default_name(path), overwrite, taken)
|
|
865
|
+
taken.add(name)
|
|
866
|
+
tasks.append((lib.root, path, fmt, name, tuple(tags), overwrite))
|
|
867
|
+
if jobs is None:
|
|
868
|
+
jobs = min(len(tasks), os.cpu_count() or 1, _MAX_JOBS)
|
|
869
|
+
if jobs <= 1 or len(tasks) <= 1:
|
|
870
|
+
for task in tasks:
|
|
871
|
+
yield _import_task(task)
|
|
872
|
+
return
|
|
873
|
+
with concurrent.futures.ProcessPoolExecutor(jobs) as pool:
|
|
874
|
+
yield from pool.map(_import_task, tasks)
|