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
from multiple threads. To run multiple threaded rollouts simultaneously, use the new class ``Rollout`` which
encapsulates the thread pool. Contribution by :github:user:`aftersomemath`.
- Fix global namespace pollution when using ``mjpython`` (:github:issue:`2265`).
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.
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);
// Wait until GLFW is initialized on macOS main thread, set up the queue and an atexit hook
// to enqueue a termination flag upon exit.
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.
with cond:
cond.wait()
import mujoco.viewer
import atexit
import threading
# Similar to a queue.Queue(maxsize=1), but where only one active task is allowed at a time.
# With queue.Queue(1), another item is allowed to be enqueued before task_done is called.
class _MjPythonImpl(mujoco.viewer._MjPythonBase):
# The mujoco.viewer module should only be imported after glfw.init() in the macOS main thread.
with cond:
cond.wait()
import mujoco.viewer
# Termination statuses
NOT_TERMINATED = 0
TERMINATION_REQUESTED = 1
TERMINATION_ACCEPTED = 2
TERMINATED = 3
# Similar to a queue.Queue(maxsize=1), but where only one active task is allowed at a time.
# With queue.Queue(1), another item is allowed to be enqueued before task_done is called.
class _MjPythonImpl(mujoco.viewer._MjPythonBase):
def __init__(self):
self._cond = threading.Condition()
self._task = None
self._termination = self.__class__.NOT_TERMINATED
self._busy = False
# Termination statuses
NOT_TERMINATED = 0
TERMINATION_REQUESTED = 1
TERMINATION_ACCEPTED = 2
TERMINATED = 3
def launch_on_ui_thread(
self,
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
):
with self._cond:
if self._busy or self._task is not None:
raise RuntimeError('another MuJoCo viewer is already open')
else:
self._task = (
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
)
def __init__(self):
self._cond = threading.Condition()
self._task = None
self._termination = self.__class__.NOT_TERMINATED
self._busy = False
def launch_on_ui_thread(
self,
model,
data,
handle_return,
key_callback,
show_left_ui,
show_right_ui,
):
with self._cond:
if self._busy or self._task is not None:
raise RuntimeError('another MuJoCo viewer is already open')
else:
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()
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)
mujoco.viewer._MJPYTHON = _MjPythonImpl()
atexit.register(mujoco.viewer._MJPYTHON.terminate)
if self._termination:
if self._termination == self.__class__.TERMINATION_REQUESTED:
self._termination = self.__class__.TERMINATION_ACCEPTED
return None
with cond:
cond.notify()
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()
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.
_mjpython_init()
)", nullptr);
// 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.
PyGILState_STATE gil = cpython.PyGILState_Ensure();
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 glfw
import mujoco.viewer
import ctypes
glfw.init()
glfw.poll_events()
ctypes.CDLL(None).mjpython_hide_dock_icon()
# GLFW must be initialized on the OS main thread (i.e. here).
import glfw
import mujoco.viewer
# Wait for Python main thread to finish setting up _MJPYTHON
with cond:
cond.notify()
cond.wait()
glfw.init()
glfw.poll_events()
ctypes.CDLL(None).mjpython_hide_dock_icon()
while True:
try:
# Wait for an incoming payload.
task = mujoco.viewer._MJPYTHON.get()
# Wait for Python main thread to finish setting up _MJPYTHON
global cond
with cond:
cond.notify()
cond.wait()
del cond
# None means that we are exiting.
if task is None:
glfw.terminate()
break
while True:
try:
# Wait for an incoming payload.
task = mujoco.viewer._MJPYTHON.get()
# Otherwise, launch the viewer.
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()
# None means that we are exiting.
if task is None:
glfw.terminate()
break
finally:
mujoco.viewer._MJPYTHON.done()
# Otherwise, launch the viewer.
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);
cpython.PyGILState_Release(gil);