From 2041db5aef21c0fd0247c6161e7cd3be2cf4d965 Mon Sep 17 00:00:00 2001 From: Dmitry Yershov Date: Thu, 2 Oct 2025 10:27:21 -0400 Subject: [PATCH] Move to unlocked middleware. --- judo/app/data/controller_data.py | 151 +++-- judo/app/data/simulation_data.py | 62 +- judo/app/data/visualization_data.py | 84 ++- judo/app/dora/__init__.py | 1 - judo/app/dora/controller.py | 134 ----- judo/app/dora/simulation.py | 86 --- judo/app/dora/visualization.py | 143 ----- judo/cli.py | 150 ++++- judo/logging_util.py | 72 +++ judo/unlocked/__init__.py | 32 ++ judo/unlocked/annotation.py | 28 + .../unlocked/benchmark/benchmark_ipc_utils.py | 148 +++++ judo/unlocked/benchmark/benchmark_shmque.py | 290 ++++++++++ judo/unlocked/const.py | 81 +++ judo/unlocked/example/fibonacci.py | 48 ++ judo/unlocked/example/loopback.py | 50 ++ judo/unlocked/example/pub_sub.py | 65 +++ judo/unlocked/input_stage.py | 90 +++ judo/unlocked/ipc_utils.py | 295 ++++++++++ judo/unlocked/node.py | 137 +++++ judo/unlocked/output_stage.py | 51 ++ judo/unlocked/policy.py | 147 +++++ judo/unlocked/queue.py | 195 +++++++ judo/unlocked/schedule.py | 69 +++ judo/unlocked/shmque.py | 536 ++++++++++++++++++ judo/unlocked/test/test_const.py | 68 +++ judo/unlocked/test/test_ipc_utils.py | 117 ++++ judo/unlocked/test/test_output_stage.py | 34 ++ judo/unlocked/test/test_queue.py | 192 +++++++ judo/unlocked/test/test_shmque.py | 250 ++++++++ pyproject.toml | 4 +- 31 files changed, 3334 insertions(+), 476 deletions(-) delete mode 100644 judo/app/dora/__init__.py delete mode 100644 judo/app/dora/controller.py delete mode 100644 judo/app/dora/simulation.py delete mode 100644 judo/app/dora/visualization.py create mode 100644 judo/logging_util.py create mode 100644 judo/unlocked/__init__.py create mode 100644 judo/unlocked/annotation.py create mode 100644 judo/unlocked/benchmark/benchmark_ipc_utils.py create mode 100644 judo/unlocked/benchmark/benchmark_shmque.py create mode 100644 judo/unlocked/const.py create mode 100644 judo/unlocked/example/fibonacci.py create mode 100644 judo/unlocked/example/loopback.py create mode 100644 judo/unlocked/example/pub_sub.py create mode 100644 judo/unlocked/input_stage.py create mode 100644 judo/unlocked/ipc_utils.py create mode 100644 judo/unlocked/node.py create mode 100644 judo/unlocked/output_stage.py create mode 100644 judo/unlocked/policy.py create mode 100644 judo/unlocked/queue.py create mode 100644 judo/unlocked/schedule.py create mode 100644 judo/unlocked/shmque.py create mode 100644 judo/unlocked/test/test_const.py create mode 100644 judo/unlocked/test/test_ipc_utils.py create mode 100644 judo/unlocked/test/test_output_stage.py create mode 100644 judo/unlocked/test/test_queue.py create mode 100644 judo/unlocked/test/test_shmque.py diff --git a/judo/app/data/controller_data.py b/judo/app/data/controller_data.py index 05206faa..38e9684b 100644 --- a/judo/app/data/controller_data.py +++ b/judo/app/data/controller_data.py @@ -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 @@ -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: @@ -34,31 +40,27 @@ 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, @@ -66,90 +68,83 @@ def _setup(self, task_name: str, optimizer_name: str) -> None: 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) diff --git a/judo/app/data/simulation_data.py b/judo/app/data/simulation_data.py index c111d588..8a97416a 100644 --- a/judo/app/data/simulation_data.py +++ b/judo/app/data/simulation_data.py @@ -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 @@ -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 @@ -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) diff --git a/judo/app/data/visualization_data.py b/judo/app/data/visualization_data.py index 4fb829b7..6b47ce2c 100644 --- a/judo/app/data/visualization_data.py +++ b/judo/app/data/visualization_data.py @@ -1,7 +1,7 @@ # Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. import threading - +import logging import mujoco import numpy as np import viser @@ -15,8 +15,12 @@ from judo.controller import ControllerConfig from judo.gui import create_gui_elements from judo.optimizers import get_registered_optimizers +from judo.optimizers.base import OptimizerConfig from judo.tasks import get_registered_tasks +from judo.tasks.base import TaskConfig from judo.visualizers.model import ViserMjModel +from judo.app.structs import MujocoState +from judo.unlocked.const import const_cast ElementType = GuiImageHandle | GuiInputHandle | GuiFolderHandle | MeshHandle | IcosphereHandle @@ -72,6 +76,82 @@ def __init__( self.set_task(init_task, init_optimizer) + def write_sim_pause(self) -> int: + while not self.sim_pause_updated.is_set(): + self.sim_pause_updated.wait(timeout=0.1) + self.sim_pause_updated.clear() + return 1 + + def write_task(self) -> str: + """Write the task name to the GUI.""" + while not self.task_updated.is_set(): + self.task_updated.wait(timeout=0.1) + self.task_updated.clear() + return self.task_name + + def write_task_reset(self) -> int: + """Write the task reset signal to the GUI.""" + while not self.task_reset_updated.is_set(): + self.task_reset_updated.wait(timeout=0.1) + self.task_reset_updated.clear() + return 1 + + def write_optimizer(self) -> str: + """Write the optimizer name to the GUI.""" + while not self.optimizer_updated.is_set(): + self.optimizer_updated.wait(timeout=0.1) + self.optimizer_updated.clear() + return self.optimizer_name + + def write_controller_config(self) -> ControllerConfig: + """Write the controller config to the GUI.""" + while not self.controller_config_updated.is_set(): + self.controller_config_updated.wait(timeout=0.1) + self.controller_config_updated.clear() + return self.controller_config + + def write_optimizer_config(self) -> OptimizerConfig: + """Write the optimizer config to the GUI.""" + while not self.optimizer_config_updated.is_set(): + self.optimizer_config_updated.wait(timeout=0.1) + self.optimizer_config_updated.clear() + return self.optimizer_config + + def write_task_config(self) -> TaskConfig: + """Write the task config to the GUI.""" + while not self.task_config_updated.is_set(): + self.task_config_updated.wait(timeout=0.1) + self.task_config_updated.clear() + return self.task_config + + def update_states(self, state_msg: MujocoState) -> None: + """Callback to update states on receiving a new state measurement.""" + if self.controller_config.spline_order == "cubic" and self.optimizer_config.num_nodes < 4: + warnings.warn("Cubic splines require at least 4 nodes. Setting num_nodes=4.", stacklevel=2) + for e in self.gui_elements["optimizer_params"]: + if e.label == "num_nodes": + e.value = 4 + break + self.optimizer_config_updated.set() + + try: + with self.task_lock: + self.data.xpos[:] = state_msg.xpos + self.data.xquat[:] = state_msg.xquat + self.viser_model.set_data(self.data) + except ValueError: + # we're switching tasks and the new task has a different number of xpos/xquat + return + + def update_traces(self, msg: tuple[np.ndarray, int]) -> None: + """Callback to update traces on receiving a new trace measurement.""" + with self.task_lock: + self.viser_model.set_traces(*const_cast(msg)) + + def update_plan_time(self, plan_time_s: float) -> None: + """Callback to update plan time on receiving a new plan time measurement.""" + self.gui_elements["plan_time_display"].value = plan_time_s * 1000 # ms + def register_controller_config_overrides(self, controller_override_cfg: DictConfig) -> None: """Register task-specific controller config overrides. @@ -140,6 +220,8 @@ def set_task(self, task: str, optimizer: str) -> None: self.setup_gui() # send the configs to the other nodes + self.task_updated.set() + self.optimizer_updated.set() self.controller_config_updated.set() self.task_config_updated.set() self.optimizer_config_updated.set() diff --git a/judo/app/dora/__init__.py b/judo/app/dora/__init__.py deleted file mode 100644 index 4d3e2daf..00000000 --- a/judo/app/dora/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. diff --git a/judo/app/dora/controller.py b/judo/app/dora/controller.py deleted file mode 100644 index 8e59cbf2..00000000 --- a/judo/app/dora/controller.py +++ /dev/null @@ -1,134 +0,0 @@ -# Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. - -import time -from threading import Lock - -import pyarrow as pa -from dora_utils.dataclasses import from_event, to_arrow -from dora_utils.node import DoraNode, on_event -from omegaconf import DictConfig - -from judo.app.data.controller_data import ControllerData -from judo.app.structs import MujocoState -from judo.controller import ControllerConfig - - -class ControllerNode(DoraNode): - """Controller node.""" - - def __init__( - self, - init_task: str = "cylinder_push", - init_optimizer: str = "cem", - node_id: str = "controller", - max_workers: int | None = None, - task_registration_cfg: DictConfig | None = None, - optimizer_registration_cfg: DictConfig | None = None, - ) -> None: - """Initialize the controller node.""" - super().__init__(node_id=node_id, max_workers=max_workers) - self._data = ControllerData( - init_task=init_task, - init_optimizer=init_optimizer, - task_registration_cfg=task_registration_cfg, - optimizer_registration_cfg=optimizer_registration_cfg, - ) - self.write_controls() - self.lock = Lock() - - @on_event("INPUT", "task") - def update_task(self, event: dict) -> None: - """Updates the task type.""" - new_task = event["value"].to_numpy(zero_copy_only=False)[0] - task_entry = self._data.available_tasks.get(new_task) - if task_entry is not None: - task_cls, task_config_cls = task_entry - task = task_cls() - task_config = task_config_cls() - self._data.optimizer_config.set_override(new_task) - optimizer = self._data.optimizer_cls(self._data.optimizer_config, task.nu) - with self.lock: - self._data.update_task(task, task_config, optimizer) - self.write_controls() - else: - raise ValueError(f"Task {new_task} not found in task registry.") - - @on_event("INPUT", "task_reset") - def reset_task(self, event: dict) -> None: - """Resets the task.""" - with self.lock: - self._data.reset_task() - self.write_controls() - - @on_event("INPUT", "sim_pause") - def set_paused_status(self, event: dict) -> None: - """Event handler for processing pause status updates.""" - self._data.pause() - - @on_event("INPUT", "optimizer") - def update_optimizer(self, event: dict) -> None: - """Updates the optimizer type.""" - new_optimizer = event["value"].to_numpy(zero_copy_only=False)[0] - optimizer_entry = self._data.available_optimizers.get(new_optimizer) - if optimizer_entry is not None: - optimizer_cls, optimizer_config_cls = optimizer_entry - optimizer_config = optimizer_config_cls() - optimizer = optimizer_cls(optimizer_config, self._data.task.nu) - with self.lock: - self._data.update_optimizer(optimizer, optimizer_config_cls, optimizer_config, optimizer_cls) - else: - raise ValueError(f"Optimizer {new_optimizer_name} not found in optimizer registry.") - - @on_event("INPUT", "controller_config") - def update_controller_config(self, event: dict) -> None: - """Callback to update controller config on receiving a new config message.""" - self._data.controller_config = from_event(event, ControllerConfig) - self._data.controller.controller_cfg = self._data.controller_config - - @on_event("INPUT", "optimizer_config") - def update_optimizer_config(self, event: dict) -> None: - """Callback to update optimizer config on receiving a new config message.""" - self._data.update_optimizer_config(from_event(event, self._data.optimizer_config_cls)) - - @on_event("INPUT", "task_config") - def update_task_config(self, event: dict) -> None: - """Callback to update optimizer task config on receiving a new config message.""" - self._data.task_config = from_event(event, type(self._data.task_config)) - self._data.task_config = self._data.task_config - - def write_controls(self) -> None: - """Util that publishes the current controller spline.""" - # send control action - arr, metadata = to_arrow(self._data.spline_data) - self.node.send_output("controls", arr, metadata) - - # send traces - if self._data.controller.traces is not None and len(self._data.controller.traces) > 0: - metadata = { - "all_traces_rollout_size": str(self._data.controller.all_traces_rollout_size), - "shape": self._data.controller.traces.shape, - } - self.node.send_output("traces", pa.array(self._data.controller.traces.flatten()), metadata=metadata) - - @on_event("INPUT", "states") - def update_states(self, event: dict) -> None: - """Callback to update states on receiving a new state measurement.""" - state_msg = from_event(event, MujocoState) - self._data.update_states(state_msg) - - def step(self) -> None: - """Updates controls using current state info, and writes to /controls.""" - self._data.step() - self.node.send_output("plan_time", pa.array([self._data.last_plan_time])) - self.write_controls() - - def spin(self) -> None: - """Spin logic for the controller node.""" - while True: - start_time = time.time() - self.parse_messages() - self.step() - - # Force controller to run at fixed rate specified by control_freq. - sleep_dt = 1 / self._data.controller_config.control_freq - (time.time() - start_time) - time.sleep(max(0, sleep_dt)) diff --git a/judo/app/dora/simulation.py b/judo/app/dora/simulation.py deleted file mode 100644 index ad1ec422..00000000 --- a/judo/app/dora/simulation.py +++ /dev/null @@ -1,86 +0,0 @@ -# Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. - -import threading -import time -import warnings - -from dora_utils.dataclasses import from_arrow, to_arrow -from dora_utils.node import DoraNode, on_event -from omegaconf import DictConfig - -from judo.app.data.simulation_data import SimulationData -from judo.app.structs import SplineData - - -class SimulationNode(DoraNode): - """The simulation node.""" - - def __init__( - self, - node_id: str = "simulation", - init_task: str = "cylinder_push", - max_workers: int | None = None, - task_registration_cfg: DictConfig | None = None, - ) -> None: - """Initialize the simulation node.""" - super().__init__(node_id=node_id, max_workers=max_workers) - - self._data = SimulationData(init_task=init_task, task_registration_cfg=task_registration_cfg) - - self.task_reset_lock = threading.Lock() - self.config_lock = threading.Lock() - self.control_lock = threading.Lock() - self.write_states() - - @on_event("INPUT", "task") - def update_task(self, event: dict) -> None: - """Event handler for processing task updates.""" - new_task = event["value"].to_numpy(zero_copy_only=False)[0] - self._data.set_task(new_task) - - def step(self) -> None: - """Step the simulation forward by one timestep.""" - self._data.step() - - def spin(self) -> None: - """Spin logic for the simulation node.""" - while True: - start_time = time.time() - self.parse_messages() - self.step() - self.write_states() - - # Force controller to run at fixed rate specified by model dt. - dt_des = self._data.task.sim_model.opt.timestep - dt_elapsed = time.time() - start_time - if dt_elapsed < dt_des: - time.sleep(dt_des - dt_elapsed) - else: - warnings.warn( - f"Sim step {dt_elapsed:.3f} longer than desired step {dt_des:.3f}!", - stacklevel=2, - ) - - def write_states(self) -> None: - """Reads data from simulation and writes to output topic.""" - arr, metadata = to_arrow(self._data.sim_state) - self.node.send_output("states", arr, metadata) - - @on_event("INPUT", "sim_pause") - def set_paused_status(self, event: dict) -> None: - """Event handler for processing pause status updates.""" - self._data.pause() - - @on_event("INPUT", "task_reset") - def reset_task(self, event: dict) -> None: - """Resets the task.""" - with self.task_reset_lock: - self._data.reset_task() - - @on_event("INPUT", "controls") - def update_control(self, event: dict) -> None: - """Event handler for processing controls received from controller node.""" - spline_data = from_arrow(event["value"], event["metadata"], SplineData) - control = spline_data.spline() - with self.control_lock: - self._data.update_control(control) diff --git a/judo/app/dora/visualization.py b/judo/app/dora/visualization.py deleted file mode 100644 index e69a1b36..00000000 --- a/judo/app/dora/visualization.py +++ /dev/null @@ -1,143 +0,0 @@ -# Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. - -import warnings - -import pyarrow as pa -from dora_utils.dataclasses import from_arrow, to_arrow -from dora_utils.node import DoraNode, on_event -from omegaconf import DictConfig -from viser import GuiFolderHandle, GuiImageHandle, GuiInputHandle, IcosphereHandle, MeshHandle - -from judo.app.data.visualization_data import VisualizationData -from judo.app.structs import MujocoState - -ElementType = GuiImageHandle | GuiInputHandle | GuiFolderHandle | MeshHandle | IcosphereHandle - - -class VisualizationNode(DoraNode): - """The visualization node.""" - - def __init__( - self, - node_id: str = "visualization", - max_workers: int | None = None, - init_task: str = "cylinder_push", - init_optimizer: str = "cem", - task_registration_cfg: DictConfig | None = None, - optimizer_registration_cfg: DictConfig | None = None, - controller_override_cfg: DictConfig | None = None, - optimizer_override_cfg: DictConfig | None = None, - sim_pause_button: bool = True, - geom_exclude_substring: str = "collision", - ) -> None: - """Initialize the visualization node.""" - super().__init__(node_id=node_id, max_workers=max_workers) - self._data = VisualizationData( - init_task=init_task, - init_optimizer=init_optimizer, - task_registration_cfg=task_registration_cfg, - optimizer_registration_cfg=optimizer_registration_cfg, - controller_override_cfg=controller_override_cfg, - optimizer_override_cfg=optimizer_override_cfg, - sim_pause_button=sim_pause_button, - geom_exclude_substring=geom_exclude_substring, - ) - - def write_sim_pause(self) -> None: - """Write the sim pause signal to the GUI.""" - with self._data.sim_pause_lock: - self.node.send_output("sim_pause", pa.array([1])) # dummy value - self._data.sim_pause_updated.clear() - - def write_task(self) -> None: - """Write the task name to the GUI.""" - with self._data.task_lock: - self.node.send_output("task", pa.array([self._data.task_name])) - self._data.task_updated.clear() - - def write_task_reset(self) -> None: - """Write the task reset signal to the GUI.""" - with self._data.task_lock: - self.node.send_output("task_reset", pa.array([1])) # dummy value - self._data.task_reset_updated.clear() - - def write_optimizer(self) -> None: - """Write the optimizer name to the GUI.""" - with self._data.optimizer_lock: - self.node.send_output("optimizer", pa.array([self._data.optimizer_name])) - self._data.optimizer_updated.clear() - - def write_controller_config(self) -> None: - """Write the controller config to the GUI.""" - with self._data.controller_config_lock: - self.node.send_output("controller_config", *to_arrow(self._data.controller_config)) - self._data.controller_config_updated.clear() - - def write_optimizer_config(self) -> None: - """Write the optimizer config to the GUI.""" - with self._data.optimizer_config_lock: - self.node.send_output("optimizer_config", *to_arrow(self._data.optimizer_config)) - self._data.optimizer_config_updated.clear() - - def write_task_config(self) -> None: - """Write the task config to the GUI.""" - with self._data.task_config_lock: - self.node.send_output("task_config", *to_arrow(self._data.task_config)) - self._data.task_config_updated.clear() - - @on_event("INPUT", "states") - def update_states(self, event: dict) -> None: - """Callback to update states on receiving a new state measurement.""" - if self._data.controller_config.spline_order == "cubic" and self._data.optimizer_config.num_nodes < 4: - warnings.warn("Cubic splines require at least 4 nodes. Setting num_nodes=4.", stacklevel=2) - for e in self._data.gui_elements["optimizer_params"]: - if e.label == "num_nodes": - e.value = 4 - break - self._data.optimizer_config_updated.set() - - state_msg = from_arrow(event["value"], event["metadata"], MujocoState) - try: - with self._data.task_lock: - self._data.data.xpos[:] = state_msg.xpos - self._data.data.xquat[:] = state_msg.xquat - self._data.viser_model.set_data(self._data.data) - except ValueError: - # we're switching tasks and the new task has a different number of xpos/xquat - return - - @on_event("INPUT", "traces") - def update_traces(self, event: dict) -> None: - """Callback to update traces on receiving a new trace measurement.""" - traces_flat = event["value"].to_numpy() - all_traces_rollout_size = int(event["metadata"]["all_traces_rollout_size"]) - shape = event["metadata"]["shape"] - traces = traces_flat.reshape(*shape) - with self._data.task_lock: - self._data.viser_model.set_traces(traces, all_traces_rollout_size) - - @on_event("INPUT", "plan_time") - def update_plan_time(self, event: dict) -> None: - """Callback to update plan time on receiving a new plan time measurement.""" - plan_time_s = event["value"].to_numpy(zero_copy_only=False)[0] - self._data.gui_elements["plan_time_display"].value = plan_time_s * 1000 # ms - - def spin(self) -> None: - """Spin logic for the visualization node.""" - for event in self.node: - if self._data.sim_pause_updated.is_set(): - self.write_sim_pause() - if self._data.task_updated.is_set(): - self.write_task() - if self._data.task_reset_updated.is_set(): - self.write_task_reset() - if self._data.optimizer_updated.is_set(): - self.write_optimizer() - if self._data.controller_config_updated.is_set(): - self.write_controller_config() - if self._data.optimizer_config_updated.is_set(): - self.write_optimizer_config() - if self._data.task_config_updated.is_set(): - self.write_task_config() - - self.handle(event) diff --git a/judo/cli.py b/judo/cli.py index 26ffdcf6..e3b4e751 100644 --- a/judo/cli.py +++ b/judo/cli.py @@ -1,13 +1,26 @@ # Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. +import logging +import colorlog import warnings from pathlib import Path +import sys +from typing import Any, Iterable +from functools import partial +import signal +import multiprocessing as mp import hydra -from dora_utils.launch.run import run from hydra import compose, initialize_config_dir from hydra.core.config_store import ConfigStore from omegaconf import DictConfig +from judo.app.data.visualization_data import VisualizationData +from judo.app.data.simulation_data import SimulationData +from judo.app.data.controller_data import ControllerData +from judo.logging_util import setup_logger +from judo.unlocked import Node +from judo.unlocked.policy import Latest, Lossless, Optional +from judo.unlocked.schedule import Threaded # suppress annoying warning from hydra warnings.filterwarnings( @@ -23,11 +36,144 @@ # app # # ### # +class StopAll: + """Terminate all schedulers.""" + + def __init__(self, stop_it: Iterable) -> None: + self._stop = list(stop_it) + + def __call__(self, signum: Any, _: Any) -> None: + for s in self._stop: + s.set() + +def clear_dora_cfg(cfg: DictConfig) -> dict: + exclude_keys = ["_target_", "node_id", "max_workers"] + return {k: v for k, v in cfg.items() if k not in exclude_keys} + +def log_str_sink(level, msg): + def log_fn(l: int, m: str, d: Any) -> None: + logging.log(l, m, d) + return partial(log_fn, level, msg + "%s") @hydra.main(config_path=str(CONFIG_PATH), config_name="judo_dora_default", version_base="1.3") def main_app(cfg: DictConfig) -> None: """Main function to run judo via a hydra configuration yaml file.""" - run(cfg) + setup_logger(log_level=logging.INFO, include_thread=True) + + controller_config = clear_dora_cfg(cfg.node_definitions.controller) + logging.debug("Got controller config:\n%s", controller_config) + controller = ControllerData(**controller_config) + + simulation_config = clear_dora_cfg(cfg.node_definitions.simulation) + logging.debug("Got simulation config:\n%s", simulation_config) + simulation = SimulationData(**simulation_config) + + visualization_config = clear_dora_cfg(cfg.node_definitions.visualization) + logging.debug("Got visualization config:\n%s", visualization_config) + visualization = VisualizationData(**visualization_config) + + controller_update_task = Node("cont_update_task", controller.update_task) + controller_reset_task = Node("cont_reset_task", controller.reset_task) + controller_update_optimizer = Node("cont_update_optimizer", controller.update_optimizer) + controller_update_optimizer_config = Node("cont_update_optimizer_config", controller.update_optimizer_config) + controller_update_controller_config = Node("cont_update_controller_config", controller.update_controller_config) + controller_update_task_config = Node("cont_update_task_config", controller.update_task_config) + controller_step = Node("cont_step", controller.step) + + simulation_pause = Node("sim_pause", simulation.pause) + simulation_reset_task = Node("sim_reset_task", simulation.reset_task) + simulation_update_task = Node("sim_update_task", simulation.update_task) + simulation_step = Node("sim_step", simulation.step, frequency=60) + + visualization_sim_pause = Node("vis_sim_pause", visualization.write_sim_pause, frequency=1000) + visualization_task_name = Node("vis_task_name", visualization.write_task, frequency=1) + visualization_task_reset = Node("vis_task_reset", visualization.write_task_reset, frequency=10) + visualization_optimizer_name = Node("vis_optimizer_name", visualization.write_optimizer, frequency=1) + visualization_controller_config = Node("vis_controller_config", visualization.write_controller_config, frequency=100) + visualization_optimizer_config = Node("vis_optimizer_config", visualization.write_optimizer_config, frequency=100) + visualization_task_config = Node("vis_task_config", visualization.write_task_config, frequency=100) + visualization_update_states = Node("vis_update_states", visualization.update_states) + visualization_update_traces = Node("vis_update_traces", visualization.update_traces) + visualization_update_plan_time = Node("vis_update_plan_time", visualization.update_plan_time) + + controller_update_task.input_stage.connect(0, visualization_task_name.output_stage[0], Lossless()) + controller_reset_task.input_stage.connect(0, visualization_task_reset.output_stage[0], Lossless()) + controller_update_optimizer.input_stage.connect(0, visualization_optimizer_name.output_stage[0], Lossless()) + controller_update_optimizer_config.input_stage.connect(0, visualization_optimizer_config.output_stage[0], Lossless()) + controller_update_controller_config.input_stage.connect(0, visualization_controller_config.output_stage[0], Lossless()) + controller_update_task_config.input_stage.connect(0, visualization_task_config.output_stage[0], Lossless()) + controller_step.input_stage.connect(0, simulation_step.output_stage[0], Latest()) + + simulation_pause.input_stage.connect(0, visualization_sim_pause.output_stage[0], Lossless()) + simulation_reset_task.input_stage.connect(0, visualization_task_reset.output_stage[0], Lossless()) + simulation_update_task.input_stage.connect(0, visualization_task_name.output_stage[0], Lossless()) + simulation_step.input_stage.connect(0, controller_step.output_stage[0], Optional()) + + visualization_update_states.input_stage.connect(0, simulation_step.output_stage[0], Latest()) + visualization_update_traces.input_stage.connect(0, controller_step.output_stage[1], Latest()) + visualization_update_plan_time.input_stage.connect(0, controller_step.output_stage[2], Latest()) + + log_str = [] + # log_str.append(Node("log_str_1", log_str_sink(logging.DEBUG, "Got from sim "))) + # log_str[-1].input_stage.connect(0, simulation_step.output_stage[0], Latest()) + # + # log_str.append(Node("log_str_2", log_str_sink(logging.DEBUG, "Got from cont"))) + # log_str[-1].input_stage.connect(0, controller_step.output_stage[0], Latest()) + # + # log_str.append(Node("log_str_3", log_str_sink(logging.WARN, "Got from task name "))) + # log_str[-1].input_stage.connect(0, visualization_task_name.output_stage[0], Latest()) + # + # log_str.append(Node("log_str_4", log_str_sink(logging.WARN, "Got from optimizer name "))) + # log_str[-1].input_stage.connect(0, visualization_optimizer_name.output_stage[0], Latest()) + + controller_scheduler = Threaded( + "cont_sch", ( + controller_update_task, + controller_reset_task, + controller_update_optimizer, + controller_update_optimizer_config, + controller_update_controller_config, + controller_update_task_config, + controller_step, + ), + mp.Event()) + simulation_scheduler = Threaded( + "sim_sch", ( + simulation_pause, + simulation_reset_task, + simulation_update_task, + simulation_step, + ), + mp.Event()) + visualization_scheduler = Threaded( + "vis_sch", ( + visualization_sim_pause, + visualization_task_name, + visualization_task_reset, + visualization_optimizer_name, + visualization_controller_config, + visualization_optimizer_config, + visualization_task_config, + visualization_update_states, + visualization_update_traces, + visualization_update_plan_time, + *log_str, + ), + mp.Event()) + controller_scheduler.start() + simulation_scheduler.start() + # visualization_scheduler.start() + signal.signal(signal.SIGINT, StopAll( + ( + controller_scheduler._stop, + simulation_scheduler._stop, + visualization_scheduler._stop, + ) + )) + visualization_scheduler.spin() + controller_scheduler.join() + simulation_scheduler.join() + # visualization_scheduler.join() def app() -> None: diff --git a/judo/logging_util.py b/judo/logging_util.py new file mode 100644 index 00000000..992e4e39 --- /dev/null +++ b/judo/logging_util.py @@ -0,0 +1,72 @@ +# Copyright (c) 2025 Robotics and AI Institute LLC. All rights reserved. + +# written by Duy +import logging +import sys + +import colorlog + +DEFAULT_DATE_FORMAT = "%H:%M:%S" +DEFAULT_FORMAT = "%(log_color)s[%(levelname)s][%(asctime)s][%(processName)s][%(module)s]: %(message)s" +DEFAULT_MCAP_FORMAT = "[%(module)s] %(message)s" +DEFAULT_THREAD_FORMAT = ( + "%(log_color)s[%(levelname)s][%(asctime)s][%(processName)s][%(threadName)s][%(module)s]: %(message)s" +) + +LOG_COLORS = { + "DEBUG": "light_black", + "INFO": "white", + "WARNING": "yellow", + "ERROR": "red", + "CRITICAL": "red", +} + + +def get_format(include_thread: bool) -> str: + return DEFAULT_THREAD_FORMAT if include_thread else DEFAULT_FORMAT + + +def setup_logger(log_level: str | int = logging.INFO, include_thread: bool = False) -> None: + """Sets root logger. You should only run this once per node/runtime.""" + colorlog.basicConfig( + level=log_level, + format=get_format(include_thread), + datefmt=DEFAULT_DATE_FORMAT, + # Force is required to make sure that any other instantiations/setting + # of a logger don't prevent this from going into effect. Even just calling + # logging.getLogger() will prevent subsequent logging.basicConfig() calls to be + # ignored unless the "force" flag it True. + force=True, + ) + + # Set custom colors so all the INFO messages aren't obnoxiously green. + # For some reason the custom colors don't set properly when done as part of the basicConfig call + formatter = colorlog.ColoredFormatter( + get_format(include_thread), + log_colors=LOG_COLORS, + reset=True, + ) + handler = colorlog.StreamHandler(sys.stdout) + handler.setFormatter(formatter) + logger = colorlog.getLogger() + logger.handlers.clear() + logger.addHandler(handler) + + +# We may want this to take in a custom format string eventually +def configure_local_logger_format(logger: logging.Logger, include_thread: bool = False) -> None: + """Configures a specific logger instance to be different from the root logger.""" + root_logger = colorlog.getLogger() + # Because the root logger may be using force=True, you need to clear associatings of this + # logger from the root logger + logger.handlers.clear() + handler = colorlog.StreamHandler(sys.stdout) + formatter = colorlog.ColoredFormatter( + fmt=get_format(include_thread), + log_colors=LOG_COLORS, + datefmt=DEFAULT_DATE_FORMAT, + ) + handler.setFormatter(formatter) + logger.addHandler(handler) + logger.propagate = False + logger.setLevel(root_logger.getEffectiveLevel()) diff --git a/judo/unlocked/__init__.py b/judo/unlocked/__init__.py new file mode 100644 index 00000000..f3312ebe --- /dev/null +++ b/judo/unlocked/__init__.py @@ -0,0 +1,32 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from .const import Const, const, const_cast +from .input_stage import InputStage +from .ipc_utils import Barrier, BroadcastEvent, Event, Mutex, NamedSemaphore, SharedTimeStat +from .node import Node, NodeStop +from .output_stage import OutputStage +from .queue import Queue, QueueView +from .shmque import SIZEOF_SIZE_T, Frame, Memory, SharedMemoryQueue, SharedMemoryQueueView + +__all__: list = [ + "Barrier", + "BroadcastEvent", + "Const", + "const", + "const_cast", + "Event", + "Frame", + "InputStage", + "Memory", + "Mutex", + "NamedSemaphore", + "Node", + "NodeStop", + "OutputStage", + "Queue", + "QueueView", + "SharedMemoryQueue", + "SharedMemoryQueueView", + "SharedTimeStat", + "SIZEOF_SIZE_T", +] diff --git a/judo/unlocked/annotation.py b/judo/unlocked/annotation.py new file mode 100644 index 00000000..c017ce6c --- /dev/null +++ b/judo/unlocked/annotation.py @@ -0,0 +1,28 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from collections.abc import Iterator as IteratorI +from inspect import Parameter +from typing import Any, Mapping +from typing import Iterator as IteratorT + + +def is_tuple_t(annotation: Any) -> bool: + """Check if annotation is of tuple type.""" + return annotation is tuple or (hasattr(annotation, "__origin__") and annotation.__origin__ is tuple) + + +def is_iterator_t(annotation: Any) -> bool: + """Check if annotation is of iterable type.""" + return ( + annotation is IteratorT + or annotation is IteratorI + or (hasattr(annotation, "__origin__") and annotation.__origin__ is IteratorI) + ) + + +def positional(parameters: Mapping, index: int) -> Parameter: + name = list(parameters)[index] + param = parameters[name] + if param.kind >= Parameter.KEYWORD_ONLY: + raise IndexError(f"Parameter {name} is not positional") + return param diff --git a/judo/unlocked/benchmark/benchmark_ipc_utils.py b/judo/unlocked/benchmark/benchmark_ipc_utils.py new file mode 100644 index 00000000..c7c62a15 --- /dev/null +++ b/judo/unlocked/benchmark/benchmark_ipc_utils.py @@ -0,0 +1,148 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import multiprocessing as mp +import traceback +from ctypes import c_bool + +from core.unlocked.ipc_utils import Barrier, BroadcastEvent, Event, Mutex, SharedTimeStat + + +def benchmark_mutex() -> None: + mutex = Mutex("benchmark_mutex") + lock_time = SharedTimeStat("mutex_lock_time", lock=False) + total_time = SharedTimeStat("mutex_lock_time", lock=False) + + n_trials = 10_000 + + for _ in range(n_trials): + lock_time.tic() + total_time.tic() + with mutex: + lock_time.toc() + total_time.toc() + + print(f" Mutex lock time: {lock_time}") + print(f" Mutex round-trip time: {total_time}") + + +def benchmark_event() -> None: + event = Event("benchmark_event") + notify_time = SharedTimeStat("event_notify_time", lock=False) + ping_pong_time = SharedTimeStat("event_ping_pong_time", lock=False) + + n_trials = 10_000 + + def ping() -> None: + for _ in range(n_trials): + ping_pong_time.tic() + notify_time.tic() + event.notify() + ping_pong_time.toc() + + def pong() -> None: + for _ in range(n_trials): + with event: + notify_time.toc() + + ping_proc = mp.Process(target=ping, name="ping") + pong_proc = mp.Process(target=pong, name="pong") + + ping_proc.start() + pong_proc.start() + + ping_proc.join(timeout=30) + pong_proc.join(timeout=30) + + print(f" Event notify time: {notify_time}") + print(f" Ping-pong time: {ping_pong_time}") + + +def benchmark_barrier() -> None: + n_sub = 10 + barrier = Barrier("benchmark_barrier", n_sub + 1) + notify_time = SharedTimeStat("barrier_notify_time", lock=False) + ping_pong_time = SharedTimeStat("barrier_ping_pong_time", lock=False) + + n_trials = 10_000 + + def first() -> None: + for _ in range(n_trials): + ping_pong_time.tic() + notify_time.tic() + with barrier: + notify_time.toc() + ping_pong_time.toc() + + def rest() -> None: + for _ in range(n_trials): + with barrier: + pass + + procs = [mp.Process(target=first, name="first")] + [mp.Process(target=rest, name=f"rest_{i}") for i in range(n_sub)] + + for p in procs: + p.start() + for p in procs: + p.join(timeout=30) + + print(f" Barrier notify time: {notify_time}") + print(f" Ping-pong time: {ping_pong_time}") + + +def benchmark_broadcast() -> None: + broadcast = BroadcastEvent("benchmark_broadcast") + notify_time = SharedTimeStat("broadcast_notify_time", lock=False) + ping_pong_time = SharedTimeStat("broadcast_ping_pong_time", lock=False) + + n_trials = 10_000 + n_sub = 10 + + done = mp.RawArray(c_bool, [False] * n_sub) + + def first() -> None: + try: + while not all(done): + ping_pong_time.tic() + notify_time.tic() + broadcast.notify() + ping_pong_time.toc() + except BaseException: + pass + + def rest(ind: int) -> None: + try: + for _ in range(n_trials): + with broadcast: + if ind == 0: + notify_time.toc() + except BaseException: + if ind == 0: + traceback.print_exc() + finally: + done[ind] = True + + procs = [mp.Process(target=first, name="first")] + [ + mp.Process(target=rest, args=(i,), name=f"rest_{i}") for i in range(n_sub) + ] + for p in procs[1:]: + p.start() + procs[0].start() + try: + for p in procs: + p.join(timeout=30) + except BaseException: + pass + + print(f"Broadcast Event notify time: {notify_time}") + print(f" Ping-pong time: {ping_pong_time}") + + +def main() -> None: + benchmark_mutex() + benchmark_event() + benchmark_barrier() + benchmark_broadcast() + + +if __name__ == "__main__": + main() diff --git a/judo/unlocked/benchmark/benchmark_shmque.py b/judo/unlocked/benchmark/benchmark_shmque.py new file mode 100644 index 00000000..e6671af8 --- /dev/null +++ b/judo/unlocked/benchmark/benchmark_shmque.py @@ -0,0 +1,290 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import multiprocessing as mp +from ctypes import c_bool +from time import perf_counter, sleep + +from core.logging_util import setup_logger +from core.unlocked.ipc_utils import Event, SharedTimeStat +from core.unlocked.shmque import SIZEOF_SIZE_T, Frame, Memory, SharedMemoryQueue + +setup_logger(log_level=logging.INFO) + +n_trials = 1000 +n_sub = 30 +payload = b"\x00" * 3 * 1920 * 1080 # expected image payload size +payload = b"\x00" * 8 * 100 # expected state payload size + + +def benchmark_perf_counter() -> None: + perf_time = SharedTimeStat("perf_time", lock=False) + + for _ in range(n_trials): # publish more not to deprive views of data + with perf_time: + perf_counter() + + print(f" Perf counter time: {perf_time}") + + +def benchmark_sleep() -> None: + sleep_time = SharedTimeStat("sleep_time", lock=False) + + for _ in range(n_trials): # publish more not to deprive views of data + with sleep_time: + sleep(1.0e-300) + + print(f" Sleep time: {sleep_time}") + + +def benchmark_event() -> None: + sync = Event("benchmarked") + + event_time = SharedTimeStat("event_time", lock=False) + + def push_fn() -> None: + for _ in range(n_trials): # publish more not to deprive views of data + event_time.tic() + sync.notify() + + push_process = mp.Process(target=push_fn) + + def view_fn() -> None: + for _ in range(n_trials): + with sync: + event_time.toc() + + view_process = mp.Process(target=view_fn) + view_process.start() + push_process.start() + push_process.join() + view_process.join() + + print(f" Event time: {event_time}") + + +def benchmark_memory_init() -> None: + init_time = SharedTimeStat("init_time", lock=False) + + for i in range(n_trials): + with init_time: + Memory((i, payload)) + + print(f" Memory init time: {init_time}") + + +def benchmark_memory_encode_decode() -> None: + memory: Memory = Memory((n_trials, payload)) + buffer = bytearray(len(memory)) + + encode_time = SharedTimeStat("encode_time", lock=False) + decode_time = SharedTimeStat("decode_time", lock=False) + + for _ in range(n_trials): + with encode_time: + memory.encode(buffer) + with decode_time: + value, _ = Memory.decode(buffer) + assert value == n_trials + + print(f" Memory encode time: {encode_time}") + print(f" Memory decode time: {decode_time}") + + +def benchmark_frame_init() -> None: + size = 400_000_000 # expected memory rate for serializing large objects + + init_time = SharedTimeStat("init_time", lock=False) + + for _ in range(n_trials): + with init_time: + Frame(name="benchmark_frame_init", create=True, size=size) + + print(f" Frame init time: {init_time}") + + +def benchmark_frame_push_memory() -> None: + memory: Memory = Memory((n_trials, payload)) + + frame: Frame = Frame(name="benchmark_frame_push_memory", create=True, size=(len(memory) + SIZEOF_SIZE_T) * n_trials) + + push_time = SharedTimeStat("push_time", lock=False) + + for i in range(n_trials): + memory = Memory((i, payload)) + with push_time: + frame.push_memory(memory) + + print(f"Frame push memory time: {push_time}") + + +def benchmark_frame_push() -> None: + memory: Memory = Memory((n_trials, payload)) + + frame: Frame = Frame(name="benchmark_frame_push", create=True, size=(len(memory) + SIZEOF_SIZE_T) * n_trials) + + push_time = SharedTimeStat("push_time", lock=False) + + for i in range(n_trials): + with push_time: + frame.push((i, payload)) + + print(f" Frame push time: {push_time}") + + +def benchmark_frame_get_item() -> None: + memory: Memory = Memory((n_trials, payload)) + + frame: Frame = Frame(name="benchmark_frame_get_item", create=True, size=(len(memory) + SIZEOF_SIZE_T) * n_trials) + + for i in range(n_trials): + memory = Memory((i, payload)) + frame.push_memory(memory) + + get_time = SharedTimeStat("get_time", lock=False) + + success = mp.RawValue(c_bool, False) + + def get_process() -> None: + frame_view: Frame = Frame(name="benchmark_frame_get_item", create=False) + for i in range(n_trials): + with get_time: + value, _ = frame_view[i] + if value != i: + logging.error("Expected %d; instead, got %d", i, value) + return + success.value = True + + get_proc = mp.Process(target=get_process) + get_proc.start() + get_proc.join() + + assert success.value + + print(f" Frame get item time: {get_time}") + + +def benchmark_queue_push() -> None: + queue: SharedMemoryQueue = SharedMemoryQueue(name="benchmark_queue_push") + queue.open() + queue_push_time = SharedTimeStat("queue_push_time", lock=False) + + for i in range(n_trials): + with queue_push_time: + queue.push((i, payload)) + + queue.close() + + print(f" Queue push time: {queue_push_time}") + + +def benchmark_queue_view() -> None: + queue: SharedMemoryQueue = SharedMemoryQueue(name="benchmark_queue_view") + view = queue.make_view() + + queue.open() + for i in range(n_trials): + queue.push((i, payload)) + + wait_time = SharedTimeStat("wait_time", lock=False) + top_time = SharedTimeStat("top_time", lock=False) + pop_time = SharedTimeStat("pop_time", lock=False) + + success = mp.RawValue(c_bool, False) + + def view_process() -> None: + view.open() + for i in range(n_trials): + with wait_time: + while not view.wait(): + pass + with top_time: + value, _ = view.top() + if value != i: + logging.error("Expected %d; instead, got %d", i, value) + view.close() + return + with pop_time: + view.pop() + view.close() + success.value = True + + view_proc = mp.Process(target=view_process) + view_proc.start() + view_proc.join() + + assert success.value + queue.close() + + print(f" Queue view wait time: {wait_time}") + print(f" Queue view top time: {top_time}") + print(f" Queue view pop time: {pop_time}") + + +def benchmark_queue_transit() -> None: + queue: SharedMemoryQueue = SharedMemoryQueue(name="benchmark_queue_transit") + views = [queue.make_view() for _ in range(n_sub)] + sync = mp.Barrier(n_sub + 1) + + transit_time = SharedTimeStat("transit_time", lock=False) + view_wait_time = [SharedTimeStat(f"view_{sub_ind}_wait_time", lock=False) for sub_ind in range(n_sub)] + + def push_fn() -> None: + queue.open() + for i in range(n_trials): # publish more not to deprive views of data + # logging.info("Publisher iteration %d", i) + queue.push((i, payload, transit_time.tic())) + sync.wait() + queue.close() + + push_process = mp.Process(target=push_fn, name="Push") + + def view_fn(sub_ind: int) -> None: + view = views[sub_ind] + view.open() + for _ in range(n_trials): + # logging.info("Subscriber %d iteration %d", sub_ind, i) + with view_wait_time[sub_ind]: + while not view.wait(): + pass + value, _, now = view.top() + # logging.info("Subscriber %d value %d", sub_ind, value) + view.pop() + if sub_ind == 0: + transit_time.tic(now) + transit_time.toc() + sync.wait() + view.close() + + view_processes = [mp.Process(target=view_fn, args=(i,), name=f"View_{i}") for i in range(n_sub)] + + push_process.start() + for view_proc in view_processes: + view_proc.start() + push_process.join() + for view_proc in view_processes: + view_proc.join() + + print(f" Queue transit time: {transit_time}") + # for sub_ind in range(n_sub): + # print( f" View {sub_ind} wait time: {view_wait_time[sub_ind]}") + + +def main() -> None: + benchmark_perf_counter() + benchmark_sleep() + benchmark_event() + benchmark_memory_init() + benchmark_memory_encode_decode() + benchmark_frame_init() + benchmark_frame_push_memory() + benchmark_frame_push() + benchmark_frame_get_item() + benchmark_queue_push() + benchmark_queue_view() + benchmark_queue_transit() + + +if __name__ == "__main__": + mp.set_start_method("fork") + main() diff --git a/judo/unlocked/const.py b/judo/unlocked/const.py new file mode 100644 index 00000000..7f894f5c --- /dev/null +++ b/judo/unlocked/const.py @@ -0,0 +1,81 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from functools import partial +from types import MethodType +from typing import Any, Generic, TypeVar +from copy import copy, deepcopy + +T = TypeVar("T") + + +def is_builtin_class_instance(obj: Any) -> bool: + return obj.__class__.__module__ == "builtins" and not isinstance(obj, (list, tuple, dict, set)) + + +class Const(Generic[T]): + """Turn any python object into a runtime constant object. + + This wrapper allows access to all member values, and non mutating member functions. + Modifying member values or calling a modifying member funtion will raise RuntimeError. + """ + + def __init__(self, obj: T): + self.__dict__["__obj"] = obj + + def __getattr__(self, name: str, /) -> Any: + result = object.__getattribute__(self.__dict__["__obj"], name) + if isinstance(result, MethodType): + result = partial(result.__func__, self) + return result + return const(result) + + def __setattr__(self, name: str, value: Any, /) -> None: + raise RuntimeError("Setting attribute to read-only object") + + def __eq__(self, obj: object) -> bool: + return self.__dict__["__obj"] == const_cast(obj) + + def __ne__(self, obj: object) -> bool: + return self.__dict__["__obj"] != const_cast(obj) + + def __len__(self) -> int: + return len(self.__dict__["__obj"]) + + def __getitem__(self, index: Any) -> Any: + return const(self.__dict__["__obj"][index]) + + def __contains__(self, el: Any) -> bool: + return el in self.__dict__["__obj"] + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + """Const call method""" + return self.__dict__["__obj"].__class__.__call__(self, *args, **kwargs) + + def __iter__(self) -> Any: + for val in self.__dict__["__obj"]: + yield const(val) + + def __str__(self) -> Any: + return "const " + str(self.__dict__["__obj"]) + + def __copy__(self) -> Any: + return copy(self.__dict__["__obj"]) + + def __deepcopy__(self) -> Any: + return deepcopy(self.__dict__["__obj"]) + + +def const_cast(obj: Any) -> Any: + """A convenience function to remove constant modifyer.""" + if isinstance(obj, Const): + return obj.__dict__["__obj"] + return obj + + +def const(obj: Any) -> Any: + """A convenience function to make a constant object.""" + if is_builtin_class_instance(obj): + return obj + if isinstance(obj, Const): + return obj + return Const(obj) diff --git a/judo/unlocked/example/fibonacci.py b/judo/unlocked/example/fibonacci.py new file mode 100644 index 00000000..79c7970c --- /dev/null +++ b/judo/unlocked/example/fibonacci.py @@ -0,0 +1,48 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import multiprocessing as mp +import signal +from typing import Any, Iterable + +# from operator import add +from core.logging_util import setup_logger +from core.unlocked import Node +from core.unlocked.policy import Lossless +from core.unlocked.schedule import Threaded + +setup_logger(log_level=logging.INFO) + + +class StopAll: + """Terminate all schedulers.""" + + def __init__(self, stop_it: Iterable) -> None: + self._stop = list(stop_it) + + def __call__(self, signum: Any, _: Any) -> None: + for s in self._stop: + s.set() + + +def fibonacci(a: int, b: int) -> tuple[int, int]: + """Subscriber printout.""" + next = a + b + logging.info("Next fibonacci number is %d", next) + return b, next + + +def main() -> None: + f_node = Node("fibonacci", fibonacci, frequency=10, warmup=[(0, 1)]) + f_node.input_stage.connect(0, f_node.output_stage[0], Lossless()) + f_node.input_stage.connect(1, f_node.output_stage[1], Lossless()) + + stop = mp.Event() + schedule1 = Threaded("sch_1", (f_node,), stop) + schedule1.start() + signal.signal(signal.SIGINT, StopAll((stop,))) + schedule1.join() + + +if __name__ == "__main__": + main() diff --git a/judo/unlocked/example/loopback.py b/judo/unlocked/example/loopback.py new file mode 100644 index 00000000..9d744609 --- /dev/null +++ b/judo/unlocked/example/loopback.py @@ -0,0 +1,50 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import multiprocessing as mp +import signal +from typing import Any, Iterable + +# from operator import add +from core.logging_util import setup_logger +from core.unlocked import Node +from core.unlocked.policy import Latest, Lossless +from core.unlocked.schedule import Threaded + +setup_logger(log_level=logging.ERROR, include_thread=True) + + +class StopAll: + """Terminate all schedulers.""" + + def __init__(self, stop_it: Iterable) -> None: + self._stop = list(stop_it) + + def __call__(self, signum: Any, _: Any) -> None: + for s in self._stop: + s.set() + + +def loop_the_loop(msg: int) -> int: + """Subscriber printout.""" + return msg + 1 + + +def subscriber(msg: int) -> None: + logging.error("Received %d", msg) + + +def main() -> None: + loop = Node("loop", loop_the_loop, frequency=5000, warmup=[(0,)]) # 5 kHz! + subs = Node("subs", subscriber, frequency=1) + loop.input_stage.connect(0, loop.output_stage[0], Lossless()) + subs.input_stage.connect("msg", loop.output_stage[0], Latest()) + + schedule1 = Threaded("sch_1", (loop, subs), mp.Event()) + schedule1.start() + signal.signal(signal.SIGINT, StopAll((schedule1._stop,))) + schedule1.join() + + +if __name__ == "__main__": + main() diff --git a/judo/unlocked/example/pub_sub.py b/judo/unlocked/example/pub_sub.py new file mode 100644 index 00000000..b32be34a --- /dev/null +++ b/judo/unlocked/example/pub_sub.py @@ -0,0 +1,65 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import multiprocessing as mp +import signal +from typing import Any, Iterable + +# from operator import add +from core.logging_util import setup_logger +from core.unlocked import Node +from core.unlocked.policy import Latest +from core.unlocked.schedule import Threaded + +setup_logger(log_level=logging.ERROR) + + +class StopAll: + """Terminate all schedulers.""" + + def __init__(self, stop_it: Iterable) -> None: + self._stop = list(stop_it) + + def __call__(self, signum: Any, _: Any) -> None: + for s in self._stop: + s.set() + + +class Publisher: + """Publishing generator.""" + + def __init__(self) -> None: + self._counter = 0 + + def __call__(self) -> int: + self._counter += 1 + # logging.info("Sending %d", self._counter) + return self._counter + + +def subscriber(msg: int) -> None: + """Subscriber printout.""" + logging.error("Received %d", msg) + + +def main() -> None: + pub = Node("pub", Publisher(), frequency=10000) # 10 kHz! + sub = Node("sub", subscriber, frequency=1) + sub.input_stage.connect("msg", pub.output_stage[0], Latest()) + + stop = mp.Event() + schedule1 = Threaded( + "sch_1", + ( + pub, + sub, + ), + stop, + ) + schedule1.start() + signal.signal(signal.SIGINT, StopAll((stop,))) + schedule1.join() + + +if __name__ == "__main__": + main() diff --git a/judo/unlocked/input_stage.py b/judo/unlocked/input_stage.py new file mode 100644 index 00000000..30be6f45 --- /dev/null +++ b/judo/unlocked/input_stage.py @@ -0,0 +1,90 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from inspect import Parameter +from typing import Any, Mapping +from typing import Any as Policy + +from .shmque import SharedMemoryQueueView + + +class InputStage: + """Node output stage.""" + + def __init__(self, name: str, parameters: Mapping): + self._name = name + self._parameters = parameters + self._inputs: dict[str | int, tuple[SharedMemoryQueueView, Policy]] = {} + self._arg_len = 0 + self._arg_inputs: list[tuple[SharedMemoryQueueView, Policy]] = [] + self._kwarg_inputs: list[tuple[str, SharedMemoryQueueView, Policy]] = [] + + def __len__(self) -> int: + """Number of inputs.""" + return len(self._parameters) + + @property + def ready(self) -> bool: + """Check if data is available.""" + return all((policy.ready(view) for view, policy in self._inputs.values())) + + def wait(self, timeout: float | None = None) -> bool: + """Wait for the data.""" + for view, policy in self._inputs.values(): + if not policy.wait(view, timeout=timeout): + return False + return True + + def open(self) -> None: + """Open input shared memory views.""" + self._assert_contiguous_positional_arguments() + self._arg_inputs = [self._inputs[arg] for arg in range(self._arg_len)] + self._kwarg_inputs = [(key, *value) for key, value in self._inputs.items() if not isinstance(key, int)] + for view, _ in self._inputs.values(): + view.open() + + def close(self) -> None: + """Close input shared memory views.""" + for view, _ in self._inputs.values(): + view.close() + + def connect(self, arg: int | str, queue_view: SharedMemoryQueueView, policy: Policy) -> None: + """Connect a view to a positional of keyword input argument.""" + if ( + not isinstance(arg, int) + and arg in self._parameters + and self._parameters[arg].kind == Parameter.POSITIONAL_OR_KEYWORD + ): + arg = list(self._parameters.keys()).index(arg) + + if arg in self._inputs: + raise ValueError("Input stage cannot subscribe to multiple output stages.") + + self._inputs[arg] = (queue_view, policy) + + def input_args(self) -> tuple[list[Any], dict[str, Any]]: + """Get input arguments from the views.""" + args = [policy.get(view) for view, policy in self._arg_inputs] + kwargs = {name: policy.get(view) for name, view, policy in self._kwarg_inputs} + return args, kwargs + + def next(self) -> None: + """Pop data from the views.""" + for view, policy in self._inputs.values(): + policy.next(view) + + def _assert_contiguous_positional_arguments(self) -> None: + if len(self._inputs) == 0: + return + + is_min_zero = False + for arg in self._inputs.keys(): + if not isinstance(arg, int): + continue + if arg == 0: + is_min_zero = True + self._arg_len = max(self._arg_len, arg + 1) + + assert is_min_zero, "Missing the first positional argument" + + for arg in range(self._arg_len): + assert arg in self._inputs, f"Missing the {arg}th positional argument" diff --git a/judo/unlocked/ipc_utils.py b/judo/unlocked/ipc_utils.py new file mode 100644 index 00000000..8e2e4fe5 --- /dev/null +++ b/judo/unlocked/ipc_utils.py @@ -0,0 +1,295 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from contextlib import nullcontext +from ctypes import c_double, c_uint64 +from math import sqrt +from multiprocessing import Semaphore +from multiprocessing.sharedctypes import RawValue +from time import perf_counter_ns +from typing import Any, Generic, TypeVar + +T = TypeVar("T") + + +def _f(nanoseconds: float) -> str: + if nanoseconds >= 1.0e9: + return f"{(nanoseconds*1.e-9):6.2f} s" + if nanoseconds >= 1.0e6: + return f"{(nanoseconds*1.e-6):6.2f} ms" + if nanoseconds >= 1.0e3: + return f"{(nanoseconds*1.e-3):6.2f} us" + return f"{(nanoseconds):6.2f} ns" + + +class SharedTimeStat: + """Accumulate time statistics across multiple processes.""" + + def __init__(self, name: str, *, lock: bool = True): + self._mutex: Mutex | nullcontext = Mutex(name + "_mutex") if lock else nullcontext() + self._n = RawValue(c_uint64, 0) + self._min = RawValue(c_uint64, -1) + self._max = RawValue(c_uint64, 0) + self._mu = RawValue(c_double, 0.0) + self._s2 = RawValue(c_double, 0.0) + self._tic = RawValue(c_uint64, 0) + + def __enter__(self) -> "SharedTimeStat": + self.tic() + return self + + def __exit__(self, *args: Any) -> None: + self.toc() + + def __str__(self) -> str: + with self._mutex: + n = self._n.value + if n == 0: + return "================================= No Data =================================" + return ( + f"n = {n}, " + f"min = {_f(self._min.value)}, " + f"max = {_f(self._max.value)}, " + f"avg = {_f(self._mu.value)}, " + f"std = {_f(sqrt(self._s2.value / n))}" + ) + + def clear(self) -> None: + """Clear accumulated statistics.""" + with self._mutex: + self._n.value = 0 + self._min.value = -1 + self._max.value = 0 + self._mu.value = 0 + self._s2.value = 0 + self._tic.value = 0 + + def tic(self, now: int | None = None) -> int: + """Start time tracking.""" + if now is None: + now = perf_counter_ns() + with self._mutex: + self._tic.value = now + return now + + def toc(self, now: int | None = None) -> None: + """Stop time tracking and accumulated staticstics.""" + if now is None: + now = perf_counter_ns() + with self._mutex: + self._n.value += 1 + x = now - self._tic.value + self._min.value = min(self._min.value, x) + self._max.value = max(self._max.value, x) + delta = x - self._mu.value + self._mu.value = self._mu.value + delta / self._n.value + delta2 = x - self._mu.value + self._s2.value = self._s2.value + delta * delta2 + + @property + def n(self) -> int: + """Number of records.""" + with self._mutex: + return self._n.value + + @property + def mu(self) -> float: + """Avarage time in seconds.""" + with self._mutex: + return self._mu.value * 1.0e-9 + + @property + def sigma(self) -> float: + """Time standard deviation in seconds.""" + with self._mutex: + n = self._n.value + if n == 0: + return float("nan") + return sqrt(self._s2.value / self._n.value) * 1.0e-9 + + @property + def min(self) -> float: + """Minimum time in seconds.""" + with self._mutex: + return self._min.value * 1.0e-9 + + @property + def max(self) -> float: + """Maximum time in seconds.""" + with self._mutex: + return self._max.value * 1.0e-9 + + +# We are preparing for using POSIX Semaphores to be able to synch with C++/Rust/Zig code. +# (POSIX Semaphores)[https://github.com/osvenskan/posix_ipc/blob/develop/USAGE.md#the-semaphore-class] +# We also will create custom Test Semaphores that will simulate random arrival times. +class NamedSemaphore: + """Posix semaphore.""" + + def __init__(self, name: str, value: int = 1): + self.__name__ = name + self._semaphore = Semaphore(value) + + def acquire(self) -> None: + """Acquire semaphore.""" + self._semaphore.acquire() + + def release(self, n: int = 1) -> None: + """Release semaphore.""" + if n == 1: # it's a bit faster to do if then a for loop of 1 + self._semaphore.release() + else: + for _ in range(n): + self._semaphore.release() + + +class UseGenericSemaphore(Generic[T]): + """Templated semaphore.""" + + __SemaphoreT__: type = NamedSemaphore + + def __class_getitem__(cls, SemaphoreT: type) -> type: + UseSemaphore = cls + UseSemaphore.__SemaphoreT__ = SemaphoreT + return UseSemaphore + + +class Mutex(UseGenericSemaphore): + """Mutex.""" + + def __init__(self, name: str): + self._name = name + self._s = self.__SemaphoreT__(name + "_s", 1) + + def __enter__(self) -> None: + self._s.acquire() + + def __exit__(self, *args: Any) -> None: + self._s.release() + + def acquire(self) -> None: + """Manually acquire mutex.""" + self._s.acquire() + + def release(self) -> None: + """Manually release mutex.""" + self._s.release() + + +class Event(UseGenericSemaphore): + """Event syncronizes a calling process with one and only one of waiting processes.""" + + def __init__(self, name: str): + self._name = name + self._s1 = self.__SemaphoreT__(name + "_s1", 0) + self._s2 = self.__SemaphoreT__(name + "_s2", 0) + + def __enter__(self) -> None: + self._s1.acquire() + + def __exit__(self, *args: Any) -> None: + self._s2.release() + + def notify(self) -> None: + """Notify one waiting process and wait for it to wake up.""" + self._s1.release() + self._s2.acquire() + + +class BroadcastEvent(UseGenericSemaphore): + """Broadcast event syncronizes a calling process with all (any number) of waiting processes.""" + + def __init__(self, name: str): + self._name = name + self._counter = RawValue(c_uint64, 0) + self._mutex = Mutex[self.__SemaphoreT__](name + "_mutex") # type: ignore[misc] + self._t = Mutex[self.__SemaphoreT__](name + "_turn") # type: ignore[misc] + + self._s1 = self.__SemaphoreT__(name + "_s1", 0) + self._s2 = self.__SemaphoreT__(name + "_s2", 1) + self._s3 = self.__SemaphoreT__(name + "_s3", 0) + + def __enter__(self) -> None: + # Block until notifier is done + with self._t: + # Count number of listeners + with self._mutex: + self._counter.value += 1 + + # Wait for the notice + self._s1.acquire() + + def __exit__(self, *args: Any) -> None: + # Release listeners + with self._mutex: + self._counter.value -= 1 + if self._counter.value == 0: + self._s2.release() + self._s3.release() + + self._s2.acquire() + self._s2.release() + + def notify(self) -> None: + """Notify all waiting process and wait for all of them to wake up.""" + # Notifier block + with self._t: + with self._mutex: + n = self._counter.value + # if no listeners are waiting---leave + if n == 0: + return + # notify + self._s2.acquire() + self._s1.release(n) + self._s3.acquire() + + +class Barrier(UseGenericSemaphore): + """Barrier syncronizes exactly n processes.""" + + def __init__(self, name: str, n: int): + self._name = name + self._n = n + self._counter = RawValue(c_uint64, 0) + self._mutex = Mutex[self.__SemaphoreT__](name + "_mutex") # type: ignore[misc] + self._s1 = self.__SemaphoreT__(name + "_s1", 0) + self._s2 = self.__SemaphoreT__(name + "_s2", 1) + + def __enter__(self) -> None: + with self._mutex: + self._counter.value += 1 + if self._counter.value == self._n: + self._s2.acquire() + self._s1.release(self._n + 1) + + self._s1.acquire() + + def __exit__(self, *args: Any) -> None: + with self._mutex: + self._counter.value -= 1 + if self._counter.value == 0: + self._s1.acquire() + self._s2.release(self._n + 1) + + self._s2.acquire() + + # A slower, but perhaps a safer solution, in which threads are unlocked one-by-one. + # def __enter__(self) -> None: + # with self._mutex: + # self._counter.value += 1 + # if self._counter.value == self._n: + # self._s2.acquire() + # self._s1.release() + # + # self._s1.acquire() + # self._s1.release() + # + # def __exit__(self, *args) -> None: + # with self._mutex: + # self._counter.value -= 1 + # if self._counter.value == 0: + # self._s1.acquire() + # self._s2.release() + # + # self._s2.acquire() + # self._s2.release() diff --git a/judo/unlocked/node.py b/judo/unlocked/node.py new file mode 100644 index 00000000..f1b22e5d --- /dev/null +++ b/judo/unlocked/node.py @@ -0,0 +1,137 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import traceback +from inspect import signature +from time import perf_counter, sleep +from typing import Any, Callable, Iterable + +from .input_stage import InputStage +from .output_stage import OutputStage + + +class NodeStop(Exception): + """Exception to stop node execution.""" + + +_ADDED_SLEEP_DURATION = 5.41e-5 + + +class Node: + """Node is a smallest execution unit in the pipeline.""" + + def __init__(self, name: str, target: Callable, *, frequency: float | None = None, warmup: Iterable = []): + self._name = name + self._target = target + self._period = None if frequency is None else 1.0 / frequency + self._warmup = list(warmup) + self._signature = signature(self._target) + self._output_stage = OutputStage(self._name, self._signature.return_annotation) + self._input_stage = InputStage(self._name, self._signature.parameters) + self._stop: NodeStop | None = None + self._in_exec = False + self._count = 0 + self._last_update = perf_counter() + + assert frequency is None or frequency > 0, "Frequency must be positive" + assert not (len(self.input_stage) == 0 and frequency is None), "Update frequency is required for generators" + + @property + def name(self) -> str: + """Check if node has not been terminated.""" + return self._name + + @property + def live(self) -> bool: + """Check if node has not been terminated.""" + return self._stop is None + + @property + def in_exec(self) -> bool: + """Check if node function is currently running.""" + return self._in_exec + + @property + def output_stage(self) -> OutputStage: + """Get output stage.""" + return self._output_stage + + @property + def input_stage(self) -> InputStage: + """Get input stage.""" + return self._input_stage + + @property + def ready(self) -> bool: + """Get input stage.""" + if self._period is not None and perf_counter() - self._last_update < self._period: + return False + return self._input_stage.ready + + def wait(self, timeout: int | None = None) -> bool: + """Wait for the input data to arrive.""" + if self._period is None: + return self._input_stage.wait(timeout=timeout) + time_left = self._period + self._last_update - _ADDED_SLEEP_DURATION - perf_counter() + if timeout is not None and time_left > timeout - _ADDED_SLEEP_DURATION: + sleep(timeout - _ADDED_SLEEP_DURATION) + return False + if time_left > 0: + sleep(time_left) + if not self.live: + return False + return self._input_stage.wait(timeout=0.5) + + def open(self) -> None: + """Open node shared memory communications.""" + self._output_stage.open() + for w in self._warmup: + self._output_stage.push(w) + self._input_stage.open() + + def close(self) -> None: + """Close node shared memory communications.""" + self._stop = NodeStop("Closed nominally.") + self._input_stage.close() + self._output_stage.close() + + def exec(self) -> None: + """Run the node function.""" + if self._period is not None: + time_passed = perf_counter() - self._last_update + if time_passed <= 1.5 * self._period: + self._last_update += self._period + else: + delay_percent = int(100 * (time_passed / self._period - 1.0)) + # logging.warning("Node %s execution is delayed by %d %%", self._name, delay_percent) + self._last_update += time_passed + assert self.live, "The node has been stopped." + self._in_exec = True + assert self._input_stage.ready, "Input stage is not ready." + args, kwargs = self._input_stage.input_args() + try: + result = self._target(*args, **kwargs) + except BaseException as e: + self._set_error(e) + return + + self._push_result(result) + + def _push_result(self, result: Any) -> None: + if result is None: + result = tuple() + elif not isinstance(result, tuple): + result = (result,) + self._input_stage.next() + self._output_stage.push(result) + self._in_exec = False + self._count += 1 + + def _set_error(self, error: Any) -> None: + assert error is not None + if isinstance(error, NodeStop): + self._stop = error + logging.info("Node %s stopped with %s", self._name, error) + else: + logging.error("Node %s failed with %s\n === traceback ===\n%s =================", self._name, error, traceback.format_exc()) + self._in_exec = False diff --git a/judo/unlocked/output_stage.py b/judo/unlocked/output_stage.py new file mode 100644 index 00000000..0aca9da5 --- /dev/null +++ b/judo/unlocked/output_stage.py @@ -0,0 +1,51 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from typing import Any + +from .annotation import is_tuple_t +from .shmque import SharedMemoryQueue, SharedMemoryQueueView + + +class OutputStage: + """Node output stage.""" + + def __init__(self, name: str, return_annotation: Any): + self._name = name + if return_annotation is None: + return_annotation = tuple() + elif is_tuple_t(return_annotation): + return_annotation = return_annotation.__args__ + else: + return_annotation = (return_annotation,) + + # the following was meant to satisfy mypy, but it actually made matters worse + # queue_annotations = tuple((Queue[t] for t in return_annotation)) + # queue_tuple_annotation = tuple[queue_annotations] + # self._output_queues: queue_tuple_annotation = tuple((Queue[t]() for t in return_annotation)) + self._output_queues: tuple = tuple( + (SharedMemoryQueue(f"{name}_output_{i}") for i, t in enumerate(return_annotation)) + ) + self._len = len(self._output_queues) + + def __len__(self) -> int: + return self._len + + def __getitem__(self, channel: int) -> SharedMemoryQueueView: + if channel < 0 or channel >= self._len: + raise IndexError("Channel index is out of range") + return self._output_queues[channel].make_view() + + def open(self) -> None: + """Open output shared memory queue.""" + for queue in self._output_queues: + queue.open() + + def close(self) -> None: + """Close output shared memory queue.""" + for queue in self._output_queues: + queue.close() + + def push(self, result: tuple) -> None: + """Push elelemtns to the output stage.""" + for r, q in zip(result, self._output_queues, strict=True): + q.push(r) diff --git a/judo/unlocked/policy.py b/judo/unlocked/policy.py new file mode 100644 index 00000000..587c869e --- /dev/null +++ b/judo/unlocked/policy.py @@ -0,0 +1,147 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from ctypes import c_double, c_uint64 +from multiprocessing.sharedctypes import RawValue +from typing import Any + +from .shmque import SharedMemoryQueueView + + +class Lossless: + """Lossless policy""" + + def ready(self, view: SharedMemoryQueueView) -> bool: + """Check if the input is ready.""" + return len(view) > 0 + + def wait(self, view: SharedMemoryQueueView, timeout: float | None = None) -> bool: + """Wait for the input.""" + return view.wait(timeout=timeout) + + def get(self, view: SharedMemoryQueueView) -> Any: + """Get the next element.""" + return view.top() + + def next(self, view: SharedMemoryQueueView) -> None: + """Advance the input.""" + view.pop() + + +class Latest: + """Lossless policy""" + + def __init__(self) -> None: + self._drop = 0 + self._n = RawValue(c_uint64, 0) + self._min = RawValue(c_uint64, -1) + self._max = RawValue(c_uint64, 0) + self._mu = RawValue(c_double, 0.0) + self._s2 = RawValue(c_double, 0.0) + + def ready(self, view: SharedMemoryQueueView) -> bool: + """Check if the input is ready.""" + return len(view) > 0 + + def wait(self, view: SharedMemoryQueueView, timeout: float | None = None) -> bool: + """Wait for the input.""" + return view.wait(timeout=timeout) + + def get(self, view: SharedMemoryQueueView) -> Any: + """Get the next eleent.""" + self._drop = len(view) - 1 + assert self._drop >= 0 + # accumulate dropped message statistics + self._n.value += 1 + self._min.value = min(self._min.value, self._drop) + self._max.value = max(self._max.value, self._drop) + delta = self._drop - self._mu.value + self._mu.value = self._mu.value + delta / self._n.value + delta2 = self._drop - self._mu.value + self._s2.value = self._s2.value + delta * delta2 + # return the latest message + return view[view.begin + self._drop] + + def next(self, view: SharedMemoryQueueView) -> None: + """Advance the input.""" + view.pop(self._drop + 1) + +class Optional: + """Lossless policy""" + + def __init__(self) -> None: + self._drop = 0 + self._n = RawValue(c_uint64, 0) + self._min = RawValue(c_uint64, -1) + self._max = RawValue(c_uint64, 0) + self._mu = RawValue(c_double, 0.0) + self._s2 = RawValue(c_double, 0.0) + + def ready(self, view: SharedMemoryQueueView) -> bool: + """Check if the input is ready.""" + return True + + def wait(self, view: SharedMemoryQueueView, timeout: float | None = None) -> bool: + """Wait for the input.""" + return True + + def get(self, view: SharedMemoryQueueView) -> Any: + """Get the next eleent.""" + self._drop = len(view) - 1 + if self._drop < 0: + return None + # accumulate dropped message statistics + self._n.value += 1 + self._min.value = min(self._min.value, self._drop) + self._max.value = max(self._max.value, self._drop) + delta = self._drop - self._mu.value + self._mu.value = self._mu.value + delta / self._n.value + delta2 = self._drop - self._mu.value + self._s2.value = self._s2.value + delta * delta2 + # return the latest message + return view[view.begin + self._drop] + + def next(self, view: SharedMemoryQueueView) -> None: + """Advance the input.""" + if self._drop < 0: + return + view.pop(self._drop + 1) + +class KeepLatest: + """Lossless policy""" + + def __init__(self) -> None: + self._drop = 0 + self._n = RawValue(c_uint64, 0) + self._min = RawValue(c_uint64, -1) + self._max = RawValue(c_uint64, 0) + self._mu = RawValue(c_double, 0.0) + self._s2 = RawValue(c_double, 0.0) + + def ready(self, view: SharedMemoryQueueView) -> bool: + """Check if the input is ready.""" + return len(view) > 0 + + def wait(self, view: SharedMemoryQueueView, timeout: float | None = None) -> bool: + """Wait for the input.""" + return view.wait(timeout=timeout) + + def get(self, view: SharedMemoryQueueView) -> Any: + """Get the next eleent.""" + self._drop = len(view) - 1 + assert self._drop >= 0 + # accumulate dropped message statistics + self._n.value += 1 + self._min.value = min(self._min.value, self._drop) + self._max.value = max(self._max.value, self._drop) + delta = self._drop - self._mu.value + self._mu.value = self._mu.value + delta / self._n.value + delta2 = self._drop - self._mu.value + self._s2.value = self._s2.value + delta * delta2 + # return the latest message + return view[view.begin + self._drop] + + def next(self, view: SharedMemoryQueueView) -> None: + """Advance the input.""" + if self._drop <= 0: + return + view.pop(self._drop) diff --git a/judo/unlocked/queue.py b/judo/unlocked/queue.py new file mode 100644 index 00000000..33f9a872 --- /dev/null +++ b/judo/unlocked/queue.py @@ -0,0 +1,195 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from collections import deque +from multiprocessing import Event +from typing import Callable, Generator, Generic, TypeVar + +from .const import Const, const + +T = TypeVar("T") + + +class Queue(Generic[T]): + """Shared queue.""" + + def __init__(self) -> None: + self._read: deque[T] = deque() + self._writable: tuple[deque[T], deque[T]] = (deque(), deque()) + self._write_index = 0 + self._offset = 0 + self._views: list[QueueView[T]] = [] + + @property + def _write(self) -> deque[T]: + return self._writable[self._write_index] + + @property + def _buffer(self) -> deque[T]: + return self._writable[self._write_index ^ 1] + + # @property.setter + # def _buffer(self, value: deque[T]) -> None: + # self._writable[self._write_index ^ 1] = value + + @property + def begin(self) -> int: + """The first available data index.""" + return self._offset + + @property + def end(self) -> int: + """Past last available data index.""" + return self._offset + len(self) + + def __getitem__(self, index: int) -> Const | T: + index -= self._offset + if index < 0: + raise IndexError("Out of bounds") + if index >= len(self._read): + self._swap() + return const(self._read[index]) + + def _swap(self) -> None: + self._write_index ^= 1 + self._read.extend(self._buffer) + self._buffer.clear() + + def __len__(self) -> int: + return len(self._read) + len(self._write) + + def flush(self) -> None: + """Manually flush unused elements.""" + min_begin = min(self.end, *(view.begin for view in self._views)) + if min_begin == self._offset: + return + if min_begin - self._offset > len(self._read): + self._swap() + + assert min_begin - self._offset <= len(self._read) + [self._read.popleft() for _ in range(min_begin - self._offset)] + self._offset = min_begin + + def push(self, value: T) -> None: + """Push an element into a queue.""" + self._write.append(value) + # With small number of subscribers (views) it is much faster to + # i) notify subscriber views directly from publish thread, e.g., + for view in self._views: + view.notify() + # , than ii) spin a separate thread and notify views from there, e.g., + # Thread(target=lambda : [view.notify() for view in self._views]).begin() + + def pop(self) -> T: + """Remove top element from the queue.""" + if not self._read: + self._swap() + result = self._read.popleft() + self._offset += 1 + return result + + def top(self) -> Const | T: + """Get top element from the queue.""" + if not self._read: + self._swap() + return const(self._read[0]) + + def clear(self) -> None: + """Clear the queue.""" + for view in self._views: + view.clear() + self._write_index ^= 1 + self._buffer.clear() + self._read.clear() + self._offset = 0 + + def make_view(self) -> "QueueView": + """Make a queue view.""" + queue_view = QueueView(self, begin=self.begin, end=self.end) + self._views.append(queue_view) + return queue_view + + +class QueueView(Generic[T]): + """The view fot the shared queue.""" + + def __init__(self, queue: Queue[T], begin: int = 0, end: int = 0, max_len: None | int = None): + self._queue = queue + self._begin = begin + self._end = end + self._max_len = max_len + self._has_data = Event() + self._on_ready: list[Callable] = [] + + @property + def begin(self) -> int: + """The first available data index.""" + return self._begin + + @property + def end(self) -> int: + """Past last available data index.""" + return self._end + + def __getitem__(self, index: int) -> Const | T | tuple: + if index < self._begin or index >= self._end: + raise IndexError("Out of bounds") + return self._queue[index] + + def notify(self) -> None: + """Notify view that new data is available.""" + if not self.ready: + [callback() for callback in self._on_ready] + self._has_data.set() + + def wait(self, timeout: float | None = None) -> None: + """Wait for new data.""" + self._has_data.wait(timeout) + + @property + def ready(self) -> bool: + """Check if data is available.""" + return len(self) > 0 or self._has_data.is_set() + + def pop(self) -> Const | T | tuple: + """Remove top element from the view.""" + if self._begin == self._end: + raise IndexError("Queue is empty") + self._begin += 1 + return self._queue[self._begin - 1] + + def top(self) -> Const | T | tuple: + """Get view top element.""" + if self._begin == self._end: + raise IndexError("Queue is empty") + return self._queue[self._begin] + + def last(self) -> Const | T | tuple: + """Get view last element.""" + if self._begin == self._end: + raise IndexError("Queue is empty") + return self._queue[self._end - 1] + + def flush(self, *, flush_queue: bool = False) -> None: + """Manually flush the view.""" + self._has_data.clear() + self._begin = self._end + self._end = self._queue.end + if self._max_len is not None: + self._begin = max(self._begin, self._end - self._max_len) + if flush_queue: + self._queue.flush() + + def __iter__(self) -> Generator: + return (self._queue[i] for i in range(self.begin, self._end)) + + def __len__(self) -> int: + return self._end - self._begin + + def clear(self) -> None: + """Clear the view.""" + self._begin = self._queue.begin + self._end = self._begin + + def register(self, on_ready: Callable) -> None: + """Register the callback on available data.""" + self._on_ready.append(on_ready) diff --git a/judo/unlocked/schedule.py b/judo/unlocked/schedule.py new file mode 100644 index 00000000..0f485061 --- /dev/null +++ b/judo/unlocked/schedule.py @@ -0,0 +1,69 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import multiprocessing as mp +import multiprocessing.synchronize +import signal +from threading import Thread +from time import sleep +from typing import Iterable + +from .node import Node + + +def _node(node: Node) -> None: + while node.live: + if node.wait(): + node.exec() + + +class Threaded: + """Threaded scheduler.""" + + def __init__(self, name: str, nodes: Iterable[Node], stop: mp.synchronize.Event): + self._name = name + self._nodes = list(nodes) + self._proc = mp.Process(target=self.exec, name=name) + self._stop = stop + + def start(self) -> None: + """Start scheduler.""" + self._proc.start() + + def join(self) -> None: + """Join scheduler.""" + self._proc.join() + + def exec(self) -> None: + """Run execution loop.""" + signal.signal(signal.SIGINT, signal.SIG_IGN) + self.spin() + + def spin(self) -> None: + logging.debug("Scheduler %s has started", self._name) + + for node in self._nodes: + node.open() + + self._start_node_threads() + + while not self._stop.is_set(): + sleep(0.1) + + for node in self._nodes: + node.close() + + logging.debug("Scheduler %s has stopped", self._name) + + def _start_node_threads(self) -> None: + self.node_threads = [ + Thread( + target=_node, + args=(node,), + name=node.name, + daemon=True, + ) + for node_ind, node in enumerate(self._nodes) + ] + for t in self.node_threads: + t.start() diff --git a/judo/unlocked/shmque.py b/judo/unlocked/shmque.py new file mode 100644 index 00000000..339b0c4e --- /dev/null +++ b/judo/unlocked/shmque.py @@ -0,0 +1,536 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import logging +import multiprocessing as mp +import multiprocessing.synchronize +import pickle +import struct +import time +from bisect import bisect +from ctypes import c_size_t, sizeof +from itertools import accumulate, pairwise +from multiprocessing.shared_memory import SharedMemory +from threading import Event, Thread +from typing import Any, Generic, TypeVar + +from .const import Const, const + +SIZEOF_SIZE_T = sizeof(c_size_t) +SIZE_T_MAX = c_size_t(-1).value +MIN_FRANE_SIZE = 64 * 1024 # 64 kb +T = TypeVar("T") + + +# Due to a bug in Python 3.10 (https://bugs.python.org/issue39959) +# we need to patch resource tracker for shared memory resource tracking +# TODO (dmitry): remove this hack when we switch to Python 3.13 or later versions +def _patch(name: str, rtype: str): + '''NoOp to patch resource tracker.''' + pass +mp.resource_tracker.register = _patch +mp.resource_tracker.unregister = _patch + +# inspired by https://github.com/joblib/joblib/issues/1094 +class Memory(Generic[T]): + """Memory class for serialization and deserialization of python objects into and from a memory buffer.""" + + def __init__(self, obj: Any): + self._buffers: list[memoryview] = [] + dump = pickle.dumps(obj, protocol=5, buffer_callback=self._on_buffer) + self._buffers.append(memoryview(dump)) + self._preamble = b"".join( + [struct.pack("N", len(self._buffers))] + [struct.pack("N", len(buffer)) for buffer in self._buffers] + ) + self._buffers = [memoryview(self._preamble), *self._buffers] + self._size = sum(map(len, self._buffers)) + + def _on_buffer(self, buffer: pickle.PickleBuffer) -> None: + self._buffers.append(buffer.raw()) + + def __len__(self) -> int: + return self._size + + def encode(self, memory: bytearray | memoryview) -> None: + """Serialize held object into a buffer""" + if not isinstance(memory, memoryview): + memory = memoryview(memory) + assert self._size <= len(memory) + offset = 0 + for buffer in self._buffers: + size = len(buffer) + memory[offset : offset + size] = buffer + offset += size + assert offset == self._size + + @classmethod + def decode(cls, memory: bytearray | memoryview) -> Any: + """Deserialize from a buffer into an object view (object data may be shared between multiple views).""" + if not isinstance(memory, memoryview): + memory = memoryview(memory) + buffers_num = struct.unpack("N", memory[:SIZEOF_SIZE_T])[0] + buffers_size: list[int] = [ + struct.unpack("N", memory[(offset + 1) * SIZEOF_SIZE_T : (offset + 2) * SIZEOF_SIZE_T])[0] + for offset in range(buffers_num) + ] + buffers_offset = accumulate(buffers_size, initial=(buffers_num + 1) * SIZEOF_SIZE_T) + + buffers: list[memoryview] = [memory[begin:end] for begin, end in pairwise(buffers_offset)] + + return pickle.loads(buffers[-1], buffers=buffers) + + +class Frame(Generic[T]): + """A frame is a fixed size shared memory block that is used by variable size shared memory queue. + + Frame layout + +--------+--------+-------+--------+------------------------+----------+-------+----------+----------+---------+ + | data_1 | data_2 | ... | data_n | //////// free //////// | offset_n | ... | offset_3 | offset_2 | n | + +--------+--------+-------+--------+------------------------+----------+-------+----------+----------+---------+ + ^ ^ ^ ^ ^ ^ ^ ^ ^ ^ ^ + 0 | offset_3 | free_offset | size - 4(n-1) | size - 8 | size + offset_2 offset_n size - 4n size - 12 size - 4 + """ + + def __init__(self, *, name: str, create: bool, size: int | None = None): + logging.debug("Opening frame%s %s", "" if create else " view", name) + assert not create or (size is not None), "Cannot specify size for shared memory that has been already created." + self.name = name + self._shm: SharedMemory + self._time = time.perf_counter() + if size is None: + self._shm = SharedMemory(name=name, create=create) + else: + self._shm = SharedMemory(name=name, create=create, size=size + SIZEOF_SIZE_T) + self._buffer = self._shm.buf + self._is_view = not create + + self._len = 0 + self._size = len(self._buffer) + self._free = self._size - SIZEOF_SIZE_T + self._used = SIZEOF_SIZE_T + self._data_offset = [0] + self._free_offset = 0 + + if self._is_view: + self.free = self.__deleted__ # type: ignore[assignment, method-assign] + self.used = self.__deleted__ # type: ignore[assignment, method-assign] + self.push = self.__deleted__ # type: ignore[assignment, method-assign] + self.push_memory = self.__deleted__ # type: ignore[assignment, method-assign] + self.data_offset = self._view_data_offset # type: ignore[method-assign] + + def __deleted__(self, *args: Any, **kwargs: Any) -> None: + raise RuntimeError("Not implemented for views") + + def __len__(self) -> int: + return struct.unpack("N", self._buffer[-SIZEOF_SIZE_T:])[0] + + def __del__(self) -> None: + logging.debug("Closing frame%s %s", " view" if self._is_view else "", self.name) + self._shm.close() + if not self._is_view: + logging.debug("Unlinking frame%s %s", " view" if self._is_view else "", self.name) + self._shm.unlink() + + def free(self) -> int: + """Free memory size in bytes.""" + return self._free + + def fits_size(self, size: int) -> bool: + """Check if an object of a given size fits into the frame.""" + return size + SIZEOF_SIZE_T <= self._free + + def used(self) -> int: + """Used memory size in bytes.""" + return self._used + + def size(self) -> int: + """Frame total memory size in bytes.""" + return self._size + + def data_offset(self, index: int) -> int: + """Memory offset in bytes where an [index] object is stored.""" + if index < 0 or index > len(self): + raise IndexError("Out of bounds") + return self._data_offset[index] + + def _view_data_offset(self, index: int) -> int: + if index < 0 or index >= len(self): + raise IndexError("Out of bounds") + if index == 0: + return 0 + return struct.unpack("N", self._buffer[-(index + 1) * SIZEOF_SIZE_T : -index * SIZEOF_SIZE_T])[0] + + def time(self) -> float: + """Time at which frame has been created.""" + return self._time + + def __getitem__(self, index: int) -> T: + if index < 0 or index >= len(self): + raise IndexError("Out of bounds") + return Memory.decode(self._buffer[self.data_offset(index) :]) + + def push(self, el: T) -> None: + """Push elelement into the frame.""" + el_mem: Memory = Memory(el) + self.push_memory(el_mem) + + def push_memory(self, mem: Memory) -> None: + """Push memory object into the frame.n + + This function saves on serialization time. + """ + assert not ((self._len == 0) ^ (self._free_offset == 0)), f"What?!?!? {self._len} {self._free_offset}" + record_size = len(mem) + (0 if self._len == 0 else SIZEOF_SIZE_T) + if record_size > self._free: + raise MemoryError( + f"Frame is out of memory: Record size {record_size} is greater than remaining buffer size {self._free}" + ) + mem.encode(self._buffer[self._free_offset :]) + if self._len > 0: + self._buffer[-(self._len + 1) * SIZEOF_SIZE_T : -self._len * SIZEOF_SIZE_T] = struct.pack( + "N", self._free_offset + ) + self._data_offset.append(self._free_offset) + self._len += 1 + self._free_offset += len(mem) + + self._buffer[-SIZEOF_SIZE_T:] = struct.pack("N", self._len) + self._free -= record_size + self._used += record_size + + +class SharedMemoryQueueView(Generic[T]): + """A view into a shared memory queue. + + A view can read from shared memory frames. + """ + + def __init__(self, ind: int, name: str, *, has_data: mp.synchronize.Condition, shared_frame_count: c_size_t): + self._ind = ind + self._name = name + self._view_name = name + f"_view_{ind}" + self._has_data = has_data + self._shared_frame_count = shared_frame_count + + self.index_begin = mp.RawValue(c_size_t, 0) + + self._frames: list[Frame] = [] + self._frame_count = 0 + self._frame_index = [0] + self._last_frame_length = 0 + self._data_available = Event() + self._continue_manage = Event() + # self._data_available = mp.Queue() + # self._data_available.cancel_join_thread() + + self._view_state = 0 + + def open(self) -> None: + """Open a view. + + A view should be opened before reading the data. + """ + if self._view_state > 0: + raise RuntimeError("Shared memory queue view has been opened in this process") + logging.debug("[%s] Opening", self._view_name) + self._view_state = 1 + self._managing_thread = Thread(target=self._manage, name=self._view_name + "_manage", daemon=True) + self._managing_thread.start() + + def close(self) -> None: + """Close a view. + + A view must be closed in order to clear all resources. + """ + logging.debug("[%s] Closing %d frames", self._view_name, len(self._frames)) + if self._view_state != 1: + raise RuntimeError("Cannot close the queue view that has not been opened") + self._view_state = 2 + logging.debug("[%s] Waiting for manageing thread to finish", self._view_name) + self._managing_thread.join() + del self._managing_thread + + self._frame_index = [0] + logging.debug("[%s] Closing %d frames", self._view_name, len(self._frames)) + del self._frames[:] + + logging.debug("[%s] Setting begin index to max value", self._view_name) + self.index_begin.value = SIZE_T_MAX + logging.debug("[%s] Closed", self._view_name) + + def _manage(self) -> None: + logging.debug("[%s] Starting management loop", self._view_name) + while self._view_state == 1: + try: + if self._sync_frames(): + self._data_available.set() + continue + with self._has_data: + self._has_data.wait(timeout=0.01) + except BaseException as e: + logging.error("[%s] Management loop encountered exception:\n%s", self._view_name, str(e)) + logging.debug("[%s] Management loop has finished", self._view_name) + + def _sync_frames(self) -> bool: + # open frames + end = self._shared_frame_count.value + if end != 0: + if self._frame_count != end: + logging.debug( + "[%s] Adding %d to %d frames", self._view_name, end - self._frame_count, len(self._frames) + ) + self._frames.extend( + [ + Frame(name=f"{self._name}_frame_{frame_count}", create=False) + for frame_count in range(self._frame_count, end) + ] + ) + self._frame_index = list( + accumulate((len(frame) for frame in self._frames), initial=self._frame_index[0]) + ) + self._last_frame_length = self._frame_index[-1] - self._frame_index[-2] + self._frame_count = end + if len(self) == 0: + logging.error("[%s] The queue is empty after new frame has arrived", self._view_name) + return True + last_frame_length = len(self._frames[-1]) + if self._last_frame_length != last_frame_length: + self._frame_index[-1] = self._frame_index[-2] + last_frame_length + self._last_frame_length = last_frame_length + if len(self) == 0: + logging.error("[%s] The queue is empty after new data has arrived", self._view_name) + return True + + # close frames + n = min(bisect(self._frame_index, self.begin) - 1, len(self._frames) - 1) + if n > 0: + logging.debug("[%s] Removing %d out of %d frames", self._view_name, n, len(self._frames)) + del self._frames[:n] + del self._frame_index[:n] + + return False + + def wait(self, timeout: float | None = None) -> bool: + """Wait for next data.""" + if self._view_state != 1: + raise RuntimeError("Queue view is not open.") + if len(self) > 0: + return True + # try: + # self._data_available.get(timeout=timeout) + # except queue.Empty: + # return Fasle + if not self._data_available.wait(timeout=timeout): + return False + else: + self._data_available.clear() + if len(self) == 0: + logging.error("[%s] The view is empty after data_available event", self._view_name) + return False + + return True + + @property + def begin(self) -> int: + """The first available data index.""" + if self._view_state != 1: + raise RuntimeError("Queue view is not open.") + return self.index_begin.value + + @property + def end(self) -> int: + """Past last available data index.""" + if self._view_state != 1: + raise RuntimeError("Queue view is not open.") + return self._frame_index[-1] + + def __len__(self) -> int: + return self.end - self.begin + + def __getitem__(self, index: int) -> T | Const[T]: + if index < self.begin or index >= self.end: + raise IndexError(f"Index {index} is out of bounds [{self.begin}, {self.end})") + i = bisect(self._frame_index, index) - 1 + index_in_frame = index - self._frame_index[i] + assert index_in_frame >= 0, f"index in frame is negative {index_in_frame}" + frame = self._frames[i] + assert index_in_frame < len(frame), f"index in frame {index_in_frame} is out of bounds {len(frame)}" + return const(frame[index_in_frame]) + + def top(self) -> T | Const[T]: + """The top element (FIFO) order.""" + if len(self) < 1: + raise IndexError("The queue is empty") + index = self.begin + i = 1 + while index >= self._frame_index[i]: + i += 1 + index_in_frame = index - self._frame_index[i - 1] + assert index_in_frame >= 0, f"index in frame is negative {index_in_frame}" + frame = self._frames[i - 1] + assert index_in_frame < len(frame), f"index in frame {index_in_frame} is out of bounds {len(frame)}" + return const(frame[index_in_frame]) + + def pop(self, i: int = 1) -> None: + """Remove the top element.""" + view_length = len(self) + if i > view_length: + raise IndexError(f"Attempting to remove {i} elements from the queue of length {view_length}") + self.index_begin.value += i + logging.debug( + "[%s] Advanced internal counter to %d out of %d", + self._view_name, + self.index_begin.value, + self._frame_index[-1], + ) + self._data_available.clear() + logging.debug("[%s] Cleared data_available", self._view_name) + + +class SharedMemoryQueue(Generic[T]): + """Shared memory queue is a dynamic FIFO queue. + + Data is stored in a shared memory frames, which can be used in a separate process running on the same host os. + """ + + def __init__(self, name: str, *, min_rate: int = MIN_FRANE_SIZE, fps: int = 1): + self._name = name + self._min_rate = min_rate # i.e., 1 MB/s the default value + self._fps = fps # expected new frames per second (default is 1) + self._rate_estimate = self._min_rate # B/s + + self._frame_count = 0 + self._frame_index = [0] + self._frames: list[Frame] = [] + + self._has_data = mp.Condition() + + self._lock = mp.Lock() + self._view_count = 0 + self._views_index_begin: list[c_size_t] = [] + self._shared_frame_count = mp.RawValue(c_size_t, 0) + + self._queue_state = mp.RawValue(c_size_t, 0) # 0 -- Initial, 1 -- Open, 2 -- Closed + + # def __getitem__(self, index: int) -> Const | T: + # if index < self.begin or index >= self.end: + # raise IndexError(f"Index {index} is out of bounds [{self.begin}, {self.end})") + # i = bisect(self._frame_index, index) - 1 + # index_in_frame = index - self._frame_index[i] + # assert index_in_frame >= 0, f"index in frame is negative {index_in_frame}" + # frame = self._frames[i] + # assert index_in_frame < len(frame), f"index in frame {index_in_frame} is out of bounds {len(frame)}" + # return const(frame[index_in_frame]) + # + # def __len__(self) -> int: + # return self.end - self.begin + + def __len__(self) -> int: + return self.end - self.begin + + @property + def begin(self) -> int: + """The first available data index.""" + return self._frame_index[0] + + @property + def end(self) -> int: + """Past last available data index.""" + return self._frame_index[-1] + + def open(self) -> None: + """Open a queue. + + A queue should be opened before pushing any data. + """ + logging.debug("[%s] Opening", self._name) + with self._lock: + if self._queue_state.value > 0: + raise RuntimeError("Shared memory queue can be opened only once") + self._queue_state.value = 1 + self._managing_thread = Thread(target=self._manage, name=self._name + "_queue_manage", daemon=True) + self._managing_thread.start() + + def close(self) -> None: + """Close a queue. + + A queue must be closed in order to clear all resources. + """ + logging.debug("[%s] Closing with %d frames", self._name, len(self._frames)) + with self._lock: + if self._queue_state.value != 1: + raise RuntimeError("Cannot close the queue that has not been opened") + self._queue_state.value = 2 + self._managing_thread.join() + del self._managing_thread + + self._frame_index = [0] + del self._frames[:] + + logging.debug("[%s] Closed", self._name) + + def _manage(self) -> None: + logging.debug("[%s] Starting management loop", self._name) + while self._queue_state.value == 1: + try: + self.flush() + time.sleep(0.01) + # TODO (dmitry): preallocate frames for reducing push latency + # Currently push lattency is measured ~ 50 ms in the worst case, and ~ 150 us on average + except BaseException as e: + logging.error("[%s] Management loop encountered exception:\n%s", self._name, str(e)) + logging.debug("[%s] Finished management loop", self._name) + + def flush(self) -> None: + """Manually flush unused frames.""" + if len(self._frames) == 0: + return + min_begin = min([view.value for view in self._views_index_begin] + [self.end]) + assert min_begin >= self.begin + # we do not delete the last frame for preventing frequent allocations + n = min(bisect(self._frame_index, min_begin) - 1, len(self._frames) - 1) + if n == 0: + return + logging.debug("[%s] Flushing %d out of %d frames", self._name, n, len(self._frames)) + del self._frames[:n] + del self._frame_index[:n] + + def make_view(self) -> SharedMemoryQueueView: + """Make a queue view.""" + if self._queue_state.value != 0: + raise RuntimeError("Cannot add views after queue has been opened") + view: SharedMemoryQueueView = SharedMemoryQueueView( + ind=self._view_count, + name=self._name, + has_data=self._has_data, + shared_frame_count=self._shared_frame_count, + ) + self._view_count += 1 + self._views_index_begin.append(view.index_begin) + return view + + def _add_frame(self, size: int = 0) -> None: + if self._frames: + self._rate_estimate = int( + sum((frame.used() for frame in self._frames)) / (time.perf_counter() - self._frames[0].time()) + ) + logging.debug("[%s] Current rate estimate %d B/s", self._name, self._rate_estimate) + self._rate_estimate = max(self._fps * self._rate_estimate, self._fps * self._min_rate, size) + frame: Frame = Frame(name=f"{self._name}_frame_{self._frame_count}", create=True, size=self._rate_estimate) + self._frames.append(frame) + self._frame_index.append(self._frame_index[-1]) + self._frame_count += 1 + self._shared_frame_count.value = self._frame_count + + def push(self, el: T) -> None: + """Push an element into a queue.""" + if self._queue_state.value != 1: + raise RuntimeError("Queue must be open for pushing elements") + el_mem: Memory = Memory(el) + el_size = len(el_mem) + if not self._frames or not self._frames[-1].fits_size(el_size): + self._add_frame(el_size) + self._frames[-1].push(el) + self._frame_index[-1] += 1 + with self._has_data: + self._has_data.notify_all() diff --git a/judo/unlocked/test/test_const.py b/judo/unlocked/test/test_const.py new file mode 100644 index 00000000..a7a69ff3 --- /dev/null +++ b/judo/unlocked/test/test_const.py @@ -0,0 +1,68 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from dataclasses import dataclass +from typing import Any + +import pytest + +from core.unlocked import Const as const + + +def test_simple_const() -> None: + @dataclass + class Simple: + value: int + + const_simple = const(Simple(value=5)) + + with pytest.raises(RuntimeError): + const_simple.value = -5 + assert const_simple.value == 5 + + simple = Simple(value=6) + simple.value = 6 + assert simple.value == 6 + + +def test_composed_const() -> None: + @dataclass + class Inner: + value: int + + @dataclass + class Outer: + first: Inner + second: Inner + + const_composed = const(Outer(first=Inner(value=4), second=Inner(value=2))) + + assert const_composed.first.value == 4 + with pytest.raises(RuntimeError): + const_composed.second.value = -2 + assert const_composed.second.value == 2 + + composed = Outer(first=Inner(value=6), second=Inner(value=6)) + + assert composed.first.value == 6 + composed.second.value = -6 + assert composed.second.value == -6 + + +def test_mutable_const() -> None: + class Mutable: + def __init__(self, value: Any): + self.value = value + + def add_one(self) -> None: + self.value += 1 + + const_mutable = const(Mutable(value=3)) + + with pytest.raises(RuntimeError): + const_mutable.add_one() + assert const_mutable.value == 3 + + mutable = Mutable(value=7) + + mutable.add_one() + assert mutable.value == 8 diff --git a/judo/unlocked/test/test_ipc_utils.py b/judo/unlocked/test/test_ipc_utils.py new file mode 100644 index 00000000..a52c5482 --- /dev/null +++ b/judo/unlocked/test/test_ipc_utils.py @@ -0,0 +1,117 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import multiprocessing as mp +from ctypes import c_bool, c_int +from random import uniform +from threading import Thread +from time import perf_counter, sleep + +from core.unlocked import Barrier, BroadcastEvent, Mutex + + +def test_mutex_thread() -> None: + mutex = Mutex("test_mutex") + locked = True + + def test_fn() -> None: + nonlocal locked + with mutex: + locked = False + + thread = Thread(target=test_fn) + with mutex: + thread.start() + sleep(0.1) + assert locked + thread.join() + sleep(0.1) + assert not locked + + +def test_mutex_mp() -> None: + mutex = Mutex("test_mutex") + n_proc = 10 + delay = 0.1 / n_proc + + def run() -> None: + with mutex: + sleep(delay) + + procs = [mp.Process(target=run, name=f"proc_{i}") for i in range(n_proc)] + + tic = perf_counter() + for p in procs: + p.start() + for p in procs: + p.join(timeout=1.0) + toc = perf_counter() - tic + + assert toc > 0.1 + assert toc < 0.2 + + +def test_barrier() -> None: + n_iter = 1000 + n_proc = 10 + delay = 0.1 / n_iter + barrier = Barrier("test_barrier", n_proc) + + array = mp.RawArray(c_int, n_proc) + success = mp.RawArray(c_bool, [False] * n_proc) + + def run(proc_ind: int) -> None: + scs = True + for _ in range(n_iter): + with barrier: + sleep(uniform(0, delay)) + array[proc_ind] += 1 + sleep(uniform(0, delay)) + if not all([a == array[proc_ind] for a in array]): + scs = False + success[proc_ind] = scs + + procs = [mp.Process(target=run, args=(i,), name=f"proc_{i}", daemon=True) for i in range(n_proc)] + + for p in procs: + p.start() + for p in procs: + p.join(timeout=1.0) + + assert all(success) + + +def test_broadcast_event() -> None: + event = BroadcastEvent("test_event") + + n_iter = 1000 + n_cons = 20 + delay = 0.1 / n_iter + + prod_success = mp.RawValue(c_bool, False) + cons_success = mp.RawArray(c_bool, [False] * n_cons) + + def producer() -> None: + while not all(cons_success): + sleep(delay) + event.notify() + prod_success.value = True + + def consumer(ind: int) -> None: + for _ in range(n_iter): + with event: + pass + cons_success[ind] = True + + prod_proc = mp.Process(target=producer, name="prod") + cons_proc = [mp.Process(target=consumer, args=(i,), name=f"cons_{i}") for i in range(n_cons)] + + for p in cons_proc: + p.start() + prod_proc.start() + + prod_proc.join() + for p in cons_proc: + p.join(timeout=1) + + assert prod_success.value + assert all(cons_success) diff --git a/judo/unlocked/test/test_output_stage.py b/judo/unlocked/test/test_output_stage.py new file mode 100644 index 00000000..31c3ffff --- /dev/null +++ b/judo/unlocked/test/test_output_stage.py @@ -0,0 +1,34 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +from pytest import raises + +from core.unlocked import Node + + +def test_annotation() -> None: + def no_return() -> None: + pass + + def return_int() -> int: + return 0 + + class My: + """Empty test class.""" + + def return_My_str() -> tuple[My, str]: + return My(), "test" + + stage = Node("test_None", no_return, frequency=1).output_stage + assert len(stage) == 0 + with raises(IndexError): + stage[0] + + stage = Node("test_int", return_int, frequency=1).output_stage + assert len(stage) == 1 + with raises(IndexError): + stage[1] + + stage = Node("test_tuple", return_My_str, frequency=1).output_stage + assert len(stage) == 2 + with raises(IndexError): + stage[2] diff --git a/judo/unlocked/test/test_queue.py b/judo/unlocked/test/test_queue.py new file mode 100644 index 00000000..fbee6b73 --- /dev/null +++ b/judo/unlocked/test/test_queue.py @@ -0,0 +1,192 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import random +from dataclasses import dataclass +from enum import Enum, auto +from threading import Thread +from time import sleep + +import pytest +from pytest import raises + +from core.unlocked import Queue, QueueView + + +@dataclass +class MyInt: + value: int + + +type_params = [int, MyInt] + + +@pytest.mark.parametrize("T", type_params) +def test_queue(T: type) -> None: + q: Queue = Queue() + q.push(T(8)) + q.push(T(7)) + assert q[0] == T(8) + assert q[1] == T(7) + assert q.top() == T(8) + assert q.pop() == T(8) + assert q[1] == T(7) + with raises(IndexError): + q[0] + assert q.pop() == T(7) + with raises(IndexError): + q[0] + with raises(IndexError): + q[1] + with raises(IndexError): + q.pop() + q.push(T(3)) + with raises(IndexError): + q[0] + with raises(IndexError): + q[1] + assert q[2] == T(3) + + +@pytest.mark.parametrize("T", type_params) +def test_queue_view(T: type) -> None: + q: Queue = Queue() + qw = q.make_view() + + q.push(T(8)) + q.push(T(7)) + with raises(IndexError): + qw.last() + with raises(IndexError): + qw.top() + with raises(IndexError): + qw.pop() + qw.flush() + q.flush() + assert len(q) == 2 + assert qw.top() == T(8) + assert qw.last() == T(7) + assert qw.pop() == T(8) + assert qw.top() == T(7) + assert qw.last() == T(7) + q.push(T(3)) + assert len(q) == 3 + assert qw.top() == T(7) + assert qw.last() == T(7) + q.flush() + assert len(q) == 2 + qw.flush(flush_queue=True) + assert len(q) == 1 + assert qw.pop() == T(3) + q.flush() + assert len(q) == 0 + + +@pytest.mark.parametrize("T", type_params) +def test_queue_stress(T: type) -> None: + q: Queue = Queue() + + random.seed(83285020238474093) + + class Act(Enum): + Push = auto() + Pop = auto() + Item = auto() + + test_num = 100 + act_num = 1000 + + for _ in range(test_num): + r_seq = [random.randint(0, 1_000_000_000) for _ in range(act_num)] + push_pos = 0 + pop_pos = 0 + q_len = 0 + q.clear() + for act in random.choices(list(Act.__members__.values()), k=act_num): + match act: + case Act.Push: + q.push(T(r_seq[push_pos])) + push_pos += 1 + q_len += 1 + case Act.Pop: + if q_len > 0: + assert q.pop() == T(r_seq[pop_pos]) + pop_pos += 1 + q_len -= 1 + else: + with raises(IndexError): + q.pop() + case Act.Item: + i = random.randint(0, act_num) + if q.begin <= i and i < q.end: + assert q[i] == T(r_seq[i]) + else: + with raises(IndexError): + q[i] + + assert len(q) == q_len + + +@pytest.mark.parametrize("T", type_params) +def test_queue_threaded_stress(T: type) -> None: + q: Queue = Queue() + sub_num = 10 + views = [q.make_view() for _ in range(sub_num)] + + random.seed(83285020238474093) + + test_num = 10 + act_num = 100 + test_time = 0.5 + delay = 2 * test_time / test_num / act_num + + publisher_success = False + subscriber_success = [False] * sub_num + + for _ in range(test_num): + r_seq = [random.randint(0, 1_000_000_000) for _ in range(act_num)] + + publisher_success = False + for i in range(sub_num): + subscriber_success[i] = False + + def publisher(r_seq: list) -> None: + nonlocal publisher_success + for r in r_seq: + q.push(T(r)) + sleep(random.uniform(0, delay)) + publisher_success = True + + def subscriber(r_seq: list, sub_ind: int, qw: QueueView) -> None: + nonlocal subscriber_success + for r in r_seq: + if len(qw) == 0: + while not qw.ready: + # if we use the following then the queue view is always of length 1 + # qw.wait(random.uniform(0, delay)) + # for stress testing we will let it accumulate a bit + sleep(random.uniform(0, delay)) + qw.flush() + assert qw.top() == T(r) + with raises(IndexError): + qw[qw.begin - 1] + with raises(IndexError): + qw[qw.end] + assert qw.begin < qw.end + index = random.randint(qw.begin, qw.end - 1) + assert qw[index] == T(r_seq[index]) + for el, ex in zip(qw, (r_seq[i] for i in range(qw.begin, qw.end)), strict=True): + assert el == T(ex) + assert qw.pop() == T(r) + subscriber_success[sub_ind] = True + + threads = [Thread(target=publisher, args=(r_seq,))] + for sub_ind in range(sub_num): + threads.append(Thread(target=subscriber, args=(r_seq, sub_ind, views[sub_ind]))) + + q.clear() + + for t in threads: + t.start() + for t in threads: + t.join(timeout=10) + assert publisher_success and all(subscriber_success) diff --git a/judo/unlocked/test/test_shmque.py b/judo/unlocked/test/test_shmque.py new file mode 100644 index 00000000..0a6d5461 --- /dev/null +++ b/judo/unlocked/test/test_shmque.py @@ -0,0 +1,250 @@ +# Copyright (c) 2025 Boston Dynamics AI Institute LLC. All rights reserved. + +import multiprocessing as mp +import random +from ctypes import c_bool +from dataclasses import dataclass +from operator import eq as scalar_eq +from time import sleep +from typing import Any, Callable + +import numpy as np +import pytest +from pytest import raises + +from core.unlocked import SIZEOF_SIZE_T, Frame, Memory, SharedMemoryQueue + + +@dataclass +class MyPair: + first: float + second: float + + +random.seed(785672498529) + + +def int_gen(size: int) -> list: + return [random.randint(0, 1_000_000_000) for _ in range(size)] + + +def pair_gen(size: int) -> list: + return [MyPair(first=random.uniform(0, 1), second=random.uniform(0, 1)) for _ in range(size)] + + +def array_gen(size: int) -> list: + return [np.random.randint(1_000_000_000, size=(10, 10)) for _ in range(size)] + + +type_gen_params = [(int, int_gen, scalar_eq), (MyPair, pair_gen, scalar_eq), (np.array, array_gen, np.array_equal)] + + +@pytest.mark.parametrize("T, gen, eq", type_gen_params) +def test_memory(T: type, gen: Callable, eq: Callable) -> None: + t = gen(1)[0] + t_mem: Memory = Memory(t) + t_bytes = bytearray(len(t_mem)) + t_mem.encode(t_bytes) + tt = Memory.decode(t_bytes) + assert t is not tt + assert eq(t, tt) + + +def test_shared_array() -> None: + t = np.arange(10) + t_mem: Memory = Memory(t) + t_bytes = bytearray(len(t_mem)) + t_mem.encode(t_bytes) + tt1 = Memory.decode(t_bytes) + tt2 = Memory.decode(t_bytes) + assert tt1 is not tt2 + assert np.array_equal(tt1, tt2) + tt1[...] = 42 + assert np.array_equal(tt1, tt2) + + +@pytest.mark.parametrize("T, gen, eq", type_gen_params) +def test_frame(T: type, gen: Callable, eq: Callable) -> None: + values = gen(5) + frame: Frame = Frame(name="test_frame_shm", create=True, size=1024 * 10) + frame_view: Frame = Frame(name="test_frame_shm", create=False) + + assert len(frame) == 0 + assert len(frame_view) == 0 + assert frame.size() >= 1024 * 10 + assert frame.free() == frame.size() - SIZEOF_SIZE_T + assert frame.used() == SIZEOF_SIZE_T + + for val in values: + frame.push(val) + + assert len(frame) == 5 + assert len(frame_view) == 5 + assert frame.free() + frame.used() == frame.size() + assert frame.used() > SIZEOF_SIZE_T + + for index in range(5): + assert eq(frame[index], values[index]), f"error at index {index}" + assert eq(frame_view[index], values[index]), f"view error at index {index}" + + +@pytest.mark.parametrize("T, gen, eq", type_gen_params) +def test_frame_multiprocessing(T: type, gen: Callable, eq: Callable) -> None: + def producer(name: str, values: list, cond: Any, events: list) -> None: + fail = False + frame: Frame = Frame(name=name, create=True, size=1024 * 10) + try: + for val in values: + frame.push(val) + assert len(frame) == 5 + with cond: + cond.notify_all() + while not all([event.is_set() for event in events]): + sleep(random.uniform(0.05, 0.1)) + except Exception as e: + print(e) + fail = True + if fail: + exit(1) + + def consumer(name: str, values: list, cond: Any, ind: int, events: list) -> None: + fail = False + frame: Frame = Frame(name=name, create=False) + try: + with cond: + cond.wait_for(lambda: len(frame) == 5) + assert len(frame) == 5 + for index in range(5): + assert eq(frame[index], values[index]) + events[ind].set() + while not all([event.is_set() for event in events]): + sleep(random.uniform(0.05, 0.1)) + except Exception as e: + print(e) + events[ind].set() + fail = True + if fail: + exit(1) + + name = "test_frame_multiprocessing" + values = gen(5) + sub_num = 10 + manager = mp.Manager() + cond = manager.Condition() + events = [manager.Event() for _ in range(sub_num)] + + procs = [mp.Process(target=producer, args=(name, values, cond, events), daemon=True)] + [ + mp.Process(target=consumer, args=(name, values, cond, ind, events), daemon=True) for ind in range(sub_num) + ] + + for proc in procs: + proc.start() + for proc in procs: + proc.join() + + assert all((proc.exitcode == 0 for proc in procs)) + + +@pytest.mark.parametrize("T, gen, eq", type_gen_params) +def test_shmque(T: type, gen: Callable, eq: Callable) -> None: + n = 10 + values = gen(n) + queue: SharedMemoryQueue = SharedMemoryQueue(name="test_shmque", min_rate=1) + view1 = queue.make_view() + view2 = queue.make_view() + queue.open() + view1.open() + view2.open() + + test_time = 0.1 + delay = 2 * test_time / n + + for val in values: + sleep(random.uniform(0, delay)) + queue.push(val) + + sleep(0.01) + assert len(queue) == n + assert len(view2) == n + assert len(view2) == n + + for index in range(n): + random_index = random.randint(0, n - 1) + if view1.begin <= random_index and random_index < view1.end: + assert eq(view1[random_index], values[random_index]) + else: + with raises(IndexError): + view1[random_index] + + assert eq(view1[view1.begin], values[index]) + assert eq(view1.top(), values[index]) + view1.pop() + + assert len(view1) == 0 + assert len(view2) == n + assert len(queue) == n + + view2.pop(n) + + assert len(view2) == 0 + + # wait for flush + sleep(0.01) + assert len(queue) < n + + view1.close() + view2.close() + queue.close() + + +@pytest.mark.parametrize("T, gen, eq", type_gen_params) +def test_shmque_stress(T: type, gen: Callable, eq: Callable) -> None: + + random.seed(47299561901461) + + sub_num = 10 + act_num = 100 + test_time = 0.1 + delay = 2 * test_time / act_num + + queue: SharedMemoryQueue = SharedMemoryQueue(name="test_shmque_stress", min_rate=1) + views = [queue.make_view() for _ in range(sub_num)] + + values = gen(act_num) + + success = mp.RawArray(c_bool, [False] * (sub_num + 1)) + + def publisher() -> None: + queue.open() + for val in values: + sleep(random.uniform(0, delay)) + queue.push(val) + queue.close() + success[0] = True + + def subscriber(ind: int) -> None: + view = views[ind] + view.open() + for val in values: + while not view.wait(): + pass + assert eq(view.top(), val) + with raises(IndexError): + view[view.begin - 1] + with raises(IndexError): + view[view.end] + assert view.begin < view.end + index = random.randint(view.begin, view.end - 1) + assert eq(view[index], values[index]) + view.pop() + view.close() + success[ind + 1] = True + + procs = [mp.Process(target=publisher)] + [ + mp.Process(target=subscriber, args=(sub_ind,)) for sub_ind in range(sub_num) + ] + for p in procs: + p.start() + for p in procs: + p.join(timeout=60) + assert all(success) diff --git a/pyproject.toml b/pyproject.toml index afac3772..1e04261f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,7 +29,6 @@ classifiers = [ ] dependencies = [ - "dora-utils", # dora + qol utils for writing nodes "mujoco", "numpy", "pillow", # for displaying app logo @@ -37,6 +36,7 @@ dependencies = [ "rich", # for displaying tables in benchmarking "scipy", "viser>=1.0.0", + "colorlog", ] [project.urls] @@ -178,4 +178,4 @@ judo-rai = { path = ".", editable = true } [tool.pixi.environments] default = { solve-group = "default" } docs = { features = ["docs"], solve-group = "default" } -dev = { features = ["dev", "docs"], solve-group = "default" } \ No newline at end of file +dev = { features = ["dev", "docs"], solve-group = "default" }