simscope 0.1.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (45) hide show
  1. simscope/__init__.py +6 -0
  2. simscope/__main__.py +8 -0
  3. simscope/_assets/simscope-app.css +2 -0
  4. simscope/_assets/simscope-app.js +4311 -0
  5. simscope/_assets/simscope-player.js +4325 -0
  6. simscope/_assets/simscope-web.LICENSES.txt +407 -0
  7. simscope/_icon.py +22 -0
  8. simscope/_mjviser.py +203 -0
  9. simscope/annotations.py +1132 -0
  10. simscope/cli.py +482 -0
  11. simscope/core.py +257 -0
  12. simscope/derived.py +697 -0
  13. simscope/export.py +799 -0
  14. simscope/highlights.py +947 -0
  15. simscope/importers.py +874 -0
  16. simscope/index.py +579 -0
  17. simscope/io/__init__.py +45 -0
  18. simscope/io/blockfile.py +938 -0
  19. simscope/io/cas.py +294 -0
  20. simscope/io/codecs.py +566 -0
  21. simscope/io/errors.py +9 -0
  22. simscope/io/manifest.py +358 -0
  23. simscope/io/pack.py +563 -0
  24. simscope/io/scene.py +239 -0
  25. simscope/isaaclab.py +1460 -0
  26. simscope/library.py +705 -0
  27. simscope/mujoco.py +578 -0
  28. simscope/py.typed +0 -0
  29. simscope/recorder.py +784 -0
  30. simscope/server/__init__.py +9 -0
  31. simscope/server/app.py +149 -0
  32. simscope/server/blocks.py +191 -0
  33. simscope/server/jobs.py +166 -0
  34. simscope/server/routes.py +707 -0
  35. simscope/server/security.py +218 -0
  36. simscope/server/state.py +751 -0
  37. simscope/server/static.py +84 -0
  38. simscope/transforms.py +147 -0
  39. simscope-0.1.1.dist-info/METADATA +132 -0
  40. simscope-0.1.1.dist-info/RECORD +45 -0
  41. simscope-0.1.1.dist-info/WHEEL +4 -0
  42. simscope-0.1.1.dist-info/entry_points.txt +3 -0
  43. simscope-0.1.1.dist-info/licenses/LICENSE.md +201 -0
  44. simscope-0.1.1.dist-info/licenses/THIRD_PARTY_NOTICES.md +267 -0
  45. simscope-0.1.1.dist-info/licenses/src/simscope/_assets/simscope-web.LICENSES.txt +407 -0
simscope/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)