Spaces:
Sleeping
Sleeping
Fix: render rollouts in a subprocess with a probed GL backend (EGL import broke training)
#9
by arminfg - opened
- app.py +2 -4
- packages.txt +3 -0
- render_worker.py +30 -0
- rollout.py +59 -22
app.py
CHANGED
|
@@ -17,8 +17,6 @@ from pathlib import Path
|
|
| 17 |
|
| 18 |
import gradio as gr
|
| 19 |
|
| 20 |
-
os.environ.setdefault("MUJOCO_GL", "egl") # offscreen rendering for rollouts
|
| 21 |
-
|
| 22 |
ENV_NAME = os.environ.get("ENV_NAME", "G1JoystickRoughTerrain")
|
| 23 |
NUM_TIMESTEPS = int(os.environ.get("NUM_TIMESTEPS", 200_000_000))
|
| 24 |
SEED = int(os.environ.get("SEED", 0))
|
|
@@ -340,8 +338,8 @@ def render(ckpt: str | None, vx: float, friction: float, depth: float, seconds:
|
|
| 340 |
if SNOW:
|
| 341 |
apply_snow(eval_env, float(friction), float(depth))
|
| 342 |
log(f"rollout {ckpt}: vx={vx} friction={friction} depth={depth} m, {seconds}s ...")
|
| 343 |
-
|
| 344 |
-
out = write_video(
|
| 345 |
verdict = (f"fell at {info['fell_at']:.1f}s" if info["fell_at"] is not None
|
| 346 |
else f"stayed up for {info['seconds']:.1f}s")
|
| 347 |
msg = f"{ckpt}: {verdict}, walked {info['distance_m']:.2f} m (commanded {vx} m/s)"
|
|
|
|
| 17 |
|
| 18 |
import gradio as gr
|
| 19 |
|
|
|
|
|
|
|
| 20 |
ENV_NAME = os.environ.get("ENV_NAME", "G1JoystickRoughTerrain")
|
| 21 |
NUM_TIMESTEPS = int(os.environ.get("NUM_TIMESTEPS", 200_000_000))
|
| 22 |
SEED = int(os.environ.get("SEED", 0))
|
|
|
|
| 338 |
if SNOW:
|
| 339 |
apply_snow(eval_env, float(friction), float(depth))
|
| 340 |
log(f"rollout {ckpt}: vx={vx} friction={friction} depth={depth} m, {seconds}s ...")
|
| 341 |
+
qpos, info = rollout(eval_env, policy, seconds=float(seconds), command=(float(vx), 0.0, 0.0))
|
| 342 |
+
out = write_video(eval_env, qpos, OUT / "rollouts" / f"{path.stem}.mp4", fps=1.0 / eval_env.dt)
|
| 343 |
verdict = (f"fell at {info['fell_at']:.1f}s" if info["fell_at"] is not None
|
| 344 |
else f"stayed up for {info['seconds']:.1f}s")
|
| 345 |
msg = f"{ckpt}: {verdict}, walked {info['distance_m']:.2f} m (commanded {vx} m/s)"
|
packages.txt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
libegl1
|
| 2 |
+
libgl1
|
| 3 |
+
libosmesa6
|
render_worker.py
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Render a qpos trajectory to MP4. Run as a subprocess by rollout.write_video
|
| 2 |
+
with MUJOCO_GL already set, because MuJoCo fixes its GL backend at import.
|
| 3 |
+
|
| 4 |
+
python render_worker.py model.mjb qpos.npy out.mp4 fps width height camera
|
| 5 |
+
"""
|
| 6 |
+
import sys
|
| 7 |
+
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def main(mjb, qpos_path, out, fps, width, height, camera):
|
| 12 |
+
import imageio.v2 as imageio
|
| 13 |
+
import mujoco
|
| 14 |
+
|
| 15 |
+
model = mujoco.MjModel.from_binary_path(mjb)
|
| 16 |
+
data = mujoco.MjData(model)
|
| 17 |
+
qpos = np.load(qpos_path)
|
| 18 |
+
cam = camera if mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_CAMERA, camera) >= 0 else -1
|
| 19 |
+
with mujoco.Renderer(model, height=int(height), width=int(width)) as r, \
|
| 20 |
+
imageio.get_writer(out, fps=float(fps), codec="libx264", quality=7,
|
| 21 |
+
macro_block_size=None) as w:
|
| 22 |
+
for q in qpos:
|
| 23 |
+
data.qpos[:] = q
|
| 24 |
+
mujoco.mj_forward(model, data)
|
| 25 |
+
r.update_scene(data, camera=cam)
|
| 26 |
+
w.append_data(r.render())
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
if __name__ == "__main__":
|
| 30 |
+
main(*sys.argv[1:8])
|
rollout.py
CHANGED
|
@@ -11,6 +11,7 @@ same env, then `make_inference_fn(params, deterministic=True)`.
|
|
| 11 |
|
| 12 |
from __future__ import annotations
|
| 13 |
|
|
|
|
| 14 |
import pickle
|
| 15 |
from pathlib import Path
|
| 16 |
|
|
@@ -58,11 +59,10 @@ def apply_snow(env, friction: float, depth: float) -> None:
|
|
| 58 |
env._mjx_model = mjx.put_model(mj, impl=env._config.impl)
|
| 59 |
|
| 60 |
|
| 61 |
-
def rollout(env, inference_fn, seconds: float = 8.0, command=(0.5, 0.0, 0.0),
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
rate, info has distance walked and whether it fell."""
|
| 66 |
import jax
|
| 67 |
import jax.numpy as jnp
|
| 68 |
|
|
@@ -76,33 +76,70 @@ def rollout(env, inference_fn, seconds: float = 8.0, command=(0.5, 0.0, 0.0),
|
|
| 76 |
state.info["command"] = cmd
|
| 77 |
|
| 78 |
n_steps = int(seconds / env.dt)
|
| 79 |
-
|
| 80 |
-
start = np.asarray(state.data.qpos[:3])
|
| 81 |
fell_at = None
|
| 82 |
for i in range(n_steps):
|
| 83 |
rng, key = jax.random.split(rng)
|
| 84 |
act, _ = jit_policy(state.obs, key)
|
| 85 |
state = jit_step(state, act)
|
| 86 |
state.info["command"] = cmd # joystick envs resample on reset only
|
| 87 |
-
|
| 88 |
if fell_at is None and float(state.done) > 0.5:
|
| 89 |
fell_at = (i + 1) * env.dt
|
| 90 |
break
|
| 91 |
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
"distance_m": float(np.linalg.norm((end - start)[:2])),
|
| 96 |
"fell_at": fell_at}
|
| 97 |
-
return
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
def
|
| 101 |
-
|
| 102 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
path.parent.mkdir(parents=True, exist_ok=True)
|
| 104 |
-
with
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 108 |
return path
|
|
|
|
| 11 |
|
| 12 |
from __future__ import annotations
|
| 13 |
|
| 14 |
+
import os
|
| 15 |
import pickle
|
| 16 |
from pathlib import Path
|
| 17 |
|
|
|
|
| 59 |
env._mjx_model = mjx.put_model(mj, impl=env._config.impl)
|
| 60 |
|
| 61 |
|
| 62 |
+
def rollout(env, inference_fn, seconds: float = 8.0, command=(0.5, 0.0, 0.0), seed: int = 0):
|
| 63 |
+
"""Run the policy for `seconds` with a fixed joystick command. Returns
|
| 64 |
+
(qpos, info): qpos is a (T, nq) trajectory at the env's control rate, info
|
| 65 |
+
has distance walked and when (if) it fell."""
|
|
|
|
| 66 |
import jax
|
| 67 |
import jax.numpy as jnp
|
| 68 |
|
|
|
|
| 76 |
state.info["command"] = cmd
|
| 77 |
|
| 78 |
n_steps = int(seconds / env.dt)
|
| 79 |
+
qpos = [np.asarray(state.data.qpos)]
|
|
|
|
| 80 |
fell_at = None
|
| 81 |
for i in range(n_steps):
|
| 82 |
rng, key = jax.random.split(rng)
|
| 83 |
act, _ = jit_policy(state.obs, key)
|
| 84 |
state = jit_step(state, act)
|
| 85 |
state.info["command"] = cmd # joystick envs resample on reset only
|
| 86 |
+
qpos.append(np.asarray(state.data.qpos))
|
| 87 |
if fell_at is None and float(state.done) > 0.5:
|
| 88 |
fell_at = (i + 1) * env.dt
|
| 89 |
break
|
| 90 |
|
| 91 |
+
qpos = np.stack(qpos)
|
| 92 |
+
info = {"seconds": len(qpos) * env.dt,
|
| 93 |
+
"distance_m": float(np.linalg.norm((qpos[-1, :2] - qpos[0, :2]))),
|
|
|
|
| 94 |
"fell_at": fell_at}
|
| 95 |
+
return qpos, info
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def offscreen_gl_backend() -> str | None:
|
| 99 |
+
"""Which MuJoCo offscreen backend this machine can actually load. The
|
| 100 |
+
backend is fixed when `mujoco` is imported, so rendering happens in a
|
| 101 |
+
subprocess with MUJOCO_GL set to this."""
|
| 102 |
+
import ctypes.util
|
| 103 |
+
import sys
|
| 104 |
+
|
| 105 |
+
forced = os.environ.get("ROLLOUT_GL")
|
| 106 |
+
if forced:
|
| 107 |
+
return forced
|
| 108 |
+
if sys.platform == "darwin":
|
| 109 |
+
return "glfw"
|
| 110 |
+
if ctypes.util.find_library("EGL"):
|
| 111 |
+
return "egl"
|
| 112 |
+
if ctypes.util.find_library("OSMesa"):
|
| 113 |
+
return "osmesa"
|
| 114 |
+
return None
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def write_video(env, qpos: np.ndarray, path: Path, fps: float,
|
| 118 |
+
width: int = 640, height: int = 480, camera: str = "track") -> Path:
|
| 119 |
+
"""Render a qpos trajectory of `env` to MP4 in a subprocess (see
|
| 120 |
+
render_worker.py), so a missing or different GL backend can never take the
|
| 121 |
+
training process down."""
|
| 122 |
+
import subprocess
|
| 123 |
+
import sys
|
| 124 |
+
import tempfile
|
| 125 |
+
|
| 126 |
+
backend = offscreen_gl_backend()
|
| 127 |
+
if backend is None:
|
| 128 |
+
raise RuntimeError("no offscreen GL backend (libEGL / libOSMesa) in this container -- "
|
| 129 |
+
"add `libegl1 libosmesa6` to packages.txt")
|
| 130 |
path.parent.mkdir(parents=True, exist_ok=True)
|
| 131 |
+
with tempfile.TemporaryDirectory() as tmp:
|
| 132 |
+
import mujoco
|
| 133 |
+
mjb = Path(tmp) / "model.mjb"
|
| 134 |
+
mujoco.mj_saveModel(env._mj_model, str(mjb), None)
|
| 135 |
+
traj = Path(tmp) / "qpos.npy"
|
| 136 |
+
np.save(traj, qpos)
|
| 137 |
+
worker = Path(__file__).with_name("render_worker.py")
|
| 138 |
+
env_vars = dict(os.environ, MUJOCO_GL=backend, PYOPENGL_PLATFORM=backend)
|
| 139 |
+
r = subprocess.run([sys.executable, str(worker), str(mjb), str(traj), str(path),
|
| 140 |
+
str(fps), str(width), str(height), camera],
|
| 141 |
+
env=env_vars, capture_output=True, text=True, timeout=600)
|
| 142 |
+
if r.returncode != 0:
|
| 143 |
+
tail = (r.stderr or r.stdout).strip().splitlines()[-6:]
|
| 144 |
+
raise RuntimeError(f"render ({backend}) failed: " + " | ".join(tail))
|
| 145 |
return path
|