bohrin 0.1.0__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 (85) hide show
  1. bohrin/__init__.py +37 -0
  2. bohrin/_arrays.py +26 -0
  3. bohrin/_compat.py +39 -0
  4. bohrin/_plugins.py +39 -0
  5. bohrin/adapters/__init__.py +21 -0
  6. bohrin/adapters/_arraysource.py +198 -0
  7. bohrin/adapters/_mapping.py +213 -0
  8. bohrin/adapters/_video.py +157 -0
  9. bohrin/adapters/base.py +89 -0
  10. bohrin/adapters/hdf5.py +205 -0
  11. bohrin/adapters/lerobot.py +499 -0
  12. bohrin/adapters/numpy_dir.py +107 -0
  13. bohrin/adapters/registry.py +97 -0
  14. bohrin/adapters/rlds.py +224 -0
  15. bohrin/adapters/zarr_replay.py +134 -0
  16. bohrin/analysis/__init__.py +45 -0
  17. bohrin/analysis/confident_learning.py +72 -0
  18. bohrin/analysis/embeddings.py +135 -0
  19. bohrin/analysis/neighbors.py +171 -0
  20. bohrin/analysis/robust.py +70 -0
  21. bohrin/analysis/shapes.py +186 -0
  22. bohrin/analysis/twosample.py +110 -0
  23. bohrin/api.py +86 -0
  24. bohrin/bench/__init__.py +18 -0
  25. bohrin/bench/harness.py +166 -0
  26. bohrin/calibrate/__init__.py +7 -0
  27. bohrin/calibrate/collect.py +131 -0
  28. bohrin/calibrate/conformal.py +111 -0
  29. bohrin/calibrate/corpus.py +203 -0
  30. bohrin/calibrate/dynamics_model.py +177 -0
  31. bohrin/calibrate/fdr.py +68 -0
  32. bohrin/calibrate/gate.py +229 -0
  33. bohrin/cli.py +422 -0
  34. bohrin/config.py +114 -0
  35. bohrin/detectors/__init__.py +8 -0
  36. bohrin/detectors/_common.py +166 -0
  37. bohrin/detectors/base.py +118 -0
  38. bohrin/detectors/causal.py +189 -0
  39. bohrin/detectors/consistency.py +212 -0
  40. bohrin/detectors/coverage.py +364 -0
  41. bohrin/detectors/dynamics.py +207 -0
  42. bohrin/detectors/integrity.py +475 -0
  43. bohrin/detectors/kinematics.py +467 -0
  44. bohrin/detectors/label.py +182 -0
  45. bohrin/detectors/multimodality.py +200 -0
  46. bohrin/detectors/policy_data.py +344 -0
  47. bohrin/detectors/registry.py +63 -0
  48. bohrin/detectors/scale.py +200 -0
  49. bohrin/detectors/smoothness.py +177 -0
  50. bohrin/detectors/stats.py +235 -0
  51. bohrin/detectors/temporal.py +378 -0
  52. bohrin/detectors/vision.py +418 -0
  53. bohrin/encoders/__init__.py +31 -0
  54. bohrin/encoders/base.py +35 -0
  55. bohrin/encoders/dino.py +82 -0
  56. bohrin/encoders/tiled.py +49 -0
  57. bohrin/engine.py +173 -0
  58. bohrin/hub.py +111 -0
  59. bohrin/ir/__init__.py +51 -0
  60. bohrin/ir/episode.py +133 -0
  61. bohrin/ir/schema.py +180 -0
  62. bohrin/policy/__init__.py +9 -0
  63. bohrin/policy/loader.py +321 -0
  64. bohrin/policy/target.py +56 -0
  65. bohrin/profile/__init__.py +14 -0
  66. bohrin/profile/action_space.py +132 -0
  67. bohrin/profile/dataset_profile.py +224 -0
  68. bohrin/profile/episode_reservoir.py +130 -0
  69. bohrin/profile/online.py +206 -0
  70. bohrin/py.typed +0 -0
  71. bohrin/report/__init__.py +31 -0
  72. bohrin/report/base.py +22 -0
  73. bohrin/report/html.py +201 -0
  74. bohrin/report/messages.py +134 -0
  75. bohrin/report/model.py +222 -0
  76. bohrin/report/sarif.py +231 -0
  77. bohrin/report/tty.py +101 -0
  78. bohrin/synth/__init__.py +21 -0
  79. bohrin/synth/pipeline.py +198 -0
  80. bohrin/version.py +17 -0
  81. bohrin-0.1.0.dist-info/METADATA +250 -0
  82. bohrin-0.1.0.dist-info/RECORD +85 -0
  83. bohrin-0.1.0.dist-info/WHEEL +4 -0
  84. bohrin-0.1.0.dist-info/entry_points.txt +61 -0
  85. bohrin-0.1.0.dist-info/licenses/LICENSE +202 -0
bohrin/__init__.py ADDED
@@ -0,0 +1,37 @@
1
+ """bohrin — the health check-up for robot-learning datasets (Layer 1).
2
+
3
+ Point it at your demonstrations and get a plain-language list of the hidden defects before
4
+ you train. No simulator, no training, no ground truth, no upload — the data never leaves
5
+ the machine.
6
+
7
+ import bohrin
8
+ report = bohrin.scan("./my_teleop_data")
9
+ print(report.score)
10
+
11
+ See ``docs/`` for the full architecture. Phase 1: reads local LeRobot datasets (v2.1 + v3)
12
+ and runs a conformally-calibrated detector battery over the frozen Canonical IR.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ from bohrin.api import scan
18
+ from bohrin.detectors.base import AnalysisContext, Detector, Requirements
19
+ from bohrin.detectors.registry import register
20
+ from bohrin.ir.schema import Family, Severity
21
+ from bohrin.report.model import Cluster, Finding, Report
22
+ from bohrin.version import REPORT_SCHEMA_VERSION, __version__
23
+
24
+ __all__ = [
25
+ "REPORT_SCHEMA_VERSION",
26
+ "AnalysisContext",
27
+ "Cluster",
28
+ "Detector",
29
+ "Family",
30
+ "Finding",
31
+ "Report",
32
+ "Requirements",
33
+ "Severity",
34
+ "__version__",
35
+ "register",
36
+ "scan",
37
+ ]
bohrin/_arrays.py ADDED
@@ -0,0 +1,26 @@
1
+ """Shared numpy typing aliases used across the IR and detectors.
2
+
3
+ Keeping these in one place lets ``mypy --strict`` see precise array element types instead
4
+ of a bare ``np.ndarray`` (which is generic and would leak ``Any``).
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from typing import Any
10
+
11
+ import numpy as np
12
+ import numpy.typing as npt
13
+
14
+ # A 1-D or 2-D array of float64 (the canonical dtype for actions/proprio after ingest).
15
+ FloatArray = npt.NDArray[np.float64]
16
+
17
+ # Integer arrays (episode indices, counts).
18
+ IntArray = npt.NDArray[np.int64]
19
+
20
+ # Boolean masks (anomaly flags).
21
+ BoolArray = npt.NDArray[np.bool_]
22
+
23
+ # An array of *unknown* dtype — only for the adapter boundary, where a source file may hold
24
+ # uint8 pixels, int64 indices or float32 signals and the adapter's job is to coerce them.
25
+ # Nothing past Stage ② should use this: the IR is float64 by contract.
26
+ AnyArray = npt.NDArray[Any]
bohrin/_compat.py ADDED
@@ -0,0 +1,39 @@
1
+ """Small standard-library shims so the package runs on Python 3.10.
2
+
3
+ A large share of the LeRobot / ROS / CUDA world is still pinned to 3.10-3.11, and a
4
+ `requires-python` floor they cannot meet fails at ``pip install`` -- silently, from our
5
+ point of view, because the user never files an issue. The floor is therefore 3.10 and
6
+ this module carries the (very small) cost of that.
7
+
8
+ ``StrEnum`` (3.11+) and ``typing.Self`` (3.11+) are the only two constructs from later
9
+ Python that the rest of the package uses.
10
+
11
+ The ``StrEnum`` fallback reproduces the two behaviours we actually depend on: members
12
+ compare equal to their ``str`` value, and ``str(member)`` yields that value rather than
13
+ ``"Class.MEMBER"`` (the plain ``str, Enum`` mixin does the latter on 3.10, which would
14
+ corrupt every f-string in the report layer). Explicit methods, not assigned-from-``str``
15
+ dunders, because mypy types ``str.__format__`` as unbound and rejects the assignment.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import enum
21
+ import sys
22
+
23
+ if sys.version_info >= (3, 11):
24
+ StrEnum = enum.StrEnum
25
+ from typing import Self as Self
26
+ else: # pragma: no cover - exercised by the 3.10 CI leg
27
+ from typing_extensions import Self as Self
28
+
29
+ class StrEnum(str, enum.Enum):
30
+ """Backport of :class:`enum.StrEnum` (Python 3.11+)."""
31
+
32
+ def __str__(self) -> str:
33
+ return str.__str__(self)
34
+
35
+ def __format__(self, format_spec: str) -> str:
36
+ return str.__format__(self, format_spec)
37
+
38
+
39
+ __all__ = ["Self", "StrEnum"]
bohrin/_plugins.py ADDED
@@ -0,0 +1,39 @@
1
+ """Entry-point plugin discovery (docs/02 §10).
2
+
3
+ The standard-library ``importlib.metadata`` route (non-provisional since 3.10). Both
4
+ built-in adapters/detectors and third-party ones advertise themselves through the same
5
+ entry-point groups — there is no privileged path. A plugin whose ``.load()`` raises is
6
+ skipped with a warning rather than crashing the whole scan.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import warnings
12
+ from importlib.metadata import entry_points
13
+
14
+
15
+ def load_plugin_classes(group: str) -> dict[str, type]:
16
+ """Load every *class* advertised under ``group``, keyed by entry-point name.
17
+
18
+ Entries that fail to import, or that resolve to a non-class, are warned about and
19
+ skipped so one bad plugin can never take down the tool. Callers narrow the result to
20
+ the base they expect (adapter / detector) via ``issubclass``.
21
+ """
22
+ found: dict[str, type] = {}
23
+ for ep in entry_points(group=group):
24
+ try:
25
+ obj = ep.load()
26
+ except Exception as exc: # a plugin must never crash discovery
27
+ warnings.warn(
28
+ f"bohrin: failed to load plugin {ep.name!r} from {group!r}: {exc}",
29
+ stacklevel=2,
30
+ )
31
+ continue
32
+ if not isinstance(obj, type):
33
+ warnings.warn(
34
+ f"bohrin: plugin {ep.name!r} in {group!r} is not a class; skipping",
35
+ stacklevel=2,
36
+ )
37
+ continue
38
+ found[ep.name] = obj
39
+ return found
@@ -0,0 +1,21 @@
1
+ """Stage ① — the adapter layer (docs/02 §1, docs/01_DATA_LANDSCAPE.md)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from bohrin.adapters.base import Adapter, DatasetHandle, Sampler
6
+ from bohrin.adapters.registry import (
7
+ UnknownFormatError,
8
+ discover,
9
+ register_adapter,
10
+ select_adapter,
11
+ )
12
+
13
+ __all__ = [
14
+ "Adapter",
15
+ "DatasetHandle",
16
+ "Sampler",
17
+ "UnknownFormatError",
18
+ "discover",
19
+ "register_adapter",
20
+ "select_adapter",
21
+ ]
@@ -0,0 +1,198 @@
1
+ """Shared machinery for adapters whose source is "a bag of named arrays".
2
+
3
+ HDF5, NumPy directories and Zarr replay buffers differ in *how* you list and slice arrays,
4
+ but not in what happens next: run the schema mapper, slice per episode, wrap in a
5
+ :class:`StepView`. That common half lives here so the three adapters stay thin and the IR
6
+ construction is written — and tested — exactly once.
7
+
8
+ The contract a concrete adapter implements is :class:`EpisodeArrays`: given an episode
9
+ index, hand back its arrays. Everything else (mapping, dtype coercion, camera specs,
10
+ provenance) is done for it.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ from collections.abc import Iterator, Mapping, Sequence
16
+ from dataclasses import dataclass
17
+ from typing import Protocol
18
+
19
+ import numpy as np
20
+
21
+ from bohrin._arrays import AnyArray, FloatArray
22
+ from bohrin.adapters._mapping import ArrayInfo, SchemaMapping, infer_mapping
23
+ from bohrin.adapters.base import Sampler
24
+ from bohrin.ir.episode import Episode, StepView, TaskLabel
25
+ from bohrin.ir.schema import ActionSpace, CameraSpec, DatasetSchema, Provenance, SchemaHints
26
+
27
+
28
+ class EpisodeArrays(Protocol):
29
+ """A source of per-episode named arrays."""
30
+
31
+ def episode_keys(self) -> Sequence[str]:
32
+ """Stable identifiers for each episode, in order."""
33
+ ...
34
+
35
+ def arrays(self, episode_key: str) -> Mapping[str, AnyArray]:
36
+ """The named arrays for one episode. Called lazily, one episode at a time."""
37
+ ...
38
+
39
+ def task_for(self, episode_key: str) -> str | None:
40
+ """The language instruction for this episode, if the format carries one."""
41
+ ...
42
+
43
+
44
+ @dataclass(frozen=True, slots=True)
45
+ class _ImageFrames:
46
+ """A lazy per-frame view over a ``(T, H, W[, C])`` array — satisfies ``LazyImage``."""
47
+
48
+ data: AnyArray
49
+ index: int
50
+
51
+ @property
52
+ def shape(self) -> tuple[int, int, int]:
53
+ frame = self.data[self.index]
54
+ if frame.ndim == 2:
55
+ return (int(frame.shape[0]), int(frame.shape[1]), 1)
56
+ return (int(frame.shape[0]), int(frame.shape[1]), int(frame.shape[2]))
57
+
58
+ def array(self) -> FloatArray:
59
+ frame = np.asarray(self.data[self.index], dtype=np.float64)
60
+ return frame if frame.ndim == 3 else frame[:, :, None]
61
+
62
+
63
+ def _as_2d(arr: AnyArray) -> FloatArray:
64
+ """Coerce a per-step signal to ``(T, D)`` float64 — the IR's low-dimensional shape."""
65
+ out = np.asarray(arr, dtype=np.float64)
66
+ if out.ndim == 1:
67
+ return out.reshape(-1, 1)
68
+ if out.ndim > 2: # flatten trailing dims, e.g. (T, 2, 3) poses → (T, 6)
69
+ return out.reshape(out.shape[0], -1)
70
+ return out
71
+
72
+
73
+ def _as_1d(arr: AnyArray) -> FloatArray:
74
+ out = np.asarray(arr, dtype=np.float64).ravel()
75
+ return out
76
+
77
+
78
+ def build_schema(
79
+ mapping: SchemaMapping,
80
+ sample: Mapping[str, AnyArray],
81
+ *,
82
+ control_hz: float | None,
83
+ embodiment: str | None,
84
+ ) -> DatasetSchema:
85
+ """Derive the dataset-wide schema from one representative episode."""
86
+ action = _as_2d(sample[mapping.action])
87
+ proprio = _as_2d(sample[mapping.proprio]) if mapping.proprio is not None and mapping.proprio in sample else None
88
+ cameras = tuple(
89
+ CameraSpec(key=key, height=int(sample[key].shape[1]), width=int(sample[key].shape[2]))
90
+ for key in mapping.images
91
+ if key in sample and sample[key].ndim >= 3
92
+ )
93
+ return DatasetSchema(
94
+ action_dim=int(action.shape[1]),
95
+ action_space=ActionSpace.UNKNOWN, # custom containers rarely declare it; stay honest
96
+ proprio_dim=None if proprio is None else int(proprio.shape[1]),
97
+ cameras=cameras,
98
+ control_hz=control_hz,
99
+ embodiment=embodiment,
100
+ )
101
+
102
+
103
+ def infer_control_hz(sample: Mapping[str, AnyArray], mapping: SchemaMapping) -> float | None:
104
+ """Estimate the control rate from timestamps, if the source has any."""
105
+ if mapping.timestamp is None or mapping.timestamp not in sample:
106
+ return None
107
+ ts = _as_1d(sample[mapping.timestamp])
108
+ if ts.size < 2:
109
+ return None
110
+ dt = float(np.median(np.diff(ts)))
111
+ return 1.0 / dt if dt > 0 else None
112
+
113
+
114
+ class ArraySourceHandle:
115
+ """A :class:`DatasetHandle` over any :class:`EpisodeArrays` source."""
116
+
117
+ def __init__(
118
+ self,
119
+ source: EpisodeArrays,
120
+ *,
121
+ adapter_name: str,
122
+ uri: str,
123
+ declared: Mapping[str, object] | None = None,
124
+ embodiment: str | None = None,
125
+ no_vision: bool = False,
126
+ hints: SchemaHints | None = None,
127
+ splits: Mapping[str, Sequence[str]] | None = None,
128
+ ) -> None:
129
+ self._source = source
130
+ self._adapter_name = adapter_name
131
+ self._uri = uri
132
+ self._no_vision = no_vision
133
+ # episode key -> declared split name, inverted once so lookup is O(1) per episode.
134
+ self._split_of: dict[str, str] = {key: name for name, members in (splits or {}).items() for key in members}
135
+ self._hints = hints or SchemaHints.empty()
136
+ self._keys = list(source.episode_keys())
137
+ if not self._keys:
138
+ raise ValueError(f"{adapter_name}: no episodes found in {uri}")
139
+
140
+ first = source.arrays(self._keys[0])
141
+ infos = [ArrayInfo(key=k, shape=tuple(int(d) for d in v.shape)) for k, v in first.items()]
142
+ self._mapping = infer_mapping(infos, declared)
143
+ hz = infer_control_hz(first, self._mapping)
144
+ self._schema = build_schema(self._mapping, first, control_hz=hz, embodiment=embodiment)
145
+
146
+ @property
147
+ def mapping(self) -> SchemaMapping:
148
+ """The resolved role→key mapping (exposed for ``bohrin init`` and tests)."""
149
+ return self._mapping
150
+
151
+ def schema(self) -> DatasetSchema:
152
+ return self._schema
153
+
154
+ def profile_hints(self) -> SchemaHints:
155
+ return self._hints
156
+
157
+ def episode_count(self) -> int | None:
158
+ return len(self._keys)
159
+
160
+ def iter_episodes(self, *, sample: Sampler) -> Iterator[Episode]:
161
+ keep = set(sample.plan(len(self._keys)).tolist())
162
+ for index, key in enumerate(self._keys):
163
+ if index not in keep:
164
+ continue
165
+ yield self._episode(index, key)
166
+
167
+ def _episode(self, index: int, key: str) -> Episode:
168
+ arrays = self._source.arrays(key)
169
+ m = self._mapping
170
+ action = _as_2d(arrays[m.action])
171
+ n_steps = action.shape[0]
172
+
173
+ images: dict[str, Sequence[_ImageFrames]] = {}
174
+ depth: dict[str, Sequence[_ImageFrames]] = {}
175
+ if not self._no_vision:
176
+ for cam in m.images:
177
+ if cam in arrays:
178
+ images[cam] = [_ImageFrames(arrays[cam], t) for t in range(min(n_steps, len(arrays[cam])))]
179
+ for cam in m.depth:
180
+ if cam in arrays:
181
+ depth[cam] = [_ImageFrames(arrays[cam], t) for t in range(min(n_steps, len(arrays[cam])))]
182
+
183
+ steps = StepView(
184
+ action=action,
185
+ timestamp=_as_1d(arrays[m.timestamp]) if m.timestamp and m.timestamp in arrays else None,
186
+ proprio=_as_2d(arrays[m.proprio]) if m.proprio and m.proprio in arrays else None,
187
+ reward=_as_1d(arrays[m.reward]) if m.reward and m.reward in arrays else None,
188
+ images=images,
189
+ depth=depth,
190
+ )
191
+ task = self._source.task_for(key)
192
+ return Episode(
193
+ episode_id=key,
194
+ steps=steps,
195
+ source=Provenance(adapter=self._adapter_name, uri=self._uri, locator=key),
196
+ task=None if task is None else TaskLabel(text=task),
197
+ split=self._split_of.get(key),
198
+ )
@@ -0,0 +1,213 @@
1
+ """The schema mapper — how a *custom* container becomes Canonical IR (docs/02 §1.3).
2
+
3
+ LeRobot and RLDS declare their own schema, so their adapters read it. Everything else —
4
+ raw HDF5, a directory of ``.npz``, a Zarr replay buffer — is a bag of arrays with names
5
+ chosen by whoever recorded it. This module answers one question for those formats:
6
+
7
+ given these array names and shapes, which one is the action, which the proprioception,
8
+ which the timestamps, and which are cameras?
9
+
10
+ Two sources of truth, in priority order:
11
+
12
+ 1. **Declared** — a ``bohrin.yaml`` written by ``bohrin init``. Always wins. This is the
13
+ long-tail escape hatch: if inference is wrong, the user states the mapping once.
14
+ 2. **Inferred** — name matching against the conventions actually used across the public
15
+ corpora (robomimic, Diffusion Policy/UMI, DROID, and the ad-hoc layouts in between),
16
+ then a shape-based tie-break.
17
+
18
+ Inference is deliberately **conservative**: when nothing matches confidently we return
19
+ ``None`` and let the caller raise, rather than silently profiling the wrong array. A
20
+ detector battery pointed at the wrong column produces confident nonsense, which is worse
21
+ than an error message telling the user to run ``bohrin init``.
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ from collections.abc import Mapping, Sequence
27
+ from dataclasses import dataclass
28
+
29
+ #: Candidate names for each role, most specific first. Matching is case-insensitive and
30
+ #: compares both the full key and its last path segment, so ``data/demo_0/actions`` and
31
+ #: ``actions`` both resolve. Sources: robomimic (``actions``, ``obs/``), Diffusion Policy
32
+ #: and UMI replay buffers (``action``, ``state``, ``robot0_eef_pos``), DROID/RLDS
33
+ #: (``action``, ``observation/state``).
34
+ _ACTION_NAMES = ("action", "actions", "act", "action_dict", "cmd", "command")
35
+ _PROPRIO_NAMES = (
36
+ "observation.state",
37
+ "observation/state",
38
+ "proprio",
39
+ "proprioception",
40
+ "state",
41
+ "states",
42
+ "qpos",
43
+ "joint_positions",
44
+ "joint_position",
45
+ "robot_state",
46
+ "eef_pose",
47
+ "robot0_eef_pos",
48
+ "obs",
49
+ )
50
+ _TIMESTAMP_NAMES = ("timestamp", "timestamps", "time", "t", "stamp", "time_stamp")
51
+ _REWARD_NAMES = ("reward", "rewards", "r")
52
+ #: Substrings that mark an array as pixels rather than a low-dimensional signal.
53
+ _IMAGE_HINTS = ("image", "rgb", "camera", "cam", "img", "pixels", "wrist", "front", "side", "top")
54
+ _DEPTH_HINTS = ("depth", "disparity", "pointcloud", "point_cloud", "xyz")
55
+
56
+ #: An array is treated as pixels if it has this many dims (T, H, W[, C]) and a plausible
57
+ #: channel count. Shape is the tie-break when the *name* is uninformative.
58
+ _MIN_IMAGE_NDIM = 3
59
+ _IMAGE_CHANNELS = (1, 3, 4)
60
+
61
+
62
+ @dataclass(frozen=True, slots=True)
63
+ class ArrayInfo:
64
+ """What the mapper needs to know about one candidate array: its name and shape."""
65
+
66
+ key: str
67
+ shape: tuple[int, ...]
68
+
69
+ @property
70
+ def leaf(self) -> str:
71
+ """The last path segment, lowercased — ``data/demo_0/actions`` → ``actions``."""
72
+ return self.key.replace("\\", "/").rsplit("/", 1)[-1].lower()
73
+
74
+ @property
75
+ def ndim(self) -> int:
76
+ return len(self.shape)
77
+
78
+
79
+ @dataclass(frozen=True, slots=True)
80
+ class SchemaMapping:
81
+ """The resolved answer: which key plays which role."""
82
+
83
+ action: str
84
+ proprio: str | None = None
85
+ timestamp: str | None = None
86
+ reward: str | None = None
87
+ images: tuple[str, ...] = ()
88
+ depth: tuple[str, ...] = ()
89
+
90
+ @property
91
+ def used_keys(self) -> frozenset[str]:
92
+ """Every key this mapping claims — the caller may treat the rest as unused."""
93
+ named = {self.action, self.proprio, self.timestamp, self.reward}
94
+ return frozenset({k for k in named if k} | set(self.images) | set(self.depth))
95
+
96
+
97
+ class UnmappableDatasetError(ValueError):
98
+ """Raised when no array can be identified as the action — the one required column.
99
+
100
+ Carries the candidate keys, because the fix is always "tell me which one it is".
101
+ """
102
+
103
+ def __init__(self, keys: Sequence[str]) -> None:
104
+ listed = ", ".join(sorted(keys)[:12]) or "(none)"
105
+ super().__init__(
106
+ "could not identify the action array in this dataset. "
107
+ f"Candidate arrays: {listed}. "
108
+ "Run `bohrin init <path>` to declare the mapping in a bohrin.yaml, "
109
+ "or pass --format to select a different adapter."
110
+ )
111
+
112
+
113
+ def _match(info: ArrayInfo, names: Sequence[str]) -> int | None:
114
+ """Rank of the first matching name (lower is better), or ``None``."""
115
+ leaf, full = info.leaf, info.key.lower()
116
+ for rank, name in enumerate(names):
117
+ if leaf == name or full == name or full.endswith("/" + name):
118
+ return rank
119
+ return None
120
+
121
+
122
+ def _is_image(info: ArrayInfo) -> bool:
123
+ if any(h in info.key.lower() for h in _DEPTH_HINTS):
124
+ return False
125
+ if any(h in info.key.lower() for h in _IMAGE_HINTS):
126
+ return info.ndim >= _MIN_IMAGE_NDIM
127
+ # Name says nothing — fall back to shape: (T, H, W) or (T, H, W, C).
128
+ if info.ndim == 4:
129
+ return info.shape[-1] in _IMAGE_CHANNELS
130
+ return False
131
+
132
+
133
+ def _is_depth(info: ArrayInfo) -> bool:
134
+ return any(h in info.key.lower() for h in _DEPTH_HINTS) and info.ndim >= _MIN_IMAGE_NDIM
135
+
136
+
137
+ def _pick(arrays: Sequence[ArrayInfo], names: Sequence[str], *, exclude: frozenset[str]) -> str | None:
138
+ """The best-matching key for a role, or ``None`` if nothing matches."""
139
+ ranked: list[tuple[int, str]] = []
140
+ for a in arrays:
141
+ if a.key in exclude:
142
+ continue
143
+ rank = _match(a, names)
144
+ if rank is not None:
145
+ ranked.append((rank, a.key))
146
+ return min(ranked)[1] if ranked else None
147
+
148
+
149
+ def _widest_2d(arrays: Sequence[ArrayInfo], *, exclude: frozenset[str]) -> str | None:
150
+ """The widest ``(T, D)`` array — the shape-based fallback for the action column."""
151
+ candidates = [a for a in arrays if a.ndim == 2 and a.key not in exclude and not _is_image(a)]
152
+ if not candidates:
153
+ return None
154
+ return max(candidates, key=lambda a: (a.shape[1], a.key)).key
155
+
156
+
157
+ def infer_mapping(
158
+ arrays: Sequence[ArrayInfo],
159
+ declared: Mapping[str, object] | None = None,
160
+ *,
161
+ allow_shape_fallback: bool = True,
162
+ ) -> SchemaMapping:
163
+ """Resolve array names to IR roles. Declared entries always beat inference.
164
+
165
+ ``declared`` is the ``schema_map`` section of a ``bohrin.yaml``; recognized keys are
166
+ ``action``, ``proprio``, ``timestamp``, ``reward``, ``images`` and ``depth``.
167
+
168
+ Raises :class:`UnmappableDatasetError` if the action column cannot be identified —
169
+ guessing it would mean profiling an arbitrary array and reporting the result as fact.
170
+ """
171
+ declared = declared or {}
172
+ present = {a.key for a in arrays}
173
+
174
+ def declared_str(role: str) -> str | None:
175
+ value = declared.get(role)
176
+ return str(value) if isinstance(value, str) else None
177
+
178
+ def declared_seq(role: str) -> tuple[str, ...]:
179
+ value = declared.get(role)
180
+ if isinstance(value, str):
181
+ return (value,)
182
+ if isinstance(value, (list, tuple)):
183
+ return tuple(str(v) for v in value)
184
+ return ()
185
+
186
+ action = declared_str("action") or _pick(arrays, _ACTION_NAMES, exclude=frozenset())
187
+ if action is None and allow_shape_fallback:
188
+ action = _widest_2d(arrays, exclude=frozenset())
189
+ if action is None or action not in present:
190
+ raise UnmappableDatasetError(sorted(present))
191
+
192
+ claimed = {action}
193
+ proprio = declared_str("proprio") or _pick(arrays, _PROPRIO_NAMES, exclude=frozenset(claimed))
194
+ if proprio:
195
+ claimed.add(proprio)
196
+ timestamp = declared_str("timestamp") or _pick(arrays, _TIMESTAMP_NAMES, exclude=frozenset(claimed))
197
+ if timestamp:
198
+ claimed.add(timestamp)
199
+ reward = declared_str("reward") or _pick(arrays, _REWARD_NAMES, exclude=frozenset(claimed))
200
+ if reward:
201
+ claimed.add(reward)
202
+
203
+ images = declared_seq("images") or tuple(sorted(a.key for a in arrays if a.key not in claimed and _is_image(a)))
204
+ depth = declared_seq("depth") or tuple(sorted(a.key for a in arrays if a.key not in claimed and _is_depth(a)))
205
+
206
+ return SchemaMapping(
207
+ action=action,
208
+ proprio=proprio,
209
+ timestamp=timestamp,
210
+ reward=reward,
211
+ images=images,
212
+ depth=depth,
213
+ )