Skip to content

Commit 608d075

Browse files
committed
Add entry-point adapter discovery (real plugin hub)
External pip packages can now register World Model backends under the `world_model_ros2.adapters` entry-point group; load_model() / `world-model list` discover them with no edits to this repo. Built-ins take precedence (an entry point can't hijack dummy/remote/ijepa/vjepa2); a failing entry point warns instead of breaking the registry. - registry.py: built-in + register() + cached entry-point discovery. - test_registry.py: 6 tests (builtins, register, discovery, precedence, bad EP), added to the ROS-free CI unit job. - CONTRIBUTING + CHANGELOG: document shipping an adapter from your own package.
1 parent b8a43a0 commit 608d075

5 files changed

Lines changed: 144 additions & 10 deletions

File tree

.github/workflows/ci.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@ jobs:
2323
run: >-
2424
python -m pytest
2525
test/test_adapters.py test/test_jepa.py test/test_bench.py
26-
test/test_wire.py test/test_remote_server.py -q
26+
test/test_registry.py test/test_wire.py test/test_remote_server.py -q
2727
2828
# Full ROS 2 Jazzy build + colcon test inside the official container.
2929
colcon:

CHANGELOG.md

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,9 @@ project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
77
## [Unreleased]
88

99
### Added
10+
- **Entry-point adapter discovery** — external pip packages can add World Model
11+
backends under the `world_model_ros2.adapters` group; `load_model()` /
12+
`world-model list` pick them up with no edits to this repo.
1013
- **`vjepa2` adapter** — real V-JEPA 2 video encoder via `torch.hub`
1114
(`facebookresearch/vjepa2`): rolling-clip latent + temporal surprise. Verified
1215
on a 16 GB GPU (ViT-L, fp16). torch stays optional/lazy-imported.

CONTRIBUTING.md

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,19 @@ register("mymodel", MyAdapter)
6060
Add unit tests with a fake backend so they run in CI without the real weights
6161
(see `test/test_jepa.py`), and a `gpu_verify_*.py` script for the real model.
6262

63+
**Ship an adapter from your own package** — no edits here needed. Expose it
64+
under the `world_model_ros2.adapters` entry-point group:
65+
66+
```toml
67+
# pyproject.toml in your package
68+
[project.entry-points."world_model_ros2.adapters"]
69+
mymodel = "my_pkg.adapter:make_my_adapter"
70+
```
71+
72+
Once installed, `load_model("mymodel")` and `world-model list` pick it up
73+
automatically. Built-in names take precedence, so an entry point can't hijack
74+
`dummy`/`remote`/`ijepa`/`vjepa2`.
75+
6376
## Pull requests
6477

6578
- Branch off `main`; keep PRs focused.
Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,73 @@
1+
"""Registry: built-ins, register(), and entry-point discovery. No ROS/GPU."""
2+
import numpy as np
3+
4+
from world_model_py import registry
5+
from world_model_py.adapters import DummyAdapter
6+
from world_model_py.adapters.base import FuturePrediction, Observation, WorldModelAdapter
7+
8+
9+
class _Toy(WorldModelAdapter):
10+
name = "toy"
11+
12+
def predict_future(self, obs, action=None, horizon=8):
13+
return FuturePrediction(dt=0.1, latents=[np.zeros(2, np.float32)] * horizon, risk=0.5)
14+
15+
16+
def _reset():
17+
registry._EXTRA.clear()
18+
registry._DISCOVERED = None
19+
20+
21+
def test_builtins_present():
22+
_reset()
23+
for name in ("dummy", "remote", "ijepa", "vjepa2"):
24+
assert name in registry.available_models()
25+
26+
27+
def test_register_adds_and_loads():
28+
_reset()
29+
registry.register("toy", _Toy)
30+
assert "toy" in registry.available_models()
31+
assert isinstance(registry.load_model("toy"), _Toy)
32+
33+
34+
def test_unknown_raises():
35+
_reset()
36+
try:
37+
registry.load_model("nope")
38+
assert False
39+
except KeyError:
40+
pass
41+
42+
43+
def test_entry_point_discovery(monkeypatch):
44+
_reset()
45+
monkeypatch.setattr(registry, "_discover_entry_points", lambda: {"toy_ep": _Toy})
46+
assert "toy_ep" in registry.available_models()
47+
assert isinstance(registry.load_model("toy_ep"), _Toy)
48+
49+
50+
def test_builtin_beats_entry_point(monkeypatch):
51+
_reset()
52+
# an entry point cannot hijack a built-in name
53+
monkeypatch.setattr(registry, "_discover_entry_points", lambda: {"dummy": _Toy})
54+
assert isinstance(registry.load_model("dummy"), DummyAdapter)
55+
56+
57+
def test_bad_entry_point_is_skipped(monkeypatch):
58+
_reset()
59+
from importlib import metadata
60+
61+
class _BadEP:
62+
name = "broken"
63+
64+
def load(self):
65+
raise RuntimeError("boom")
66+
67+
monkeypatch.setattr(metadata, "entry_points", lambda **k: [_BadEP()])
68+
import warnings
69+
70+
with warnings.catch_warnings():
71+
warnings.simplefilter("ignore")
72+
models = registry.available_models() # must not raise
73+
assert "dummy" in models

world_model_py/world_model_py/registry.py

Lines changed: 54 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,28 @@
11
"""Adapter registry: ``load_model("dummy")`` -> a WorldModelAdapter.
22
3-
This is the ``transformers``-style entry point for the project. New backends
4-
register a factory here (or, later, via package entry points) and become
5-
reachable by name from the CLI, the runtime node and the benchmark.
3+
This is the ``transformers``-style entry point for the project. Backends become
4+
reachable by name from the CLI, the runtime node and the benchmark in three ways:
5+
6+
1. built-in adapters (dummy / remote / ijepa / vjepa2),
7+
2. in-process ``register("name", factory)``,
8+
3. **package entry points** — an external pip package exposes adapters under the
9+
``world_model_ros2.adapters`` group, so installing it makes new World Models
10+
available with no edits here. This is what makes it an adapter *hub*.
11+
12+
In your package's pyproject.toml:
13+
14+
[project.entry-points."world_model_ros2.adapters"]
15+
mymodel = "my_pkg.adapter:make_my_adapter"
616
"""
717
from __future__ import annotations
818

19+
import warnings
920
from typing import Callable, Dict
1021

1122
from .adapters import DummyAdapter, RemoteAdapter, WorldModelAdapter
1223

24+
ENTRY_POINT_GROUP = "world_model_ros2.adapters"
25+
1326

1427
def _make_ijepa(**kwargs) -> WorldModelAdapter:
1528
# Lazy: importing jepa pulls torch/transformers only when actually used.
@@ -25,32 +38,64 @@ def _make_vjepa2(**kwargs) -> WorldModelAdapter:
2538
return make_vjepa2_adapter(**kwargs)
2639

2740

28-
# name -> factory(**kwargs) -> adapter
29-
_REGISTRY: Dict[str, Callable[..., WorldModelAdapter]] = {
41+
# built-in name -> factory(**kwargs) -> adapter
42+
_BUILTIN: Dict[str, Callable[..., WorldModelAdapter]] = {
3043
"dummy": DummyAdapter,
3144
"remote": RemoteAdapter,
3245
"ijepa": _make_ijepa,
3346
"vjepa2": _make_vjepa2,
3447
}
3548

49+
# adapters added at runtime via register()
50+
_EXTRA: Dict[str, Callable[..., WorldModelAdapter]] = {}
51+
52+
# cache of entry-point-discovered factories (None until first discovery)
53+
_DISCOVERED: Dict[str, Callable[..., WorldModelAdapter]] | None = None
54+
55+
56+
def _discover_entry_points() -> Dict[str, Callable[..., WorldModelAdapter]]:
57+
"""Load adapter factories advertised by installed packages. Cached; a single
58+
bad entry point warns rather than breaking the whole registry."""
59+
from importlib.metadata import entry_points
60+
61+
found: Dict[str, Callable[..., WorldModelAdapter]] = {}
62+
try:
63+
eps = entry_points(group=ENTRY_POINT_GROUP)
64+
except TypeError: # very old importlib.metadata
65+
eps = entry_points().get(ENTRY_POINT_GROUP, []) # type: ignore[attr-defined]
66+
for ep in eps:
67+
try:
68+
found[ep.name] = ep.load()
69+
except Exception as exc: # noqa: BLE001
70+
warnings.warn(f"failed to load adapter entry point '{ep.name}': {exc}")
71+
return found
72+
73+
74+
def _registry() -> Dict[str, Callable[..., WorldModelAdapter]]:
75+
global _DISCOVERED
76+
if _DISCOVERED is None:
77+
_DISCOVERED = _discover_entry_points()
78+
# precedence: discovered < built-in < explicit register()
79+
return {**_DISCOVERED, **_BUILTIN, **_EXTRA}
80+
3681

3782
def register(name: str, factory: Callable[..., WorldModelAdapter]) -> None:
3883
"""Register a new adapter factory under ``name`` (overwrites if present)."""
39-
_REGISTRY[name] = factory
84+
_EXTRA[name] = factory
4085

4186

4287
def available_models() -> list[str]:
43-
return sorted(_REGISTRY)
88+
return sorted(_registry())
4489

4590

4691
def load_model(name: str, **kwargs) -> WorldModelAdapter:
4792
"""Instantiate the adapter registered under ``name``.
4893
49-
Extra keyword arguments are passed to the adapter constructor, e.g.
94+
Extra keyword arguments are passed to the adapter factory, e.g.
5095
``load_model("remote", url="http://gpu-box:8080/predict_future")``.
5196
"""
5297
try:
53-
factory = _REGISTRY[name]
98+
factory = _registry()[name]
5499
except KeyError:
55100
raise KeyError(
56101
f"unknown world model '{name}'. available: {available_models()}"

0 commit comments

Comments
 (0)