84950fa371
This change introduces parallel chunked downloading for large model files (.mjb) directly into WASM linear memory, enables model loading without full page reloads and reduces memory overhead. Key changes: - Implement a chunked model endpoint in the Python web server to support range requests. - Add parallel chunked fetching in the frontend with retry logic and a single-fetch fallback. - Increase initial WASM memory to 3 GB to accommodate large models and prevent heap fragmentation. - Support in-place model reloading in the C++ client, including texture cache invalidation when the Filament context is recreated. - Display a model download progress bar and model parsing/loading banner to the UI. - Fixes model drag and drop (caused by typo in sessionId, corrected to session_id). PiperOrigin-RevId: 960249236 Change-Id: Icac89e6a4ca099882aaf9b111744c1b6ab0cc6c1
403 lines
16 KiB
Python
403 lines
16 KiB
Python
# Copyright 2026 DeepMind Technologies Limited
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# https://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
"""Simulation-agnostic web viewer for MuJoCo models.
|
|
|
|
The WebViewer streams UI and simulation state to a browser:
|
|
|
|
* The ImGui UI is built into a headless ImGui context and streamed to the
|
|
browser with the NetImgui protocol through a WebSocket-to-TCP proxy.
|
|
Input captured in the browser flows back over the same connection and is
|
|
injected into the headless context, so all viewer-side handlers work
|
|
unmodified.
|
|
* Physics state and render function state are streamed to the browser over
|
|
a WebSocket with latest-wins semantics. Note that Message types are only
|
|
streamed between the Python simulation and viewer; the data streamed to
|
|
the browser is fixed and independent of the user's custom message types.
|
|
* The browser runs the ``web_client`` WASM app, which renders the MuJoCo
|
|
scene with Filament and overlays the remote ImGui draw data.
|
|
* Everything is served through one public port (default 8080): the page and
|
|
model over HTTP, the UI stream at path /ui and the state stream at /state.
|
|
Exposing or tunneling that single port exposes the whole viewer; the other
|
|
ports are loopback-internal.
|
|
"""
|
|
|
|
import multiprocessing
|
|
import os
|
|
import queue
|
|
import shutil
|
|
import socket
|
|
import tempfile
|
|
from typing import Any
|
|
import zlib
|
|
|
|
import mujoco
|
|
from mujoco.experimental.studio import endpoints
|
|
from mujoco.experimental.studio import messages
|
|
from mujoco.experimental.studio import ux
|
|
from mujoco.experimental.studio import viewer_protocol
|
|
from mujoco.experimental.studio.web import headless_ui
|
|
from mujoco.experimental.studio.web import state_payload
|
|
from mujoco.experimental.studio.web import web_server
|
|
import numpy as np
|
|
|
|
from mujoco.experimental.dear_imgui import dear_imgui as imgui
|
|
from mujoco.experimental.implot import implot
|
|
|
|
# File extensions a dropped file may load as a model.
|
|
_MODEL_EXTENSIONS = ('.xml', '.urdf', '.mjb', '.mjz', '.zip')
|
|
|
|
|
|
def _find_assets_dir() -> str:
|
|
"""Locates the Studio assets directory inside the mujoco package."""
|
|
env_dir = os.environ.get('MUJOCO_STUDIO_ASSETS_DIR')
|
|
if env_dir and os.path.isdir(env_dir):
|
|
return env_dir
|
|
package_dir = os.path.dirname(os.path.abspath(mujoco.__file__))
|
|
assets_dir = os.path.join(package_dir, 'experimental', 'studio', 'assets')
|
|
if os.path.isdir(assets_dir):
|
|
return assets_dir
|
|
return ''
|
|
|
|
|
|
def _print_url_banner(host: str, port: int) -> None:
|
|
"""Prints a prominent boxed banner with the URLs browsers can use."""
|
|
rows = []
|
|
if host in ('::', '0.0.0.0'):
|
|
fqdn = socket.getfqdn()
|
|
rows.append(('Remote:', f'http://{fqdn}:{port}/'))
|
|
rows.append(('Local:', f'http://localhost:{port}/'))
|
|
|
|
label_width = max(len(label) for label, _ in rows)
|
|
lines_plain = ['MuJoCo Web Viewer running at:', '']
|
|
lines_plain += [f' {label.ljust(label_width)} {url}' for label, url in rows]
|
|
lines_plain += ['', 'Ctrl+C to quit']
|
|
|
|
width = max(len(line) for line in lines_plain)
|
|
|
|
def _hyperlink(url: str) -> str:
|
|
return f'\033]8;;{url}\033\\{url}\033]8;;\033\\'
|
|
|
|
lines_formatted = ['MuJoCo Web Viewer running at:', '']
|
|
lines_formatted += [
|
|
f' {label.ljust(label_width)} {_hyperlink(url)}' for label, url in rows
|
|
]
|
|
lines_formatted += ['', 'Ctrl+C to quit']
|
|
|
|
banner = ['+' + '-' * (width + 2) + '+']
|
|
for plain, formatted in zip(lines_plain, lines_formatted):
|
|
padding = ' ' * (width - len(plain))
|
|
banner.append(f'| {formatted}{padding} |')
|
|
banner.append('+' + '-' * (width + 2) + '+')
|
|
print('\n'.join(banner), flush=True)
|
|
|
|
|
|
def _pick_drop_root(paths: list[str]) -> str | None:
|
|
"""Picks the model file to load from a dropped set of files.
|
|
|
|
Rules, in order: the only loadable file wins; otherwise, among the loadable
|
|
files at the shallowest directory depth, prefer one named after its folder
|
|
(the mjz convention, e.g. cards/cards.xml), then scene.xml (the
|
|
mujoco_menagerie convention), then the alphabetically first.
|
|
|
|
Args:
|
|
paths: A list of relative file paths.
|
|
|
|
Returns:
|
|
The path to the model file to load, or None if no model file was found.
|
|
"""
|
|
candidates = [p for p in paths if p.lower().endswith(_MODEL_EXTENSIONS)]
|
|
|
|
if not candidates:
|
|
return None
|
|
|
|
if len(candidates) == 1:
|
|
return candidates[0]
|
|
|
|
depth = min(p.count('/') for p in candidates)
|
|
shallow = sorted(p for p in candidates if p.count('/') == depth)
|
|
for path in shallow:
|
|
parts = path.split('/')
|
|
stem = os.path.splitext(parts[-1])[0]
|
|
if len(parts) > 1 and stem == parts[-2]:
|
|
return path
|
|
|
|
for path in shallow:
|
|
if os.path.basename(path) == 'scene.xml':
|
|
return path
|
|
|
|
return shallow[0]
|
|
|
|
|
|
class WebViewer(viewer_protocol.Viewer):
|
|
"""Simulation-agnostic web viewer for MuJoCo models."""
|
|
|
|
def __init__(
|
|
self,
|
|
config: viewer_protocol.ViewerConfig,
|
|
endpoint: endpoints.ViewerEndpoint,
|
|
*,
|
|
model: mujoco.MjModel | None = None,
|
|
model_path: str = '',
|
|
handlers: list[Any] | None = None,
|
|
camera: mujoco.MjvCamera | None = None,
|
|
vis_options: mujoco.MjvOption | None = None,
|
|
perturb: mujoco.MjvPerturb | None = None,
|
|
render_flags: ux.RenderFlags | None = None,
|
|
extra_geoms: list[mujoco.MjvGeom] | None = None,
|
|
host: str = '::',
|
|
http_port: int | None = None,
|
|
ui_tcp_port: int = 0,
|
|
) -> None:
|
|
"""Initializes the WebViewer.
|
|
|
|
Args:
|
|
config: Viewer window configuration.
|
|
endpoint: The viewer endpoint for communication with the sim side.
|
|
model: Optional initial MjModel. Forwarded to the base Viewer.
|
|
model_path: Optional path to the model file.
|
|
handlers: Optional list of handler instances.
|
|
camera: Camera parameters. Internal object is created if None.
|
|
vis_options: Visualization options. Internal object is created if None.
|
|
perturb: Perturbation parameters. Internal object is created if None.
|
|
render_flags: Render flags. Internal object is created if None.
|
|
extra_geoms: List of extra geoms. Internal list is created if None.
|
|
host: Public interface the server binds to. The default "::" accepts both
|
|
IPv6 and IPv4 connections (IPv4-only where IPv6 is unavailable).
|
|
http_port: The single public port: page, WASM, /model, and the /ui and
|
|
/state WebSocket paths. None falls back to config.http_port; 0 picks the
|
|
first free port starting at 8080, so several viewers can run side by
|
|
side.
|
|
ui_tcp_port: Loopback TCP port the headless NetImgui client connects to. 0
|
|
(the default) uses an OS-assigned ephemeral port — both endpoints live
|
|
in this process tree, so no fixed number is needed.
|
|
"""
|
|
super().__init__(
|
|
config,
|
|
endpoint,
|
|
model=model,
|
|
model_path=model_path,
|
|
handlers=handlers,
|
|
camera=camera,
|
|
vis_options=vis_options,
|
|
perturb=perturb,
|
|
render_flags=render_flags,
|
|
extra_geoms=extra_geoms,
|
|
)
|
|
|
|
self._host = host
|
|
# Bind both listening sockets up front, in this process: bind conflicts
|
|
# surface here as one clear error, the public port stays stable across
|
|
# server restarts, and the loopback port is OS-assigned so viewer instances
|
|
# can never collide on it.
|
|
requested_port = config.http_port if http_port is None else http_port
|
|
self._http_sock = web_server.bind_public_socket(host, requested_port)
|
|
self._http_port = self._http_sock.getsockname()[1]
|
|
self._tcp_sock = web_server.bind_loopback_socket(ui_tcp_port)
|
|
self._ui_tcp_port = self._tcp_sock.getsockname()[1]
|
|
|
|
# Headless ImGui context streaming UI draw data via NetImgui.
|
|
self._headless_ui = headless_ui.HeadlessUi(
|
|
config.title or 'MuJoCo Web Viewer',
|
|
self._ui_tcp_port,
|
|
_find_assets_dir(),
|
|
)
|
|
|
|
# Point the Python Dear ImGui bindings at the headless context so that
|
|
# viewer-side handlers build their GUI into it. The ImPlot context must be
|
|
# shared in the same way.
|
|
ctx = self._headless_ui.get_context()
|
|
imgui.SetCurrentContext(ctx)
|
|
ux.set_imgui_context(ctx)
|
|
ux.set_implot_context(self._headless_ui.get_implot_context())
|
|
implot.set_imgui_context(ctx)
|
|
implot.set_implot_context(self._headless_ui.get_implot_context())
|
|
|
|
# The single-port server (HTTP + /ui + /state + /drop). This is restarted
|
|
# whenever the model changes.
|
|
self._web_server: web_server.WebServer | None = None
|
|
self._model_crc32 = 0
|
|
|
|
# Files dropped onto the browser page arrive here from the server child as a
|
|
# dict of relative path -> bytes; owned by the viewer so it survives server
|
|
# restarts.
|
|
self._drop_queue = multiprocessing.get_context('fork').Queue()
|
|
|
|
# The controlling page's session id, written by the server child. Owned by
|
|
# the viewer so a fresh server (model change restarts it) can reserve the
|
|
# controller slot for the same page instead of letting whichever page
|
|
# reconnects first win it.
|
|
self._controller_sid = multiprocessing.get_context('fork').Array('c', 64)
|
|
|
|
# Temp dir holding the most recent drop's files; removed when the next drop
|
|
# supersedes it (its model is already parsed) and on close.
|
|
self._drop_dir = None
|
|
self._start_servers()
|
|
_print_url_banner(self._host, self._http_port)
|
|
|
|
# Dispatch lifecycle event so handlers can cache the viewer reference.
|
|
self.dispatch(viewer_protocol.ViewerInitEvent(viewer=self))
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Server lifecycle.
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def _start_servers(self) -> None:
|
|
"""Starts (or restarts) the web server, serving the current model."""
|
|
self._stop_servers()
|
|
|
|
# Serialize the compiled model to MJB bytes (served as /model).
|
|
buffer = np.empty(mujoco.mj_sizeModel(self.model), np.uint8)
|
|
mujoco.mj_saveModel(self.model, None, buffer)
|
|
mjb_data = buffer.tobytes()
|
|
|
|
# Identity of the served model, included in every state payload. When it
|
|
# changes, the browser refetches /model by reloading the page.
|
|
self._model_crc32 = zlib.crc32(mjb_data)
|
|
|
|
state_sig = int(mujoco.mjtState.mjSTATE_INTEGRATION)
|
|
state_size = mujoco.mj_stateSize(self.model, state_sig)
|
|
max_payload = state_payload.max_state_payload_size(
|
|
state_size * np.float64().itemsize
|
|
)
|
|
|
|
self._web_server = web_server.WebServer(
|
|
http_sock=self._http_sock,
|
|
tcp_sock=self._tcp_sock,
|
|
mjb_data=mjb_data,
|
|
max_payload_size=max_payload,
|
|
drop_queue=self._drop_queue,
|
|
controller_sid_shared=self._controller_sid,
|
|
)
|
|
self._web_server.start()
|
|
|
|
def _stop_servers(self) -> None:
|
|
if self._web_server is not None:
|
|
self._web_server.stop()
|
|
self._web_server = None
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Message handlers.
|
|
# ---------------------------------------------------------------------------
|
|
|
|
@messages.handler(priority=messages.Priority.CRITICAL)
|
|
def _on_model(self, event: messages.ModelEvent) -> bool:
|
|
"""Loads the new model, then restarts the servers to serve it.
|
|
|
|
The handler registry discovers handlers by name, so this override replaces
|
|
the base Viewer's _on_model and must call it explicitly to load the model
|
|
before the servers serialize it. The browser reconnects to the new state
|
|
server, notices the changed model identity in the payload, and reloads
|
|
itself to fetch the new model.
|
|
"""
|
|
super()._on_model(event)
|
|
self._start_servers()
|
|
print('Model changed: the browser page reloads automatically.', flush=True)
|
|
return False # Do not consume; let other handlers see the event.
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Viewer interface.
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def prepare_next_frame(self) -> bool:
|
|
"""Advances to the next headless frame; returns False when disconnected."""
|
|
return self._headless_ui.new_frame()
|
|
|
|
def sync(self) -> None:
|
|
"""Streams state to the browser and ends the headless ImGui frame."""
|
|
if self._web_server is not None:
|
|
state_sig = int(mujoco.mjtState.mjSTATE_INTEGRATION)
|
|
state_size = mujoco.mj_stateSize(self.model, state_sig)
|
|
state = np.empty(state_size, np.float64)
|
|
mujoco.mj_getState(self.model, self.data, state, state_sig)
|
|
payload = state_payload.serialize_state_payload(
|
|
self._model_crc32,
|
|
state_sig,
|
|
state.tobytes(),
|
|
self.camera,
|
|
self.perturb,
|
|
self.vis_options,
|
|
self.model,
|
|
list(self.render_flags.flags),
|
|
self.extra_geoms[: state_payload.MAX_EXTRA_GEOMS],
|
|
)
|
|
self._web_server.update_state(payload)
|
|
|
|
# Finish the ImGui frame; NetImgui sends the draw data to the browser.
|
|
self._headless_ui.end_frame()
|
|
|
|
def close(self) -> None:
|
|
self._stop_servers()
|
|
self._http_sock.close()
|
|
self._tcp_sock.close()
|
|
if self._drop_dir is not None:
|
|
shutil.rmtree(self._drop_dir, ignore_errors=True)
|
|
self._drop_dir = None
|
|
super().close()
|
|
|
|
def get_drop_file(self) -> str:
|
|
"""Returns the path of a model file dropped onto the browser page, or ''.
|
|
|
|
The browser uploads the dropped files' bytes over the /drop WebSocket
|
|
(so drops work even when the browser runs on another machine). They are
|
|
written to a temporary directory here, preserving names and relative
|
|
paths — a dropped folder's asset references resolve, and the mjz/zip
|
|
decoder can locate an archive's root XML by the archive's own name —
|
|
and the regular drop-loading flow (ViewerApp -> parser.parse) applies.
|
|
"""
|
|
try:
|
|
files = self._drop_queue.get_nowait()
|
|
except queue.Empty:
|
|
return ''
|
|
|
|
# TODO(matijak): This could work without disk access: mjVFS can hold the
|
|
# dropped files in memory (mj_addBufferVFS) and mj_parse resolves
|
|
# includes/assets from it (ModelHolder::InitFromBuffer already covers
|
|
# the single-buffer XML/MJB/ZIP cases). Needs a buffer-based parser
|
|
# entry point and a viewer drop interface that isn't a file path.
|
|
# The previous drop's model has already been parsed, so its temp dir can
|
|
# go now; keeping only the newest avoids leaking one dir per drop.
|
|
if self._drop_dir is not None:
|
|
shutil.rmtree(self._drop_dir, ignore_errors=True)
|
|
self._drop_dir = tempfile.mkdtemp(prefix='mujoco_drop_')
|
|
drop_dir = self._drop_dir
|
|
written = []
|
|
for name, payload in files.items():
|
|
rel = name.replace('\\', '/').lstrip('/')
|
|
parts = [p for p in rel.split('/') if p not in ('', '.', '..')]
|
|
if not parts:
|
|
continue
|
|
rel = '/'.join(parts)
|
|
path = os.path.join(drop_dir, *parts)
|
|
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
with open(path, 'wb') as f:
|
|
f.write(payload)
|
|
written.append(rel)
|
|
root = _pick_drop_root(written)
|
|
if root is None:
|
|
print(
|
|
'Dropped file(s) contain no loadable model '
|
|
f'({", ".join(_MODEL_EXTENSIONS)}).',
|
|
flush=True,
|
|
)
|
|
return ''
|
|
return os.path.join(drop_dir, *root.split('/'))
|
|
|
|
def upload_image(
|
|
self, tex_id: int, img: str | bytes, width: int, height: int, bpp: int
|
|
) -> int:
|
|
"""Uploads an image to the browser over the NetImgui texture channel."""
|
|
if isinstance(img, str):
|
|
img = img.encode('latin-1')
|
|
return self._headless_ui.upload_image(tex_id, img, width, height, bpp)
|