diff --git a/src/unilab/managers/_buffers/circular_buffer.py b/src/unilab/managers/_buffers/circular_buffer.py index b0575dd50..40738a7c8 100644 --- a/src/unilab/managers/_buffers/circular_buffer.py +++ b/src/unilab/managers/_buffers/circular_buffer.py @@ -1,6 +1,6 @@ # Derived from mujocolab/mjlab v1.6.0 (0fb8a681), src/mjlab/utils/buffers/circular_buffer.py. # Copyright 2025, The mjlab Developers. -# Modified by UniLab for NumPy and UniLab contracts; licensed under Apache-2.0. +# Modified by UniLab for Torch temporal buffers and UniLab contracts; Apache-2.0. """Circular buffer for storing a history of batched tensor data. Understanding Dimensions @@ -80,6 +80,7 @@ from collections.abc import Sequence import numpy as np +import torch class CircularBuffer: @@ -130,10 +131,10 @@ def __init__(self, max_len: int, batch_size: int) -> None: self._max_len = max_len self._batch_size = batch_size self._pointer: int = -1 - self._buffer: np.ndarray | None = None - self._all_indices = np.arange(batch_size) - self._num_pushes = np.zeros(batch_size, dtype=np.int64) - self._max_len_array = np.full(batch_size, max_len, dtype=np.int64) + self._buffer: torch.Tensor | None = None + self._all_indices: torch.Tensor | None = None + self._num_pushes: torch.Tensor | None = None + self._max_len_array: torch.Tensor | None = None @property def batch_size(self) -> int: @@ -144,9 +145,12 @@ def max_length(self) -> int: return self._max_len @property - def current_length(self) -> np.ndarray: + def current_length(self) -> torch.Tensor: """Per-batch count of valid frames. Shape: (batch_size,).""" - return np.minimum(self._num_pushes, self._max_len_array) + self._initialize_counters() + assert self._num_pushes is not None + assert self._max_len_array is not None + return torch.minimum(self._num_pushes, self._max_len_array) @property def is_initialized(self) -> bool: @@ -154,7 +158,7 @@ def is_initialized(self) -> bool: return self._buffer is not None @property - def buffer(self) -> np.ndarray: + def buffer(self) -> torch.Tensor: """History in chronological order (oldest to newest). Returns: @@ -165,9 +169,8 @@ def buffer(self) -> np.ndarray: raise RuntimeError("Buffer not initialized. Call append() first.") start = (self._pointer + 1) % self._max_len - idx = (np.arange(self._max_len) + start) % self._max_len - buf = self._buffer[idx] # (max_len, batch, ...) - return np.swapaxes(buf, 0, 1) # (batch, max_len, ...) + idx = (torch.arange(self._max_len, device=self._buffer.device) + start) % self._max_len + return self._buffer[idx].transpose(0, 1) def reset(self, batch_ids: Sequence[int] | np.ndarray | slice | None = None) -> None: """Zero out values and counters for specified batch rows. @@ -175,12 +178,16 @@ def reset(self, batch_ids: Sequence[int] | np.ndarray | slice | None = None) -> Args: batch_ids: Batch indices to reset, or None to reset all. """ - ids: Sequence[int] | np.ndarray | slice = slice(None) if batch_ids is None else batch_ids + self._initialize_counters() + assert self._num_pushes is not None + ids: Sequence[int] | np.ndarray | torch.Tensor | slice = ( + slice(None) if batch_ids is None else batch_ids + ) self._num_pushes[ids] = 0 if self._buffer is not None: self._buffer[:, ids] = 0.0 - def backfill(self, data: np.ndarray, batch_ids: np.ndarray) -> None: + def backfill(self, data: np.ndarray | torch.Tensor, batch_ids: np.ndarray) -> None: """Fill the given rows' entire history with one frame, without advancing time. Unlike append, the global pointer does not move and other rows are @@ -197,11 +204,14 @@ def backfill(self, data: np.ndarray, batch_ids: np.ndarray) -> None: if self._buffer is None: raise RuntimeError("Buffer not initialized. Call append() first.") - data = np.asarray(data) - self._buffer[:, batch_ids] = np.expand_dims(data[batch_ids], axis=0) + if not isinstance(data, torch.Tensor): + data = torch.from_numpy(np.ascontiguousarray(data)) + self._initialize_counters(data.device) + assert self._num_pushes is not None + self._buffer[:, batch_ids] = data[batch_ids].unsqueeze(0) self._num_pushes[batch_ids] = 1 - def append(self, data: np.ndarray) -> None: + def append(self, data: np.ndarray | torch.Tensor) -> None: """Append a new frame for all batch elements. Args: @@ -210,11 +220,17 @@ def append(self, data: np.ndarray) -> None: if data.shape[0] != self._batch_size: raise ValueError(f"Expected batch size {self._batch_size}, got {data.shape[0]}") - data = np.asarray(data) + if not isinstance(data, torch.Tensor): + data = torch.from_numpy(np.ascontiguousarray(data)) if self._buffer is None: self._pointer = -1 - self._buffer = np.empty((self._max_len, *data.shape), dtype=data.dtype) + self._buffer = torch.empty( + (self._max_len, *data.shape), dtype=data.dtype, device=data.device + ) + self._initialize_counters(data.device) + assert self._num_pushes is not None + assert self._all_indices is not None self._pointer = (self._pointer + 1) % self._max_len self._buffer[self._pointer] = data @@ -222,13 +238,13 @@ def append(self, data: np.ndarray) -> None: # Backfill only newly initialized rows. After warm-up this branch avoids # scanning the full history buffer on every hot-path append. is_first_push = self._num_pushes == 0 - if np.any(is_first_push): - first_ids = np.flatnonzero(is_first_push) - self._buffer[:, first_ids] = np.expand_dims(data[first_ids], axis=0) + if bool(is_first_push.any()): + first_ids = torch.nonzero(is_first_push, as_tuple=False).flatten() + self._buffer[:, first_ids] = data[first_ids].unsqueeze(0) self._num_pushes += 1 - def __getitem__(self, key: np.ndarray | int) -> np.ndarray: + def __getitem__(self, key: torch.Tensor | np.ndarray | int) -> torch.Tensor: """Retrieve lagged frames per batch (LIFO). Args: @@ -236,23 +252,40 @@ def __getitem__(self, key: np.ndarray | int) -> np.ndarray: """ if self._buffer is None: raise RuntimeError("Buffer not initialized. Call append() first.") + self._initialize_counters(self._buffer.device) + assert self._num_pushes is not None + assert self._max_len_array is not None + assert self._all_indices is not None if isinstance(key, int): - key = np.full(self._batch_size, key, dtype=np.int64) + key = torch.full((self._batch_size,), key, dtype=torch.int64) else: - key = np.asarray(key, dtype=np.int64) + key = torch.as_tensor(np.asarray(key), dtype=torch.int64) if key.ndim == 0: - key = np.full(self._batch_size, key.item(), dtype=np.int64) + key = torch.full((self._batch_size,), int(key.item()), dtype=torch.int64) - if key.size != self._batch_size: - raise ValueError(f"Expected {self._batch_size} lags, got {key.size}") + if key.numel() != self._batch_size: + raise ValueError(f"Expected {self._batch_size} lags, got {key.numel()}") # Clamp to the oldest retained frame: without the max_len bound, a lag # beyond the buffer length would wrap around to a newer frame once # num_pushes exceeds max_len. - pushes = np.maximum(self._num_pushes, 1) - max_lag = np.minimum(pushes, self._max_len_array) - 1 - valid = np.maximum(np.minimum(key, max_lag), 0) + if key.device != self._num_pushes.device: + key = key.to(device=self._num_pushes.device) + pushes = torch.maximum(self._num_pushes, torch.ones_like(self._num_pushes)) + max_lag = torch.minimum(pushes, self._max_len_array) - 1 + valid = torch.clamp(torch.minimum(key, max_lag), min=0) - idx = np.remainder(self._pointer - valid, self._max_len) + idx = torch.remainder(self._pointer - valid, self._max_len) return self._buffer[idx, self._all_indices] + + def _initialize_counters(self, device: torch.device | None = None) -> None: + """Bind per-batch metadata to the first appended observation device.""" + if self._num_pushes is not None and self._max_len_array is not None: + return + resolved = torch.device("cpu") if device is None else device + self._all_indices = torch.arange(self._batch_size, device=resolved) + self._num_pushes = torch.zeros(self._batch_size, dtype=torch.int64, device=resolved) + self._max_len_array = torch.full( + (self._batch_size,), self._max_len, dtype=torch.int64, device=resolved + ) diff --git a/src/unilab/managers/_buffers/delay_buffer.py b/src/unilab/managers/_buffers/delay_buffer.py index da29bb29c..4b4236742 100644 --- a/src/unilab/managers/_buffers/delay_buffer.py +++ b/src/unilab/managers/_buffers/delay_buffer.py @@ -1,6 +1,6 @@ # Derived from mujocolab/mjlab v1.6.0 (0fb8a681), src/mjlab/utils/buffers/delay_buffer.py. # Copyright 2025, The mjlab Developers. -# Modified by UniLab for NumPy and UniLab contracts; licensed under Apache-2.0. +# Modified by UniLab for Torch temporal buffers and UniLab contracts; Apache-2.0. """Delay buffer for stochastically delayed observations.""" from __future__ import annotations @@ -8,6 +8,7 @@ from collections.abc import Sequence import numpy as np +import torch from unilab.managers._buffers import CircularBuffer @@ -204,7 +205,7 @@ def reset(self, batch_ids: Sequence[int] | np.ndarray | slice | None = None) -> ) self._phase_offsets[idx] = new_phases[idx] - def append(self, data: np.ndarray) -> None: + def append(self, data: np.ndarray | torch.Tensor) -> None: """Append new observation to buffer. Args: @@ -212,7 +213,7 @@ def append(self, data: np.ndarray) -> None: """ self._buffer.append(data) - def backfill(self, data: np.ndarray, batch_ids: np.ndarray) -> None: + def backfill(self, data: np.ndarray | torch.Tensor, batch_ids: np.ndarray) -> None: """Backfill the given rows with one frame, without advancing time. Used after a partial reset: the reset rows (whose lags and step counters @@ -225,7 +226,7 @@ def backfill(self, data: np.ndarray, batch_ids: np.ndarray) -> None: """ self._buffer.backfill(data, batch_ids) - def compute(self) -> np.ndarray: + def compute(self) -> torch.Tensor: """Compute delayed observation for current step. Advances the lag update schedule, then returns the delayed observation. @@ -239,7 +240,7 @@ def compute(self) -> np.ndarray: self._update_lags() return self.peek() - def peek(self) -> np.ndarray: + def peek(self) -> torch.Tensor: """Return the delayed observation using current lags, without advancing. Unlike compute, this neither steps the update schedule nor resamples lags, @@ -253,8 +254,14 @@ def peek(self) -> np.ndarray: # Clamp lags to valid range [0, buffer_length - 1]. # Buffer may not be full yet (e.g., only 2 frames but sampled lag=3). - valid_lags = np.minimum(self._current_lags, self._buffer.current_length - 1) - valid_lags = np.maximum(valid_lags, 0) + current_length = self._buffer.current_length + valid_lags = torch.clamp( + torch.minimum( + torch.from_numpy(self._current_lags).to(current_length.device), + current_length - 1, + ), + min=0, + ) return self._buffer[valid_lags] diff --git a/src/unilab/managers/event_manager.py b/src/unilab/managers/event_manager.py index fb98d7ea9..b477051f1 100644 --- a/src/unilab/managers/event_manager.py +++ b/src/unilab/managers/event_manager.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Literal import numpy as np +import torch from prettytable import PrettyTable from unilab.managers.manager_base import ManagerBase, ManagerTermBaseCfg @@ -76,6 +77,7 @@ class EventManager(ManagerBase): def __init__(self, cfg: dict[str, EventTermCfg | None], env: ManagerBasedRlEnv): self.cfg = deepcopy(cfg) + self._device = torch.device(getattr(env, "device", torch.device("cpu"))) self._mode_term_names: dict[EventMode, list[str]] = dict() self._mode_term_cfgs: dict[EventMode, list[EventTermCfg]] = dict() self._mode_class_term_cfgs: dict[EventMode, list[EventTermCfg]] = dict() @@ -141,7 +143,9 @@ def reset(self, env_ids: np.ndarray | None = None): assert term_cfg.interval_range_s is not None lower, upper = term_cfg.interval_range_s sampled_interval = self._env.rng.uniform(lower, upper, num_envs) - self._interval_term_time_left[index][ids] = sampled_interval + self._interval_term_time_left[index][ids] = torch.as_tensor( + sampled_interval, dtype=torch.float64, device=self._device + ) return {} def apply( @@ -179,7 +183,9 @@ def apply( assert term_cfg.interval_range_s is not None lower, upper = term_cfg.interval_range_s sampled_interval = self._env.rng.uniform(lower, upper, 1) - self._interval_term_time_left[index][:] = sampled_interval + self._interval_term_time_left[index][:] = torch.as_tensor( + sampled_interval, dtype=torch.float64, device=self._device + ) term_cfg.func(self._env, None, **term_cfg.params) else: valid_env_ids = np.flatnonzero(time_left < 1e-6) @@ -187,7 +193,9 @@ def apply( assert term_cfg.interval_range_s is not None lower, upper = term_cfg.interval_range_s sampled_time = self._env.rng.uniform(lower, upper, len(valid_env_ids)) - self._interval_term_time_left[index][valid_env_ids] = sampled_time + self._interval_term_time_left[index][valid_env_ids] = torch.as_tensor( + sampled_time, dtype=torch.float64, device=self._device + ) term_cfg.func(self._env, valid_env_ids, **term_cfg.params) elif mode == "step": term_cfg.func(self._env, None, **term_cfg.params) @@ -224,9 +232,9 @@ def apply( term_cfg.func(self._env, env_ids, **term_cfg.params) def _prepare_terms(self) -> None: - self._interval_term_time_left: list[np.ndarray] = list() - self._reset_term_last_triggered_step_id: list[np.ndarray] = list() - self._reset_term_last_triggered_once: list[np.ndarray] = list() + self._interval_term_time_left: list[torch.Tensor] = list() + self._reset_term_last_triggered_step_id: list[torch.Tensor] = list() + self._reset_term_last_triggered_once: list[torch.Tensor] = list() for term_name, term_cfg in self.cfg.items(): if term_cfg is None: @@ -257,15 +265,23 @@ def _prepare_terms(self) -> None: f"{term_cfg.interval_range_s}." ) if term_cfg.is_global_time: - time_left = self._env.rng.uniform(lower, upper, 1) + time_left = torch.as_tensor( + self._env.rng.uniform(lower, upper, 1), + dtype=torch.float64, + device=self._device, + ) self._interval_term_time_left.append(time_left) else: - time_left = self._env.rng.uniform(lower, upper, self.num_envs) + time_left = torch.as_tensor( + self._env.rng.uniform(lower, upper, self.num_envs), + dtype=torch.float64, + device=self._device, + ) self._interval_term_time_left.append(time_left) elif term_cfg.mode == "reset": - step_count = np.zeros(self.num_envs, dtype=np.int64) + step_count = torch.zeros(self.num_envs, dtype=torch.int64, device=self._device) self._reset_term_last_triggered_step_id.append(step_count) - no_trigger = np.zeros(self.num_envs, dtype=np.bool_) + no_trigger = torch.zeros(self.num_envs, dtype=torch.bool, device=self._device) self._reset_term_last_triggered_once.append(no_trigger) func = term_cfg.func diff --git a/src/unilab/managers/observation_manager.py b/src/unilab/managers/observation_manager.py index 96bf95a36..8ee03c98d 100644 --- a/src/unilab/managers/observation_manager.py +++ b/src/unilab/managers/observation_manager.py @@ -559,29 +559,24 @@ def compute_group( ) if term_cfg.delay_max_lag > 0: - if tensor_obs: - obs = cast("torch.Tensor", obs).detach().cpu().numpy() - tensor_obs = False delay_buffer = self._group_obs_term_delay_buffer[group_name][term_name] if env_ids is None or not delay_buffer.is_initialized: - delay_buffer.append(cast("np.ndarray", obs)) + delay_buffer.append(obs) obs = delay_buffer.compute() else: - delay_buffer.backfill(cast("np.ndarray", obs), env_ids) + delay_buffer.backfill(obs, env_ids) obs = delay_buffer.peek() if term_cfg.history_length > 0: - if tensor_obs: - obs = cast("torch.Tensor", obs).detach().cpu().numpy() - tensor_obs = False circular_buffer = self._group_obs_term_history_buffer[group_name][term_name] if env_ids is None or not circular_buffer.is_initialized: if update_history or not circular_buffer.is_initialized: - circular_buffer.append(cast("np.ndarray", obs)) + circular_buffer.append(obs) else: - circular_buffer.backfill(cast("np.ndarray", obs), env_ids) + circular_buffer.backfill(obs, env_ids) if term_cfg.flatten_history_dim: - group_obs[term_name] = circular_buffer.buffer.reshape(self._env.num_envs, -1) + history = circular_buffer.buffer + group_obs[term_name] = history.reshape(self._env.num_envs, -1) else: group_obs[term_name] = circular_buffer.buffer else: diff --git a/tests/envs/mdp/test_events.py b/tests/envs/mdp/test_events.py index aa8dcc537..c5a44e29e 100644 --- a/tests/envs/mdp/test_events.py +++ b/tests/envs/mdp/test_events.py @@ -8,6 +8,7 @@ import numpy as np import pytest +import torch from unisim.backend.base import BackendRootStateLayout, SimBackend from unisim.dr.types import ( RESET_TERM_BODY_INERTIA, @@ -854,7 +855,9 @@ def test_min_step_count_gating_reuses_committed_field_values() -> None: second = backend.randomization_calls[-1] assert second is not None and second.body_mass is not None np.testing.assert_allclose(second.body_mass[:2], first.body_mass) - np.testing.assert_array_equal(manager._reset_term_last_triggered_step_id[0], [0, 0, 50]) + torch.testing.assert_close( + manager._reset_term_last_triggered_step_id[0], torch.tensor([0, 0, 50], dtype=torch.int64) + ) # Fully gated resets skip the field entirely; the backend keeps the # previously applied per-env values, so no dense payload is needed. @@ -870,7 +873,10 @@ def test_min_step_count_gating_reuses_committed_field_values() -> None: manager.apply(mode="reset", env_ids=first_ids, global_env_step_count=200) fourth = backend.randomization_calls[-1] assert fourth is not None and fourth.body_mass is not None - np.testing.assert_array_equal(manager._reset_term_last_triggered_step_id[0], [200, 200, 50]) + torch.testing.assert_close( + manager._reset_term_last_triggered_step_id[0], + torch.tensor([200, 200, 50], dtype=torch.int64), + ) assert not np.allclose(fourth.body_mass, second.body_mass[:2]) @@ -901,7 +907,9 @@ def test_min_step_count_gating_applies_to_pd_gains_payload_rows() -> None: assert second is not None and second.kp is not None and second.kd is not None np.testing.assert_allclose(second.kp[:2], first.kp) np.testing.assert_allclose(second.kd[:2], first.kd) - np.testing.assert_array_equal(manager._reset_term_last_triggered_step_id[0], [0, 0, 10]) + torch.testing.assert_close( + manager._reset_term_last_triggered_step_id[0], torch.tensor([0, 0, 10], dtype=torch.int64) + ) def test_apply_body_impulse_lifecycle_stages_sustains_and_expires() -> None: @@ -1119,7 +1127,7 @@ def test_velocity_push_uses_env_rng_and_interval_subset_plan() -> None: }, env, ) - manager._interval_term_time_left[0][:] = [0.0, 1.0, 0.0] + manager._interval_term_time_left[0][:] = torch.tensor([0.0, 1.0, 0.0]) manager.apply(mode="interval", dt=0.1) @@ -1151,7 +1159,7 @@ def test_velocity_push_dispatches_angular_delta_when_supported() -> None: }, env, ) - manager._interval_term_time_left[0][:] = [0.0, 1.0, 0.0] + manager._interval_term_time_left[0][:] = torch.tensor([0.0, 1.0, 0.0]) manager.apply(mode="interval", dt=0.1) diff --git a/tests/managers/test_event_command_metrics_recorder.py b/tests/managers/test_event_command_metrics_recorder.py index 7e6d6a1d7..4cf84a16b 100644 --- a/tests/managers/test_event_command_metrics_recorder.py +++ b/tests/managers/test_event_command_metrics_recorder.py @@ -91,7 +91,14 @@ def test_event_interval_rng_is_reproducible() -> None: } left = EventManager(cfg, FakeEnv(seed=17)) right = EventManager(cfg, FakeEnv(seed=17)) - np.testing.assert_array_equal(left._interval_term_time_left, right._interval_term_time_left) + assert len(left._interval_term_time_left) == len(right._interval_term_time_left) + for left_time, right_time in zip( + left._interval_term_time_left, right._interval_term_time_left, strict=True + ): + assert isinstance(left_time, torch.Tensor) + assert left_time.dtype == torch.float64 + assert left_time.device.type == torch.device("cpu").type + torch.testing.assert_close(left_time, right_time) class DummyCommand(CommandTerm): diff --git a/tests/managers/test_observation_buffers_noise.py b/tests/managers/test_observation_buffers_noise.py index b14a11840..57ea21732 100644 --- a/tests/managers/test_observation_buffers_noise.py +++ b/tests/managers/test_observation_buffers_noise.py @@ -49,7 +49,7 @@ def test_delay_buffer_constant_delay_and_partial_backfill() -> None: outputs = [] for value in (1.0, 2.0, 3.0, 4.0): buffer.append(np.full((2, 1), value, dtype=np.float32)) - outputs.append(buffer.compute().copy()) + outputs.append(buffer.compute().clone()) np.testing.assert_array_equal(np.stack(outputs)[:, 0, 0], [1, 1, 1, 2]) buffer.reset(np.array([1])) @@ -210,6 +210,47 @@ def test_non_temporal_tensor_terms_stay_on_device_for_clip_scale_concat( assert source.isfinite().all() +def test_tensor_delay_and_history_stay_on_device_without_host_detour( + fake_env: FakeEnv, +) -> None: + device = fake_env.device + source = torch.arange(fake_env.num_envs * 2, dtype=torch.float32, device=device).reshape( + fake_env.num_envs, 2 + ) + manager = ObservationManager( + { + "policy": ObservationGroupCfg( + terms={ + "delayed": ObservationTermCfg( + func=lambda env: source, delay_min_lag=1, delay_max_lag=1 + ), + "history": ObservationTermCfg(func=lambda env: source, history_length=2), + } + ) + }, + fake_env, + ) + + first = manager.compute(update_history=True)["policy"] + assert isinstance(first, torch.Tensor) + assert first.device == device + + history_buffer = manager._group_obs_term_history_buffer["policy"]["history"] + delay_buffer = manager._group_obs_term_delay_buffer["policy"]["delayed"] + assert history_buffer.buffer.device == device + assert delay_buffer.peek().device == device + + source = source + 10 + second = manager.compute(update_history=True)["policy"] + assert second.device == device + torch.testing.assert_close(second[:, :2], first[:, :2]) + torch.testing.assert_close( + second[:, 2:4], + first[:, 2:4], + ) + torch.testing.assert_close(second[:, 4:6], source) + + def test_tensor_terms_are_row_scoped_on_reset(fake_env: FakeEnv) -> None: device = fake_env.device source = torch.arange(fake_env.num_envs * 2, dtype=torch.float32, device=device).reshape( diff --git a/tests/managers/test_observation_partial_reset.py b/tests/managers/test_observation_partial_reset.py index 584ea060c..e15fcbe7d 100644 --- a/tests/managers/test_observation_partial_reset.py +++ b/tests/managers/test_observation_partial_reset.py @@ -109,8 +109,8 @@ def test_partial_reset_temporal_group_falls_back_and_preserves_rows() -> None: ids = np.array([1], dtype=np.int32) keep_ids = np.array([0, 2, 3], dtype=np.int32) - history_before = manager._group_obs_term_history_buffer["policy"]["state"].buffer.copy() - delay_before = manager._group_obs_term_delay_buffer["policy"]["delayed"].peek().copy() + history_before = manager._group_obs_term_history_buffer["policy"]["state"].buffer.clone() + delay_before = manager._group_obs_term_delay_buffer["policy"]["delayed"].peek().clone() rng_state = env.rng.bit_generator.state rows = manager.compute(update_history=True, env_ids=ids)