Don't pollute the global namespace when using mjpython.

Fixes #2265

PiperOrigin-RevId: 713442833
Change-Id: If155d9f57d4deef8b848478ebd8077a00417d5c5
This commit is contained in:
Saran Tunyasuvunakool
2025-01-08 15:47:15 -08:00
committed by Copybara-Service
parent 357ea024c0
commit 40ef08c8ed
2 changed files with 136 additions and 108 deletions
+1
View File
@@ -12,6 +12,7 @@ Python bindings
the computation. The thread pool can be reused across calls, but then the function cannot be called simultaneously the computation. The thread pool can be reused across calls, but then the function cannot be called simultaneously
from multiple threads. To run multiple threaded rollouts simultaneously, use the new class ``Rollout`` which from multiple threads. To run multiple threaded rollouts simultaneously, use the new class ``Rollout`` which
encapsulates the thread pool. Contribution by :github:user:`aftersomemath`. encapsulates the thread pool. Contribution by :github:user:`aftersomemath`.
- Fix global namespace pollution when using ``mjpython`` (:github:issue:`2265`).
General General
^^^^^^^ ^^^^^^^
+135 -108
View File
@@ -86,95 +86,113 @@ void* mjpython_pymain(void* vargs) {
// Set up the condition variable to pass control back to the macOS main thread. // Set up the condition variable to pass control back to the macOS main thread.
gil = cpython.PyGILState_Ensure(); gil = cpython.PyGILState_Ensure();
cpython.PyRun_SimpleStringFlags("import threading; cond = threading.Condition()", nullptr); cpython.PyRun_SimpleStringFlags(R"(
def _mjpython_make_cond():
# Don't pollute the global namespace.
global _mjpython_make_cond
del _mjpython_make_cond
import threading
global cond
cond = threading.Condition()
_mjpython_make_cond()
)", nullptr);
py_initialized.store(true); py_initialized.store(true);
// Wait until GLFW is initialized on macOS main thread, set up the queue and an atexit hook // Wait until GLFW is initialized on macOS main thread, set up the queue and an atexit hook
// to enqueue a termination flag upon exit. // to enqueue a termination flag upon exit.
cpython.PyRun_SimpleStringFlags(R"( cpython.PyRun_SimpleStringFlags(R"(
import atexit def _mjpython_init():
# Don't pollute the global namespace.
global _mjpython_init
del _mjpython_init
# The mujoco.viewer module should only be imported here after glfw.init() in the macOS main thread. import atexit
with cond: import threading
cond.wait()
import mujoco.viewer
# Similar to a queue.Queue(maxsize=1), but where only one active task is allowed at a time. # The mujoco.viewer module should only be imported after glfw.init() in the macOS main thread.
# With queue.Queue(1), another item is allowed to be enqueued before task_done is called. with cond:
class _MjPythonImpl(mujoco.viewer._MjPythonBase): cond.wait()
import mujoco.viewer
# Termination statuses # Similar to a queue.Queue(maxsize=1), but where only one active task is allowed at a time.
NOT_TERMINATED = 0 # With queue.Queue(1), another item is allowed to be enqueued before task_done is called.
TERMINATION_REQUESTED = 1 class _MjPythonImpl(mujoco.viewer._MjPythonBase):
TERMINATION_ACCEPTED = 2
TERMINATED = 3
def __init__(self): # Termination statuses
self._cond = threading.Condition() NOT_TERMINATED = 0
self._task = None TERMINATION_REQUESTED = 1
self._termination = self.__class__.NOT_TERMINATED TERMINATION_ACCEPTED = 2
self._busy = False TERMINATED = 3
def launch_on_ui_thread( def __init__(self):
self, self._cond = threading.Condition()
model, self._task = None
data, self._termination = self.__class__.NOT_TERMINATED
handle_return, self._busy = False
key_callback,
show_left_ui, def launch_on_ui_thread(
show_right_ui, self,
): model,
with self._cond: data,
if self._busy or self._task is not None: handle_return,
raise RuntimeError('another MuJoCo viewer is already open') key_callback,
else: show_left_ui,
self._task = ( show_right_ui,
model, ):
data, with self._cond:
handle_return, if self._busy or self._task is not None:
key_callback, raise RuntimeError('another MuJoCo viewer is already open')
show_left_ui, else:
show_right_ui, self._task = (
) model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
)
self._cond.notify()
def terminate(self):
with self._cond:
self._termination = self.__class__.TERMINATION_REQUESTED
self._cond.notify()
self._cond.wait_for(
lambda: self._termination == self.__class__.TERMINATED)
def get(self):
with self._cond:
self._cond.wait_for(
lambda: self._task is not None or self._termination)
if self._termination:
if self._termination == self.__class__.TERMINATION_REQUESTED:
self._termination = self.__class__.TERMINATION_ACCEPTED
return None
task = self._task
self._busy = True
self._task = None
return task
def done(self):
with self._cond:
self._busy = False
if self._termination == self.__class__.TERMINATION_ACCEPTED:
self._termination = self.__class__.TERMINATED
self._cond.notify() self._cond.notify()
def terminate(self):
with self._cond:
self._termination = self.__class__.TERMINATION_REQUESTED
self._cond.notify()
self._cond.wait_for(
lambda: self._termination == self.__class__.TERMINATED)
def get(self): mujoco.viewer._MJPYTHON = _MjPythonImpl()
with self._cond: atexit.register(mujoco.viewer._MJPYTHON.terminate)
self._cond.wait_for(
lambda: self._task is not None or self._termination)
if self._termination: with cond:
if self._termination == self.__class__.TERMINATION_REQUESTED: cond.notify()
self._termination = self.__class__.TERMINATION_ACCEPTED
return None
task = self._task _mjpython_init()
self._busy = True
self._task = None
return task
def done(self):
with self._cond:
self._busy = False
if self._termination == self.__class__.TERMINATION_ACCEPTED:
self._termination = self.__class__.TERMINATED
self._cond.notify()
mujoco.viewer._MJPYTHON = _MjPythonImpl()
atexit.register(mujoco.viewer._MJPYTHON.terminate)
del _MjPythonImpl # Don't pollute globals for user script.
with cond:
cond.notify()
del cond # Don't pollute globals for user script.
)", nullptr); )", nullptr);
// Run the Python interpreter main loop. // Run the Python interpreter main loop.
@@ -283,47 +301,56 @@ int main(int argc, char** argv) {
// to finish setting up _MJPYTHON, then serve incoming viewer launch requests. // to finish setting up _MJPYTHON, then serve incoming viewer launch requests.
PyGILState_STATE gil = cpython.PyGILState_Ensure(); PyGILState_STATE gil = cpython.PyGILState_Ensure();
cpython.PyRun_SimpleStringFlags(R"( cpython.PyRun_SimpleStringFlags(R"(
import ctypes def _mjpython_main():
# Don't pollute the global namespace.
global _mjpython_main
del _mjpython_main
# GLFW must be initialized on the OS main thread (i.e. here). import ctypes
import glfw
import mujoco.viewer
glfw.init() # GLFW must be initialized on the OS main thread (i.e. here).
glfw.poll_events() import glfw
ctypes.CDLL(None).mjpython_hide_dock_icon() import mujoco.viewer
# Wait for Python main thread to finish setting up _MJPYTHON glfw.init()
with cond: glfw.poll_events()
cond.notify() ctypes.CDLL(None).mjpython_hide_dock_icon()
cond.wait()
while True: # Wait for Python main thread to finish setting up _MJPYTHON
try: global cond
# Wait for an incoming payload. with cond:
task = mujoco.viewer._MJPYTHON.get() cond.notify()
cond.wait()
del cond
# None means that we are exiting. while True:
if task is None: try:
glfw.terminate() # Wait for an incoming payload.
break task = mujoco.viewer._MJPYTHON.get()
# Otherwise, launch the viewer. # None means that we are exiting.
model, data, handle_return, key_callback, show_left_ui, show_right_ui = task if task is None:
ctypes.CDLL(None).mjpython_show_dock_icon() glfw.terminate()
mujoco.viewer._launch_internal( break
model,
data,
run_physics_thread=False,
handle_return=handle_return,
key_callback=key_callback,
show_left_ui=show_left_ui,
show_right_ui=show_right_ui,
)
ctypes.CDLL(None).mjpython_hide_dock_icon()
finally: # Otherwise, launch the viewer.
mujoco.viewer._MJPYTHON.done() model, data, handle_return, key_callback, show_left_ui, show_right_ui = task
ctypes.CDLL(None).mjpython_show_dock_icon()
mujoco.viewer._launch_internal(
model,
data,
run_physics_thread=False,
handle_return=handle_return,
key_callback=key_callback,
show_left_ui=show_left_ui,
show_right_ui=show_right_ui,
)
ctypes.CDLL(None).mjpython_hide_dock_icon()
finally:
mujoco.viewer._MJPYTHON.done()
_mjpython_main()
)", nullptr); )", nullptr);
cpython.PyGILState_Release(gil); cpython.PyGILState_Release(gil);