make every controller consume lowstate for obs

This commit is contained in:
Martino Russi
2026-07-31 10:31:11 +02:00
parent 89ac5b4408
commit 1799eceb22
2 changed files with 17 additions and 35 deletions
@@ -199,36 +199,21 @@ class SonicDecoder:
assert off == 994, f"Decoder obs mismatch: {off}" assert off == 994, f"Decoder obs mismatch: {off}"
return obs return obs
def step(self, robot_obs, token): def step(self, lowstate, token):
"""One control tick: read robot obs, decode the supplied token -> joint targets. """One control tick: read robot lowstate, decode the supplied token -> joint targets.
Args: Args:
robot_obs: dict with ``<joint>.q``/``.dq`` and ``imu.*`` fields. lowstate: Unitree lowstate with ``motor_state[i].q/.dq`` and ``imu_state`` fields.
token: 64-D latent supplied by the policy (encoder bypassed). token: 64-D latent supplied by the policy (encoder bypassed).
Returns: Returns:
dict of ``<joint>.q`` target positions (rad) in IsaacLab joint order. dict of ``<joint>.q`` target positions (rad) in IsaacLab joint order.
""" """
self.token = np.asarray(token, np.float32) self.token = np.asarray(token, np.float32)
jnames = [m.name for m in G1_29_JointIndex] q = np.array([lowstate.motor_state[m.value].q for m in G1_29_JointIndex], np.float32)
q = np.array( dq = np.array([lowstate.motor_state[m.value].dq for m in G1_29_JointIndex], np.float32)
[ quat = np.array(lowstate.imu_state.quaternion, np.float32) # (w, x, y, z)
robot_obs.get(f"{n}.q", self.default_angles[m.value]) ang = np.array(lowstate.imu_state.gyroscope, np.float32)
for m, n in zip(G1_29_JointIndex, jnames, strict=False)
],
np.float32,
)
dq = np.array([robot_obs.get(f"{n}.dq", 0.0) for n in jnames], np.float32)
quat = np.array(
[
robot_obs.get("imu.quat.w", 1),
robot_obs.get("imu.quat.x", 0),
robot_obs.get("imu.quat.y", 0),
robot_obs.get("imu.quat.z", 0),
],
np.float32,
)
ang = np.array([robot_obs.get(f"imu.gyro.{a}", 0) for a in "xyz"], np.float32)
self.update_history(q, dq, ang, quat) self.update_history(q, dq, ang, quat)
action_mj = ( action_mj = (
self.decoder.run(None, {self.decoder_input: self.build_decoder_obs().reshape(1, -1)})[0] self.decoder.run(None, {self.decoder_input: self.build_decoder_obs().reshape(1, -1)})[0]
@@ -294,7 +279,7 @@ class SonicWholeBodyController:
token = self._last_token if self._last_token is not None else np.zeros(TOKEN_DIM, dtype=np.float32) token = self._last_token if self._last_token is not None else np.zeros(TOKEN_DIM, dtype=np.float32)
return {token_state_key(i): float(v) for i, v in enumerate(token)} return {token_state_key(i): float(v) for i, v in enumerate(token)}
def _startup_blend(self, obs: dict, out: dict) -> dict: def _startup_blend(self, lowstate, out: dict) -> dict:
"""Ease into policy control at startup: for the first ``INIT_RAMP_S`` seconds, """Ease into policy control at startup: for the first ``INIT_RAMP_S`` seconds,
interpolate between the robot's pose captured on the first tick and the policy's interpolate between the robot's pose captured on the first tick and the policy's
live commanded target, so the handoff has no snap. live commanded target, so the handoff has no snap.
@@ -307,8 +292,7 @@ class SonicWholeBodyController:
if self._init_step == 0: if self._init_step == 0:
# Capture the robot's actual pose as the interpolation start point. # Capture the robot's actual pose as the interpolation start point.
self._start_pose = { self._start_pose = {
f"{m.name}.q": float(obs.get(f"{m.name}.q", self._default_angles[m.value])) f"{m.name}.q": float(lowstate.motor_state[m.value].q) for m in G1_29_JointIndex
for m in G1_29_JointIndex
} }
self._init_step += 1 self._init_step += 1
ratio = min(1.0, self._init_step / self._init_ramp_steps) ratio = min(1.0, self._init_step / self._init_ramp_steps)
@@ -320,8 +304,8 @@ class SonicWholeBodyController:
logger.info("SONIC startup blend complete -> full policy control") logger.info("SONIC startup blend complete -> full policy control")
return blended return blended
def run_step(self, action: dict, obs: dict) -> dict: def run_step(self, action: dict, lowstate) -> dict:
if not obs: if lowstate is None:
return {} return {}
# Token-only interface (token-output VLA): a dense 64-D ``motion_token.{i}`` command # Token-only interface (token-output VLA): a dense 64-D ``motion_token.{i}`` command
@@ -335,7 +319,7 @@ class SonicWholeBodyController:
self._last_token = self._neutral_token.copy() self._last_token = self._neutral_token.copy()
# Either a fresh token this tick or the last one received (held between the ~30 Hz # Either a fresh token this tick or the last one received (held between the ~30 Hz
# token stream and the ~50 Hz control loop). # token stream and the ~50 Hz control loop).
return self._startup_blend(obs, self.controller.step(obs, self._last_token)) return self._startup_blend(lowstate, self.controller.step(lowstate, self._last_token))
def reset(self): def reset(self):
self.controller.reset() self.controller.reset()
+5 -7
View File
@@ -241,12 +241,14 @@ class UnitreeG1(Robot):
if self.controller is None: if self.controller is None:
return {f"{G1_29_JointIndex(motor).name}.q": float for motor in G1_29_JointIndex} return {f"{G1_29_JointIndex(motor).name}.q": float for motor in G1_29_JointIndex}
# Controllers may define their own action space (e.g. SONIC's 64-D latent token). # Whole-body controllers (SONIC): 64-D latent token.
controller_ft = getattr(self.controller, "action_features", None) controller_ft = getattr(self.controller, "action_features", None)
if controller_ft is not None: if controller_ft is not None:
return dict(controller_ft) return dict(controller_ft)
# Locomotion controllers (GR00T / Holosoma): arm joint targets + joystick axes. # Locomotion controllers (GR00T / Holosoma): arm joint targets + joystick axes.
# TODO: have GR00T/Holosoma advertise their own action_features too, so every
# controller declares its action space and this fallthrough can be dropped.
arm_features = {f"{G1_29_JointArmIndex(motor).name}.q": float for motor in G1_29_JointArmIndex} arm_features = {f"{G1_29_JointArmIndex(motor).name}.q": float for motor in G1_29_JointArmIndex}
remote_features = dict.fromkeys(REMOTE_AXES, float) remote_features = dict.fromkeys(REMOTE_AXES, float)
return {**arm_features, **remote_features} return {**arm_features, **remote_features}
@@ -278,12 +280,8 @@ class UnitreeG1(Robot):
with self._controller_action_lock: with self._controller_action_lock:
controller_input = dict(self.controller_input) controller_input = dict(self.controller_input)
# Full-body controllers (SONIC) consume the full observation dict; others # Run controller step
# take the raw lowstate. get_observation() is the single lowstate -> obs builder. controller_action = self.controller.run_step(controller_input, lowstate)
controller_state = (
self.get_observation() if getattr(self.controller, "full_body", False) else lowstate
)
controller_action = self.controller.run_step(controller_input, controller_state)
# Write controller output snapshot # Write controller output snapshot
with self._controller_action_lock: with self._controller_action_lock: