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/__init__.py +82 -0
- simrig/_version.py +3 -0
- simrig/browser_render.py +145 -0
- simrig/browser_shell.py +130 -0
- simrig/cli.py +412 -0
- simrig/core.py +144 -0
- simrig/custom_env.py +109 -0
- simrig/huggingface.py +78 -0
- simrig/io.py +49 -0
- simrig/live_view.py +553 -0
- simrig/model_view.py +944 -0
- simrig/mujoco_backend.py +197 -0
- simrig/paths.py +54 -0
- simrig/playground_backend.py +603 -0
- simrig/presets.py +93 -0
- simrig/preview.py +956 -0
- simrig/rendering.py +127 -0
- simrig/scaffold.py +127 -0
- simrig/three_scene.py +107 -0
- simrig/validate_env.py +211 -0
- simrig-0.2.2.dist-info/METADATA +238 -0
- simrig-0.2.2.dist-info/RECORD +26 -0
- simrig-0.2.2.dist-info/WHEEL +5 -0
- simrig-0.2.2.dist-info/entry_points.txt +2 -0
- simrig-0.2.2.dist-info/licenses/LICENSE +21 -0
- simrig-0.2.2.dist-info/top_level.txt +1 -0
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
|