Source code for bigym.loco.controller
"""The lower-body controller contract (thin required core).
``LowerBodyController`` is the *entire* surface the BiGym env integration
relies on. Anything else a concrete adapter exposes is an implementation
detail. New backends normally subclass :class:`bigym.loco.base.LowerBodyBase`
(the fat base class that absorbs joint addressing, reset pose application,
failure detection and declarative replay state) and provide a policy loader,
an obs builder and a policy step — a few hundred lines for a real backend —
but any object satisfying this protocol plugs in.
Contract summary:
- ``controlled_joints`` joints whose position targets this backend owns.
- ``command_spec`` typed command declaration (see loco.command).
- ``control_dt`` seconds between ``step()`` calls (1/control_hz).
- ``output_spec`` names + bounds of the produced targets.
- ``reset()`` re-anchor to the backend's init pose, clear state.
- ``set_command(...)`` latch the current command (VELOCITY kind channels).
- ``step()`` run the policy once -> joint position targets, in
``controlled_joints`` order.
- ``is_failed()`` has the robot fallen / left the recoverable region.
- ``get_state()/set_state()`` every mutable field that shapes future
targets, for bit-exact mid-episode save/restore
(demo replay). MuJoCo qpos/qvel/ctrl/qacc_warmstart
are snapshotted by the caller; this covers the
controller-side remainder (obs histories, last
actions, rate-limiter anchors, gait clocks...).
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional, Protocol, runtime_checkable
import numpy as np
from bigym.loco.command import CommandSpec
[docs]
@dataclass(frozen=True)
class OutputSpec:
"""Names and bounds of the joint position targets a backend produces."""
joint_names: tuple[str, ...]
low: np.ndarray
high: np.ndarray
@property
def dim(self) -> int:
"""Number of joints in this output spec."""
return len(self.joint_names)
[docs]
@runtime_checkable
class LowerBodyController(Protocol):
"""Thin required core every lower-body backend implements."""
@property
def controlled_joints(self) -> tuple[str, ...]:
"""Joints whose position targets this backend owns (incl. waist)."""
...
@property
def command_spec(self) -> CommandSpec:
"""Typed command declaration and its bounds."""
...
@property
def control_dt(self) -> float:
"""Seconds between step() calls."""
...
@property
def output_spec(self) -> OutputSpec:
"""Names/bounds of the produced targets."""
...
[docs]
def reset(self) -> None:
"""Re-apply the anchor pose and clear mutable state."""
...
[docs]
def set_command(
self,
cmd_vx: float,
cmd_vy: float,
cmd_wz: float,
*,
height: Optional[float] = None,
torso_pitch: Optional[float] = None,
) -> None:
"""Latch the VELOCITY-kind command channels.
Keyword names match the ``command_spec`` field names (``height`` /
``torso_pitch``). Channels absent from ``command_spec`` are ignored.
Non-velocity command kinds (EE_POSE / MOTION_REF) will extend this
surface as a pure addition; velocity backends stay untouched.
"""
...
[docs]
def step(self) -> np.ndarray:
"""Run the policy once; return targets in controlled_joints order."""
...
[docs]
def is_failed(self) -> bool:
"""True when the base has fallen / tipped beyond recovery."""
...
[docs]
def get_state(self) -> dict[str, np.ndarray]:
"""Snapshot every mutable field that shapes future targets."""
...
[docs]
def set_state(self, state: dict[str, np.ndarray]) -> None:
"""Restore a snapshot produced by get_state()."""
...
# ------------------------------------------------------------------
# Introspection the env integration also relies on. LowerBodyBase
# implements all of these (get_base_obs excepted — frame conventions
# are backend-specific), so subclassing the base satisfies the whole
# protocol; a from-scratch implementation must provide them too.
# ------------------------------------------------------------------
[docs]
def get_base_obs(self) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""(base_lin_vel, base_ang_vel, projected_gravity) in the base frame."""
...
[docs]
def get_command(self) -> np.ndarray:
"""Current [vx, vy, wz] command (clipped view)."""
...
[docs]
def get_height_command(self) -> float:
"""Current height command in meters."""
...
[docs]
def get_last_action(self) -> np.ndarray:
"""Raw policy action from the most recent step()."""
...