From 359e0b2723c8d415e88ca5707d88d9b80190155b Mon Sep 17 00:00:00 2001 From: KraHsu Date: Wed, 13 May 2026 21:42:11 +0800 Subject: [PATCH 1/5] Add backend-agnostic Sensor abstraction and obs noise pipeline --- src/genelab/configs.py | 3 + src/genelab/envs/manager_based_rl_env.py | 17 +++ src/genelab/lab.py | 7 + src/genelab/managers/observation_manager.py | 7 +- src/genelab/mdp/__init__.py | 6 + src/genelab/mdp/noise.py | 36 +++++ src/genelab/mdp/observations.py | 5 + src/genelab/sensor/__init__.py | 4 + src/genelab/sensor/sensor.py | 75 ++++++++++ tests/test_sensor.py | 144 ++++++++++++++++++++ 10 files changed, 303 insertions(+), 1 deletion(-) create mode 100644 src/genelab/mdp/noise.py create mode 100644 src/genelab/sensor/sensor.py create mode 100644 tests/test_sensor.py diff --git a/src/genelab/configs.py b/src/genelab/configs.py index 07861aac..4647ccb8 100644 --- a/src/genelab/configs.py +++ b/src/genelab/configs.py @@ -5,6 +5,8 @@ from types import UnionType from typing import Any, cast, get_args, get_origin, get_type_hints +from genelab.sensor import SensorCfg + type _Annotation = object @@ -19,6 +21,7 @@ class SceneCfg: substeps: int = 4 num_envs: int = 1 env_spacing: tuple[float, float] = (2.0, 2.0) + sensors: tuple[SensorCfg, ...] = field(default_factory=tuple) @dataclass diff --git a/src/genelab/envs/manager_based_rl_env.py b/src/genelab/envs/manager_based_rl_env.py index 9c780c4c..4993c2ce 100644 --- a/src/genelab/envs/manager_based_rl_env.py +++ b/src/genelab/envs/manager_based_rl_env.py @@ -28,6 +28,7 @@ TerminationManager, TerminationTermCfg, ) +from genelab.sensor import Sensor if TYPE_CHECKING: pass @@ -140,6 +141,14 @@ def __init__(self, cfg: ManagerBasedRlEnvCfg) -> None: self.event_manager = EventManager(cfg.events_cfg, self) self.curriculum_manager = CurriculumManager(cfg.curriculum_cfg, self) + # Build sensors after managers (so cfg parsing is settled) but before reset + # (so observation_manager.compute can read sensor.data on the first frame). + self._sensors: dict[str, Sensor[Any]] = {} + for sensor_cfg in cfg.scene.sensors: + sensor = sensor_cfg.build() + sensor.bind(self) + self._sensors[sensor_cfg.name] = sensor + # Apply PD gains, default pose, then run startup events. self._apply_default_gains() self.event_manager.apply("startup") @@ -189,6 +198,10 @@ def robot(self) -> Any: def robot_state(self) -> RobotState: return self._robot_state + @property + def sensors(self) -> dict[str, Sensor[Any]]: + return self._sensors + @property def joint_names(self) -> list[str]: return list(self._joint_names) @@ -482,6 +495,8 @@ def _reset_idx(self, env_ids: torch.Tensor) -> None: self.event_manager.apply("reset", env_ids) self.command_manager.reset(env_ids) self.action_manager.reset(env_ids) + for sensor in self._sensors.values(): + sensor.reset(env_ids) reward_extras = self.reward_manager.reset(env_ids) term_extras = self.termination_manager.reset(env_ids) curr_extras = self.curriculum_manager.compute(env_ids) @@ -497,6 +512,8 @@ def step( self.action_manager.apply_action() self._scene.step() self._refresh_robot_state() + for sensor in self._sensors.values(): + sensor.update(self._step_dt) self._episode_length_buf += 1 # Resample commands and trigger interval-mode events. self.command_manager.compute(self._step_dt) diff --git a/src/genelab/lab.py b/src/genelab/lab.py index 3d29c857..2c01440e 100644 --- a/src/genelab/lab.py +++ b/src/genelab/lab.py @@ -4,6 +4,7 @@ from typing import Protocol from genelab.configs import ManagerBasedEnvCfg, TaskCfg, apply_overrides +from genelab.mdp.noise import Gnoise, NoiseCfg, Unoise from genelab.registry import ( ENVS, ROBOTS, @@ -14,6 +15,7 @@ load_entrypoint_extensions, load_extension_module, ) +from genelab.sensor import Sensor, SensorCfg class ManagerBasedEnv(Protocol): @@ -39,11 +41,16 @@ class GenesisBackendCfg: "ROBOTS", "TASKS", "GenesisBackendCfg", + "Gnoise", "ManagerBasedEnv", "ManagerBasedEnvCfg", + "NoiseCfg", "Registry", "RegistryEntry", + "Sensor", + "SensorCfg", "TaskCfg", + "Unoise", "apply_overrides", "load_builtin_registries", "load_entrypoint_extensions", diff --git a/src/genelab/managers/observation_manager.py b/src/genelab/managers/observation_manager.py index 1609c57b..fe914677 100644 --- a/src/genelab/managers/observation_manager.py +++ b/src/genelab/managers/observation_manager.py @@ -11,20 +11,23 @@ if TYPE_CHECKING: from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + from genelab.mdp.noise import NoiseCfg @dataclass class ObservationTermCfg(ManagerTermBaseCfg): - """Observation term: optional scale + clip applied to the raw tensor.""" + """Observation term: optional noise / scale / clip applied to the raw tensor.""" scale: float | None = None clip: tuple[float, float] | None = None + noise: "NoiseCfg | None" = None @dataclass class ObservationGroupCfg: terms: dict[str, ObservationTermCfg] = field(default_factory=dict) concatenate_terms: bool = True + enable_corruption: bool = False class ObservationManager: @@ -67,6 +70,8 @@ def compute(self) -> dict[str, torch.Tensor]: value = term_cfg.func(self._env, **term_cfg.params) if value.dim() == 1: value = value.unsqueeze(-1) + if term_cfg.noise is not None and group_cfg.enable_corruption: + value = term_cfg.noise.apply(value) if term_cfg.scale is not None: value = value * term_cfg.scale if term_cfg.clip is not None: diff --git a/src/genelab/mdp/__init__.py b/src/genelab/mdp/__init__.py index c3eddaa6..3d84c289 100644 --- a/src/genelab/mdp/__init__.py +++ b/src/genelab/mdp/__init__.py @@ -13,6 +13,7 @@ reset_joints_to_default, reset_root_state_uniform, ) +from genelab.mdp.noise import Gnoise, NoiseCfg, Unoise from genelab.mdp.observations import ( base_ang_vel, base_lin_vel, @@ -25,6 +26,7 @@ projected_gravity, robot_body_ori_b, robot_body_pos_b, + sensor_data, ) from genelab.mdp.rewards import ( action_rate_l2, @@ -51,13 +53,16 @@ ) __all__ = [ + "Gnoise", "JointPositionAction", "JointPositionActionCfg", "MotionCommand", "MotionCommandCfg", "MotionLoader", + "NoiseCfg", "UniformVelocityCommand", "UniformVelocityCommandCfg", + "Unoise", "action_rate_l2", "bad_anchor_ori", "bad_anchor_pos_z_only", @@ -88,6 +93,7 @@ "robot_body_ori_b", "robot_body_pos_b", "root_height_below", + "sensor_data", "time_out", "track_angular_velocity_z_exp", "track_linear_velocity_xy_exp", diff --git a/src/genelab/mdp/noise.py b/src/genelab/mdp/noise.py new file mode 100644 index 00000000..bc6a4eae --- /dev/null +++ b/src/genelab/mdp/noise.py @@ -0,0 +1,36 @@ +"""Additive noise models for observation corruption (matches mjlab's ``Unoise`` / ``Gnoise``).""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass + +import torch + + +@dataclass +class NoiseCfg(ABC): + """Base config for additive noise injected into an observation term.""" + + @abstractmethod + def apply(self, data: torch.Tensor) -> torch.Tensor: ... + + +@dataclass +class Unoise(NoiseCfg): + """Uniform additive noise in ``[n_min, n_max]``.""" + + n_min: float = -1.0 + n_max: float = 1.0 + + def apply(self, data: torch.Tensor) -> torch.Tensor: + return data + torch.empty_like(data).uniform_(self.n_min, self.n_max) + + +@dataclass +class Gnoise(NoiseCfg): + """Gaussian additive noise: ``data + N(mean, std^2)``.""" + + mean: float = 0.0 + std: float = 1.0 + + def apply(self, data: torch.Tensor) -> torch.Tensor: + return data + torch.randn_like(data) * self.std + self.mean diff --git a/src/genelab/mdp/observations.py b/src/genelab/mdp/observations.py index f6e565a7..850bd6ee 100644 --- a/src/genelab/mdp/observations.py +++ b/src/genelab/mdp/observations.py @@ -44,6 +44,11 @@ def generated_commands(env: "ManagerBasedRlEnv", command_name: str) -> torch.Ten return env.command_manager.get_command(command_name) +def sensor_data(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: + """Return the per-step cached tensor of the named sensor.""" + return env.sensors[sensor_name].data + + # --------------------------------------------------------------------- motion imitation diff --git a/src/genelab/sensor/__init__.py b/src/genelab/sensor/__init__.py index e1db484b..fdb9236b 100644 --- a/src/genelab/sensor/__init__.py +++ b/src/genelab/sensor/__init__.py @@ -1 +1,5 @@ """Sensor abstractions for cameras, contacts, tactile data, and future streams.""" + +from genelab.sensor.sensor import Sensor, SensorCfg + +__all__ = ["Sensor", "SensorCfg"] diff --git a/src/genelab/sensor/sensor.py b/src/genelab/sensor/sensor.py new file mode 100644 index 00000000..a2bfbf08 --- /dev/null +++ b/src/genelab/sensor/sensor.py @@ -0,0 +1,75 @@ +"""Backend-agnostic sensor abstraction. + +The shape mirrors mjlab's ``SensorCfg`` / ``Sensor[T]`` so configs and observation terms can move +between the two backends without ceremony, but the lifecycle drops mjlab's ``edit_spec`` / +``initialize`` pair — Genesis has no MJCF spec to rewrite, so ``bind(env)`` is the single hook. +""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, cast + +import torch + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class SensorCfg(ABC): + """Backend-agnostic sensor configuration. Subclasses describe what to sense and how.""" + + name: str = "" + + @abstractmethod + def build(self) -> "Sensor[Any]": ... + + +class Sensor[T](ABC): + """Per-step cached sensor. Subclasses implement ``_compute_data``. + + Lifecycle: ``__init__(cfg)`` → ``bind(env)`` once → per-step ``update(dt)`` (invalidates the + cache) → ``data`` triggers a lazy ``_compute_data`` on first access → ``reset(env_ids)`` + invalidates and lets subclasses clear stateful buffers. + """ + + def __init__(self, cfg: SensorCfg) -> None: + self._cfg = cfg + self._env: "ManagerBasedRlEnv | None" = None + self._cached_data: T | None = None + self._cache_valid: bool = False + + @property + def cfg(self) -> SensorCfg: + return self._cfg + + @property + def name(self) -> str: + return self._cfg.name + + def bind(self, env: "ManagerBasedRlEnv") -> None: + """Called once during env construction. Subclasses may cache link indices etc.""" + self._env = env + + @property + def data(self) -> T: + if not self._cache_valid: + self._cached_data = self._compute_data() + self._cache_valid = True + return cast(T, self._cached_data) + + def _invalidate_cache(self) -> None: + self._cache_valid = False + + def update(self, dt: float) -> None: + """Called each control step after ``_refresh_robot_state``. Override + ``super().update(dt)``.""" + del dt + self._invalidate_cache() + + def reset(self, env_ids: torch.Tensor | None = None) -> None: + """Called per-env on reset. Override + ``super().reset(env_ids)`` to clear state buffers.""" + del env_ids + self._invalidate_cache() + + @abstractmethod + def _compute_data(self) -> T: ... diff --git a/tests/test_sensor.py b/tests/test_sensor.py new file mode 100644 index 00000000..ad55c649 --- /dev/null +++ b/tests/test_sensor.py @@ -0,0 +1,144 @@ +"""Sensor abstraction + obs-noise pipeline. Uses a fake env so Genesis is not required.""" + +from dataclasses import dataclass +from typing import Any + +import torch + +from genelab.managers import ( + ObservationGroupCfg, + ObservationManager, + ObservationTermCfg, +) +from genelab.mdp.noise import Gnoise, Unoise +from genelab.sensor import Sensor, SensorCfg + + +@dataclass +class _ConstSensorCfg(SensorCfg): + value: float = 1.0 + + def build(self) -> "_ConstSensor": + return _ConstSensor(self) + + +class _ConstSensor(Sensor[torch.Tensor]): + def __init__(self, cfg: _ConstSensorCfg) -> None: + super().__init__(cfg) + self._cfg_typed = cfg + self.compute_calls: int = 0 + + def _compute_data(self) -> torch.Tensor: + assert self._env is not None + self.compute_calls += 1 + return torch.full((self._env.num_envs, 3), self._cfg_typed.value, device=self._env.device) + + +class _FakeEnv: + """Just enough surface for sensors and the observation manager.""" + + def __init__(self, num_envs: int = 4, device: str = "cpu") -> None: + self.num_envs = num_envs + self.device = device + self.sensors: dict[str, Sensor[Any]] = {} + + +def test_sensor_caches_until_invalidated() -> None: + env = _FakeEnv(num_envs=2) + sensor = _ConstSensorCfg(name="c", value=1.5).build() + sensor.bind(env) + first = sensor.data + assert first.shape == (2, 3) + assert torch.allclose(first, torch.full((2, 3), 1.5)) + assert sensor.compute_calls == 1 + _ = sensor.data + assert sensor.compute_calls == 1 + sensor.update(0.02) + _ = sensor.data + assert sensor.compute_calls == 2 + sensor.reset(torch.tensor([0, 1])) + _ = sensor.data + assert sensor.compute_calls == 3 + + +def test_unoise_stays_within_bounds_and_is_additive() -> None: + torch.manual_seed(0) + data = torch.zeros(1024, 3) + noisy = Unoise(n_min=-0.2, n_max=0.2).apply(data) + delta = noisy - data + assert (delta >= -0.2 - 1e-6).all() + assert (delta <= 0.2 + 1e-6).all() + assert delta.abs().mean() > 1e-3 + + +def test_gnoise_has_expected_std() -> None: + torch.manual_seed(0) + data = torch.zeros(4096, 1) + noisy = Gnoise(mean=0.0, std=0.5).apply(data) + assert abs(noisy.std().item() - 0.5) < 0.05 + + +def test_observation_pipeline_noise_only_when_corruption_enabled() -> None: + env = _FakeEnv(num_envs=8) + + def const_two(_env: Any) -> torch.Tensor: + return torch.full((_env.num_envs, 3), 2.0, device=_env.device) + + torch.manual_seed(0) + cfg = { + "policy": ObservationGroupCfg( + enable_corruption=True, + terms={"x": ObservationTermCfg(func=const_two, noise=Unoise(-0.1, 0.1))}, + ), + "critic": ObservationGroupCfg( + enable_corruption=False, + terms={"x": ObservationTermCfg(func=const_two, noise=Unoise(-0.1, 0.1))}, + ), + } + mgr = ObservationManager(cfg, env) + obs = mgr.compute() + assert torch.equal(obs["critic"], torch.full((8, 3), 2.0)) + assert not torch.equal(obs["policy"], obs["critic"]) + assert ((obs["policy"] - 2.0).abs() <= 0.1 + 1e-6).all() + + +def test_observation_pipeline_applies_noise_scale_clip_in_order() -> None: + env = _FakeEnv(num_envs=1) + + def const_ten(_env: Any) -> torch.Tensor: + return torch.full((_env.num_envs, 1), 10.0, device=_env.device) + + # zero-magnitude noise so we can pin the math: 10 -> +0 noise -> *0.5 scale -> clip to [0, 3] + cfg = { + "policy": ObservationGroupCfg( + enable_corruption=True, + terms={ + "x": ObservationTermCfg( + func=const_ten, + noise=Unoise(0.0, 0.0), + scale=0.5, + clip=(0.0, 3.0), + ) + }, + ) + } + mgr = ObservationManager(cfg, env) + obs = mgr.compute() + assert torch.allclose(obs["policy"], torch.tensor([[3.0]])) + + +def test_observation_manager_skips_corruption_when_noise_unset() -> None: + env = _FakeEnv(num_envs=2) + + def const_one(_env: Any) -> torch.Tensor: + return torch.full((_env.num_envs, 2), 1.0, device=_env.device) + + cfg = { + "g": ObservationGroupCfg( + enable_corruption=True, + terms={"x": ObservationTermCfg(func=const_one)}, + ) + } + mgr = ObservationManager(cfg, env) + obs = mgr.compute() + assert torch.equal(obs["g"], torch.full((2, 2), 1.0)) From 14472eb7926a0fd0de5c65dfb26746409cd4ff62 Mon Sep 17 00:00:00 2001 From: KraHsu Date: Wed, 13 May 2026 21:45:00 +0800 Subject: [PATCH 2/5] Wire G1 IMU obs through BodyVelocitySensor with Unoise corruption --- .../unitree/src/genelab_unitree/g1/env_cfg.py | 63 +++++++++---- src/genelab/sensor/__init__.py | 8 +- src/genelab/sensor/body_velocity.py | 89 +++++++++++++++++++ tests/test_sensor.py | 85 +++++++++++++++++- 4 files changed, 228 insertions(+), 17 deletions(-) create mode 100644 src/genelab/sensor/body_velocity.py diff --git a/examples/unitree/src/genelab_unitree/g1/env_cfg.py b/examples/unitree/src/genelab_unitree/g1/env_cfg.py index 41570dda..e72ac3e4 100644 --- a/examples/unitree/src/genelab_unitree/g1/env_cfg.py +++ b/examples/unitree/src/genelab_unitree/g1/env_cfg.py @@ -17,29 +17,48 @@ ) from genelab.mdp.actions.joint_position import JointPositionActionCfg from genelab.mdp.commands.velocity_command import UniformVelocityCommandCfg +from genelab.mdp.noise import Unoise +from genelab.sensor import BodyVelocitySensorCfg from genelab_unitree.g1.constants import G1_ACTION_SCALE from genelab_unitree.g1.robot import get_g1_robot_cfg +# IMU site offset from the pelvis link origin; matches g1.xml's . +_IMU_OFFSET = (0.04525, 0.0, -0.08339) + + +def _obs_terms() -> dict[str, ObservationTermCfg]: + return { + "base_lin_vel": ObservationTermCfg( + func=mdp.sensor_data, + params={"sensor_name": "imu_lin_vel"}, + noise=Unoise(-0.5, 0.5), + ), + "base_ang_vel": ObservationTermCfg( + func=mdp.sensor_data, + params={"sensor_name": "imu_ang_vel"}, + noise=Unoise(-0.2, 0.2), + ), + "projected_gravity": ObservationTermCfg( + func=mdp.projected_gravity, noise=Unoise(-0.05, 0.05) + ), + "velocity_commands": ObservationTermCfg( + func=mdp.generated_commands, params={"command_name": "twist"} + ), + "joint_pos": ObservationTermCfg(func=mdp.joint_pos_rel, noise=Unoise(-0.01, 0.01)), + # Unoise lives in raw rad/s; the existing scale=0.05 brings it to ±0.075 final. + "joint_vel": ObservationTermCfg( + func=mdp.joint_vel_rel, scale=0.05, noise=Unoise(-1.5, 1.5) + ), + "actions": ObservationTermCfg(func=mdp.last_action), + } + def _policy_obs_group() -> ObservationGroupCfg: - return ObservationGroupCfg( - terms={ - "base_lin_vel": ObservationTermCfg(func=mdp.base_lin_vel), - "base_ang_vel": ObservationTermCfg(func=mdp.base_ang_vel), - "projected_gravity": ObservationTermCfg(func=mdp.projected_gravity), - "velocity_commands": ObservationTermCfg( - func=mdp.generated_commands, params={"command_name": "twist"} - ), - "joint_pos": ObservationTermCfg(func=mdp.joint_pos_rel), - "joint_vel": ObservationTermCfg(func=mdp.joint_vel_rel, scale=0.05), - "actions": ObservationTermCfg(func=mdp.last_action), - } - ) + return ObservationGroupCfg(enable_corruption=True, terms=_obs_terms()) def _critic_obs_group() -> ObservationGroupCfg: - # Same terms as policy in v1 — privileged signals will land here once contact sensors exist. - return _policy_obs_group() + return ObservationGroupCfg(enable_corruption=False, terms=_obs_terms()) def unitree_g1_velocity_env_cfg(play: bool = False) -> ManagerBasedRlEnvCfg: @@ -53,6 +72,20 @@ def unitree_g1_velocity_env_cfg(play: bool = False) -> ManagerBasedRlEnvCfg: substeps=1, env_spacing=(2.5, 2.5), vis=play, + sensors=( + BodyVelocitySensorCfg( + name="imu_lin_vel", + link_name="pelvis", + offset=_IMU_OFFSET, + measure="lin_vel", + ), + BodyVelocitySensorCfg( + name="imu_ang_vel", + link_name="pelvis", + offset=_IMU_OFFSET, + measure="ang_vel", + ), + ), ), decimation=10, episode_length_s=20.0, diff --git a/src/genelab/sensor/__init__.py b/src/genelab/sensor/__init__.py index fdb9236b..859feebf 100644 --- a/src/genelab/sensor/__init__.py +++ b/src/genelab/sensor/__init__.py @@ -1,5 +1,11 @@ """Sensor abstractions for cameras, contacts, tactile data, and future streams.""" +from genelab.sensor.body_velocity import BodyVelocitySensor, BodyVelocitySensorCfg from genelab.sensor.sensor import Sensor, SensorCfg -__all__ = ["Sensor", "SensorCfg"] +__all__ = [ + "BodyVelocitySensor", + "BodyVelocitySensorCfg", + "Sensor", + "SensorCfg", +] diff --git a/src/genelab/sensor/body_velocity.py b/src/genelab/sensor/body_velocity.py new file mode 100644 index 00000000..f637030b --- /dev/null +++ b/src/genelab/sensor/body_velocity.py @@ -0,0 +1,89 @@ +"""Velocity sensor at a site rigidly attached to a robot link. + +Mirrors MuJoCo's ```` and ````: the value is reported in the +link's body frame, with the velocimeter accounting for the lever arm from link origin to site. +""" + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Literal + +import torch + +from genelab.sensor.sensor import Sensor, SensorCfg +from genelab.utils.math import matrix_from_quat, quat_apply_inverse + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class BodyVelocitySensorCfg(SensorCfg): + """Configuration for ``BodyVelocitySensor``. + + ``offset`` is the site position in the link's local frame; ignored when ``measure="ang_vel"`` + since angular velocity does not depend on the lever arm. ``bias_range``, when set, samples a + constant per-env bias on every ``reset(env_ids)`` to emulate sensor offset randomization. + """ + + link_name: str = "" + offset: tuple[float, float, float] = (0.0, 0.0, 0.0) + measure: Literal["lin_vel", "ang_vel"] = "lin_vel" + bias_range: tuple[float, float] | None = None + + def build(self) -> "BodyVelocitySensor": + return BodyVelocitySensor(self) + + +class BodyVelocitySensor(Sensor[torch.Tensor]): + def __init__(self, cfg: BodyVelocitySensorCfg) -> None: + super().__init__(cfg) + self._cfg_typed = cfg + self._link_idx: int = -1 + self._offset_local: torch.Tensor | None = None + self._bias: torch.Tensor | None = None + + def bind(self, env: "ManagerBasedRlEnv") -> None: + super().bind(env) + if not self._cfg_typed.link_name: + raise ValueError(f"BodyVelocitySensorCfg(name={self._cfg.name!r}) requires link_name") + try: + self._link_idx = env.link_names.index(self._cfg_typed.link_name) + except ValueError as exc: + raise ValueError( + f"sensor {self._cfg.name!r}: link {self._cfg_typed.link_name!r} not found in " + f"env.link_names={env.link_names!r}" + ) from exc + self._offset_local = torch.tensor( + self._cfg_typed.offset, dtype=torch.float32, device=env.device + ) + self._bias = torch.zeros(env.num_envs, 3, device=env.device) + if self._cfg_typed.bias_range is not None: + self._resample_bias(torch.arange(env.num_envs, device=env.device)) + + def reset(self, env_ids: torch.Tensor | None = None) -> None: + super().reset(env_ids) + if self._cfg_typed.bias_range is None or self._bias is None or self._env is None: + return + if env_ids is None: + env_ids = torch.arange(self._env.num_envs, device=self._env.device) + self._resample_bias(env_ids) + + def _resample_bias(self, env_ids: torch.Tensor) -> None: + assert self._bias is not None and self._cfg_typed.bias_range is not None + lo, hi = self._cfg_typed.bias_range + self._bias[env_ids] = ( + torch.rand(env_ids.numel(), 3, device=self._bias.device) * (hi - lo) + lo + ) + + def _compute_data(self) -> torch.Tensor: + assert self._env is not None and self._offset_local is not None and self._bias is not None + rs = self._env.robot_state + link_quat = rs.link_quat_w[:, self._link_idx] + link_ang_vel_w = rs.link_ang_vel_w[:, self._link_idx] + if self._cfg_typed.measure == "ang_vel": + return quat_apply_inverse(link_quat, link_ang_vel_w) + self._bias + link_lin_vel_w = rs.link_lin_vel_w[:, self._link_idx] + R = matrix_from_quat(link_quat) + r_world = torch.matmul(R, self._offset_local) + v_site_w = link_lin_vel_w + torch.cross(link_ang_vel_w, r_world, dim=-1) + return quat_apply_inverse(link_quat, v_site_w) + self._bias diff --git a/tests/test_sensor.py b/tests/test_sensor.py index ad55c649..5eb75e73 100644 --- a/tests/test_sensor.py +++ b/tests/test_sensor.py @@ -11,7 +11,7 @@ ObservationTermCfg, ) from genelab.mdp.noise import Gnoise, Unoise -from genelab.sensor import Sensor, SensorCfg +from genelab.sensor import BodyVelocitySensorCfg, Sensor, SensorCfg @dataclass @@ -142,3 +142,86 @@ def const_one(_env: Any) -> torch.Tensor: mgr = ObservationManager(cfg, env) obs = mgr.compute() assert torch.equal(obs["g"], torch.full((2, 2), 1.0)) + + +# --------------------------------------------------------------------- BodyVelocitySensor + + +class _FakeRobotState: + def __init__(self, num_envs: int, num_links: int, device: str) -> None: + self.link_quat_w = torch.zeros(num_envs, num_links, 4, device=device) + self.link_quat_w[..., 0] = 1.0 + self.link_lin_vel_w = torch.zeros(num_envs, num_links, 3, device=device) + self.link_ang_vel_w = torch.zeros(num_envs, num_links, 3, device=device) + + +class _FakeRobotEnv: + def __init__(self, num_envs: int = 2, link_names: tuple[str, ...] = ("pelvis",)) -> None: + self.num_envs = num_envs + self.device = "cpu" + self.link_names = list(link_names) + self.robot_state = _FakeRobotState(num_envs, len(link_names), self.device) + + +def test_body_velocity_sensor_gyro_rotates_to_body_frame() -> None: + env = _FakeRobotEnv() + # 90° rotation about +z: q = (cos45°, 0, 0, sin45°). World ω = +x → body ω = +y after inverse. + s2 = 2.0**0.5 / 2 + env.robot_state.link_quat_w[:, 0] = torch.tensor([s2, 0.0, 0.0, s2]) + env.robot_state.link_ang_vel_w[:, 0] = torch.tensor([1.0, 0.0, 0.0]) + sensor = BodyVelocitySensorCfg(name="g", link_name="pelvis", measure="ang_vel").build() + sensor.bind(env) + out = sensor.data + assert out.shape == (2, 3) + assert torch.allclose(out[0], torch.tensor([0.0, -1.0, 0.0]), atol=1e-6) + + +def test_body_velocity_sensor_velocimeter_lever_arm_cancels() -> None: + # Pure rotation ω = +z, site offset r = +x → v_site_world = ω × r = +y. + # Then rotate into body frame (identity quat here) → still +y. + env = _FakeRobotEnv() + env.robot_state.link_ang_vel_w[:, 0] = torch.tensor([0.0, 0.0, 1.0]) + sensor = BodyVelocitySensorCfg( + name="v", link_name="pelvis", offset=(1.0, 0.0, 0.0), measure="lin_vel" + ).build() + sensor.bind(env) + out = sensor.data + assert torch.allclose(out[0], torch.tensor([0.0, 1.0, 0.0]), atol=1e-6) + + +def test_body_velocity_sensor_velocimeter_with_translation_only() -> None: + # Pure translation, no rotation: world vel directly returned in body frame (identity quat). + env = _FakeRobotEnv() + env.robot_state.link_lin_vel_w[:, 0] = torch.tensor([2.0, -1.0, 0.5]) + sensor = BodyVelocitySensorCfg( + name="v", link_name="pelvis", offset=(0.05, 0.0, -0.08), measure="lin_vel" + ).build() + sensor.bind(env) + out = sensor.data + assert torch.allclose(out[0], torch.tensor([2.0, -1.0, 0.5]), atol=1e-6) + + +def test_body_velocity_sensor_bias_randomizes_on_reset() -> None: + torch.manual_seed(0) + env = _FakeRobotEnv(num_envs=64) + sensor = BodyVelocitySensorCfg( + name="g", link_name="pelvis", measure="ang_vel", bias_range=(-0.1, 0.1) + ).build() + sensor.bind(env) + first = sensor.data.clone() + assert ((first.abs() <= 0.1 + 1e-6).all()) + assert first.std() > 0.0 # initial bias should already vary across envs + sensor.reset(torch.arange(64)) + second = sensor.data + assert not torch.equal(first, second) + + +def test_body_velocity_sensor_rejects_unknown_link() -> None: + env = _FakeRobotEnv() + sensor = BodyVelocitySensorCfg(name="x", link_name="nonexistent").build() + try: + sensor.bind(env) + except ValueError as exc: + assert "nonexistent" in str(exc) + else: + raise AssertionError("expected ValueError for unknown link_name") From acda61ee19818d0004c45eb864ce63c31ac9f23d Mon Sep 17 00:00:00 2001 From: KraHsu Date: Wed, 13 May 2026 21:56:35 +0800 Subject: [PATCH 3/5] Add ContactSensor with air-time tracking for G1 critic obs --- .../unitree/src/genelab_unitree/g1/env_cfg.py | 20 +- src/genelab/lab.py | 5 +- src/genelab/mdp/__init__.py | 6 + src/genelab/mdp/observations.py | 24 ++ src/genelab/sensor/__init__.py | 4 + src/genelab/sensor/contact.py | 205 ++++++++++++++++++ tests/test_sensor.py | 139 +++++++++++- 7 files changed, 399 insertions(+), 4 deletions(-) create mode 100644 src/genelab/sensor/contact.py diff --git a/examples/unitree/src/genelab_unitree/g1/env_cfg.py b/examples/unitree/src/genelab_unitree/g1/env_cfg.py index e72ac3e4..2daebe5b 100644 --- a/examples/unitree/src/genelab_unitree/g1/env_cfg.py +++ b/examples/unitree/src/genelab_unitree/g1/env_cfg.py @@ -18,12 +18,13 @@ from genelab.mdp.actions.joint_position import JointPositionActionCfg from genelab.mdp.commands.velocity_command import UniformVelocityCommandCfg from genelab.mdp.noise import Unoise -from genelab.sensor import BodyVelocitySensorCfg +from genelab.sensor import BodyVelocitySensorCfg, ContactSensorCfg from genelab_unitree.g1.constants import G1_ACTION_SCALE from genelab_unitree.g1.robot import get_g1_robot_cfg # IMU site offset from the pelvis link origin; matches g1.xml's . _IMU_OFFSET = (0.04525, 0.0, -0.08339) +_G1_FOOT_LINKS = ("left_ankle_roll_link", "right_ankle_roll_link") def _obs_terms() -> dict[str, ObservationTermCfg]: @@ -58,7 +59,17 @@ def _policy_obs_group() -> ObservationGroupCfg: def _critic_obs_group() -> ObservationGroupCfg: - return ObservationGroupCfg(enable_corruption=False, terms=_obs_terms()) + terms = _obs_terms() + terms["foot_air_time"] = ObservationTermCfg( + func=mdp.foot_air_time, params={"sensor_name": "feet_ground_contact"} + ) + terms["foot_contact"] = ObservationTermCfg( + func=mdp.foot_contact, params={"sensor_name": "feet_ground_contact"} + ) + terms["foot_contact_forces"] = ObservationTermCfg( + func=mdp.foot_contact_forces, params={"sensor_name": "feet_ground_contact"} + ) + return ObservationGroupCfg(enable_corruption=False, terms=terms) def unitree_g1_velocity_env_cfg(play: bool = False) -> ManagerBasedRlEnvCfg: @@ -85,6 +96,11 @@ def unitree_g1_velocity_env_cfg(play: bool = False) -> ManagerBasedRlEnvCfg: offset=_IMU_OFFSET, measure="ang_vel", ), + ContactSensorCfg( + name="feet_ground_contact", + link_names=_G1_FOOT_LINKS, + track_air_time=True, + ), ), ), decimation=10, diff --git a/src/genelab/lab.py b/src/genelab/lab.py index 2c01440e..b3a6e36a 100644 --- a/src/genelab/lab.py +++ b/src/genelab/lab.py @@ -15,7 +15,7 @@ load_entrypoint_extensions, load_extension_module, ) -from genelab.sensor import Sensor, SensorCfg +from genelab.sensor import ContactData, ContactSensor, ContactSensorCfg, Sensor, SensorCfg class ManagerBasedEnv(Protocol): @@ -40,6 +40,9 @@ class GenesisBackendCfg: "ENVS", "ROBOTS", "TASKS", + "ContactData", + "ContactSensor", + "ContactSensorCfg", "GenesisBackendCfg", "Gnoise", "ManagerBasedEnv", diff --git a/src/genelab/mdp/__init__.py b/src/genelab/mdp/__init__.py index 3d84c289..8ed9fd9e 100644 --- a/src/genelab/mdp/__init__.py +++ b/src/genelab/mdp/__init__.py @@ -17,6 +17,9 @@ from genelab.mdp.observations import ( base_ang_vel, base_lin_vel, + foot_air_time, + foot_contact, + foot_contact_forces, generated_commands, joint_pos_rel, joint_vel_rel, @@ -72,6 +75,9 @@ "base_lin_vel", "feet_air_time", "flat_orientation_l2", + "foot_air_time", + "foot_contact", + "foot_contact_forces", "generated_commands", "joint_acc_l2", "joint_pos_limits", diff --git a/src/genelab/mdp/observations.py b/src/genelab/mdp/observations.py index 850bd6ee..b3972295 100644 --- a/src/genelab/mdp/observations.py +++ b/src/genelab/mdp/observations.py @@ -5,6 +5,7 @@ import torch from genelab.mdp.commands.motion_command import MotionCommand +from genelab.sensor.contact import ContactSensor from genelab.utils.math import matrix_from_quat, subtract_frame_transforms if TYPE_CHECKING: @@ -49,6 +50,29 @@ def sensor_data(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: return env.sensors[sensor_name].data +def _contact_sensor(env: "ManagerBasedRlEnv", sensor_name: str) -> ContactSensor: + sensor = env.sensors[sensor_name] + if not isinstance(sensor, ContactSensor): + raise TypeError(f"sensor {sensor_name!r} is not a ContactSensor (got {type(sensor).__name__})") + return sensor + + +def foot_air_time(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: + """Current air time per foot (zero while in contact).""" + return _contact_sensor(env, sensor_name).data.current_air_time + + +def foot_contact(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: + """Binary contact mask per foot as a float tensor.""" + return _contact_sensor(env, sensor_name).data.found.float() + + +def foot_contact_forces(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: + """Per-foot contact force, compressed via ``sign(f) * log1p(|f|)`` and flattened to ``(B, N*3)``.""" + force = _contact_sensor(env, sensor_name).data.force + return (force.sign() * torch.log1p(force.abs())).reshape(force.shape[0], -1) + + # --------------------------------------------------------------------- motion imitation diff --git a/src/genelab/sensor/__init__.py b/src/genelab/sensor/__init__.py index 859feebf..a3aef660 100644 --- a/src/genelab/sensor/__init__.py +++ b/src/genelab/sensor/__init__.py @@ -1,11 +1,15 @@ """Sensor abstractions for cameras, contacts, tactile data, and future streams.""" from genelab.sensor.body_velocity import BodyVelocitySensor, BodyVelocitySensorCfg +from genelab.sensor.contact import ContactData, ContactSensor, ContactSensorCfg from genelab.sensor.sensor import Sensor, SensorCfg __all__ = [ "BodyVelocitySensor", "BodyVelocitySensorCfg", + "ContactData", + "ContactSensor", + "ContactSensorCfg", "Sensor", "SensorCfg", ] diff --git a/src/genelab/sensor/contact.py b/src/genelab/sensor/contact.py new file mode 100644 index 00000000..5f4601e0 --- /dev/null +++ b/src/genelab/sensor/contact.py @@ -0,0 +1,205 @@ +"""Per-link contact sensor with optional air-time / contact-time state machine. + +The shape mirrors mjlab's ``ContactSensor`` — same ``found`` / ``force`` / ``current_air_time`` / +``last_air_time`` semantics so obs / reward terms transfer 1:1 — but the backend reads Genesis's +``robot.get_links_net_contact_force()`` (per-link external-contact aggregate) instead of MuJoCo +contact pairs. A contact is "found" when the force magnitude exceeds ``force_threshold`` (N). +""" + +import re +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import torch + +from genelab.sensor.sensor import Sensor, SensorCfg + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class ContactData: + """Per-step contact snapshot. All tensors are shaped ``(num_envs, num_links)`` unless noted.""" + + force: torch.Tensor # (B, N, 3) world-frame net contact force per selected link + force_norm: torch.Tensor # (B, N) magnitude of ``force`` + found: torch.Tensor # (B, N) bool: force_norm > force_threshold + current_air_time: torch.Tensor # (B, N) seconds since last lift-off; 0 while in contact + last_air_time: torch.Tensor # (B, N) duration of the most recently completed air phase + current_contact_time: torch.Tensor # (B, N) seconds since last landing; 0 while airborne + last_contact_time: torch.Tensor # (B, N) duration of the most recently completed contact + + +@dataclass +class ContactSensorCfg(SensorCfg): + """Configuration for ``ContactSensor``. + + Provide either an explicit tuple of ``link_names`` or a regex ``link_names_expr`` (matched + against ``env.link_names`` with ``re.search``). ``track_air_time`` toggles allocation of the + air-time / contact-time state machine; turn it off to skip the per-step state update on + sensors that only need ``force`` / ``found``. + """ + + link_names: tuple[str, ...] = () + link_names_expr: str | None = None + force_threshold: float = 1.0 + track_air_time: bool = True + + def build(self) -> "ContactSensor": + return ContactSensor(self) + + +@dataclass +class _AirTimeState: + current_air_time: torch.Tensor + last_air_time: torch.Tensor + current_contact_time: torch.Tensor + last_contact_time: torch.Tensor + + +def _resolve_link_indices( + cfg: ContactSensorCfg, link_names: list[str] +) -> tuple[list[int], list[str]]: + if cfg.link_names_expr is not None: + pattern = re.compile(cfg.link_names_expr) + matched = [(i, n) for i, n in enumerate(link_names) if pattern.search(n)] + elif cfg.link_names: + missing = [n for n in cfg.link_names if n not in link_names] + if missing: + raise ValueError( + f"ContactSensorCfg(name={cfg.name!r}): link(s) {missing!r} not in env.link_names" + ) + order = {n: i for i, n in enumerate(link_names)} + matched = [(order[n], n) for n in cfg.link_names] + else: + raise ValueError( + f"ContactSensorCfg(name={cfg.name!r}) requires either link_names or link_names_expr" + ) + if not matched: + raise ValueError( + f"ContactSensorCfg(name={cfg.name!r}): " + f"link_names_expr={cfg.link_names_expr!r} matched nothing" + ) + return [i for i, _ in matched], [n for _, n in matched] + + +class ContactSensor(Sensor[ContactData]): + def __init__(self, cfg: ContactSensorCfg) -> None: + super().__init__(cfg) + self._cfg_typed = cfg + self._link_idx_tensor: torch.Tensor | None = None + self._resolved_link_names: list[str] = [] + self._air_state: _AirTimeState | None = None + + @property + def link_names(self) -> list[str]: + return list(self._resolved_link_names) + + def bind(self, env: "ManagerBasedRlEnv") -> None: + super().bind(env) + indices, names = _resolve_link_indices(self._cfg_typed, env.link_names) + self._resolved_link_names = names + self._link_idx_tensor = torch.tensor(indices, dtype=torch.long, device=env.device) + if self._cfg_typed.track_air_time: + shape = (env.num_envs, len(indices)) + self._air_state = _AirTimeState( + current_air_time=torch.zeros(*shape, device=env.device), + last_air_time=torch.zeros(*shape, device=env.device), + current_contact_time=torch.zeros(*shape, device=env.device), + last_contact_time=torch.zeros(*shape, device=env.device), + ) + + def update(self, dt: float) -> None: + super().update(dt) + if self._air_state is None or self._env is None: + return + force = self._read_link_forces() + in_contact = force.norm(dim=-1) > self._cfg_typed.force_threshold + state = self._air_state + # Detect just-landed / just-lifted-off transitions before mutating the timers. + first_contact = in_contact & (state.current_air_time > 0) + first_detached = (~in_contact) & (state.current_contact_time > 0) + state.last_air_time = torch.where( + first_contact, state.current_air_time + dt, state.last_air_time + ) + state.last_contact_time = torch.where( + first_detached, state.current_contact_time + dt, state.last_contact_time + ) + state.current_air_time = torch.where( + in_contact, + torch.zeros_like(state.current_air_time), + state.current_air_time + dt, + ) + state.current_contact_time = torch.where( + in_contact, + state.current_contact_time + dt, + torch.zeros_like(state.current_contact_time), + ) + + def reset(self, env_ids: torch.Tensor | None = None) -> None: + super().reset(env_ids) + if self._air_state is None or self._env is None: + return + if env_ids is None: + env_ids = torch.arange(self._env.num_envs, device=self._env.device) + for buf in ( + self._air_state.current_air_time, + self._air_state.last_air_time, + self._air_state.current_contact_time, + self._air_state.last_contact_time, + ): + buf[env_ids] = 0.0 + + def _compute_data(self) -> ContactData: + force = self._read_link_forces() + force_norm = force.norm(dim=-1) + found = force_norm > self._cfg_typed.force_threshold + if self._air_state is None: + zeros = torch.zeros_like(force_norm) + return ContactData( + force=force, + force_norm=force_norm, + found=found, + current_air_time=zeros, + last_air_time=zeros, + current_contact_time=zeros.clone(), + last_contact_time=zeros.clone(), + ) + s = self._air_state + return ContactData( + force=force, + force_norm=force_norm, + found=found, + current_air_time=s.current_air_time, + last_air_time=s.last_air_time, + current_contact_time=s.current_contact_time, + last_contact_time=s.last_contact_time, + ) + + def _read_link_forces(self) -> torch.Tensor: + """Return per-selected-link external contact force ``(num_envs, num_links, 3)``.""" + assert self._env is not None and self._link_idx_tensor is not None + robot = self._env.robot + getter = getattr(robot, "get_links_net_contact_force", None) + if getter is None: + # Genesis fake / not built — return zeros so the abstraction stays usable in tests. + return torch.zeros( + self._env.num_envs, + self._link_idx_tensor.numel(), + 3, + device=self._env.device, + ) + try: + forces = getter() + except Exception: + return torch.zeros( + self._env.num_envs, + self._link_idx_tensor.numel(), + 3, + device=self._env.device, + ) + forces = forces.to(self._env.device) + if forces.dim() == 2: + forces = forces.unsqueeze(0).expand(self._env.num_envs, -1, -1) + return forces.index_select(-2, self._link_idx_tensor) diff --git a/tests/test_sensor.py b/tests/test_sensor.py index 5eb75e73..b3a4e86a 100644 --- a/tests/test_sensor.py +++ b/tests/test_sensor.py @@ -11,7 +11,7 @@ ObservationTermCfg, ) from genelab.mdp.noise import Gnoise, Unoise -from genelab.sensor import BodyVelocitySensorCfg, Sensor, SensorCfg +from genelab.sensor import BodyVelocitySensorCfg, ContactSensorCfg, Sensor, SensorCfg @dataclass @@ -225,3 +225,140 @@ def test_body_velocity_sensor_rejects_unknown_link() -> None: assert "nonexistent" in str(exc) else: raise AssertionError("expected ValueError for unknown link_name") + + +# --------------------------------------------------------------------- ContactSensor + + +class _FakeRobot: + def __init__(self, contact_force: torch.Tensor) -> None: + # contact_force shape: (num_envs, num_links, 3) + self._contact_force = contact_force + + def get_links_net_contact_force(self) -> torch.Tensor: + return self._contact_force + + def set_contact_force(self, contact_force: torch.Tensor) -> None: + self._contact_force = contact_force + + +class _FakeContactEnv: + def __init__(self, num_envs: int, link_names: tuple[str, ...]) -> None: + self.num_envs = num_envs + self.device = "cpu" + self.link_names = list(link_names) + self.robot = _FakeRobot(torch.zeros(num_envs, len(link_names), 3)) + + +def test_contact_sensor_explicit_link_names_resolves_indices() -> None: + env = _FakeContactEnv(num_envs=2, link_names=("base", "left_foot", "right_foot", "head")) + sensor = ContactSensorCfg( + name="feet", link_names=("left_foot", "right_foot"), track_air_time=False + ).build() + sensor.bind(env) + assert sensor.link_names == ["left_foot", "right_foot"] + env.robot.set_contact_force( + torch.tensor( + [ + [[0, 0, 0], [0, 0, 50.0], [0, 0, 0], [0, 0, 0]], + [[0, 0, 0], [0, 0, 0], [0, 0, 30.0], [0, 0, 0]], + ] + ) + ) + sensor.update(0.02) + data = sensor.data + assert data.force_norm.shape == (2, 2) + assert torch.allclose(data.force_norm[0], torch.tensor([50.0, 0.0])) + assert torch.allclose(data.force_norm[1], torch.tensor([0.0, 30.0])) + assert data.found.dtype == torch.bool + assert torch.equal(data.found, torch.tensor([[True, False], [False, True]])) + + +def test_contact_sensor_regex_match() -> None: + env = _FakeContactEnv(num_envs=1, link_names=("base", "left_foot", "right_foot")) + sensor = ContactSensorCfg( + name="f", link_names_expr=r"_foot$", track_air_time=False + ).build() + sensor.bind(env) + assert sensor.link_names == ["left_foot", "right_foot"] + + +def test_contact_sensor_air_time_state_machine() -> None: + env = _FakeContactEnv(num_envs=1, link_names=("foot",)) + sensor = ContactSensorCfg(name="c", link_names=("foot",), track_air_time=True).build() + sensor.bind(env) + + in_air = torch.zeros(1, 1, 3) + in_contact = torch.tensor([[[0.0, 0.0, 100.0]]]) + dt = 0.05 + + env.robot.set_contact_force(in_air) + sensor.update(dt) # tick 1 in air + assert torch.allclose(sensor.data.current_air_time, torch.tensor([[dt]])) + assert torch.allclose(sensor.data.current_contact_time, torch.tensor([[0.0]])) + + sensor._invalidate_cache() + sensor.update(dt) # tick 2 in air + assert torch.allclose(sensor.data.current_air_time, torch.tensor([[2 * dt]])) + + env.robot.set_contact_force(in_contact) + sensor._invalidate_cache() + sensor.update(dt) # landing tick + d = sensor.data + assert torch.allclose(d.current_air_time, torch.tensor([[0.0]])) + assert torch.allclose(d.last_air_time, torch.tensor([[3 * dt]])) + assert torch.allclose(d.current_contact_time, torch.tensor([[dt]])) + + sensor._invalidate_cache() + sensor.update(dt) # continued contact + assert torch.allclose(sensor.data.current_contact_time, torch.tensor([[2 * dt]])) + + env.robot.set_contact_force(in_air) + sensor._invalidate_cache() + sensor.update(dt) # lift-off + d = sensor.data + assert torch.allclose(d.current_contact_time, torch.tensor([[0.0]])) + assert torch.allclose(d.last_contact_time, torch.tensor([[3 * dt]])) + assert torch.allclose(d.current_air_time, torch.tensor([[dt]])) + + +def test_contact_sensor_reset_clears_state_for_env_ids_only() -> None: + env = _FakeContactEnv(num_envs=2, link_names=("foot",)) + sensor = ContactSensorCfg(name="c", link_names=("foot",), track_air_time=True).build() + sensor.bind(env) + env.robot.set_contact_force(torch.zeros(2, 1, 3)) + for _ in range(3): + sensor._invalidate_cache() + sensor.update(0.1) + assert torch.allclose(sensor.data.current_air_time, torch.tensor([[0.3], [0.3]])) + sensor.reset(torch.tensor([0])) + sensor._invalidate_cache() + sensor.update(0.1) + out = sensor.data.current_air_time + assert torch.allclose(out[0], torch.tensor([0.1])) + assert torch.allclose(out[1], torch.tensor([0.4])) + + +def test_contact_sensor_force_threshold_controls_found_bit() -> None: + env = _FakeContactEnv(num_envs=1, link_names=("foot",)) + sensor = ContactSensorCfg( + name="c", link_names=("foot",), force_threshold=10.0, track_air_time=False + ).build() + sensor.bind(env) + env.robot.set_contact_force(torch.tensor([[[0.0, 0.0, 5.0]]])) + sensor.update(0.02) + assert sensor.data.found.item() is False + env.robot.set_contact_force(torch.tensor([[[0.0, 0.0, 12.0]]])) + sensor._invalidate_cache() + sensor.update(0.02) + assert sensor.data.found.item() is True + + +def test_contact_sensor_rejects_unresolved_links() -> None: + env = _FakeContactEnv(num_envs=1, link_names=("a", "b")) + try: + ContactSensorCfg(name="x", link_names=("c",)).build().bind(env) + except ValueError as exc: + assert "'c'" in str(exc) + else: + raise AssertionError("expected ValueError for unknown link") From 48fd3bddb8e6cbb4fba3c00269f13aefd6908449 Mon Sep 17 00:00:00 2001 From: KraHsu Date: Wed, 13 May 2026 22:01:23 +0800 Subject: [PATCH 4/5] Add RayCastSensor + TerrainHeightSensor with flat-plane backend --- src/genelab/lab.py | 20 +++- src/genelab/mdp/__init__.py | 2 + src/genelab/mdp/observations.py | 10 ++ src/genelab/sensor/__init__.py | 13 +++ src/genelab/sensor/ray_cast.py | 158 +++++++++++++++++++++++++++ src/genelab/sensor/terrain_height.py | 63 +++++++++++ tests/test_sensor.py | 120 +++++++++++++++++++- 7 files changed, 384 insertions(+), 2 deletions(-) create mode 100644 src/genelab/sensor/ray_cast.py create mode 100644 src/genelab/sensor/terrain_height.py diff --git a/src/genelab/lab.py b/src/genelab/lab.py index b3a6e36a..eac36350 100644 --- a/src/genelab/lab.py +++ b/src/genelab/lab.py @@ -15,7 +15,19 @@ load_entrypoint_extensions, load_extension_module, ) -from genelab.sensor import ContactData, ContactSensor, ContactSensorCfg, Sensor, SensorCfg +from genelab.sensor import ( + ContactData, + ContactSensor, + ContactSensorCfg, + GridPattern, + RayCastData, + RayCastSensor, + RayCastSensorCfg, + Sensor, + SensorCfg, + TerrainHeightSensor, + TerrainHeightSensorCfg, +) class ManagerBasedEnv(Protocol): @@ -45,14 +57,20 @@ class GenesisBackendCfg: "ContactSensorCfg", "GenesisBackendCfg", "Gnoise", + "GridPattern", "ManagerBasedEnv", "ManagerBasedEnvCfg", "NoiseCfg", + "RayCastData", + "RayCastSensor", + "RayCastSensorCfg", "Registry", "RegistryEntry", "Sensor", "SensorCfg", "TaskCfg", + "TerrainHeightSensor", + "TerrainHeightSensorCfg", "Unoise", "apply_overrides", "load_builtin_registries", diff --git a/src/genelab/mdp/__init__.py b/src/genelab/mdp/__init__.py index 8ed9fd9e..0982e402 100644 --- a/src/genelab/mdp/__init__.py +++ b/src/genelab/mdp/__init__.py @@ -21,6 +21,7 @@ foot_contact, foot_contact_forces, generated_commands, + height_scan, joint_pos_rel, joint_vel_rel, last_action, @@ -79,6 +80,7 @@ "foot_contact", "foot_contact_forces", "generated_commands", + "height_scan", "joint_acc_l2", "joint_pos_limits", "joint_pos_rel", diff --git a/src/genelab/mdp/observations.py b/src/genelab/mdp/observations.py index b3972295..cab0015f 100644 --- a/src/genelab/mdp/observations.py +++ b/src/genelab/mdp/observations.py @@ -73,6 +73,16 @@ def foot_contact_forces(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Ten return (force.sign() * torch.log1p(force.abs())).reshape(force.shape[0], -1) +def height_scan(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: + """Per-ray heights from a ``TerrainHeightSensor`` (positive = above terrain).""" + out = env.sensors[sensor_name].data + if not isinstance(out, torch.Tensor): + raise TypeError( + f"sensor {sensor_name!r} does not return a height tensor (got {type(out).__name__})" + ) + return out + + # --------------------------------------------------------------------- motion imitation diff --git a/src/genelab/sensor/__init__.py b/src/genelab/sensor/__init__.py index a3aef660..708698b0 100644 --- a/src/genelab/sensor/__init__.py +++ b/src/genelab/sensor/__init__.py @@ -2,7 +2,14 @@ from genelab.sensor.body_velocity import BodyVelocitySensor, BodyVelocitySensorCfg from genelab.sensor.contact import ContactData, ContactSensor, ContactSensorCfg +from genelab.sensor.ray_cast import ( + GridPattern, + RayCastData, + RayCastSensor, + RayCastSensorCfg, +) from genelab.sensor.sensor import Sensor, SensorCfg +from genelab.sensor.terrain_height import TerrainHeightSensor, TerrainHeightSensorCfg __all__ = [ "BodyVelocitySensor", @@ -10,6 +17,12 @@ "ContactData", "ContactSensor", "ContactSensorCfg", + "GridPattern", + "RayCastData", + "RayCastSensor", + "RayCastSensorCfg", "Sensor", "SensorCfg", + "TerrainHeightSensor", + "TerrainHeightSensorCfg", ] diff --git a/src/genelab/sensor/ray_cast.py b/src/genelab/sensor/ray_cast.py new file mode 100644 index 00000000..b4db654e --- /dev/null +++ b/src/genelab/sensor/ray_cast.py @@ -0,0 +1,158 @@ +"""Ray-cast sensor with a configurable grid pattern, attached to a robot link. + +The flat-plane backend assumes the ground is a horizontal infinite plane at ``ground_height`` — +matches GeneLab's current ``gs.morphs.Plane()`` setup. Subclasses override +``_intersect_world_rays`` to plug in terrain queries, BVH raycasting, or sloped grounds. + +Pattern shape mirrors mjlab's ``GridPatternCfg``; for now only a 2D grid is provided. +""" + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +import torch + +from genelab.sensor.sensor import Sensor, SensorCfg +from genelab.utils.math import quat_apply, yaw_quat + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class GridPattern: + """2D grid of ray origins centred at the sensor frame, all pointing along ``direction``.""" + + resolution: float = 0.1 + size: tuple[float, float] = (1.6, 1.0) + direction: tuple[float, float, float] = (0.0, 0.0, -1.0) + + def num_rays(self) -> int: + nx = int(round(self.size[0] / self.resolution)) + 1 + ny = int(round(self.size[1] / self.resolution)) + 1 + return nx * ny + + def generate(self, device: str) -> tuple[torch.Tensor, torch.Tensor]: + """Return ``(starts_local, dirs_local)``, each shape ``(M, 3)``.""" + nx = int(round(self.size[0] / self.resolution)) + 1 + ny = int(round(self.size[1] / self.resolution)) + 1 + xs = torch.linspace(-self.size[0] / 2, self.size[0] / 2, nx, device=device) + ys = torch.linspace(-self.size[1] / 2, self.size[1] / 2, ny, device=device) + xx, yy = torch.meshgrid(xs, ys, indexing="ij") + starts = torch.stack( + [xx.reshape(-1), yy.reshape(-1), torch.zeros_like(xx).reshape(-1)], dim=-1 + ) + dir_tensor = torch.tensor(self.direction, dtype=torch.float32, device=device) + # Normalise so distances match world units even if direction wasn't unit length. + dir_tensor = dir_tensor / dir_tensor.norm().clamp_min(1e-9) + dirs = dir_tensor.unsqueeze(0).expand(starts.shape[0], -1).contiguous() + return starts, dirs + + +@dataclass +class RayCastData: + distances: torch.Tensor # (B, M) — clamped to max_distance; equals max_distance on a miss + hit_pos_w: torch.Tensor # (B, M, 3) — world-frame intersection points + normals_w: torch.Tensor # (B, M, 3) — surface normals at hit + ray_starts_w: torch.Tensor # (B, M, 3) — world-frame ray origins + ray_dirs_w: torch.Tensor # (B, M, 3) — world-frame ray directions + + +@dataclass +class RayCastSensorCfg(SensorCfg): + """Configuration for ``RayCastSensor``. + + ``link_name`` anchors the pattern. ``attach_yaw_only`` keeps the grid axis-aligned to the + horizon when the link rolls / pitches (typical for terrain height scans on a torso link); + set ``False`` to follow the full link orientation. ``ground_height`` is the z of the flat + plane the default backend intersects against; subclasses ignore it. + """ + + link_name: str = "" + pattern: GridPattern = field(default_factory=GridPattern) + attach_yaw_only: bool = True + max_distance: float = 10.0 + ground_height: float = 0.0 + + def build(self) -> "RayCastSensor": + return RayCastSensor(self) + + +class RayCastSensor(Sensor[RayCastData]): + def __init__(self, cfg: RayCastSensorCfg) -> None: + super().__init__(cfg) + self._cfg_typed = cfg + self._link_idx: int = -1 + self._ray_starts_local: torch.Tensor | None = None + self._ray_dirs_local: torch.Tensor | None = None + + @property + def num_rays(self) -> int: + return self._cfg_typed.pattern.num_rays() + + def bind(self, env: "ManagerBasedRlEnv") -> None: + super().bind(env) + if not self._cfg_typed.link_name: + raise ValueError(f"RayCastSensorCfg(name={self._cfg.name!r}) requires link_name") + try: + self._link_idx = env.link_names.index(self._cfg_typed.link_name) + except ValueError as exc: + raise ValueError( + f"sensor {self._cfg.name!r}: link {self._cfg_typed.link_name!r} not in " + f"env.link_names={env.link_names!r}" + ) from exc + self._ray_starts_local, self._ray_dirs_local = self._cfg_typed.pattern.generate(env.device) + + def _compute_data(self) -> RayCastData: + assert ( + self._env is not None + and self._ray_starts_local is not None + and self._ray_dirs_local is not None + ) + rs = self._env.robot_state + link_pos = rs.link_pos[:, self._link_idx] + link_quat = rs.link_quat_w[:, self._link_idx] + rot_q = yaw_quat(link_quat) if self._cfg_typed.attach_yaw_only else link_quat + # Project the local pattern into world coordinates per env. + b = link_pos.shape[0] + m = self._ray_starts_local.shape[0] + q_expanded = rot_q.unsqueeze(1).expand(b, m, 4).contiguous() + starts_local_b = self._ray_starts_local.unsqueeze(0).expand(b, m, 3).contiguous() + dirs_local_b = self._ray_dirs_local.unsqueeze(0).expand(b, m, 3).contiguous() + starts_w = quat_apply(q_expanded, starts_local_b) + link_pos.unsqueeze(1) + dirs_w = quat_apply(q_expanded, dirs_local_b) + distances, hit_pos_w, normals_w = self._intersect_world_rays(starts_w, dirs_w) + return RayCastData( + distances=distances, + hit_pos_w=hit_pos_w, + normals_w=normals_w, + ray_starts_w=starts_w, + ray_dirs_w=dirs_w, + ) + + def _intersect_world_rays( + self, starts_w: torch.Tensor, dirs_w: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Default backend: intersect every ray with the horizontal plane ``z = ground_height``. + + Subclasses may override to query terrain height fields or call into a Genesis raycaster. + Returns ``(distances, hit_pos_w, normals_w)`` each shaped ``(B, M, ...)``. + """ + assert self._env is not None + ground_z = self._cfg_typed.ground_height + max_dist = self._cfg_typed.max_distance + dir_z = dirs_w[..., 2] + # Avoid division-by-zero / wrong-side hits (ray going up or parallel) by masking them + # to a "miss" with distance = max_dist and hit at start + max_dist * dir. + valid = (dir_z < -1e-6) & (starts_w[..., 2] > ground_z) + t = torch.where( + valid, (ground_z - starts_w[..., 2]) / dir_z, torch.full_like(dir_z, max_dist) + ) + t = t.clamp(min=0.0, max=max_dist) + hit_pos_w = starts_w + t.unsqueeze(-1) * dirs_w + # When invalid, place hit_pos at the end of the unit-length ray. + miss_pos = starts_w + max_dist * dirs_w + hit_pos_w = torch.where(valid.unsqueeze(-1), hit_pos_w, miss_pos) + normals_w = torch.zeros_like(hit_pos_w) + normals_w[..., 2] = 1.0 + return t, hit_pos_w, normals_w diff --git a/src/genelab/sensor/terrain_height.py b/src/genelab/sensor/terrain_height.py new file mode 100644 index 00000000..0d1c07c3 --- /dev/null +++ b/src/genelab/sensor/terrain_height.py @@ -0,0 +1,63 @@ +"""Terrain height-scan sensor: per-ray ``frame_z - hit_z`` (positive = above terrain). + +Composes a ``RayCastSensor`` internally rather than subclassing — keeps the Sensor[T] type +parameter clean (this sensor returns a ``torch.Tensor``, the underlying ray cast returns a +``RayCastData``). With the default flat-plane backend, every output entry is +``link_z - ground_height``; for non-flat ground, swap in a ``RayCastSensor`` subclass that +overrides ``_intersect_world_rays``. +""" + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +import torch + +from genelab.sensor.ray_cast import GridPattern, RayCastSensor, RayCastSensorCfg +from genelab.sensor.sensor import Sensor, SensorCfg + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class TerrainHeightSensorCfg(SensorCfg): + link_name: str = "" + pattern: GridPattern = field(default_factory=GridPattern) + attach_yaw_only: bool = True + max_distance: float = 10.0 + ground_height: float = 0.0 + + def build(self) -> "TerrainHeightSensor": + return TerrainHeightSensor(self) + + +class TerrainHeightSensor(Sensor[torch.Tensor]): + def __init__(self, cfg: TerrainHeightSensorCfg) -> None: + super().__init__(cfg) + self._cfg_typed = cfg + self._inner = RayCastSensor( + RayCastSensorCfg( + name=f"{cfg.name}__inner", + link_name=cfg.link_name, + pattern=cfg.pattern, + attach_yaw_only=cfg.attach_yaw_only, + max_distance=cfg.max_distance, + ground_height=cfg.ground_height, + ) + ) + + def bind(self, env: "ManagerBasedRlEnv") -> None: + super().bind(env) + self._inner.bind(env) + + def update(self, dt: float) -> None: + super().update(dt) + self._inner.update(dt) + + def reset(self, env_ids: torch.Tensor | None = None) -> None: + super().reset(env_ids) + self._inner.reset(env_ids) + + def _compute_data(self) -> torch.Tensor: + raw = self._inner.data + return raw.ray_starts_w[..., 2] - raw.hit_pos_w[..., 2] diff --git a/tests/test_sensor.py b/tests/test_sensor.py index b3a4e86a..2f30050a 100644 --- a/tests/test_sensor.py +++ b/tests/test_sensor.py @@ -11,7 +11,15 @@ ObservationTermCfg, ) from genelab.mdp.noise import Gnoise, Unoise -from genelab.sensor import BodyVelocitySensorCfg, ContactSensorCfg, Sensor, SensorCfg +from genelab.sensor import ( + BodyVelocitySensorCfg, + ContactSensorCfg, + GridPattern, + RayCastSensorCfg, + Sensor, + SensorCfg, + TerrainHeightSensorCfg, +) @dataclass @@ -362,3 +370,113 @@ def test_contact_sensor_rejects_unresolved_links() -> None: assert "'c'" in str(exc) else: raise AssertionError("expected ValueError for unknown link") + + +# --------------------------------------------------------------------- RayCast / TerrainHeight + + +class _FakeTerrainState: + def __init__( + self, + num_envs: int, + num_links: int, + device: str, + link_pos: torch.Tensor, + ) -> None: + self.link_pos = link_pos + self.link_quat_w = torch.zeros(num_envs, num_links, 4, device=device) + self.link_quat_w[..., 0] = 1.0 + self.link_lin_vel_w = torch.zeros(num_envs, num_links, 3, device=device) + self.link_ang_vel_w = torch.zeros(num_envs, num_links, 3, device=device) + + +class _FakeTerrainEnv: + def __init__( + self, + num_envs: int, + link_names: tuple[str, ...], + link_pos: torch.Tensor, + ) -> None: + self.num_envs = num_envs + self.device = "cpu" + self.link_names = list(link_names) + self.robot_state = _FakeTerrainState(num_envs, len(link_names), self.device, link_pos) + + +def test_grid_pattern_generates_expected_count_and_shape() -> None: + pattern = GridPattern(resolution=0.5, size=(2.0, 1.0), direction=(0.0, 0.0, -1.0)) + starts, dirs = pattern.generate("cpu") + assert pattern.num_rays() == 5 * 3 + assert starts.shape == (15, 3) + assert dirs.shape == (15, 3) + assert torch.allclose(dirs[0], torch.tensor([0.0, 0.0, -1.0])) + # Outer corners are at the configured extent. + assert torch.allclose(starts[..., 0].max(), torch.tensor(1.0)) + assert torch.allclose(starts[..., 0].min(), torch.tensor(-1.0)) + + +def test_raycast_sensor_returns_link_z_distance_on_flat_plane() -> None: + link_pos = torch.tensor([[[0.0, 0.0, 1.2]]]) + env = _FakeTerrainEnv(1, ("torso",), link_pos) + sensor = RayCastSensorCfg( + name="r", + link_name="torso", + pattern=GridPattern(resolution=0.5, size=(1.0, 1.0)), + max_distance=5.0, + ).build() + sensor.bind(env) + data = sensor.data + # 9 rays in a 3x3 grid; each starts at link_pos.z=1.2 and hits z=0 → distance 1.2. + assert data.distances.shape == (1, 9) + assert torch.allclose(data.distances, torch.full((1, 9), 1.2)) + assert torch.allclose(data.hit_pos_w[..., 2], torch.zeros(1, 9)) + assert torch.allclose(data.normals_w[..., 2], torch.ones(1, 9)) + + +def test_raycast_sensor_clamps_at_max_distance_when_no_intersection() -> None: + # Link is below ground → upward-only hit; default backend treats as miss → max_distance. + link_pos = torch.tensor([[[0.0, 0.0, -1.0]]]) + env = _FakeTerrainEnv(1, ("torso",), link_pos) + sensor = RayCastSensorCfg( + name="r", + link_name="torso", + pattern=GridPattern(resolution=1.0, size=(0.0, 0.0)), + max_distance=3.0, + ).build() + sensor.bind(env) + data = sensor.data + assert torch.allclose(data.distances, torch.full_like(data.distances, 3.0)) + + +def test_terrain_height_sensor_returns_link_height_on_flat_plane() -> None: + # Two envs at different heights — height_scan output is constant across rays. + link_pos = torch.tensor([[[0.0, 0.0, 0.8]], [[1.0, 1.0, 1.5]]]) + env = _FakeTerrainEnv(2, ("torso",), link_pos) + sensor = TerrainHeightSensorCfg( + name="h", + link_name="torso", + pattern=GridPattern(resolution=0.5, size=(1.0, 1.0)), + ).build() + sensor.bind(env) + heights = sensor.data + assert heights.shape == (2, 9) + assert torch.allclose(heights[0], torch.full((9,), 0.8)) + assert torch.allclose(heights[1], torch.full((9,), 1.5)) + + +def test_terrain_height_sensor_lifecycle_propagates_to_inner_sensor() -> None: + link_pos = torch.tensor([[[0.0, 0.0, 1.0]]]) + env = _FakeTerrainEnv(1, ("torso",), link_pos) + sensor = TerrainHeightSensorCfg( + name="h", + link_name="torso", + pattern=GridPattern(resolution=1.0, size=(0.0, 0.0)), + ).build() + sensor.bind(env) + _ = sensor.data + assert sensor._cache_valid is True + sensor.update(0.02) + assert sensor._cache_valid is False + _ = sensor.data + sensor.reset(torch.tensor([0])) + assert sensor._cache_valid is False From a57c89d1ff71eb3161ace94a75d8c60306f33ae9 Mon Sep 17 00:00:00 2001 From: KraHsu Date: Wed, 13 May 2026 22:09:45 +0800 Subject: [PATCH 5/5] Document the sensor abstraction in concepts/sensors --- docs/concepts/sensors.en.md | 185 ++++++++++++++++++++++++++++++++ docs/concepts/sensors.zh.md | 177 ++++++++++++++++++++++++++++++ mkdocs.yml | 2 + src/genelab/mdp/observations.py | 4 +- tests/test_sensor.py | 6 +- 5 files changed, 369 insertions(+), 5 deletions(-) create mode 100644 docs/concepts/sensors.en.md create mode 100644 docs/concepts/sensors.zh.md diff --git a/docs/concepts/sensors.en.md b/docs/concepts/sensors.en.md new file mode 100644 index 00000000..aa8831ff --- /dev/null +++ b/docs/concepts/sensors.en.md @@ -0,0 +1,185 @@ +# Sensors + +Genesis does not parse the MJCF `` block, so GeneLab introduces a backend-agnostic +sensor abstraction. The interface mirrors mjlab's `SensorCfg` / `Sensor[T]`, but every concrete +sensor reads from the env's `RobotState` instead of MuJoCo sensordata — observation and reward +terms transfer between backends without ceremony. + +## Lifecycle + +`bind(env)` runs once at construction. Each control step, `update(dt)` invalidates the cache; +the first access to `data` triggers `_compute_data` lazily. `reset(env_ids)` invalidates the +cache and lets stateful sensors clear per-env buffers. + +```python +class Sensor[T](ABC): + def bind(self, env: "ManagerBasedRlEnv") -> None: ... + @property + def data(self) -> T: ... + def update(self, dt: float) -> None: ... + def reset(self, env_ids: torch.Tensor | None = None) -> None: ... + @abstractmethod + def _compute_data(self) -> T: ... +``` + +The env wires the lifecycle automatically: sensors built from `SceneCfg.sensors` get `update` +called after `_refresh_robot_state` and `reset` called from inside `_reset_idx`, so reward and +observation terms always see fresh sensor data. + +## Registering on a scene + +`SceneCfg.sensors` is a tuple of `SensorCfg`. `ManagerBasedRlEnv.__init__` calls `build()` on +each cfg and binds the resulting sensor to the env. Access at runtime is via +`env.sensors[name].data`. + +```python +from genelab.configs import SceneCfg +from genelab.sensor import BodyVelocitySensorCfg, ContactSensorCfg + +scene = SceneCfg( + num_envs=4096, + sensors=( + BodyVelocitySensorCfg( + name="imu_lin_vel", + link_name="pelvis", + offset=(0.04525, 0.0, -0.08339), + measure="lin_vel", + ), + ContactSensorCfg( + name="feet_ground_contact", + link_names=("left_ankle_roll_link", "right_ankle_roll_link"), + track_air_time=True, + ), + ), +) +``` + +## Built-in sensors + +### BodyVelocitySensor + +Mirrors a MuJoCo `` / `` on a site rigidly attached to a robot link. Returns +the linear or angular velocity of the site in the link's body frame, with a configurable +lever-arm offset and optional per-env uniform bias. + +| Field | Type | Meaning | +|-------|------|---------| +| `link_name` | `str` | Link the site is rigidly attached to. | +| `offset` | `tuple[float, float, float]` | Site position in the link's local frame. Ignored for `ang_vel`. | +| `measure` | `Literal["lin_vel", "ang_vel"]` | Velocimeter (linear) or gyro (angular). | +| `bias_range` | `tuple[float, float] \| None` | Uniform per-env bias resampled on `reset`. | + +The velocimeter math is `v_site = v_link + ω × (R_link · offset)` rotated into the body frame — +matching MuJoCo's lever-arm convention so a GeneLab-trained policy reads the same signal as the +mjlab reference. + +### ContactSensor + +Per-link aggregate of `robot.get_links_net_contact_force()`. With `track_air_time=True`, an +internal state machine advances `current_air_time` / `current_contact_time` per env, and snaps +the completed durations into `last_air_time` / `last_contact_time` at the contact transition. + +| Field | Type | Meaning | +|-------|------|---------| +| `link_names` | `tuple[str, ...]` | Explicit list of link names to monitor. | +| `link_names_expr` | `str \| None` | Regex matched against `env.link_names`. | +| `force_threshold` | `float` | Force magnitude (N) above which `found` is true. | +| `track_air_time` | `bool` | Allocate the air-time / contact-time state machine. | + +`data` is a `ContactData` dataclass with `force`, `force_norm`, `found`, plus the four air-time +buffers. The matching obs terms — `mdp.foot_air_time`, `mdp.foot_contact`, +`mdp.foot_contact_forces` — read straight off this dataclass. + +### TerrainHeightSensor + +2D grid of downward rays anchored to a robot link. Output is per-ray height above the terrain +(positive = above), useful as a privileged `height_scan` critic observation. The default +backend intersects every ray against a horizontal plane at `ground_height`; subclassing +`RayCastSensor` and overriding `_intersect_world_rays` is the extension point for non-flat +terrain. + +| Field | Type | Meaning | +|-------|------|---------| +| `link_name` | `str` | Anchor link for the grid origin. | +| `pattern` | `GridPattern` | Grid resolution / size / direction. | +| `attach_yaw_only` | `bool` | Rotate the grid by yaw only so it stays horizon-aligned. | +| `max_distance` | `float` | Distance clamp for the ray cast. | +| `ground_height` | `float` | Plane height used by the default flat-plane backend. | + +## Observations with noise + +`ObservationTermCfg.noise` accepts an additive noise model (`Unoise(n_min, n_max)` or +`Gnoise(mean, std)`). `ObservationGroupCfg.enable_corruption` gates whether the noise is +applied — disabled by default to keep the critic on ground truth. The per-term pipeline order +is **noise → scale → clip**, so noise magnitudes live in raw signal space: `Unoise(-1.5, 1.5)` +on a raw `joint_vel` term with `scale=0.05` ends up as ±0.075 final jitter. + +The canonical pattern shares terms between policy and critic and differs only on +`enable_corruption`: + +```python +from genelab import mdp +from genelab.managers import ObservationGroupCfg, ObservationTermCfg +from genelab.mdp.noise import Unoise + + +def _obs_terms() -> dict[str, ObservationTermCfg]: + return { + "base_lin_vel": ObservationTermCfg( + func=mdp.sensor_data, + params={"sensor_name": "imu_lin_vel"}, + noise=Unoise(-0.5, 0.5), + ), + "joint_vel": ObservationTermCfg( + func=mdp.joint_vel_rel, + scale=0.05, + noise=Unoise(-1.5, 1.5), + ), + } + + +policy = ObservationGroupCfg(enable_corruption=True, terms=_obs_terms()) +critic = ObservationGroupCfg(enable_corruption=False, terms=_obs_terms()) +``` + +## Writing a custom sensor + +Subclass `Sensor[T]` with the desired return type and implement `_compute_data`. Override +`bind` to cache link indices once, `update` to advance integrators, and `reset` to clear per-env +state — always calling `super()` first so the cache-invalidation chain stays intact. + +```python +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import torch + +from genelab.sensor import Sensor, SensorCfg + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class JointTorqueSensorCfg(SensorCfg): + def build(self) -> "JointTorqueSensor": + return JointTorqueSensor(self) + + +class JointTorqueSensor(Sensor[torch.Tensor]): + def bind(self, env: "ManagerBasedRlEnv") -> None: + super().bind(env) + # Cache anything that depends on env.link_names / env.joint_names here. + + def _compute_data(self) -> torch.Tensor: + assert self._env is not None + rs = self._env.robot_state + return self._env.joint_kp * (rs.joint_pos - self._env.default_joint_pos) +``` + +Add the cfg to `SceneCfg.sensors` and the sensor is reachable as `env.sensors[name].data`. + +## See also + +- [Configs](configs.md) +- [API Reference](../api/reference.md) diff --git a/docs/concepts/sensors.zh.md b/docs/concepts/sensors.zh.md new file mode 100644 index 00000000..e29eb050 --- /dev/null +++ b/docs/concepts/sensors.zh.md @@ -0,0 +1,177 @@ +# 传感器 + +Genesis 不解析 MJCF 中的 `` 段,因此 GeneLab 提供了一套 backend-agnostic 的传感器抽象。 +接口与 mjlab 的 `SensorCfg` / `Sensor[T]` 对齐,每个具体传感器从 env 的 `RobotState` 而非 +MuJoCo sensordata 取数 —— 观测与奖励 term 可在两套后端之间无缝迁移。 + +## 生命周期 + +`bind(env)` 在构造时调用一次。每个控制步 `update(dt)` 失效缓存;首次访问 `data` 触发 +`_compute_data` 惰性求值;`reset(env_ids)` 同样失效缓存,并允许有状态传感器清空对应 env 的 +缓冲区。 + +```python +class Sensor[T](ABC): + def bind(self, env: "ManagerBasedRlEnv") -> None: ... + @property + def data(self) -> T: ... + def update(self, dt: float) -> None: ... + def reset(self, env_ids: torch.Tensor | None = None) -> None: ... + @abstractmethod + def _compute_data(self) -> T: ... +``` + +env 自动接好生命周期:`SceneCfg.sensors` 里构造出的传感器,`update` 在 `_refresh_robot_state` +之后被调用,`reset` 在 `_reset_idx` 内部被调用 —— 奖励和观测 term 总能读到当前步的数据。 + +## 注册到 scene + +`SceneCfg.sensors` 是一个 `SensorCfg` 元组。`ManagerBasedRlEnv.__init__` 对每个 cfg 调用 +`build()` 并把生成的传感器绑定到 env。运行时通过 `env.sensors[name].data` 访问。 + +```python +from genelab.configs import SceneCfg +from genelab.sensor import BodyVelocitySensorCfg, ContactSensorCfg + +scene = SceneCfg( + num_envs=4096, + sensors=( + BodyVelocitySensorCfg( + name="imu_lin_vel", + link_name="pelvis", + offset=(0.04525, 0.0, -0.08339), + measure="lin_vel", + ), + ContactSensorCfg( + name="feet_ground_contact", + link_names=("left_ankle_roll_link", "right_ankle_roll_link"), + track_air_time=True, + ), + ), +) +``` + +## 内置传感器 + +### BodyVelocitySensor + +对应 MuJoCo 的 `` / ``:在与 link 刚性连接的 site 上读取线速度或角速度, +结果旋转到 link body frame。支持杠杆臂偏移与可选的每 env 均匀分布偏置。 + +| 字段 | 类型 | 含义 | +|------|------|------| +| `link_name` | `str` | site 刚性连接到的 link。 | +| `offset` | `tuple[float, float, float]` | site 在 link 局部系下的位置,`ang_vel` 模式忽略。 | +| `measure` | `Literal["lin_vel", "ang_vel"]` | 选择 velocimeter(线速度)或 gyro(角速度)。 | +| `bias_range` | `tuple[float, float] \| None` | 每 env 均匀偏置,`reset` 时重采样。 | + +velocimeter 计算公式为 `v_site = v_link + ω × (R_link · offset)`,再旋转到 body frame —— +与 MuJoCo 的杠杆臂约定一致,使得 GeneLab 训出的策略读到的信号与 mjlab 参考实现完全相同。 + +### ContactSensor + +按 link 名汇总 `robot.get_links_net_contact_force()` 输出。开启 `track_air_time=True` 时, +一个内部状态机会推进每 env 的 `current_air_time` / `current_contact_time`,并在接触状态翻转 +的瞬间把已完成的时长快照到 `last_air_time` / `last_contact_time`。 + +| 字段 | 类型 | 含义 | +|------|------|------| +| `link_names` | `tuple[str, ...]` | 显式列出需要监控的 link。 | +| `link_names_expr` | `str \| None` | 对 `env.link_names` 做正则匹配。 | +| `force_threshold` | `float` | 力幅值(N)大于此值时 `found` 为真。 | +| `track_air_time` | `bool` | 是否分配 air-time / contact-time 状态机。 | + +`data` 是 `ContactData` 数据类,包含 `force`、`force_norm`、`found` 以及四个 air-time 缓冲。 +配套的观测 term —— `mdp.foot_air_time`、`mdp.foot_contact`、`mdp.foot_contact_forces` —— +直接读取这个数据类。 + +### TerrainHeightSensor + +锚定在某个 link 上的 2D 下射光线网格,输出每条光线相对地形的高度(正值表示在地形上方), +适合作为 critic 的 privileged `height_scan` 观测。默认 backend 把每条光线打到位于 +`ground_height` 的水平面上;非平地场景的扩展点是继承 `RayCastSensor` 并重载 +`_intersect_world_rays`。 + +| 字段 | 类型 | 含义 | +|------|------|------| +| `link_name` | `str` | 网格原点锚定的 link。 | +| `pattern` | `GridPattern` | 网格分辨率 / 尺寸 / 方向。 | +| `attach_yaw_only` | `bool` | 仅按 yaw 旋转网格,使其保持水平。 | +| `max_distance` | `float` | 光线距离上限。 | +| `ground_height` | `float` | 默认平面 backend 使用的地面高度。 | + +## 给观测加噪 + +`ObservationTermCfg.noise` 接收加性噪声模型(`Unoise(n_min, n_max)` 或 +`Gnoise(mean, std)`)。`ObservationGroupCfg.enable_corruption` 控制是否真正加噪 —— 默认关闭, +让 critic 看到 ground truth。逐 term 的管线顺序是 **noise → scale → clip**,因此噪声幅值 +定义在原始信号空间:在带 `scale=0.05` 的原始 `joint_vel` 上加 `Unoise(-1.5, 1.5)`,最终的 +抖动量为 ±0.075。 + +惯例做法是 policy / critic 共享 term,只在 `enable_corruption` 上分叉: + +```python +from genelab import mdp +from genelab.managers import ObservationGroupCfg, ObservationTermCfg +from genelab.mdp.noise import Unoise + + +def _obs_terms() -> dict[str, ObservationTermCfg]: + return { + "base_lin_vel": ObservationTermCfg( + func=mdp.sensor_data, + params={"sensor_name": "imu_lin_vel"}, + noise=Unoise(-0.5, 0.5), + ), + "joint_vel": ObservationTermCfg( + func=mdp.joint_vel_rel, + scale=0.05, + noise=Unoise(-1.5, 1.5), + ), + } + + +policy = ObservationGroupCfg(enable_corruption=True, terms=_obs_terms()) +critic = ObservationGroupCfg(enable_corruption=False, terms=_obs_terms()) +``` + +## 自定义传感器 + +继承 `Sensor[T]` 并指定返回类型,实现 `_compute_data`。如需缓存 link 索引、推进积分器或 +清空 env 状态,重载 `bind` / `update` / `reset` 并先调用 `super()`,让缓存失效链路保持完整。 + +```python +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import torch + +from genelab.sensor import Sensor, SensorCfg + +if TYPE_CHECKING: + from genelab.envs.manager_based_rl_env import ManagerBasedRlEnv + + +@dataclass +class JointTorqueSensorCfg(SensorCfg): + def build(self) -> "JointTorqueSensor": + return JointTorqueSensor(self) + + +class JointTorqueSensor(Sensor[torch.Tensor]): + def bind(self, env: "ManagerBasedRlEnv") -> None: + super().bind(env) + # 在这里缓存依赖 env.link_names / env.joint_names 的索引。 + + def _compute_data(self) -> torch.Tensor: + assert self._env is not None + rs = self._env.robot_state + return self._env.joint_kp * (rs.joint_pos - self._env.default_joint_pos) +``` + +把 cfg 加入 `SceneCfg.sensors` 后,传感器即可通过 `env.sensors[name].data` 访问。 + +## See also + +- [配置系统](configs.md) +- [API 参考](../api/reference.md) diff --git a/mkdocs.yml b/mkdocs.yml index 5cc3fd6f..dc596d37 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -97,6 +97,7 @@ plugins: Concepts: 核心概念 Registry: 注册表 Configs: 配置系统 + Sensors: 传感器 Extensions: 扩展加载 Examples: 示例 API Reference: API 参考 @@ -158,6 +159,7 @@ nav: - Concepts: - Registry: concepts/registry.md - Configs: concepts/configs.md + - Sensors: concepts/sensors.md - Extensions: concepts/extensions.md - Examples: - Overview: examples/overview.md diff --git a/src/genelab/mdp/observations.py b/src/genelab/mdp/observations.py index cab0015f..de7fe74a 100644 --- a/src/genelab/mdp/observations.py +++ b/src/genelab/mdp/observations.py @@ -53,7 +53,9 @@ def sensor_data(env: "ManagerBasedRlEnv", sensor_name: str) -> torch.Tensor: def _contact_sensor(env: "ManagerBasedRlEnv", sensor_name: str) -> ContactSensor: sensor = env.sensors[sensor_name] if not isinstance(sensor, ContactSensor): - raise TypeError(f"sensor {sensor_name!r} is not a ContactSensor (got {type(sensor).__name__})") + raise TypeError( + f"sensor {sensor_name!r} is not a ContactSensor (got {type(sensor).__name__})" + ) return sensor diff --git a/tests/test_sensor.py b/tests/test_sensor.py index 2f30050a..6cf4372b 100644 --- a/tests/test_sensor.py +++ b/tests/test_sensor.py @@ -217,7 +217,7 @@ def test_body_velocity_sensor_bias_randomizes_on_reset() -> None: ).build() sensor.bind(env) first = sensor.data.clone() - assert ((first.abs() <= 0.1 + 1e-6).all()) + assert (first.abs() <= 0.1 + 1e-6).all() assert first.std() > 0.0 # initial bias should already vary across envs sensor.reset(torch.arange(64)) second = sensor.data @@ -284,9 +284,7 @@ def test_contact_sensor_explicit_link_names_resolves_indices() -> None: def test_contact_sensor_regex_match() -> None: env = _FakeContactEnv(num_envs=1, link_names=("base", "left_foot", "right_foot")) - sensor = ContactSensorCfg( - name="f", link_names_expr=r"_foot$", track_air_time=False - ).build() + sensor = ContactSensorCfg(name="f", link_names_expr=r"_foot$", track_air_time=False).build() sensor.bind(env) assert sensor.link_names == ["left_foot", "right_foot"]