simrig 0.2.2__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.
simrig/cli.py ADDED
@@ -0,0 +1,412 @@
1
+ """Command-line interface for SimRig."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ import sys
8
+ from pathlib import Path
9
+ from typing import Any
10
+
11
+ from simrig._version import __version__
12
+ from simrig.core import report_markdown, to_dict
13
+ from simrig.huggingface import resolve_policy_checkpoint
14
+ from simrig.io import save_json, save_report_pair, slugify
15
+ from simrig.mujoco_backend import inspect_model, list_models
16
+ from simrig.paths import ensure_project_dirs
17
+ from simrig.playground_backend import (
18
+ demo_policy,
19
+ eval_policy,
20
+ inspect_env,
21
+ list_envs,
22
+ smoke_env,
23
+ train_ppo,
24
+ )
25
+ from simrig.scaffold import new_env
26
+ from simrig.validate_env import validate_env
27
+
28
+
29
+ def main(argv: list[str] | None = None) -> int:
30
+ parser = build_parser()
31
+ args = parser.parse_args(argv)
32
+ try:
33
+ result = args.func(args)
34
+ except Exception as exc:
35
+ print(f"simrig: error: {exc}", file=sys.stderr)
36
+ return 1
37
+ return 0 if result is None else int(result)
38
+
39
+
40
+ def build_parser() -> argparse.ArgumentParser:
41
+ parser = argparse.ArgumentParser(
42
+ prog="simrig",
43
+ description="Physical AI simulation training starter workflows.",
44
+ )
45
+ parser.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
46
+ sub = parser.add_subparsers(dest="command", required=True)
47
+
48
+ init = sub.add_parser("init", help="Create local SimRig output folders.")
49
+ init.add_argument("--root", type=Path, default=Path("."))
50
+ init.set_defaults(func=_cmd_init)
51
+
52
+ list_models_parser = sub.add_parser("list-models", help="List MuJoCo Menagerie models.")
53
+ list_models_parser.add_argument("--menagerie", type=Path)
54
+ list_models_parser.add_argument("--json", action="store_true")
55
+ list_models_parser.set_defaults(func=_cmd_list_models)
56
+
57
+ inspect_model_parser = sub.add_parser("inspect-model", help="Inspect a MuJoCo XML/model.")
58
+ inspect_model_parser.add_argument("model_or_xml")
59
+ inspect_model_parser.add_argument("--menagerie", type=Path)
60
+ inspect_model_parser.add_argument("--steps", type=int, default=25)
61
+ inspect_model_parser.add_argument("--json", action="store_true")
62
+ inspect_model_parser.add_argument("--save-report", action="store_true")
63
+ inspect_model_parser.set_defaults(func=_cmd_inspect_model)
64
+
65
+ list_envs_parser = sub.add_parser("list-envs", help="List trainable backend envs.")
66
+ list_envs_parser.add_argument("--backend", default="mujoco-playground")
67
+ list_envs_parser.add_argument("--json", action="store_true")
68
+ list_envs_parser.set_defaults(func=_cmd_list_envs)
69
+
70
+ inspect_env_parser = sub.add_parser("inspect-env", help="Inspect a Playground or custom env.")
71
+ inspect_env_parser.add_argument(
72
+ "env_name",
73
+ help="Playground env name or path to a custom *.py env module.",
74
+ )
75
+ inspect_env_parser.add_argument("--backend", default="mujoco-playground")
76
+ inspect_env_parser.add_argument("--json", action="store_true")
77
+ inspect_env_parser.add_argument("--save-report", action="store_true")
78
+ inspect_env_parser.set_defaults(func=_cmd_inspect_env)
79
+
80
+ smoke_parser = sub.add_parser("smoke", help="Run a short env reset/step smoke test.")
81
+ smoke_parser.add_argument(
82
+ "env_name",
83
+ help="Playground env name or path to a custom *.py env module.",
84
+ )
85
+ smoke_parser.add_argument("--backend", default="mujoco-playground")
86
+ smoke_parser.add_argument("--steps", type=int, default=10)
87
+ smoke_parser.add_argument("--json", action="store_true")
88
+ smoke_parser.set_defaults(func=_cmd_smoke)
89
+
90
+ train_parser = sub.add_parser(
91
+ "train",
92
+ help="Train a Playground env or custom *.py module with Brax PPO.",
93
+ )
94
+ train_parser.add_argument(
95
+ "env_name",
96
+ help="Playground env name or path to a custom *.py env module.",
97
+ )
98
+ train_parser.add_argument("--backend", default="mujoco-playground")
99
+ train_parser.add_argument("--preset", choices=("smoke", "local", "cloud"), default="smoke")
100
+ train_parser.add_argument("--output", type=Path)
101
+ train_parser.add_argument("--timesteps", type=int)
102
+ train_parser.add_argument("--num-envs", type=int)
103
+ train_parser.add_argument("--batch-size", type=int)
104
+ train_parser.set_defaults(func=_cmd_train)
105
+
106
+ eval_parser = sub.add_parser("eval", help="Headless policy eval.")
107
+ eval_parser.add_argument("checkpoint")
108
+ eval_parser.add_argument(
109
+ "--env",
110
+ dest="env_name",
111
+ required=True,
112
+ help="Playground env name or path to a custom *.py env module.",
113
+ )
114
+ eval_parser.add_argument("--backend", default="mujoco-playground")
115
+ eval_parser.add_argument("--steps", type=int, default=500)
116
+ eval_parser.add_argument(
117
+ "--seed",
118
+ type=int,
119
+ default=0,
120
+ help="Deterministic environment and policy rollout seed.",
121
+ )
122
+ eval_parser.add_argument(
123
+ "--command",
124
+ type=float,
125
+ nargs="+",
126
+ help="Fix command-like environment state, for example X Y YAW.",
127
+ )
128
+ eval_parser.add_argument("--small-network", action=argparse.BooleanOptionalAction, default=None)
129
+ eval_parser.add_argument("--hf-revision", help="Revision for hf:// policy checkpoints.")
130
+ eval_parser.add_argument("--hf-token", help="Hugging Face token for private policy repos.")
131
+ eval_parser.add_argument("--json", action="store_true")
132
+ eval_parser.set_defaults(func=_cmd_eval)
133
+
134
+ demo_parser = sub.add_parser("demo", help="Run a trained policy in a desktop MuJoCo viewer.")
135
+ demo_parser.add_argument("checkpoint")
136
+ demo_parser.add_argument("--env", dest="env_name", required=True)
137
+ demo_parser.add_argument("--backend", default="mujoco-playground")
138
+ demo_parser.add_argument("--steps", type=int, default=5000)
139
+ demo_parser.add_argument("--small-network", action=argparse.BooleanOptionalAction, default=None)
140
+ demo_parser.add_argument("--hf-revision", help="Revision for hf:// policy checkpoints.")
141
+ demo_parser.add_argument("--hf-token", help="Hugging Face token for private policy repos.")
142
+ demo_parser.add_argument("--command", type=float, nargs="+")
143
+ demo_parser.add_argument("--speed", type=float, default=1.0)
144
+ demo_parser.add_argument("--camera-distance", type=float)
145
+ demo_parser.add_argument("--json", action="store_true")
146
+ demo_parser.set_defaults(func=_cmd_demo)
147
+
148
+ preview_parser = sub.add_parser("preview", help="Serve a trained policy preview in the browser.")
149
+ preview_parser.add_argument("checkpoint")
150
+ preview_parser.add_argument("--env", dest="env_name", required=True)
151
+ preview_parser.add_argument("--backend", default="mujoco-playground")
152
+ preview_parser.add_argument("--host", default="127.0.0.1")
153
+ preview_parser.add_argument("--port", type=int, default=8765)
154
+ preview_parser.add_argument("--width", type=int, default=960)
155
+ preview_parser.add_argument("--height", type=int, default=540)
156
+ preview_parser.add_argument("--frame-skip", type=int, default=1)
157
+ preview_parser.add_argument("--fps", type=int, default=24, help="Browser render loop target FPS.")
158
+ preview_parser.add_argument("--small-network", action=argparse.BooleanOptionalAction, default=None)
159
+ preview_parser.add_argument("--hf-revision", help="Revision for hf:// policy checkpoints.")
160
+ preview_parser.add_argument("--hf-token", help="Hugging Face token for private policy repos.")
161
+ preview_parser.add_argument("--command", type=float, nargs="+")
162
+ preview_parser.add_argument("--camera")
163
+ preview_parser.add_argument("--paused", action="store_true", help="Start the browser preview paused.")
164
+ preview_parser.add_argument(
165
+ "--render-mode",
166
+ choices=("threejs", "mujoco", "topdown"),
167
+ default="threejs",
168
+ help=(
169
+ "Browser render mode. threejs renders rollout geometry interactively in WebGL; "
170
+ "mujoco streams offscreen frames; topdown is a schematic debug view."
171
+ ),
172
+ )
173
+ preview_parser.set_defaults(func=_cmd_preview)
174
+
175
+ view_model_parser = sub.add_parser(
176
+ "view-model",
177
+ help="Serve a MuJoCo model in the browser with per-joint controls.",
178
+ )
179
+ view_model_parser.add_argument("model_or_xml")
180
+ view_model_parser.add_argument("--menagerie", type=Path)
181
+ view_model_parser.add_argument("--host", default="127.0.0.1")
182
+ view_model_parser.add_argument("--port", type=int, default=8766)
183
+ view_model_parser.add_argument("--width", type=int, default=960)
184
+ view_model_parser.add_argument("--height", type=int, default=540)
185
+ view_model_parser.add_argument("--fps", type=int, default=24, help="Browser render loop target FPS.")
186
+ view_model_parser.add_argument("--camera")
187
+ view_model_parser.add_argument(
188
+ "--render-mode",
189
+ choices=("threejs", "mujoco", "topdown"),
190
+ default="threejs",
191
+ help=(
192
+ "Browser render mode. threejs renders geometry interactively in WebGL; "
193
+ "mujoco streams offscreen frames; topdown is a schematic debug view."
194
+ ),
195
+ )
196
+ view_model_parser.set_defaults(func=_cmd_view_model)
197
+
198
+ new_env_parser = sub.add_parser("new-env", help="Create an editable custom env starter.")
199
+ new_env_parser.add_argument("name")
200
+ new_env_parser.add_argument("--model", required=True)
201
+ new_env_parser.add_argument("--template", default="mjx")
202
+ new_env_parser.add_argument("--root", type=Path, default=Path("envs"))
203
+ new_env_parser.set_defaults(func=_cmd_new_env)
204
+
205
+ validate_env_parser = sub.add_parser(
206
+ "validate-env",
207
+ help="Validate a custom env module (static checklist; optional runtime).",
208
+ )
209
+ validate_env_parser.add_argument("path", type=Path)
210
+ validate_env_parser.add_argument(
211
+ "--runtime",
212
+ action="store_true",
213
+ help="Import the module and run construct/reset/step checks when possible.",
214
+ )
215
+ validate_env_parser.add_argument("--json", action="store_true")
216
+ validate_env_parser.set_defaults(func=_cmd_validate_env)
217
+
218
+ return parser
219
+
220
+
221
+ def _cmd_init(args: argparse.Namespace) -> None:
222
+ paths = ensure_project_dirs(args.root)
223
+ for path in paths:
224
+ print(path)
225
+
226
+
227
+ def _cmd_list_models(args: argparse.Namespace) -> None:
228
+ entries = list_models(args.menagerie)
229
+ _print(entries, as_json=args.json)
230
+
231
+
232
+ def _cmd_inspect_model(args: argparse.Namespace) -> None:
233
+ report = inspect_model(args.model_or_xml, menagerie=args.menagerie, steps=args.steps)
234
+ if args.save_report:
235
+ md_path, json_path = save_report_pair(
236
+ report,
237
+ name=report.name,
238
+ title=f"Model Inspection: {report.name}",
239
+ )
240
+ print(f"saved {md_path}")
241
+ print(f"saved {json_path}")
242
+ _print_report(f"Model Inspection: {report.name}", report, as_json=args.json)
243
+
244
+
245
+ def _cmd_list_envs(args: argparse.Namespace) -> None:
246
+ _print(list_envs(args.backend), as_json=args.json)
247
+
248
+
249
+ def _cmd_inspect_env(args: argparse.Namespace) -> None:
250
+ report = inspect_env(args.env_name, backend=args.backend)
251
+ if args.save_report:
252
+ md_path, json_path = save_report_pair(
253
+ report,
254
+ name=report.name,
255
+ title=f"Env Inspection: {report.name}",
256
+ )
257
+ print(f"saved {md_path}")
258
+ print(f"saved {json_path}")
259
+ _print_report(f"Env Inspection: {report.name}", report, as_json=args.json)
260
+
261
+
262
+ def _cmd_smoke(args: argparse.Namespace) -> int:
263
+ result = smoke_env(args.env_name, backend=args.backend, steps=args.steps)
264
+ _print(result, as_json=args.json)
265
+ return 0 if result.passed else 1
266
+
267
+
268
+ def _cmd_train(args: argparse.Namespace) -> None:
269
+ overrides = _training_overrides(args)
270
+ run_config = train_ppo(
271
+ args.env_name,
272
+ backend=args.backend,
273
+ preset_name=args.preset,
274
+ output=args.output,
275
+ overrides=overrides,
276
+ )
277
+ print(f"saved run: {run_config.output_dir}")
278
+
279
+
280
+ def _cmd_eval(args: argparse.Namespace) -> None:
281
+ command = tuple(args.command) if args.command is not None else None
282
+ checkpoint = resolve_policy_checkpoint(
283
+ args.checkpoint,
284
+ hf_revision=args.hf_revision,
285
+ hf_token=args.hf_token,
286
+ )
287
+ result = eval_policy(
288
+ checkpoint,
289
+ env_name=args.env_name,
290
+ backend=args.backend,
291
+ steps=args.steps,
292
+ small_network=args.small_network,
293
+ seed=args.seed,
294
+ command=command,
295
+ )
296
+ save_json(Path("reports") / f"{slugify(args.env_name)}_eval.json", result)
297
+ _print(result, as_json=args.json)
298
+
299
+
300
+ def _cmd_demo(args: argparse.Namespace) -> None:
301
+ command = tuple(args.command) if args.command is not None else None
302
+ checkpoint = resolve_policy_checkpoint(
303
+ args.checkpoint,
304
+ hf_revision=args.hf_revision,
305
+ hf_token=args.hf_token,
306
+ )
307
+ result = demo_policy(
308
+ checkpoint,
309
+ env_name=args.env_name,
310
+ backend=args.backend,
311
+ steps=args.steps,
312
+ small_network=args.small_network,
313
+ command=command,
314
+ speed=args.speed,
315
+ camera_distance=args.camera_distance,
316
+ )
317
+ _print(result, as_json=args.json)
318
+
319
+
320
+ def _cmd_view_model(args: argparse.Namespace) -> None:
321
+ from simrig.model_view import serve_model_view
322
+
323
+ serve_model_view(
324
+ args.model_or_xml,
325
+ menagerie=args.menagerie,
326
+ host=args.host,
327
+ port=args.port,
328
+ width=args.width,
329
+ height=args.height,
330
+ render_mode=args.render_mode,
331
+ camera=args.camera,
332
+ fps=args.fps,
333
+ )
334
+
335
+
336
+ def _cmd_preview(args: argparse.Namespace) -> None:
337
+ from simrig.preview import serve_policy_preview
338
+
339
+ command = tuple(args.command) if args.command is not None else None
340
+ checkpoint = resolve_policy_checkpoint(
341
+ args.checkpoint,
342
+ hf_revision=args.hf_revision,
343
+ hf_token=args.hf_token,
344
+ )
345
+ serve_policy_preview(
346
+ checkpoint,
347
+ env_name=args.env_name,
348
+ backend=args.backend,
349
+ host=args.host,
350
+ port=args.port,
351
+ width=args.width,
352
+ height=args.height,
353
+ frame_skip=args.frame_skip,
354
+ small_network=args.small_network,
355
+ command=command,
356
+ camera=args.camera,
357
+ render_mode=args.render_mode,
358
+ paused=args.paused,
359
+ fps=args.fps,
360
+ )
361
+
362
+
363
+ def _cmd_new_env(args: argparse.Namespace) -> None:
364
+ path = new_env(args.name, args.model, template=args.template, root=args.root)
365
+ print(path)
366
+
367
+
368
+ def _cmd_validate_env(args: argparse.Namespace) -> int:
369
+ result = validate_env(args.path, runtime=bool(args.runtime))
370
+ _print(result, as_json=args.json)
371
+ if not args.json:
372
+ status = "passed" if result.passed else "failed"
373
+ print(f"validate-env: {status} (trainable={result.trainable})")
374
+ return 0 if result.passed else 1
375
+
376
+
377
+ def _training_overrides(args: argparse.Namespace) -> dict[str, Any]:
378
+ overrides: dict[str, Any] = {}
379
+ for cli_name, key in (
380
+ ("timesteps", "timesteps"),
381
+ ("num_envs", "num_envs"),
382
+ ("batch_size", "batch_size"),
383
+ ):
384
+ value = getattr(args, cli_name)
385
+ if value is not None:
386
+ overrides[key] = value
387
+ return overrides
388
+
389
+
390
+ def _print_report(title: str, report: Any, *, as_json: bool) -> None:
391
+ if as_json:
392
+ _print(report, as_json=True)
393
+ else:
394
+ print(report_markdown(title, report), end="")
395
+
396
+
397
+ def _print(value: Any, *, as_json: bool) -> None:
398
+ if as_json:
399
+ print(json.dumps(to_dict(value), indent=2, sort_keys=True))
400
+ return
401
+ if isinstance(value, list):
402
+ for item in value:
403
+ if isinstance(item, dict) and "name" in item:
404
+ print(item["name"])
405
+ else:
406
+ print(item)
407
+ return
408
+ print(json.dumps(to_dict(value), indent=2, sort_keys=True))
409
+
410
+
411
+ if __name__ == "__main__":
412
+ raise SystemExit(main())
simrig/core.py ADDED
@@ -0,0 +1,144 @@
1
+ """Backend-neutral SimRig data contracts."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import asdict, dataclass, field, is_dataclass
6
+ from enum import Enum
7
+ from pathlib import Path
8
+ from typing import Any
9
+
10
+
11
+ class TrainabilityStatus(str, Enum):
12
+ """Coarse status for what SimRig can safely do with an asset."""
13
+
14
+ UNKNOWN = "unknown"
15
+ INSPECTABLE = "inspectable"
16
+ SIMULATABLE = "simulatable"
17
+ TRAINABLE_EXISTING_ENV = "trainable_existing_env"
18
+ NEEDS_CUSTOM_ENV = "needs_custom_env"
19
+ FAILED = "failed"
20
+
21
+
22
+ @dataclass(frozen=True)
23
+ class BackendInfo:
24
+ """Information about a simulation/training backend."""
25
+
26
+ name: str
27
+ available: bool
28
+ version: str | None = None
29
+ detail: str | None = None
30
+
31
+
32
+ @dataclass(frozen=True)
33
+ class ModelInspectionReport:
34
+ """Summary of a MuJoCo model inspection."""
35
+
36
+ name: str
37
+ path: str
38
+ backend: str
39
+ status: TrainabilityStatus
40
+ compiled: bool
41
+ stepped: bool
42
+ bodies: int = 0
43
+ joints: int = 0
44
+ dofs: int = 0
45
+ actuators: int = 0
46
+ sensors: int = 0
47
+ keyframes: int = 0
48
+ has_freejoint: bool = False
49
+ has_mjx_hint: bool = False
50
+ warnings: list[str] = field(default_factory=list)
51
+ errors: list[str] = field(default_factory=list)
52
+ notes: list[str] = field(default_factory=list)
53
+
54
+
55
+ @dataclass(frozen=True)
56
+ class EnvInspectionReport:
57
+ """Summary of a training environment inspection."""
58
+
59
+ name: str
60
+ backend: str
61
+ status: TrainabilityStatus
62
+ available: bool
63
+ loaded: bool
64
+ observation_size: Any = None
65
+ action_size: int | None = None
66
+ xml_path: str | None = None
67
+ model_bodies: int | None = None
68
+ model_actuators: int | None = None
69
+ has_domain_randomizer: bool | None = None
70
+ warnings: list[str] = field(default_factory=list)
71
+ errors: list[str] = field(default_factory=list)
72
+ notes: list[str] = field(default_factory=list)
73
+
74
+
75
+ @dataclass(frozen=True)
76
+ class RunConfig:
77
+ """Resolved training or evaluation run metadata."""
78
+
79
+ env_name: str
80
+ backend: str
81
+ preset: str
82
+ output_dir: str
83
+ config: dict[str, Any] = field(default_factory=dict)
84
+ command: list[str] = field(default_factory=list)
85
+
86
+
87
+ @dataclass(frozen=True)
88
+ class SmokeResult:
89
+ """Result of a short environment smoke test."""
90
+
91
+ env_name: str
92
+ backend: str
93
+ steps_requested: int
94
+ steps_completed: int
95
+ passed: bool
96
+ action_size: int | None = None
97
+ observation_size: Any = None
98
+ final_reward: float | None = None
99
+ final_done: bool | None = None
100
+ errors: list[str] = field(default_factory=list)
101
+
102
+
103
+ def to_dict(value: Any) -> Any:
104
+ """Convert dataclasses, enums, paths, and containers into JSONable values."""
105
+ if isinstance(value, Enum):
106
+ return value.value
107
+ if isinstance(value, Path):
108
+ return str(value)
109
+ if is_dataclass(value):
110
+ return {key: to_dict(item) for key, item in asdict(value).items()}
111
+ if isinstance(value, dict):
112
+ return {str(key): to_dict(item) for key, item in value.items()}
113
+ if isinstance(value, (list, tuple)):
114
+ return [to_dict(item) for item in value]
115
+ if isinstance(value, (str, int, float, bool)) or value is None:
116
+ return value
117
+
118
+ # JAX/NumPy scalars and arrays.
119
+ item = getattr(value, "item", None)
120
+ if callable(item):
121
+ shape = getattr(value, "shape", None)
122
+ if shape == ():
123
+ return to_dict(item())
124
+ tolist = getattr(value, "tolist", None)
125
+ if callable(tolist):
126
+ return to_dict(tolist())
127
+
128
+ return value
129
+
130
+
131
+ def report_markdown(title: str, report: Any) -> str:
132
+ """Render a compact Markdown report for humans and agents."""
133
+ data = to_dict(report)
134
+ lines = [f"# {title}", ""]
135
+ for key, value in data.items():
136
+ label = key.replace("_", " ").title()
137
+ if isinstance(value, list):
138
+ rendered = ", ".join(str(item) for item in value) if value else "none"
139
+ else:
140
+ rendered = str(value)
141
+ lines.append(f"- **{label}:** {rendered}")
142
+ lines.append("")
143
+ return "\n".join(lines)
144
+
simrig/custom_env.py ADDED
@@ -0,0 +1,109 @@
1
+ """Load user-authored custom environment modules."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import importlib.util
6
+ import sys
7
+ from pathlib import Path
8
+ from types import ModuleType
9
+ from typing import Any
10
+
11
+
12
+ def is_env_module_path(env_ref: str | Path) -> bool:
13
+ """Return True when env_ref points at a custom *.py env module."""
14
+ text = str(env_ref).strip()
15
+ if not text.endswith(".py"):
16
+ return False
17
+ return True
18
+
19
+
20
+ def resolve_env_label(env_ref: str | Path) -> str:
21
+ """Stable label for reports/run dirs (basename without .py for modules)."""
22
+ if is_env_module_path(env_ref):
23
+ return Path(str(env_ref)).expanduser().stem
24
+ return str(env_ref)
25
+
26
+
27
+ def import_env_module(path: Path | str) -> ModuleType:
28
+ """Import a custom env module from a filesystem path."""
29
+ env_path = Path(path).expanduser().resolve()
30
+ if not env_path.is_file():
31
+ raise FileNotFoundError(f"Custom env module not found: {env_path}")
32
+ if env_path.suffix != ".py":
33
+ raise ValueError(f"Custom env module must be a .py file: {env_path}")
34
+
35
+ module_name = f"simrig_custom_env_{env_path.stem}_{abs(hash(str(env_path)))}"
36
+ spec = importlib.util.spec_from_file_location(module_name, env_path)
37
+ if spec is None or spec.loader is None:
38
+ raise ImportError(f"Could not load custom env module: {env_path}")
39
+ module = importlib.util.module_from_spec(spec)
40
+ sys.modules[module_name] = module
41
+ spec.loader.exec_module(module)
42
+ return module
43
+
44
+
45
+ def load_custom_env(
46
+ path: Path | str,
47
+ *,
48
+ class_name: str = "CustomEnv",
49
+ config_overrides: dict[str, Any] | None = None,
50
+ ) -> Any:
51
+ """Instantiate a custom env from a module path.
52
+
53
+ Supports:
54
+ - make_env(config_overrides=...) factory if present
55
+ - CustomEnv(config=...) / CustomEnv(config=..., config_overrides=...)
56
+ """
57
+ module = import_env_module(path)
58
+ overrides = dict(config_overrides or {})
59
+
60
+ make_env = getattr(module, "make_env", None)
61
+ if callable(make_env):
62
+ try:
63
+ return make_env(config_overrides=overrides or None)
64
+ except TypeError:
65
+ return make_env()
66
+
67
+ cls = getattr(module, class_name, None)
68
+ if cls is None:
69
+ raise AttributeError(
70
+ f"Custom env module {path} must define make_env() or class {class_name}."
71
+ )
72
+
73
+ config = None
74
+ default_config = getattr(module, "default_config", None)
75
+ if callable(default_config):
76
+ config = default_config()
77
+
78
+ if isinstance(config, dict):
79
+ merged = dict(config)
80
+ merged.update(overrides)
81
+ return _construct_env(cls, config=merged, config_overrides=None)
82
+
83
+ if config is not None and overrides:
84
+ return _construct_env(cls, config=config, config_overrides=overrides)
85
+ if config is not None:
86
+ return _construct_env(cls, config=config, config_overrides=None)
87
+ if overrides:
88
+ return _construct_env(cls, config=overrides, config_overrides=None)
89
+ return _construct_env(cls, config=None, config_overrides=None)
90
+
91
+
92
+ def _construct_env(cls: type, *, config: Any, config_overrides: dict[str, Any] | None) -> Any:
93
+ if config is None and config_overrides is None:
94
+ return cls()
95
+ if config_overrides is None:
96
+ try:
97
+ return cls(config=config)
98
+ except TypeError:
99
+ return cls(config)
100
+ try:
101
+ return cls(config=config, config_overrides=config_overrides)
102
+ except TypeError:
103
+ try:
104
+ return cls(config, config_overrides)
105
+ except TypeError as exc:
106
+ raise TypeError(
107
+ f"Could not construct {cls.__name__} with config/config_overrides. "
108
+ "Prefer make_env(config_overrides=...) in the module."
109
+ ) from exc