diff --git a/judo/app/simulation.py b/judo/app/simulation.py index cb577689..b25e7b79 100644 --- a/judo/app/simulation.py +++ b/judo/app/simulation.py @@ -6,7 +6,6 @@ from dora_utils.dataclasses import from_arrow, to_arrow from dora_utils.node import DoraNode, on_event -from mujoco import mj_step from omegaconf import DictConfig from judo.app.structs import MujocoState, SplineData @@ -50,6 +49,7 @@ def set_task(self, task_name: str) -> None: self.task: Task = task_cls() self.task_config = task_config_cls() + self.sim_backend = self.task.SimBackend() self.task.reset() @on_event("INPUT", "task") @@ -61,14 +61,13 @@ def update_task(self, event: dict) -> None: def step(self) -> None: """Step the simulation forward by one timestep.""" if self.control is not None and not self.paused: - try: - self.task.data.ctrl[:] = self.control(self.task.data.time) - self.task.pre_sim_step() - mj_step(self.task.sim_model, self.task.data) - self.task.post_sim_step() - except ValueError: + self.sim_controls = self.control(self.task.data.time) + if self.sim_controls.shape != (self.task.nu,): # we're switching tasks and the new task has a different number of actuators return + self.task.pre_sim_step() + self.sim_backend.step(self.task.sim_model, self.task.data, self.sim_controls) + self.task.post_sim_step() def spin(self) -> None: """Spin logic for the simulation node.""" diff --git a/judo/controller/controller.py b/judo/controller/controller.py index f261baf7..465154fa 100644 --- a/judo/controller/controller.py +++ b/judo/controller/controller.py @@ -11,7 +11,7 @@ from judo.gui import slider from judo.optimizers import Optimizer, OptimizerConfig from judo.tasks.base import Task, TaskConfig -from judo.utils.mujoco import RolloutBackend, make_model_data_pairs +from judo.utils.mujoco import make_model_data_pairs from judo.utils.normalization import ( IdentityNormalizer, Normalizer, @@ -69,7 +69,7 @@ def __init__( self.model = task.model self.model_data_pairs = make_model_data_pairs(self.model, self.optimizer_cfg.num_rollouts) - self.rollout_backend = RolloutBackend(num_threads=self.optimizer_cfg.num_rollouts, backend=rollout_backend) + self.rollout_backend = task.RolloutBackend(num_threads=self.optimizer_cfg.num_rollouts, backend=rollout_backend) self.action_normalizer = self._init_action_normalizer() @@ -78,7 +78,7 @@ def __init__( self.states = np.zeros((self.optimizer_cfg.num_rollouts, self.num_timesteps, self.model.nq + self.model.nv)) self.sensors = np.zeros((self.optimizer_cfg.num_rollouts, self.num_timesteps, self.model.nsensordata)) - self.rollout_controls = np.zeros((self.optimizer_cfg.num_rollouts, self.num_timesteps, self.model.nu)) + self.rollout_controls = np.zeros((self.optimizer_cfg.num_rollouts, self.num_timesteps, self.task.nu)) self.rewards = np.zeros((self.optimizer_cfg.num_rollouts,)) self.reset() @@ -173,8 +173,8 @@ def update_action(self, curr_state: np.ndarray, curr_time: float) -> None: candidate_knots_normalized = self.optimizer.sample_control_knots(nominal_knots_normalized) candidate_knots_normalized = np.clip( candidate_knots_normalized, - self.action_normalizer.normalize(self.task.actuator_ctrlrange[:, 0]), - self.action_normalizer.normalize(self.task.actuator_ctrlrange[:, 1]), + self.action_normalizer.normalize(self.task.ctrlrange[:, 0]), + self.action_normalizer.normalize(self.task.ctrlrange[:, 1]), ) self.candidate_knots = self.action_normalizer.denormalize(candidate_knots_normalized) @@ -283,11 +283,11 @@ def _init_action_normalizer(self) -> Normalizer: """Initialize the action normalizer.""" action_normalizer_kwargs = {} if self.action_normalizer_type == "min_max": - action_normalizer_kwargs["min"] = self.task.actuator_ctrlrange[:, 0] - action_normalizer_kwargs["max"] = self.task.actuator_ctrlrange[:, 1] + action_normalizer_kwargs["min"] = self.task.ctrlrange[:, 0] + action_normalizer_kwargs["max"] = self.task.ctrlrange[:, 1] elif self.action_normalizer_type == "running": action_normalizer_kwargs["init_std"] = 1.0 # TODO(yunhai): make this configurable - return make_normalizer(self.action_normalizer_type, self.model.nu, **action_normalizer_kwargs) + return make_normalizer(self.action_normalizer_type, self.task.nu, **action_normalizer_kwargs) def make_spline(times: np.ndarray, controls: np.ndarray, spline_order: str) -> interp1d: diff --git a/judo/tasks/base.py b/judo/tasks/base.py index 63bc7aee..1c7e3298 100644 --- a/judo/tasks/base.py +++ b/judo/tasks/base.py @@ -9,6 +9,8 @@ import numpy as np from mujoco import MjData, MjModel, MjSpec +from judo.utils.mujoco import RolloutBackend, SimBackend + @dataclass class TaskConfig: @@ -31,6 +33,9 @@ def __init__(self, model_path: Path | str = "", sim_model_path: Path | str | Non self.model_path = model_path self.sim_model = self.model if sim_model_path is None else MjModel.from_xml_path(str(sim_model_path)) + self.RolloutBackend = RolloutBackend + self.SimBackend = SimBackend + @property def time(self) -> float: """Returns the current simulation time.""" @@ -71,8 +76,8 @@ def nu(self) -> int: return self.model.nu @property - def actuator_ctrlrange(self) -> np.ndarray: - """Mujoco actuator limits for this task.""" + def ctrlrange(self) -> np.ndarray: + """Mujoco actuator limits for this task. Same as actuator limits for this task.""" limits = self.model.actuator_ctrlrange limited: np.ndarray = self.model.actuator_ctrllimited.astype(bool) # type: ignore limits[~limited] = np.array([-np.inf, np.inf], dtype=limits.dtype) # if not limited, set to inf diff --git a/judo/utils/mujoco.py b/judo/utils/mujoco.py index fe0f7d3a..ee6f3daf 100644 --- a/judo/utils/mujoco.py +++ b/judo/utils/mujoco.py @@ -5,7 +5,7 @@ from typing import Literal import numpy as np -from mujoco import MjData, MjModel +from mujoco import MjData, MjModel, mj_step from mujoco.rollout import Rollout @@ -79,3 +79,12 @@ def update(self, num_threads: int) -> None: self.setup_mujoco_backend(num_threads) else: raise ValueError(f"Unknown backend: {self.backend}") + + +class SimBackend: + """The backend for conducting simulation.""" + + def step(self, sim_model: MjModel, sim_data: MjData, sim_controls: np.ndarray) -> None: + """Conduct a simulation step.""" + sim_data.ctrl[:] = sim_controls + mj_step(sim_model, sim_data) diff --git a/tests/test_controller/test_action_normalization.py b/tests/test_controller/test_action_normalization.py index 5200da81..a193c129 100644 --- a/tests/test_controller/test_action_normalization.py +++ b/tests/test_controller/test_action_normalization.py @@ -195,8 +195,8 @@ def test_min_max_normalizer_with_task_control_ranges() -> None: assert isinstance(controller.action_normalizer, MinMaxNormalizer) # Check that the normalizer is initialized with correct control ranges - np.testing.assert_array_almost_equal(controller.action_normalizer.min, controller.task.actuator_ctrlrange[:, 0]) - np.testing.assert_array_almost_equal(controller.action_normalizer.max, controller.task.actuator_ctrlrange[:, 1]) + np.testing.assert_array_almost_equal(controller.action_normalizer.min, controller.task.ctrlrange[:, 0]) + np.testing.assert_array_almost_equal(controller.action_normalizer.max, controller.task.ctrlrange[:, 1]) # Run optimization loop curr_state = np.random.rand(controller.task.model.nq + controller.task.model.nv) @@ -204,12 +204,12 @@ def test_min_max_normalizer_with_task_control_ranges() -> None: controller.update_action(curr_state, curr_time) # Check that all candidate actions are within the control range bounds - assert np.all(controller.candidate_knots >= controller.task.actuator_ctrlrange[:, 0] - 1e-6) - assert np.all(controller.candidate_knots <= controller.task.actuator_ctrlrange[:, 1] + 1e-6) + assert np.all(controller.candidate_knots >= controller.task.ctrlrange[:, 0] - 1e-6) + assert np.all(controller.candidate_knots <= controller.task.ctrlrange[:, 1] + 1e-6) # Check that normalized actions are within the normalized control range bounds - min_normalized = controller.action_normalizer.normalize(controller.task.actuator_ctrlrange[:, 0]) - max_normalized = controller.action_normalizer.normalize(controller.task.actuator_ctrlrange[:, 1]) + min_normalized = controller.action_normalizer.normalize(controller.task.ctrlrange[:, 0]) + max_normalized = controller.action_normalizer.normalize(controller.task.ctrlrange[:, 1]) candidate_knots_normalized = controller.action_normalizer.normalize(controller.candidate_knots) assert np.all(candidate_knots_normalized >= min_normalized - 1e-6) assert np.all(candidate_knots_normalized <= max_normalized + 1e-6)