Add a Viewer-to-Sim Snapshot channel and convert relevant messages.
The ViewerEndpoint can now send Snapshot messages to the SimEndpoint. The SimEndpoint has a new method, `get_viewer_snapshots`, to retrieve these. Thsi channel will be used for UI state needed for the simulation (e.g., `MjOption` and `StepControl`). PiperOrigin-RevId: 940981799 Change-Id: Ib9bd6ad848cd5479bc49492a34fddd82569cb09f
This commit is contained in:
committed by
Copybara-Service
parent
2012646711
commit
a1b0d2c933
@@ -29,10 +29,12 @@ class ViewerEndpoint:
|
||||
s2v_snapshot: vp.SnapshotChannel,
|
||||
s2v_events: vp.EventChannel,
|
||||
v2s_events: vp.EventChannel,
|
||||
v2s_snapshot: vp.SnapshotChannel,
|
||||
):
|
||||
self._s2v_snapshot = s2v_snapshot
|
||||
self._s2v_events = s2v_events
|
||||
self._v2s_events = v2s_events
|
||||
self._v2s_snapshot = v2s_snapshot
|
||||
self._is_closed = False
|
||||
|
||||
def send_to_sim(self, message: vp.Message) -> None:
|
||||
@@ -41,9 +43,10 @@ class ViewerEndpoint:
|
||||
if isinstance(message, vp.Event):
|
||||
self._v2s_events.put(message)
|
||||
elif isinstance(message, vp.Snapshot):
|
||||
raise TypeError('viewer cannot produce snapshots')
|
||||
self._v2s_snapshot.put(message)
|
||||
else:
|
||||
raise TypeError(f'expected Event, got {type(message).__name__}')
|
||||
name = type(message).__name__
|
||||
raise TypeError(f'expected Event or Snapshot, got {name}')
|
||||
|
||||
def get_sim_events(self) -> list[vp.Event]:
|
||||
"""Returns all pending events from the simulation."""
|
||||
@@ -62,6 +65,7 @@ class ViewerEndpoint:
|
||||
if not self._is_closed:
|
||||
self._is_closed = True
|
||||
self._v2s_events.close()
|
||||
self._v2s_snapshot.close()
|
||||
self._s2v_events.close()
|
||||
self._s2v_snapshot.close()
|
||||
|
||||
@@ -79,10 +83,12 @@ class SimEndpoint:
|
||||
s2v_snapshot: vp.SnapshotChannel,
|
||||
s2v_events: vp.EventChannel,
|
||||
v2s_events: vp.EventChannel,
|
||||
v2s_snapshot: vp.SnapshotChannel,
|
||||
):
|
||||
self._s2v_snapshot = s2v_snapshot
|
||||
self._s2v_events = s2v_events
|
||||
self._v2s_events = v2s_events
|
||||
self._v2s_snapshot = v2s_snapshot
|
||||
self._is_closed = False
|
||||
|
||||
def send_to_viewer(self, message: vp.Message) -> None:
|
||||
@@ -103,6 +109,12 @@ class SimEndpoint:
|
||||
raise RuntimeError('SimEndpoint is closed')
|
||||
return self._v2s_events.get()
|
||||
|
||||
def get_viewer_snapshots(self) -> list[vp.Snapshot]:
|
||||
"""Returns all pending latest snapshots from the viewer, one per type."""
|
||||
if self._is_closed:
|
||||
raise RuntimeError('SimEndpoint is closed')
|
||||
return self._v2s_snapshot.get()
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close all channels owned by this endpoint."""
|
||||
if not self._is_closed:
|
||||
@@ -110,6 +122,7 @@ class SimEndpoint:
|
||||
self._s2v_snapshot.close()
|
||||
self._s2v_events.close()
|
||||
self._v2s_events.close()
|
||||
self._v2s_snapshot.close()
|
||||
|
||||
|
||||
def make_endpoints(
|
||||
@@ -117,16 +130,19 @@ def make_endpoints(
|
||||
s2v_snapshot: vp.SnapshotChannel,
|
||||
s2v_events: vp.EventChannel,
|
||||
v2s_events: vp.EventChannel,
|
||||
v2s_snapshot: vp.SnapshotChannel,
|
||||
) -> tuple[ViewerEndpoint, SimEndpoint]:
|
||||
"""Returns viewer and simulation endpoints."""
|
||||
viewer_endpoint = ViewerEndpoint(
|
||||
s2v_snapshot=s2v_snapshot,
|
||||
s2v_events=s2v_events,
|
||||
v2s_events=v2s_events,
|
||||
v2s_snapshot=v2s_snapshot,
|
||||
)
|
||||
sim_endpoint = SimEndpoint(
|
||||
s2v_snapshot=s2v_snapshot,
|
||||
s2v_events=s2v_events,
|
||||
v2s_events=v2s_events,
|
||||
v2s_snapshot=v2s_snapshot,
|
||||
)
|
||||
return viewer_endpoint, sim_endpoint
|
||||
|
||||
@@ -113,8 +113,8 @@ class ModelEvent(Event):
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class OptionEvent(Event):
|
||||
"""An event that transports updated MuJoCo model options."""
|
||||
class MjOptionSnapshot(Snapshot):
|
||||
"""A snapshot sending mjOption state from viewer to sim each frame."""
|
||||
|
||||
opt: mujoco.MjOption
|
||||
|
||||
@@ -128,8 +128,8 @@ class PerturbEvent(Event):
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class StepControlEvent(Event):
|
||||
"""An event carrying the full step control state from the viewer to the sim."""
|
||||
class StepControlSnapshot(Snapshot):
|
||||
"""A snapshot sending step control state from viewer to sim each frame."""
|
||||
|
||||
pause_state: sim.PauseState
|
||||
speed: float
|
||||
|
||||
Reference in New Issue
Block a user