Skip to content
Open
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
151 changes: 73 additions & 78 deletions judo/app/data/controller_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,9 @@

import time
from typing import Any
import threading
import logging
from copy import copy

import numpy as np
from omegaconf import DictConfig
Expand All @@ -12,6 +15,9 @@
from judo.controller import Controller, ControllerConfig
from judo.optimizers import Optimizer, OptimizerConfig, OptimizerConfigType, OptimizerType, get_registered_optimizers
from judo.tasks import Task, TaskConfig, get_registered_tasks
from judo.tasks.base import TaskConfig
from judo.optimizers.base import OptimizerConfig
from judo.unlocked.node import _ADDED_SLEEP_DURATION


class ControllerData:
Expand All @@ -34,122 +40,111 @@ def __init__(
register_tasks_from_cfg(task_registration_cfg)
if optimizer_registration_cfg is not None:
register_optimizers_from_cfg(optimizer_registration_cfg)
self.last_plan_time = 0.0
self.paused = False
self.available_optimizers = get_registered_optimizers()
self.available_tasks = get_registered_tasks()
self._setup(init_task, init_optimizer)

def _setup(self, task_name: str, optimizer_name: str) -> None:
"""Set up the task and optimizer for the controller."""
def update_task(
self,
task_name: str,
) -> None:
"""Updates the task, task config, and optimizer."""
logging.info("Updating task to %s", task_name)
task_entry = self.available_tasks.get(task_name)
optimizer_entry = self.available_optimizers.get(optimizer_name)

assert task_entry is not None, f"Task {task_name} not found in task registry."
assert optimizer_entry is not None, f"Optimizer {optimizer_name} not found in optimizer registry."

# instantiate the task/optimizer/controller
task_cls, task_config_cls = task_entry
self.optimizer_cls, self.optimizer_config_cls = optimizer_entry

self.task = task_cls()
self.task_config = task_config_cls()
self.optimizer_config = self.optimizer_config_cls()
if task_entry is None:
raise ValueError(f"Unknown task {task_name}")
self.task_cls, self.task_config_cls = task_entry

self.task = self.task_cls()
self.task_config = self.task_config_cls()
self._wait_for("optimizer_config")
self.optimizer_config.set_override(task_name)
self._wait_for("optimizer_cls")
self.optimizer = self.optimizer_cls(self.optimizer_config, self.task.nu)

self.controller_config = ControllerConfig()
self.controller_config.set_override(task_name)
self.controller = Controller(
self.controller_config,
self.task,
self.task_config,
self.optimizer,
self.optimizer_config,
)

# Initialize the task data.
self.states = np.concatenate([self.task.data.qpos, self.task.data.qvel])
self.curr_time = self.task.data.time
logging.info("Task has been updated to %s", task_name)

def update_task(
def update_optimizer(
self,
task: Task,
task_config: TaskConfig,
optimizer: OptimizerType,
optimizer_name: str,
) -> None:
"""Updates the task, task config, and optimizer.

Args:
task_cls: The class of the task.
task_config_cls: The class of the task config.
task: The task instance.
task_config: The task config instance.
optimizer: The optimizer instance.
"""
self.task = task
self.task_config = task_config
self.optimizer = optimizer
self.controller = Controller(
self.controller_config,
self.task,
self.task_config,
self.optimizer,
self.optimizer_config,
)
self.states = np.concatenate([self.task.data.qpos, self.task.data.qvel])
self.curr_time = self.task.data.time
"""Updates the optimizer based on a given name."""
logging.info("Updating optimizer to %s", optimizer_name)
optimizer_entry = self.available_optimizers.get(optimizer_name)
if optimizer_entry is None:
raise ValueError(f"Unknown optimizer {optimizer_name}.")
self.optimizer_cls, self.optimizer_config_cls = optimizer_entry
self.optimizer_config = self.optimizer_config_cls()
self._wait_for("task")
self.optimizer = self.optimizer_cls(self.optimizer_config, self.task.nu)
self._wait_for("controller")
self.controller.optimizer = self.optimizer
logging.info("Optimizer has been updated to %s", optimizer_name)


def reset_task(self) -> None:
def reset_task(self, _: int) -> None:
"""Resets the task and controller, setting the states to the default values."""
logging.info("Resetting task")
self._wait_for("task")
self.task.reset()
self._wait_for("controller")
self.controller.reset()
self.states = np.concatenate([self.task.data.qpos, self.task.data.qvel])
self.curr_time = self.task.data.time

def update_optimizer_config(self, optimizer_cfg: Any) -> None:
def update_optimizer_config(self, optimizer_config: OptimizerConfig) -> None:
"""Updates the optimizer config."""
self.optimizer_config = optimizer_cfg
logging.info("Updating optimizer config")
self.optimizer_config = copy(optimizer_config)
self._wait_for("controller")
self.controller.optimizer.config = self.optimizer_config
self.controller.optimizer_cfg = self.optimizer_config

def update_states(self, state_msg: MujocoState) -> None:
"""Updates the states."""
def update_task_config(self, task_config: TaskConfig) -> None:
"""Callback to update optimizer task config on receiving a new config message."""
logging.info("Updating task config")
self.task_config = copy(task_config)

def update_controller_config(self, controller_config: ControllerConfig) -> None:
"""Callback to update controller config on receiving a new config message."""
logging.info("Updating controller config")
self.controller_config = copy(controller_config)
self.delay_dt = 1. / self.controller_config.control_freq
self._wait_for("controller")
self.controller.controller_cfg = controller_config

def step(self, state_msg: MujocoState) -> tuple[SplineData, tuple[np.ndarray, int], float]:
"""Updates the controls state internally."""
self.states = np.concatenate([state_msg.qpos, state_msg.qvel])
self.curr_time = state_msg.time
self._wait_for("controller")
self.controller.system_metadata = state_msg.sim_metadata
self.controller.task.time = state_msg.time

def step(self) -> None:
"""Updates the controls state internally."""
if self.states.shape != (self.controller.model.nq + self.controller.model.nv,):
return
elif self.paused:
return
raise ValueError("State and model dimension mismatch")

start = time.perf_counter()
self.controller.update_action(self.states, self.curr_time)
end = time.perf_counter()
self.last_plan_time = end - start
plan_time = end - start

def pause(self) -> None:
"""Pauses or starts the controller."""
self.paused = not self.paused
if plan_time + _ADDED_SLEEP_DURATION < self.delay_dt:
time.sleep(self.delay_dt - plan_time - _ADDED_SLEEP_DURATION)

control = SplineData(self.controller.times, self.controller.nominal_knots)

return control, (self.controller.traces, self.controller.all_traces_rollout_size), plan_time

def _wait_for(self, attr_name: str) -> None:
while not hasattr(self, attr_name):
time.sleep(0.1)

def update_optimizer(
self,
optimizer: Optimizer,
optimizer_config_cls: OptimizerConfigType,
optimizer_config: OptimizerConfig,
optimizer_cls: OptimizerType,
) -> None:
"""Updates the optimizer based on a given name."""
self.optimizer = optimizer
self.controller.optimizer = optimizer
self.optimizer_config_cls = optimizer_config_cls
self.optimizer_config = optimizer_config
self.optimizer_cls = optimizer_cls

@property
def spline_data(self) -> np.ndarray:
"""Returns the spline data for the current state."""
return SplineData(self.controller.times, self.controller.nominal_knots)
62 changes: 33 additions & 29 deletions judo/app/data/simulation_data.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,13 @@
# Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved.

from typing import Callable
import logging
import time

from mujoco import mj_step
from omegaconf import DictConfig

from judo.app.structs import MujocoState
from judo.app.structs import MujocoState, SplineData
from judo.app.utils import register_tasks_from_cfg
from judo.tasks import get_registered_tasks
from judo.tasks.base import Task, TaskConfig
Expand All @@ -31,39 +33,49 @@ def __init__(

self.control = None
self.paused = False
self.set_task(init_task)

def set_task(self, task_name: str) -> None:
def update_task(self, task_name: str) -> None:
"""Helper to initialize task from task name."""
logging.info("Updating task to %s", task_name)
task_entry = get_registered_tasks().get(task_name)
if task_entry is None:
raise ValueError(f"Init task {task_name} not found in task registry")
raise ValueError(f"Unknown task {task_name}")

task_cls, task_config_cls = task_entry

self.task: Task = task_cls()
self.task_config: TaskConfig = task_config_cls()
self.task.reset()
self._set_state()
logging.info("Task has been updated to %s", task_name)

def step(self) -> None:
def reset_task(self, _: int) -> None:
"""Resets the task."""
logging.info("Resetting task")
self.task.reset()

def pause(self, _: int) -> None:
"""Event handler for processing pause status updates."""
self.paused = not self.paused

def step(self, control: SplineData | None) -> MujocoState:
"""Step the simulation forward by one timestep."""
if self.control is not None and not self.paused:
if self.paused:
raise ValueError("Simulation is paused.")
if control is not None:
self.control = control

if self.control is not None:
self._wait_for("task")
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()
self.task.data.ctrl[:] = self.control.spline()(self.task.data.time)
except ValueError:
# we're switching tasks and the new task has a different number of actuators
pass
self._wait_for("task")
self.task.pre_sim_step()
mj_step(self.task.sim_model, self.task.data)
self.task.post_sim_step()

# Sets the internal state message based on the control and simuilation output
self._set_state()

def _set_state(self) -> None:
"""Set the state of the simulation."""
self.sim_state = MujocoState(
return MujocoState(
time=self.task.data.time,
qpos=self.task.data.qpos, # type: ignore
qvel=self.task.data.qvel, # type: ignore
Expand All @@ -74,14 +86,6 @@ def _set_state(self) -> None:
sim_metadata=self.task.get_sim_metadata(),
)

def pause(self) -> None:
"""Event handler for processing pause status updates."""
self.paused = not self.paused

def reset_task(self) -> None:
"""Resets the task."""
self.task.reset()

def update_control(self, control_spline: Callable) -> None:
"""Event handler for processing controls received from controller node."""
self.control = control_spline
def _wait_for(self, attr_name: str) -> None:
while not hasattr(self, attr_name):
time.sleep(0.1)
Loading