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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 64 additions & 31 deletions src/unilab/managers/_buffers/circular_buffer.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -80,6 +80,7 @@
from collections.abc import Sequence

import numpy as np
import torch


class CircularBuffer:
Expand Down Expand Up @@ -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:
Expand All @@ -144,17 +145,20 @@ 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:
"""Check if the buffer has been initialized with at least one append."""
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:
Expand All @@ -165,22 +169,25 @@ 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.

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
Expand All @@ -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:
Expand All @@ -210,49 +220,72 @@ 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

# 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:
key: Per-batch lags (Tensor) or shared lag (int). Shape (batch_size,) or scalar.
"""
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
)
21 changes: 14 additions & 7 deletions src/unilab/managers/_buffers/delay_buffer.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,14 @@
# 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

from collections.abc import Sequence

import numpy as np
import torch

from unilab.managers._buffers import CircularBuffer

Expand Down Expand Up @@ -204,15 +205,15 @@ 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:
data: Observation tensor of shape (batch_size, ...).
"""
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
Expand All @@ -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.
Expand All @@ -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,
Expand All @@ -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]

Expand Down
36 changes: 26 additions & 10 deletions src/unilab/managers/event_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -179,15 +183,19 @@ 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)
if len(valid_env_ids) > 0:
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)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
17 changes: 6 additions & 11 deletions src/unilab/managers/observation_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Loading