Improve async example and remove Network demo code

- Move SimToView and ViewToSim dataclasses from async.py to viewer_protocol.py (adding a user_data field) to consolidate IPC protocol definitions.
  - Remove the simulated Network class and latency UI controls from async.py to simplify the example and focus on core multiprocessing communication.
  - Use parser.parse() in main() for cleaner and more direct MuJoCo model loading.
  - Add cancel_join_thread() to multiprocessing queues and a join timeout to prevent atexit handlers from blocking during process shutdown.

PiperOrigin-RevId: 932327625
Change-Id: If91777dba934814b3f9faa17f9055869920575d7
This commit is contained in:
Matija Kecman
2026-06-15 02:30:15 -07:00
committed by Copybara-Service
parent 2a210da0f7
commit 7c22920720
2 changed files with 69 additions and 73 deletions
@@ -14,13 +14,11 @@
"""This script runs a simulation and viewer in separate processes communicating asynchronously.
In this example, we will run the viewer in an independent process communicating
via multiprocessing queues. Controls are provided to simulate network transit
latency and adjust the communication rates.
via multiprocessing queues. A slider is provided to adjust the state send rate.
You must provide a mjcf model file via the first command-line argument.
"""
import dataclasses
import multiprocessing
import os
import sys
@@ -30,79 +28,41 @@ from absl import app as absl_app
from absl import flags as absl_flags
import mujoco
from mujoco.experimental.studio import native_viewer as _viewer
from mujoco.experimental.studio import parser
from mujoco.experimental.studio import sim as _sim
from mujoco.experimental.studio import studio_app
from mujoco.experimental.studio import ux
from mujoco.experimental.studio import viewer_protocol
from mujoco.experimental.studio import viewer_protocol as vp
import numpy as np
from mujoco.experimental.dear_imgui import dear_imgui as imgui
_GFX = absl_flags.DEFINE_enum(
'gfx', None, viewer_protocol.GFX_MODES, 'Rendering graphics mode.'
'gfx', None, vp.GFX_MODES, 'Rendering graphics mode.'
)
_WIDTH = absl_flags.DEFINE_integer('width', 1200, 'Width of the output image.')
_HEIGHT = absl_flags.DEFINE_integer('height', 800, 'Height of the output image')
@dataclasses.dataclass
class SimToView:
"""A message sent from the simulation process to the viewer process."""
model: mujoco.MjModel | None = None
data: mujoco.MjData | None = None
state: np.ndarray | None = None
state_sig: int = 0
send_time: float = 0.0
@dataclasses.dataclass
class ViewToSim:
"""A message sent from the viewer process to the simulation process."""
state: np.ndarray | None = None
state_sig: int = 0
reset: bool = False
send_rate: float = 60.0
class Network:
"""Simulated networking parameters."""
def __init__(self) -> None:
self.transit_buffer = []
self.send_rate = 60.0
self.network_delay = 0.2
def get_arrived(self, q: multiprocessing.Queue) -> SimToView | None:
now = time.time()
while not q.empty():
self.transit_buffer.append(q.get())
arrived = None
while (
self.transit_buffer
and now >= self.transit_buffer[0].send_time + self.network_delay
):
arrived = self.transit_buffer.pop(0)
return arrived
def view(
sim_to_view: multiprocessing.Queue,
view_to_sim: multiprocessing.Queue,
) -> None:
"""Entry-point for process that renders the simulation."""
# Block until the first message (containing the model) arrives.
msg = sim_to_view.get()
assert msg.model is not None, 'First message must contain the MuJoCo model.'
first_msg = sim_to_view.get()
assert (
first_msg.model is not None
), 'First message must contain the MuJoCo model.'
title = os.path.basename(sys.argv[0])
xfrc_sig = int(mujoco.mjtState.mjSTATE_XFRC_APPLIED)
xfrc_size = mujoco.mj_stateSize(msg.model, xfrc_sig)
xfrc_size = mujoco.mj_stateSize(first_msg.model, xfrc_sig)
xfrc_state = np.zeros(xfrc_size, np.float64)
app = studio_app.StudioApp(msg.model, msg.data)
network = Network()
data = mujoco.MjData(first_msg.model)
app = studio_app.StudioApp(first_msg.model, data)
send_rate = 60.0
viewer = _viewer.NativeViewer(
app.model,
title=title,
@@ -113,26 +73,28 @@ def view(
while viewer.is_running() and app.is_running():
# Determine which messages have arrived through the simulated network.
arrived = network.get_arrived(sim_to_view)
# Drain the queue, keeping only the latest message.
msg = None
while not sim_to_view.empty():
msg = sim_to_view.get()
# Update the camera and compute the perturbation.
app.handle_mouse_events(viewer.camera, viewer.vis_options, viewer.perturb)
# Sync state from the backend if a new payload actually arrived.
if arrived is not None and arrived.state is not None:
mujoco.mj_setState(app.model, app.data, arrived.state, arrived.state_sig)
if msg is not None and msg.state is not None:
mujoco.mj_setState(app.model, app.data, msg.state, msg.state_sig)
mujoco.mj_forward(app.model, app.data)
# Always apply the perturbation forces from the viewer.
app.apply_perturb(viewer.perturb)
# Transmit user interaction when we get a new state
if arrived is not None:
if msg is not None:
mujoco.mj_getState(app.model, app.data, xfrc_state, xfrc_sig)
view_to_sim.put(
ViewToSim(
send_rate=network.send_rate, state=xfrc_state, state_sig=xfrc_sig
vp.ViewToSim(
send_rate=send_rate, state=xfrc_state, state_sig=xfrc_sig
)
)
@@ -145,24 +107,25 @@ def view(
| int(imgui.WindowFlags.NoCollapse),
):
imgui.PushItemWidth(200.0)
_, network.network_delay = imgui.SliderFloat(
'Network Latency (s)', network.network_delay, 0.0, 2.0
)
updated, network.send_rate = imgui.SliderFloat(
'Send Rate (Hz)', network.send_rate, 1.0, 120.0
updated, send_rate = imgui.SliderFloat(
'Send Rate (Hz)', send_rate, 1.0, 120.0
)
if updated:
view_to_sim.put(ViewToSim(send_rate=network.send_rate))
view_to_sim.put(vp.ViewToSim(send_rate=send_rate))
imgui.SetNextItemWidth(-1)
if imgui.Button('Reset Simulation'):
view_to_sim.put(ViewToSim(reset=True, send_rate=network.send_rate))
view_to_sim.put(vp.ViewToSim(reset=True, send_rate=send_rate))
imgui.PopItemWidth()
imgui.End()
viewer.sync(app.model, app.data)
# Prevent multiprocessing.Queue atexit handler from blocking on exit.
sim_to_view.cancel_join_thread()
view_to_sim.cancel_join_thread()
def sim(
data: mujoco.MjData,
@@ -173,14 +136,14 @@ def sim(
) -> None:
"""Entry-point for process that runs the simulation."""
sim_to_view.put(SimToView(model=model, data=data))
sim_to_view.put(vp.SimToView(model=model))
step_control = _sim.StepControl()
integration_sig = int(mujoco.mjtState.mjSTATE_INTEGRATION)
integration_size = mujoco.mj_stateSize(model, integration_sig)
integration_state = np.empty(integration_size, np.float64)
msg = ViewToSim()
msg = vp.ViewToSim()
last_send_time = time.time()
while view_process.is_alive():
@@ -204,17 +167,21 @@ def sim(
if now - last_send_time >= 1.0 / max(1.0, msg.send_rate):
mujoco.mj_getState(model, data, integration_state, integration_sig)
sim_to_view.put(
SimToView(
vp.SimToView(
state=integration_state,
state_sig=integration_sig,
send_time=now,
)
)
last_send_time = now
def main(argv: list[str]) -> None:
app = studio_app.StudioApp.from_argv(argv)
if len(argv) < 2:
print('Usage: async <model_path.xml>')
sys.exit(1)
data = parser.parse(argv[1])
model = data.model
# Queues for communication between the simulation and viewer processes.
sim_to_view = multiprocessing.Queue()
@@ -227,7 +194,13 @@ def main(argv: list[str]) -> None:
view_process.start()
# Start the simulation in the main process.
sim(app.data, app.model, sim_to_view, view_to_sim, view_process)
sim(data, model, sim_to_view, view_to_sim, view_process)
# Prevent multiprocessing.Queue atexit handler from blocking on exit.
sim_to_view.cancel_join_thread()
view_to_sim.cancel_join_thread()
view_process.join(timeout=5.0)
if __name__ == '__main__':
@@ -16,10 +16,12 @@
StudioApp uses the protocol for convenience methods that accept any viewer.
"""
import dataclasses
from typing import Any
from typing import Protocol
import mujoco
from mujoco.experimental.studio import ux
import numpy as np
GFX_MODES = (
'classic',
@@ -33,6 +35,27 @@ GFX_MODES = (
)
@dataclasses.dataclass
class SimToView:
"""A message sent from the simulation to the viewer."""
model: mujoco.MjModel | None = None
state: np.ndarray | None = None
state_sig: int = 0
user_data: dict[str, Any] = dataclasses.field(default_factory=dict)
@dataclasses.dataclass
class ViewToSim:
"""A message sent from the viewer to the simulation."""
state: np.ndarray | None = None
state_sig: int = 0
reset: bool = False
send_rate: float = 60.0
user_data: dict[str, Any] = dataclasses.field(default_factory=dict)
class Viewer(Protocol):
"""Structural interface for any viewer."""