Fix: render rollouts in a subprocess with a probed GL backend (EGL import broke training)

#9
Files changed (4) hide show
  1. app.py +2 -4
  2. packages.txt +3 -0
  3. render_worker.py +30 -0
  4. 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
- frames, info = rollout(eval_env, policy, seconds=float(seconds), command=(float(vx), 0.0, 0.0))
344
- out = write_video(frames, OUT / "rollouts" / f"{path.stem}.mp4", fps=1.0 / eval_env.dt)
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
- seed: int = 0, width: int = 640, height: int = 480, camera: str = "track"):
63
- """Run the policy for `seconds` with a fixed joystick command and return
64
- (frames, info): frames is a list of HxWx3 uint8 images at the env's control
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
- states = [state]
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
- states.append(state)
88
  if fell_at is None and float(state.done) > 0.5:
89
  fell_at = (i + 1) * env.dt
90
  break
91
 
92
- end = np.asarray(states[-1].data.qpos[:3])
93
- frames = env.render(states, width=width, height=height, camera=camera)
94
- info = {"seconds": len(states) * env.dt,
95
- "distance_m": float(np.linalg.norm((end - start)[:2])),
96
  "fell_at": fell_at}
97
- return frames, info
98
-
99
-
100
- def write_video(frames, path: Path, fps: float) -> Path:
101
- import imageio.v2 as imageio
102
-
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
103
  path.parent.mkdir(parents=True, exist_ok=True)
104
- with imageio.get_writer(str(path), fps=fps, codec="libx264", quality=7,
105
- macro_block_size=None) as w:
106
- for fr in frames:
107
- w.append_data(np.asarray(fr))
 
 
 
 
 
 
 
 
 
 
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