Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 6 additions & 7 deletions judo/app/simulation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand All @@ -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,):

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This change is necessary because for different backends (eg cpp) the try-catch does not work

# 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."""
Expand Down
16 changes: 8 additions & 8 deletions judo/controller/controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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()

Expand All @@ -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()

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand Down
9 changes: 7 additions & 2 deletions judo/tasks/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
import numpy as np
from mujoco import MjData, MjModel, MjSpec

from judo.utils.mujoco import RolloutBackend, SimBackend


@dataclass
class TaskConfig:
Expand All @@ -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."""
Expand Down Expand Up @@ -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
Expand Down
11 changes: 10 additions & 1 deletion judo/utils/mujoco.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
12 changes: 6 additions & 6 deletions tests/test_controller/test_action_normalization.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,21 +195,21 @@ 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)
curr_time = 0.0
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)
Expand Down
Loading