Source code for bigym.loco.base

"""LowerBodyBase — the fat base class implementing the shared pipeline.

Every concrete backend deals with the same plumbing: resolving joint
addresses by name, reading the floating base and IMU-style quantities,
applying an anchor pose at reset, clipping commands to training ranges, and
snapshotting mutable state for bit-exact demo replay. This class implements
all of it so a new backend only provides:

- a policy loader (``__init__``),
- ``step()`` (build obs -> run policy -> joint position targets),
- declarations: ``controlled_joints``, ``command_spec``, and either the
  declarative ``STATEFUL`` mapping or custom ``get_state``/``set_state``.

Where an adapter's behavior deviates from the base defaults (e.g. a
controller that stores raw commands and clips on read), the adapter
overrides the base method rather than the base absorbing the quirk.
"""

from __future__ import annotations

from typing import TYPE_CHECKING, Any, Optional

import mujoco
import numpy as np

from bigym.loco.command import CommandSpec
from bigym.loco.controller import OutputSpec


def mujoco_basename(full_identifier: str | None) -> str:
    """Strip the attachment namespace from a scene element name."""
    if not full_identifier:
        return ""
    return full_identifier.split("/")[-1]


def quat_rotate_inverse_wxyz(q_wxyz: np.ndarray, v_xyz: np.ndarray) -> np.ndarray:
    """Rotate vector ``v`` by the inverse of quaternion ``q`` (wxyz order)."""
    w, x, y, z = (
        float(q_wxyz[0]),
        float(q_wxyz[1]),
        float(q_wxyz[2]),
        float(q_wxyz[3]),
    )
    rx, ry, rz = -x, -y, -z

    tx = 2.0 * (ry * v_xyz[2] - rz * v_xyz[1])
    ty = 2.0 * (rz * v_xyz[0] - rx * v_xyz[2])
    tz = 2.0 * (rx * v_xyz[1] - ry * v_xyz[0])

    out_x = v_xyz[0] + w * tx + (ry * tz - rz * ty)
    out_y = v_xyz[1] + w * ty + (rz * tx - rx * tz)
    out_z = v_xyz[2] + w * tz + (rx * ty - ry * tx)
    return np.array([out_x, out_y, out_z], dtype=np.float32)


def yaw_from_quat_wxyz(quat: np.ndarray) -> float:
    """Yaw angle in radians from a wxyz quaternion."""
    w, x, y, z = (float(quat[0]), float(quat[1]), float(quat[2]), float(quat[3]))
    siny_cosp = 2.0 * (w * z + x * y)
    cosy_cosp = 1.0 - 2.0 * (y * y + z * z)
    return float(np.arctan2(siny_cosp, cosy_cosp))


def yaw_from_qpos(qpos: np.ndarray) -> float:
    """Yaw from a floating-base qpos: free joint (>=7) or slide/hinge stack."""
    if qpos.size >= 7:
        return yaw_from_quat_wxyz(qpos[3:7])
    if qpos.size:
        return float(qpos[-1])
    return 0.0


[docs] class LowerBodyBase: """Shared pipeline for lower-body backends. Subclasses set in ``__init__``: ``controlled_joints``, ``command_spec``, ``control_dt``, ``controlled_range_low``/``controlled_range_high`` (from :meth:`build_joint_ranges`), ``velocity_clip``/``yaw_rate_clip``, and the latched ``command``, ``height_command`` and ``last_action``. """ # Declarative replay state: snapshot key -> attribute name. The base # get_state()/set_state() stack and restore them, so a forgotten field is # a loud KeyError instead of a silent divergence. STATEFUL: dict[str, str] = {} # Post-reset settle steps the env runs before handing control to the # agent when the config does not set reset_warmup_steps. Measured per # checkpoint; adapters override. Demos are recorded from the post-settle # engage moment, so eval must settle too — 0 is only correct for callers # that manage settling themselves. recommended_reset_warmup_steps: int = 0 def weight_files(self) -> tuple[str, ...]: """Absolute paths of the policy weight files this controller loaded. Used by ``substrate_fingerprint()`` to record a SHA-256 per weight file, so a leaderboard record pins the exact lower-body policy. Adapters override; the default (no files) is only right for analytic controllers. """ return () _env: Any _pelvis_body_id: Optional[int] = None #: Joints whose position targets this backend owns. controlled_joints: tuple[str, ...] #: The backend's typed command declaration and its bounds. command_spec: CommandSpec #: Seconds between ``step()`` calls. control_dt: float #: Symmetric clip on ``vx``/``vy``; 0 disables it. velocity_clip: float #: Symmetric clip on ``wz``; 0 disables it. yaw_rate_clip: float #: Lower joint limits of ``controlled_joints``. controlled_range_low: np.ndarray #: Upper joint limits of ``controlled_joints``. controlled_range_high: np.ndarray #: The latched ``[vx, vy, wz]`` command. command: np.ndarray #: The latched height command in meters. height_command: float #: The policy's most recent raw action. last_action: np.ndarray if TYPE_CHECKING: def step(self) -> np.ndarray: """Advance the policy one control step; backends implement it.""" ... # ------------------------------------------------------------------ # Contract properties # ------------------------------------------------------------------ @property def output_spec(self) -> OutputSpec: """Names and bounds of the joint targets produced by ``step()``.""" return OutputSpec( joint_names=self.controlled_joints, low=np.asarray(self.controlled_range_low, dtype=np.float32).copy(), high=np.asarray(self.controlled_range_high, dtype=np.float32).copy(), ) def get_command(self) -> np.ndarray: """A copy of the latched twist command ``[vx, vy, wz]``.""" return self.command.astype(np.float32, copy=True) def get_height_command(self) -> float: """The latched height command in meters.""" return float(self.height_command) def get_last_action(self) -> np.ndarray: """A copy of the lower-body policy's most recent raw action.""" return self.last_action.astype(np.float32, copy=True) # ------------------------------------------------------------------ # Failure detection # ------------------------------------------------------------------ # Upright projected gravity is [0, 0, -1]; proj_g_z above this means the # base tilted past ~53 degrees — backend-independent (fallen is fallen). _FAILED_MAX_PROJ_G_Z = -0.6 # Height backstop: a commanded squat at the bottom of the backend's # training height range is NOT a failure, so the threshold derives from # command_spec (height.low - margin) rather than one fixed constant. # The margin absorbs pelvis-vs-commanded-height tracking slack; an # actually fallen robot is far below it and also trips the tilt check. _FAILED_HEIGHT_MARGIN = 0.10 # Fallback for specs without a height field (twist-only backends). _FAILED_MIN_BASE_HEIGHT = 0.35 def _failed_min_base_height(self) -> float: spec = self.command_spec if spec.has("height"): return float(spec.field("height").low) - self._FAILED_HEIGHT_MARGIN return self._FAILED_MIN_BASE_HEIGHT
[docs] def is_failed(self) -> bool: """Fallen / tipped-over detection from base height and tilt.""" _, _, proj_g = self.get_base_obs() if float(proj_g[2]) > self._FAILED_MAX_PROJ_G_Z: return True pelvis_z = self._pelvis_height() if pelvis_z is not None and pelvis_z < self._failed_min_base_height(): return True return False
def _pelvis_height(self) -> Optional[float]: pelvis_body_id = self._pelvis_body_id if pelvis_body_id is None: model = self._env.model pelvis_body_id = -1 for candidate in range(int(model.nbody)): name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_BODY, candidate) if mujoco_basename(name) == "pelvis": pelvis_body_id = int(candidate) break self._pelvis_body_id = pelvis_body_id if pelvis_body_id < 0: return None return float(self._env.data.xpos[pelvis_body_id][2]) # ------------------------------------------------------------------ # Command handling # ------------------------------------------------------------------ def _clip_twist( self, cmd_vx: float, cmd_vy: float, cmd_wz: float ) -> tuple[float, float, float]: cmd_clip = float(self.velocity_clip) wz_clip = float(self.yaw_rate_clip) vx = ( float(np.clip(cmd_vx, -cmd_clip, cmd_clip)) if cmd_clip > 0 else float(cmd_vx) ) vy = ( float(np.clip(cmd_vy, -cmd_clip, cmd_clip)) if cmd_clip > 0 else float(cmd_vy) ) wz = float(np.clip(cmd_wz, -wz_clip, wz_clip)) if wz_clip > 0 else float(cmd_wz) return vx, vy, wz
[docs] def set_command( self, cmd_vx: float, cmd_vy: float, cmd_wz: float, *, height: Optional[float] = None, torso_pitch: Optional[float] = None, ) -> None: """Default clip-on-set semantics (groot_wbc family). Keyword names match the ``command_spec`` field names. """ vx, vy, wz = self._clip_twist(cmd_vx, cmd_vy, cmd_wz) self.command = np.asarray([vx, vy, wz], dtype=np.float32) if height is not None: self.height_command = float(height) if torso_pitch is not None: self._set_pitch_command(float(torso_pitch))
def _set_pitch_command(self, pitch_cmd: float) -> None: """Hook for backends with a torso-pitch channel; default: ignore.""" # ------------------------------------------------------------------ # Joint / sensor resolution # ------------------------------------------------------------------ def _joint_name_to_id(self) -> dict[str, int]: model = self._env.model out: dict[str, int] = {} for jid in range(int(model.njnt)): name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_JOINT, jid) out[mujoco_basename(name)] = jid return out
[docs] def build_joint_addresses( self, joint_names: tuple[str, ...] ) -> tuple[np.ndarray, np.ndarray]: """The qpos and dof addresses of ``joint_names``, in order. Raises: ValueError: A joint is missing from the model. """ model = self._env.model joint_name_to_id = self._joint_name_to_id() qposadr = [] dofadr = [] for joint_name in joint_names: if joint_name not in joint_name_to_id: raise ValueError( f"Missing joint '{joint_name}' required by {type(self).__name__}." ) jid = int(joint_name_to_id[joint_name]) qposadr.append(int(model.jnt_qposadr[jid])) dofadr.append(int(model.jnt_dofadr[jid])) return np.asarray(qposadr, dtype=np.int32), np.asarray(dofadr, dtype=np.int32)
[docs] def build_joint_ranges( self, joint_names: tuple[str, ...] ) -> tuple[np.ndarray, np.ndarray]: """The lower and upper limits of ``joint_names``; unlimited joints get +-inf.""" model = self._env.model joint_name_to_id = self._joint_name_to_id() lows = [] highs = [] for joint_name in joint_names: jid = int(joint_name_to_id[joint_name]) if int(model.jnt_limited[jid]): lows.append(float(model.jnt_range[jid, 0])) highs.append(float(model.jnt_range[jid, 1])) else: lows.append(-np.inf) highs.append(np.inf) return np.asarray(lows, dtype=np.float32), np.asarray(highs, dtype=np.float32)
[docs] def find_sensor(self, sensor_name: str, *, dim: int) -> Optional[int]: """The ``sensordata`` address of the ``dim``-wide sensor ``sensor_name``, or None.""" model = self._env.model for sid in range(int(model.nsensor)): name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_SENSOR, sid) if mujoco_basename(name) != sensor_name: continue if int(model.sensor_dim[sid]) != int(dim): continue return int(model.sensor_adr[sid]) return None
# ------------------------------------------------------------------ # Reset pose application # ------------------------------------------------------------------
[docs] def apply_pose( self, qpos_addresses: np.ndarray, dof_addresses: np.ndarray, targets: np.ndarray, joint_names: tuple[str, ...], ) -> None: """Write a joint pose + matching actuator targets and re-forward.""" model = self._env.model data = self._env.data for i, qadr in enumerate(qpos_addresses): data.qpos[int(qadr)] = float(targets[i]) for adr in dof_addresses: adr = int(adr) if 0 <= adr < int(model.nv): data.qvel[adr] = 0.0 data.qacc[adr] = 0.0 target_by_joint = { name: float(targets[i]) for i, name in enumerate(joint_names) } for aid in range(int(model.nu)): act_name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_ACTUATOR, aid) joint_name = mujoco_basename(act_name) if joint_name not in target_by_joint: continue target = target_by_joint[joint_name] if int(model.actuator_ctrllimited[aid]): low = float(model.actuator_ctrlrange[aid, 0]) high = float(model.actuator_ctrlrange[aid, 1]) target = float(np.clip(target, min(low, high), max(low, high))) data.ctrl[aid] = target mujoco.mj_forward(model, data)
# ------------------------------------------------------------------ # Declarative replay state # ------------------------------------------------------------------
[docs] def get_state(self) -> dict[str, np.ndarray]: """Snapshot every attribute named in ``STATEFUL``. Scalars become 0-d float32 arrays, arrays are float32 copies — matching the recorded snapshot format so demo npz files stay interchangeable. """ state: dict[str, np.ndarray] = {} for key, attribute in self.STATEFUL.items(): value = getattr(self, attribute) if isinstance(value, np.ndarray): state[key] = value.astype(np.float32, copy=True) else: state[key] = np.asarray(value, dtype=np.float32) return state
[docs] def set_state(self, state: dict[str, np.ndarray]) -> None: """Restore a snapshot; missing keys fall back to _state_default(). Restoration follows ``STATEFUL`` order, an ordering contract adapters may rely on: a ``_state_default`` for a derived field (e.g. a slewed height that equals the raw command) must come after the field it derives from. """ for key, attribute in self.STATEFUL.items(): if key in state: raw = state[key] else: raw = self._state_default(key) current = getattr(self, attribute) if isinstance(current, np.ndarray): setattr( self, attribute, np.asarray(raw, dtype=np.float32) .reshape(np.asarray(current).shape) .copy(), ) elif isinstance(current, bool): setattr(self, attribute, bool(np.asarray(raw).reshape(-1)[0])) else: setattr( self, attribute, float(np.asarray(raw, dtype=np.float64).reshape(-1)[0]), )
def _state_default(self, key: str): """Value for a snapshot key the snapshot lacks. Default: a KeyError, since a silently defaulted stateful field makes replay diverge. Adapters override it for keys whose absence has one meaning (e.g. a snapshot without a pitch channel restores pitch 0). """ raise KeyError( f"Snapshot is missing stateful key {key!r} required by " f"{type(self).__name__}." ) # ------------------------------------------------------------------ # Base observation (adapter-specific frames; must be provided) # ------------------------------------------------------------------
[docs] def get_base_obs(self) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Return (base_lin_vel, base_ang_vel, projected_gravity), base frame.""" raise NotImplementedError