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
96 changes: 96 additions & 0 deletions src/unilab/base/entity.py
Original file line number Diff line number Diff line change
Expand Up @@ -474,6 +474,7 @@ def __init__(
actuator_ids = self._resolve_enumerated_ids(
"actuator", self._actuator_names, backend.get_actuator_names
)
self._actuator_ids = actuator_ids

self._validate_joint_state(backend, joint_pos_ids, joint_vel_ids)
self._validate_body_state(backend, root_body_ids, body_ids)
Expand Down Expand Up @@ -1025,6 +1026,68 @@ def write_root_state_to_sim(
term_name=f"{self.name}.write_root_state_to_sim",
)

def bind_actuator_gain_write(
self,
actuator_ids: np.ndarray | Sequence[int] | slice | None = None,
*,
term_name: str,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Bind selected actuator columns and immutable gain defaults on the cold path."""
if self._reset_state is None:
raise self._capability_error(
"reset actuator-gain write",
"EntityScene was materialized without an env-owned reset transaction",
)
if self._actuator_ids is None:
raise self._capability_error(
"reset actuator-gain write",
"actuator_names were not declared in EntityCfg",
)
local_ids = self._normalize_local_actuator_ids(
actuator_ids,
capability="reset actuator-gain write",
)
if local_ids.size == 0:
raise ValueError(
f"Entity '{self.name}' reset actuator-gain write selected no actuators"
)
backend_ids = self._actuator_ids[local_ids]
_, default_kp, default_kd = self._reset_state.bind_actuator_gain_write(
backend_ids,
term_name=f"{term_name}:{self.name}",
)
bound_local_ids = np.array(local_ids, copy=True)
bound_local_ids.setflags(write=False)
return bound_local_ids, default_kp, default_kd

def write_actuator_gains_to_sim(
self,
kp: np.ndarray,
kd: np.ndarray,
actuator_ids: np.ndarray | Sequence[int] | slice | None = None,
env_ids: np.ndarray | slice | None = None,
*,
term_name: str = "pd_gains",
) -> None:
"""Stage entity-local actuator gains in the active reset transaction."""
if self._reset_state is None or self._actuator_ids is None:
raise self._capability_error(
"reset actuator-gain write",
"actuator metadata or the env-owned reset transaction was not materialized",
)
local_ids = self._normalize_local_actuator_ids(
actuator_ids,
capability="reset actuator-gain write",
)
resolved_env_ids = self._normalize_reset_env_ids(env_ids)
self._reset_state.write_actuator_gains(
resolved_env_ids,
self._actuator_ids[local_ids],
kp,
kd,
term_name=f"{term_name}:{self.name}",
)

def write_root_link_pose_to_sim(
self,
root_pose: np.ndarray,
Expand Down Expand Up @@ -1160,6 +1223,39 @@ def _normalize_local_joint_ids(
)
return ids

def _normalize_local_actuator_ids(
self,
actuator_ids: np.ndarray | Sequence[int] | slice | None,
*,
capability: str,
) -> np.ndarray:
if actuator_ids is None:
ids = np.arange(self.num_actuators, dtype=np.intp)
elif isinstance(actuator_ids, slice):
ids = np.arange(self.num_actuators, dtype=np.intp)[actuator_ids]
else:
raw = np.asarray(actuator_ids)
if (
raw.ndim != 1
or not np.issubdtype(raw.dtype, np.integer)
or np.issubdtype(raw.dtype, np.bool_)
):
raise TypeError(
f"Entity '{self.name}' {capability} actuator_ids must be a 1-D "
"integer array or slice"
)
ids = np.asarray(raw, dtype=np.intp)
if np.any(ids < 0) or np.any(ids >= self.num_actuators):
raise IndexError(
f"Entity '{self.name}' {capability} actuator_ids out of range for "
f"{self.num_actuators} actuators: {ids.tolist()}"
)
if np.unique(ids).size != ids.size:
raise ValueError(
f"Entity '{self.name}' {capability} actuator_ids contain duplicates: {ids.tolist()}"
)
return ids

def _normalize_reset_env_ids(self, env_ids: np.ndarray | slice | None) -> np.ndarray:
if env_ids is None:
return np.arange(self._backend.num_envs, dtype=np.int32)
Expand Down
153 changes: 153 additions & 0 deletions src/unilab/base/reset_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
import numpy as np

from unilab.base.backend.base import BackendRootStateLayout, SimBackend
from unilab.dr.types import RESET_TERM_KD, RESET_TERM_KP, ResetRandomizationPayload
from unilab.utils.rotation import np_quat_apply_inverse


Expand All @@ -35,6 +36,11 @@ def __init__(
self._default_qvel: np.ndarray | None = None
self._qpos: np.ndarray | None = None
self._qvel: np.ndarray | None = None
self._default_kp: np.ndarray | None = None
self._default_kd: np.ndarray | None = None
self._kp: np.ndarray | None = None
self._kd: np.ndarray | None = None
self._gain_dirty_mask = np.zeros(self._num_envs, dtype=np.bool_)
self._requesting_terms: set[str] = set()

@property
Expand Down Expand Up @@ -62,9 +68,83 @@ def begin(self, env_ids: np.ndarray) -> None:
self._active_mask.fill(False)
self._active_mask[ids] = True
self._dirty_mask.fill(False)
self._gain_dirty_mask.fill(False)
self._requesting_terms.clear()
self._active = True

def bind_actuator_gain_write(
self,
actuator_ids: np.ndarray,
*,
term_name: str,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Resolve gain mutation capability and immutable defaults on the cold path."""
columns = self._validate_columns(
actuator_ids,
width=self._backend.num_actuators,
capability="actuator IDs",
term_name=term_name,
)
self._materialize_default_actuator_gains(term_name)
assert self._default_kp is not None
assert self._default_kd is not None
selected_kp = np.array(self._default_kp[columns], copy=True)
selected_kd = np.array(self._default_kd[columns], copy=True)
selected_kp.setflags(write=False)
selected_kd.setflags(write=False)
bound_columns = np.array(columns, copy=True)
bound_columns.setflags(write=False)
return bound_columns, selected_kp, selected_kd

def write_actuator_gains(
self,
env_ids: np.ndarray,
actuator_ids: np.ndarray,
kp: np.ndarray,
kd: np.ndarray,
*,
term_name: str,
) -> None:
"""Stage selected per-environment actuator gains in the reset transaction."""
ids = self._prepare_state_write(
env_ids,
capability="actuator-gain",
term_name=term_name,
)
columns = self._validate_columns(
actuator_ids,
width=self._backend.num_actuators,
capability="actuator IDs",
term_name=term_name,
)
self._materialize_default_actuator_gains(term_name)
gains_shape = (ids.size, columns.size)
kp_values = self._validate_values(
kp,
shape=gains_shape,
capability="actuator kp",
term_name=term_name,
)
kd_values = self._validate_values(
kd,
shape=gains_shape,
capability="actuator kd",
term_name=term_name,
)
assert self._default_kp is not None
assert self._default_kd is not None
assert self._kp is not None
assert self._kd is not None
uninitialized = ids[~self._gain_dirty_mask[ids]]
if uninitialized.size:
self._kp[uninitialized] = self._default_kp
self._kd[uninitialized] = self._default_kd
if ids.size and columns.size:
self._kp[ids[:, None], columns[None, :]] = kp_values
self._kd[ids[:, None], columns[None, :]] = kd_values
self._gain_dirty_mask[ids] = True
self._dirty_mask[ids] = True

def reset_to_default(self, env_ids: np.ndarray, *, term_name: str) -> None:
"""Stage backend default qpos/qvel for a subset of the active reset."""
self._require_active()
Expand Down Expand Up @@ -243,11 +323,13 @@ def commit(self) -> dict | None:
return None
assert self._qpos is not None
assert self._qvel is not None
randomization = self._build_randomization_payload(dirty_ids)
try:
return self._backend.set_state(
dirty_ids,
self._qpos[dirty_ids],
self._qvel[dirty_ids],
randomization=randomization,
)
except (AttributeError, NotImplementedError) as exc:
terms = ", ".join(sorted(self._requesting_terms))
Expand Down Expand Up @@ -284,6 +366,61 @@ def _materialize_default_state(self, term_name: str) -> None:
self._qpos = np.empty((self._num_envs, default_qpos.size), dtype=default_qpos.dtype)
self._qvel = np.empty((self._num_envs, default_qvel.size), dtype=default_qvel.dtype)

def _materialize_default_actuator_gains(self, term_name: str) -> None:
if self._default_kp is not None:
return
try:
capabilities = self._backend.get_dr_capabilities()
except (AttributeError, NotImplementedError) as exc:
raise self._capability_error(term_name, "actuator gain randomization", exc) from exc
required = frozenset((RESET_TERM_KP, RESET_TERM_KD))
unsupported = capabilities.get_unsupported_reset_terms(required)
if unsupported:
detail = ", ".join(sorted(unsupported))
raise self._capability_error(
term_name,
"actuator gain randomization",
NotImplementedError(f"unsupported reset payload fields: {detail}"),
)
try:
kp, kd = self._backend.get_actuator_gains()
except (AttributeError, NotImplementedError) as exc:
raise self._capability_error(term_name, "default actuator gains", exc) from exc
default_kp = self._validate_gain_vector(kp, "default actuator kp", term_name)
default_kd = self._validate_gain_vector(kd, "default actuator kd", term_name)
self._default_kp = default_kp
self._default_kd = default_kd
self._kp = np.empty(
(self._num_envs, self._backend.num_actuators),
dtype=default_kp.dtype,
)
self._kd = np.empty(
(self._num_envs, self._backend.num_actuators),
dtype=default_kd.dtype,
)

def _build_randomization_payload(
self,
dirty_ids: np.ndarray,
) -> ResetRandomizationPayload | None:
gain_ids = np.flatnonzero(self._gain_dirty_mask).astype(np.int32, copy=False)
if gain_ids.size == 0:
return None
missing = dirty_ids[~self._gain_dirty_mask[dirty_ids]]
if missing.size:
terms = ", ".join(sorted(self._requesting_terms))
raise RuntimeError(
"EventManager reset actuator-gain payload cannot represent sparse rows in "
f"one SimBackend.set_state call for term(s) [{terms}] on backend "
f"'{self._backend.backend_type}'; missing env IDs {missing.tolist()}"
)
assert self._kp is not None
assert self._kd is not None
return ResetRandomizationPayload(
kp=np.array(self._kp[dirty_ids], copy=True),
kd=np.array(self._kd[dirty_ids], copy=True),
)

def _prepare_state_write(
self,
env_ids: np.ndarray,
Expand Down Expand Up @@ -378,6 +515,21 @@ def _validate_state_vector(
result.setflags(write=False)
return result

def _validate_gain_vector(
self,
value: np.ndarray,
capability: str,
term_name: str,
) -> np.ndarray:
result = self._validate_state_vector(value, capability, term_name)
expected = (self._backend.num_actuators,)
if result.shape != expected:
raise ValueError(
f"EventManager term '{term_name}' capability '{capability}' on backend "
f"'{self._backend.backend_type}' returned shape {result.shape}; expected {expected}"
)
return result

def _validate_ids(self, env_ids: np.ndarray, *, capability: str) -> np.ndarray:
if not isinstance(env_ids, np.ndarray):
raise TypeError(
Expand Down Expand Up @@ -484,6 +636,7 @@ def _finish(self) -> None:
self._active = False
self._active_mask.fill(False)
self._dirty_mask.fill(False)
self._gain_dirty_mask.fill(False)
self._requesting_terms.clear()


Expand Down
2 changes: 2 additions & 0 deletions src/unilab/envs/mdp/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from unilab.envs.mdp.actions import JointPositionActionCfg as JointPositionActionCfg
from unilab.envs.mdp.commands import UniformVelocityCommand as UniformVelocityCommand
from unilab.envs.mdp.commands import UniformVelocityCommandCfg as UniformVelocityCommandCfg
from unilab.envs.mdp.events import pd_gains as pd_gains
from unilab.envs.mdp.events import reset_root_state_uniform as reset_root_state_uniform
from unilab.envs.mdp.events import reset_scene_to_default as reset_scene_to_default
from unilab.envs.mdp.events import resolve_env_ids as resolve_env_ids
Expand Down Expand Up @@ -53,6 +54,7 @@
"joint_vel_rel",
"joint_vel_l2",
"last_action",
"pd_gains",
"is_alive",
"is_terminated",
"projected_gravity",
Expand Down
Loading