From 3b23fea2ef9592b6e9ba37f6d313dcc67bfe3cc5 Mon Sep 17 00:00:00 2001 From: KraHsu Date: Thu, 14 May 2026 14:17:28 +0800 Subject: [PATCH] Add mouse-interaction plugin hook to ManagerBasedRlEnv Expose `SceneCfg.mouse_interaction` (default False) and, when enabled together with `scene.vis`, attach `GeneLabMouseInteractionPlugin` before `scene.build()` so the raycaster sees the fully constructed rigid solver. The subclass squeezes the leading length-1 batch dim from Genesis' new batched link APIs so the upstream spring-force math keeps working at `num_envs=1`. --- src/genelab/configs.py | 1 + src/genelab/envs/manager_based_rl_env.py | 8 ++ src/genelab/sim/mouse_interaction.py | 133 +++++++++++++++++++++++ 3 files changed, 142 insertions(+) create mode 100644 src/genelab/sim/mouse_interaction.py diff --git a/src/genelab/configs.py b/src/genelab/configs.py index 4647ccb8..d35ca50d 100644 --- a/src/genelab/configs.py +++ b/src/genelab/configs.py @@ -22,6 +22,7 @@ class SceneCfg: num_envs: int = 1 env_spacing: tuple[float, float] = (2.0, 2.0) sensors: tuple[SensorCfg, ...] = field(default_factory=tuple) + mouse_interaction: bool = False @dataclass diff --git a/src/genelab/envs/manager_based_rl_env.py b/src/genelab/envs/manager_based_rl_env.py index 4993c2ce..4b0e3106 100644 --- a/src/genelab/envs/manager_based_rl_env.py +++ b/src/genelab/envs/manager_based_rl_env.py @@ -278,6 +278,14 @@ def _build_scene(self) -> None: quat=tuple(robot_cfg.init_quat), ) self._robot: Any = self._scene.add_entity(morph) + if self.cfg.scene.vis and self.cfg.scene.mouse_interaction: + from genelab.sim.mouse_interaction import GeneLabMouseInteractionPlugin + + # MouseInteractionPlugin must be attached BEFORE ``scene.build()`` — pre-build registration + # routes through ``viewer.build`` so the plugin's raycaster sees the fully constructed + # rigid solver. Post-build registration deadlocks the sim/viewer loop in some Genesis builds. + # ``GeneLabMouseInteractionPlugin`` patches the upstream plugin for batched link APIs. + self._scene.viewer.add_plugin(GeneLabMouseInteractionPlugin(use_force=True)) self._scene.build( n_envs=self._num_envs, env_spacing=tuple(self.cfg.scene.env_spacing), diff --git a/src/genelab/sim/mouse_interaction.py b/src/genelab/sim/mouse_interaction.py new file mode 100644 index 00000000..6032b9a5 --- /dev/null +++ b/src/genelab/sim/mouse_interaction.py @@ -0,0 +1,133 @@ +"""Subclass of ``MouseInteractionPlugin`` patched for Genesis' batched tensor outputs. + +The upstream plugin (genesis 0.4.6 ``vis.viewer_plugins.MouseInteractionPlugin``) assumes +``link.get_pos()`` / ``get_quat()`` / ``get_vel()`` / ``get_ang()`` return un-batched +``(3,)`` / ``(4,)`` arrays, while the link's ``inertial_pos`` / ``inertial_quat`` are +already un-batched. With the current Genesis (post-build), the dynamic getters return +``(num_envs, ...)`` tensors — even at ``num_envs=1`` — so the inertial / link arrays no +longer share shapes and ``_np_quat_mul`` asserts inside ``_apply_spring_force``. + +This subclass squeezes the leading batch dim from the dynamic per-link arrays so the +rest of the spring-force math (which operates on a single env's COM frame) works. +""" + +from collections.abc import Callable +from typing import Any, cast + +import numpy as np + +import genesis.utils.geom as gu +from genesis.utils.misc import tensor_to_array +from genesis.vis.viewer_plugins import MouseInteractionPlugin + + +def _squeeze_env(arr: np.ndarray) -> np.ndarray: + """Drop a leading length-1 batch dim if present, leaving ``(3,)`` / ``(4,)``.""" + if arr.ndim >= 2 and arr.shape[0] == 1: + return arr[0] + return arr + + +class GeneLabMouseInteractionPlugin(MouseInteractionPlugin): + """Mouse-interaction plugin compatible with Genesis' batched link APIs.""" + + def on_mouse_press(self, x: int, y: int, button: int, modifiers: int) -> Any: + # The base class computes ``_held_point_local`` via + # ``inv_transform_by_trans_quat(ray_hit.position, link_pos, link_quat)``; + # with batched link state ``(1, 3)`` / ``(1, 4)`` this broadcasts into a + # ``(1, 3)`` held point, which then breaks every later use. Squeeze after. + result = super().on_mouse_press(x, y, button, modifiers) + if self._held_point_local is not None: + self._held_point_local = _squeeze_env(np.asarray(self._held_point_local)) + return result + + def _apply_spring_force(self, control_point: np.ndarray, dt: float) -> None: + if not self._held_link: + return + + link_pos = _squeeze_env(tensor_to_array(self._held_link.get_pos())) + link_quat = _squeeze_env(tensor_to_array(self._held_link.get_quat())) + lin_vel = _squeeze_env(tensor_to_array(self._held_link.get_vel())) + ang_vel = _squeeze_env(tensor_to_array(self._held_link.get_ang())) + + held_point_world = gu.transform_by_trans_quat(self._held_point_local, link_pos, link_quat) + + inertial_pos = tensor_to_array(cast(Any, self._held_link.inertial_pos)) + inertial_quat = tensor_to_array(cast(Any, self._held_link.inertial_quat)) + world_principal_quat = gu.transform_quat_by_quat(inertial_quat, link_quat) + + arm_in_principal = gu.inv_transform_by_trans_quat( + self._held_point_local, inertial_pos, inertial_quat + ) + arm_in_world = gu.transform_by_quat(arm_in_principal, world_principal_quat) + + R_world = gu.quat_to_R(world_principal_quat) + inertia_world = R_world @ self._held_link.inertial_i @ R_world.T + inv_inertia_world = np.linalg.inv(inertia_world) + + pos_err_v = control_point - held_point_world + mass = self._held_link.get_mass() + if mass is None or mass <= 0: + return + inv_mass = float(1.0 / mass) + + total_impulse = np.zeros(3, dtype=control_point.dtype) + total_torque_impulse = np.zeros(3, dtype=control_point.dtype) + + for i in range(3): + body_point_vel = lin_vel + np.cross(ang_vel, arm_in_world) + vel_err_v = -body_point_vel + + direction = np.zeros(3, dtype=control_point.dtype) + direction[i] = 1.0 + + pos_err = float(np.dot(direction, pos_err_v)) + vel_err = float(np.dot(direction, vel_err_v)) + + arm_x_dir = np.cross(arm_in_world, direction) + rot_mass = float(np.dot(arm_x_dir, inv_inertia_world @ arm_x_dir)) + virtual_mass = 1.0 / (inv_mass + rot_mass + 1e-12) + + damping_coeff = 2.0 * np.sqrt(self.spring_const * virtual_mass) + impulse = (self.spring_const * pos_err + damping_coeff * vel_err) * dt + + lin_vel = lin_vel + direction * impulse * inv_mass + ang_vel = ang_vel + inv_inertia_world @ (arm_x_dir * impulse) + + total_impulse[i] += impulse + total_torque_impulse += arm_x_dir * impulse + + self._held_link.solver.apply_links_external_force( + total_impulse / dt, (self._held_link.idx,), ref="link_com", local=False + ) + self._held_link.solver.apply_links_external_torque( + total_torque_impulse / dt, (self._held_link.idx,), ref="link_com", local=False + ) + + def on_draw(self) -> None: + # The base class's on_draw also dereferences batched ``link.get_pos`` / ``get_quat`` + # outputs through ``gu.transform_by_trans_quat`` to position the debug sphere/line. + # Pre-squeeze by temporarily wrapping the link's accessors. + link = self._held_link + if link is None: + super().on_draw() + return + + orig_get_pos = link.get_pos + orig_get_quat = link.get_quat + + def _wrap(fn: Callable[[], Any]) -> Callable[[], Any]: + def _wrapped() -> np.ndarray: + out = fn() + arr = tensor_to_array(out) + return _squeeze_env(arr) + + return _wrapped + + try: + link.get_pos = _wrap(orig_get_pos) # type: ignore[method-assign] + link.get_quat = _wrap(orig_get_quat) # type: ignore[method-assign] + super().on_draw() + finally: + link.get_pos = orig_get_pos # type: ignore[method-assign] + link.get_quat = orig_get_quat # type: ignore[method-assign]