mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 13:09:40 +00:00
Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 40a5e70352 | |||
| 0cef9cd197 | |||
| 643ffb4785 | |||
| d59505a735 | |||
| 6ac95363b0 | |||
| ede1fc2978 | |||
| 49d5ea49bc | |||
| d23b65416f |
@@ -59,6 +59,7 @@ The `lerobot-rollout --strategy.type=dagger` mode requires **teleoperators with
|
|||||||
|
|
||||||
- `bi_openarm_mini` - Bimanual OpenArm Mini
|
- `bi_openarm_mini` - Bimanual OpenArm Mini
|
||||||
- `so_leader` - SO100 / SO101 leader arm
|
- `so_leader` - SO100 / SO101 leader arm
|
||||||
|
- `bi_so_leader` - Bimanual SO100 / SO101 leader arms
|
||||||
|
|
||||||
> [!IMPORTANT]
|
> [!IMPORTANT]
|
||||||
> The provided commands default to `bi_openarm_follower` + `bi_openarm_mini`.
|
> The provided commands default to `bi_openarm_follower` + `bi_openarm_mini`.
|
||||||
|
|||||||
@@ -338,7 +338,7 @@ It is advisable to install one 3-pin cable in the motor after placing them befor
|
|||||||
<hfoption id="Leader">
|
<hfoption id="Leader">
|
||||||
|
|
||||||
- Mount the leader holder onto the wrist and secure it with 4 M3x6mm screws.
|
- Mount the leader holder onto the wrist and secure it with 4 M3x6mm screws.
|
||||||
- Attach the handle to motor 5 using 1 M2x6mm screw.
|
- Attach the handle to the leader holder using 1 M2x6mm screw.
|
||||||
- Insert the gripper motor, secure it with 2 M2x6mm screws on each side, attach a motor horn using a M3x6mm horn screw.
|
- Insert the gripper motor, secure it with 2 M2x6mm screws on each side, attach a motor horn using a M3x6mm horn screw.
|
||||||
- Attach the follower trigger with 4 M3x6mm screws.
|
- Attach the follower trigger with 4 M3x6mm screws.
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -67,7 +67,7 @@ dependencies = [
|
|||||||
"einops>=0.8.0,<0.9.0",
|
"einops>=0.8.0,<0.9.0",
|
||||||
|
|
||||||
# Config & Hub
|
# Config & Hub
|
||||||
"draccus==0.10.0", # TODO: Relax version constraint
|
"draccus>=0.11.6,<0.12.0",
|
||||||
"huggingface-hub>=1.0.0,<2.0.0",
|
"huggingface-hub>=1.0.0,<2.0.0",
|
||||||
"requests>=2.32.0,<3.0.0",
|
"requests>=2.32.0,<3.0.0",
|
||||||
|
|
||||||
|
|||||||
@@ -163,8 +163,10 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
def _save_pretrained(self, save_directory: Path) -> None:
|
def _save_pretrained(self, save_directory: Path) -> None:
|
||||||
with open(save_directory / CONFIG_NAME, "w") as f, draccus.config_type("json"):
|
# Encode against the base class so draccus includes the choice "type" key,
|
||||||
draccus.dump(self, f, indent=4)
|
# which `from_pretrained` needs to resolve the concrete subclass.
|
||||||
|
with open(save_directory / CONFIG_NAME, "w") as f:
|
||||||
|
json.dump(draccus.encode(self, PreTrainedConfig), f, indent=4)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(
|
def from_pretrained(
|
||||||
|
|||||||
@@ -103,8 +103,10 @@ class RewardModelConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def _save_pretrained(self, save_directory: Path) -> None:
|
def _save_pretrained(self, save_directory: Path) -> None:
|
||||||
with open(save_directory / CONFIG_NAME, "w") as f, draccus.config_type("json"):
|
# Encode against the base class so draccus includes the choice "type" key,
|
||||||
draccus.dump(self, f, indent=4)
|
# which `from_pretrained` needs to resolve the concrete subclass.
|
||||||
|
with open(save_directory / CONFIG_NAME, "w") as f:
|
||||||
|
json.dump(draccus.encode(self, RewardModelConfig), f, indent=4)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(
|
def from_pretrained(
|
||||||
|
|||||||
@@ -194,7 +194,11 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if Path(config_path).resolve().exists():
|
if Path(config_path).resolve().exists():
|
||||||
policy_dir = Path(config_path).parent
|
# `config_path` may point at the checkpoint's train_config.json or at its
|
||||||
|
# pretrained_model/ directory (both documented above) — resolve either to
|
||||||
|
# the pretrained_model/ directory.
|
||||||
|
config_path_obj = Path(config_path)
|
||||||
|
policy_dir = config_path_obj.parent if config_path_obj.is_file() else config_path_obj
|
||||||
self.checkpoint_path = policy_dir.parent
|
self.checkpoint_path = policy_dir.parent
|
||||||
elif self.job.is_remote:
|
elif self.job.is_remote:
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -42,6 +42,9 @@ class Evo1Policy(PreTrainedPolicy):
|
|||||||
config_class = Evo1Config
|
config_class = Evo1Config
|
||||||
name = "evo1"
|
name = "evo1"
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def __init__(self, config: Evo1Config, *, vlm_hub_kwargs: dict | None = None, **kwargs):
|
def __init__(self, config: Evo1Config, *, vlm_hub_kwargs: dict | None = None, **kwargs):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
config.validate_features()
|
config.validate_features()
|
||||||
|
|||||||
@@ -68,6 +68,9 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
name = "groot"
|
name = "groot"
|
||||||
config_class = GrootConfig
|
config_class = GrootConfig
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def __init__(self, config: GrootConfig, **kwargs):
|
def __init__(self, config: GrootConfig, **kwargs):
|
||||||
"""Initialize Groot policy wrapper."""
|
"""Initialize Groot policy wrapper."""
|
||||||
require_package("transformers", extra="groot")
|
require_package("transformers", extra="groot")
|
||||||
|
|||||||
@@ -520,6 +520,9 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
config_class = MolmoAct2Config
|
config_class = MolmoAct2Config
|
||||||
name = "molmoact2"
|
name = "molmoact2"
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
return self.config.inference_action_mode == "continuous"
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: MolmoAct2Config,
|
config: MolmoAct2Config,
|
||||||
|
|||||||
@@ -749,6 +749,9 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
config_class = PI0Config
|
config_class = PI0Config
|
||||||
name = "pi0"
|
name = "pi0"
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: PI0Config,
|
config: PI0Config,
|
||||||
|
|||||||
@@ -714,6 +714,9 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
config_class = PI05Config
|
config_class = PI05Config
|
||||||
name = "pi05"
|
name = "pi05"
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: PI05Config,
|
config: PI05Config,
|
||||||
|
|||||||
@@ -249,6 +249,10 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
"""Whether this policy implements Real-Time Chunking inference semantics."""
|
||||||
|
return False
|
||||||
|
|
||||||
# TODO(aliberts, rcadene): split into 'forward' and 'compute_loss'?
|
# TODO(aliberts, rcadene): split into 'forward' and 'compute_loss'?
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict | None]:
|
||||||
|
|||||||
@@ -145,6 +145,9 @@ class SmolVLAPolicy(PreTrainedPolicy):
|
|||||||
config_class = SmolVLAConfig
|
config_class = SmolVLAConfig
|
||||||
name = "smolvla"
|
name = "smolvla"
|
||||||
|
|
||||||
|
def supports_rtc(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: SmolVLAConfig,
|
config: SmolVLAConfig,
|
||||||
|
|||||||
@@ -168,14 +168,23 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
last_layers.append(self.num_vlm_layers - 2)
|
last_layers.append(self.num_vlm_layers - 2)
|
||||||
frozen_layers = [
|
frozen_layers = [
|
||||||
"lm_head",
|
"lm_head",
|
||||||
"text_model.model.norm.weight",
|
"text_model.norm.weight",
|
||||||
]
|
]
|
||||||
for layer in last_layers:
|
for layer in last_layers:
|
||||||
frozen_layers.append(f"text_model.model.layers.{layer}.")
|
frozen_layers.append(f"text_model.layers.{layer}.")
|
||||||
|
|
||||||
|
unmatched_patterns = set(frozen_layers)
|
||||||
for name, params in self.vlm.named_parameters():
|
for name, params in self.vlm.named_parameters():
|
||||||
if any(k in name for k in frozen_layers):
|
matched_patterns = [k for k in frozen_layers if k in name]
|
||||||
|
if matched_patterns:
|
||||||
params.requires_grad = False
|
params.requires_grad = False
|
||||||
|
unmatched_patterns.difference_update(matched_patterns)
|
||||||
|
if unmatched_patterns:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Some frozen layer patterns matched no VLM parameters, so the corresponding layers "
|
||||||
|
"would silently remain trainable (parameter naming may have changed in transformers): "
|
||||||
|
f"{sorted(unmatched_patterns)}"
|
||||||
|
)
|
||||||
# To avoid unused params issue with distributed training
|
# To avoid unused params issue with distributed training
|
||||||
for name, params in self.lm_expert.named_parameters():
|
for name, params in self.lm_expert.named_parameters():
|
||||||
if "lm_head" in name:
|
if "lm_head" in name:
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import abc
|
import abc
|
||||||
import builtins
|
import builtins
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
@@ -78,8 +79,10 @@ class RLAlgorithmConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC):
|
|||||||
|
|
||||||
def _save_pretrained(self, save_directory: Path) -> None:
|
def _save_pretrained(self, save_directory: Path) -> None:
|
||||||
"""Serialize this config as ``config.json`` inside ``save_directory``."""
|
"""Serialize this config as ``config.json`` inside ``save_directory``."""
|
||||||
with open(save_directory / CONFIG_NAME, "w") as f, draccus.config_type("json"):
|
# Encode against the base class so draccus includes the choice "type" key,
|
||||||
draccus.dump(self, f, indent=4)
|
# which `from_pretrained` needs to resolve the concrete subclass.
|
||||||
|
with open(save_directory / CONFIG_NAME, "w") as f:
|
||||||
|
json.dump(draccus.encode(self, RLAlgorithmConfig), f, indent=4)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(
|
def from_pretrained(
|
||||||
|
|||||||
@@ -68,10 +68,6 @@ class UnitreeG1Config(RobotConfig):
|
|||||||
# Compensates for gravity on the unitree's arms using the arm ik solver
|
# Compensates for gravity on the unitree's arms using the arm ik solver
|
||||||
gravity_compensation: bool = False
|
gravity_compensation: bool = False
|
||||||
|
|
||||||
# Locomotion controller class name, e.g. "GrootLocomotionController",
|
# Lower-body controller class name, e.g. "GrootLocomotionController" or
|
||||||
# "HolosomaLocomotionController", or "SonicWholeBodyController". None disables it.
|
# "HolosomaLocomotionController". None disables it.
|
||||||
# Selecting "SonicWholeBodyController" implicitly switches the robot to the 64-D
|
|
||||||
# latent-token action/observation interface (``motion_token.{i}.pos`` action and a
|
|
||||||
# ``motion_token_state.{i}.pos`` state echo) so ``lerobot-rollout`` can drive a
|
|
||||||
# policy trained on SONIC motion tokens (e.g. nepyope/sonic_walk).
|
|
||||||
controller: str | None = None
|
controller: str | None = None
|
||||||
|
|||||||
@@ -1,27 +0,0 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""Unitree G1 locomotion controllers (Groot, Holosoma, SONIC)."""
|
|
||||||
|
|
||||||
from .gr00t_locomotion import GrootLocomotionController
|
|
||||||
from .holosoma_locomotion import HolosomaLocomotionController
|
|
||||||
from .sonic_whole_body import SonicWholeBodyController
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"GrootLocomotionController",
|
|
||||||
"HolosomaLocomotionController",
|
|
||||||
"SonicWholeBodyController",
|
|
||||||
]
|
|
||||||
@@ -1,360 +0,0 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""SONIC decoder whole-body controller for the Unitree G1 (token-only).
|
|
||||||
|
|
||||||
Pure-Python/ONNX re-implementation of the *decode* half of NVIDIA's SONIC deploy stack.
|
|
||||||
The encoder is intentionally absent: a token-output VLA (e.g. ``nepyope/sonic_walk``)
|
|
||||||
supplies the 64-D latent ``motion_token`` directly each tick, and the SONIC **decoder**
|
|
||||||
maps ``token + recent proprioception history`` to a residual action that is scaled and
|
|
||||||
added onto the standing pose (``default_angles``) to produce 50 Hz joint-position targets
|
|
||||||
for the robot's PD controller.
|
|
||||||
|
|
||||||
Index spaces: joints exist in two orderings — **IsaacLab** (policy/training order) and
|
|
||||||
**MuJoCo** (deploy order). ``ISAACLAB_TO_MUJOCO`` / ``MUJOCO_TO_ISAACLAB`` (in g1_utils)
|
|
||||||
convert between them. Quaternions are scalar-first ``(w, x, y, z)``.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
|
||||||
import logging
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import onnx
|
|
||||||
import onnxruntime as ort
|
|
||||||
from huggingface_hub import hf_hub_download
|
|
||||||
|
|
||||||
from ..g1_utils import (
|
|
||||||
ISAACLAB_TO_MUJOCO,
|
|
||||||
MUJOCO_TO_ISAACLAB,
|
|
||||||
G1_29_JointIndex,
|
|
||||||
get_gravity_orientation,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
# ── Constants (hardware-validated; see the NVIDIA SONIC deploy reference) ──────
|
|
||||||
CONTROL_DT = 0.02 # 50 Hz control period (s)
|
|
||||||
TOKEN_DIM = 64 # decoder latent size
|
|
||||||
|
|
||||||
# SONIC decoder checkpoint: NVIDIA's decoder ONNX re-packaged with its deploy constants
|
|
||||||
# (kp/kd PD gains, the standing pose default_angles, and the residual action_scale) embedded
|
|
||||||
# in the ONNX metadata; see upload_sonic_decoder.py for provisioning. The runtime loads the
|
|
||||||
# model *and* all of these straight from the checkpoint (the Holosoma convention), so no
|
|
||||||
# motor-physics math happens at deploy time.
|
|
||||||
DEFAULT_SONIC_REPO_ID = "lerobot/sonic_decoder"
|
|
||||||
DECODER_FILENAME = "model_decoder.onnx"
|
|
||||||
DECODER_INPUT_DIM = 994 # token(64) + 10-frame proprio history + gravity
|
|
||||||
|
|
||||||
|
|
||||||
def load_sonic_decoder(repo_id: str = DEFAULT_SONIC_REPO_ID):
|
|
||||||
"""Load the SONIC decoder ONNX and its baked-in deploy constants from the checkpoint.
|
|
||||||
|
|
||||||
Returns ``(decoder_session, kp, kd, default_angles, action_scale, neutral_token)``. The
|
|
||||||
gains/pose/scale are (29,) float32 in IsaacLab joint order and ``neutral_token`` is the
|
|
||||||
(64,) float32 idle latent -- all read from the ONNX ``metadata_props`` rather than
|
|
||||||
recomputed/hardcoded at deploy time (mirrors ``holosoma_locomotion.load_policy``).
|
|
||||||
"""
|
|
||||||
decoder_path = hf_hub_download(repo_id=repo_id, filename=DECODER_FILENAME)
|
|
||||||
so = ort.SessionOptions()
|
|
||||||
so.log_severity_level = 3 # quiet ORT logs
|
|
||||||
session = ort.InferenceSession(decoder_path, sess_options=so)
|
|
||||||
dec_dim = int(session.get_inputs()[0].shape[1])
|
|
||||||
if dec_dim != DECODER_INPUT_DIM:
|
|
||||||
raise RuntimeError(f"Unexpected decoder input dim {dec_dim} (expected {DECODER_INPUT_DIM})")
|
|
||||||
|
|
||||||
meta = {p.key: p.value for p in onnx.load(decoder_path, load_external_data=False).metadata_props}
|
|
||||||
required = ("kp", "kd", "default_angles", "action_scale", "neutral_token")
|
|
||||||
missing = [k for k in required if k not in meta]
|
|
||||||
if missing:
|
|
||||||
raise ValueError(
|
|
||||||
f"SONIC decoder ONNX at {repo_id} is missing metadata {missing}; "
|
|
||||||
"re-run upload_sonic_decoder.py to (re)provision the checkpoint."
|
|
||||||
)
|
|
||||||
arr = {k: np.array(json.loads(meta[k]), dtype=np.float32) for k in required}
|
|
||||||
logger.info("Loaded SONIC deploy constants from %s (%d joints)", repo_id, len(arr["kp"]))
|
|
||||||
return session, arr["kp"], arr["kd"], arr["default_angles"], arr["action_scale"], arr["neutral_token"]
|
|
||||||
|
|
||||||
|
|
||||||
# Action-feature prefix for the latent-token interface (see _extract_token_from_action).
|
|
||||||
TOKEN_ACTION_PREFIX = "motion_token" # nosec B105 - feature-key prefix, not a secret
|
|
||||||
# Proprio-state prefix for the token interface: the robot echoes the last commanded token
|
|
||||||
# here so ``lerobot-rollout`` aggregates it into a 64-D ``observation.state``.
|
|
||||||
TOKEN_STATE_PREFIX = "motion_token_state" # nosec B105 - feature-key prefix, not a secret
|
|
||||||
|
|
||||||
|
|
||||||
def token_action_key(i: int) -> str:
|
|
||||||
"""Action-dict key for the i-th component of the 64-D SONIC latent token.
|
|
||||||
|
|
||||||
The ``.pos`` suffix is required so the value flows through ``lerobot-rollout``, which
|
|
||||||
only routes ``.pos`` scalar features onto the policy action vector.
|
|
||||||
"""
|
|
||||||
return f"{TOKEN_ACTION_PREFIX}.{i}.pos"
|
|
||||||
|
|
||||||
|
|
||||||
def token_state_key(i: int) -> str:
|
|
||||||
"""Observation key for the i-th component of the 64-D SONIC latent token state."""
|
|
||||||
return f"{TOKEN_STATE_PREFIX}.{i}.pos"
|
|
||||||
|
|
||||||
|
|
||||||
# Startup blend duration: over the first control ticks, linearly interpolate every joint
|
|
||||||
# from the robot's initial measured pose into the policy's commanded target, so control
|
|
||||||
# eases in without a snap on the first command.
|
|
||||||
INIT_RAMP_S = 3.0
|
|
||||||
|
|
||||||
|
|
||||||
def _extract_token_from_action(action: dict | None) -> np.ndarray | None:
|
|
||||||
"""Reassemble a dense (64,) latent token from ``motion_token.{i}`` keys, or None.
|
|
||||||
|
|
||||||
The token-only interface: the caller supplies the 64-D encoder latent directly (e.g. a
|
|
||||||
token-output VLA's action), which the decoder consumes with the encoder bypassed.
|
|
||||||
Requires the full dense token; a partial one is ignored (returns None).
|
|
||||||
"""
|
|
||||||
if not action:
|
|
||||||
return None
|
|
||||||
keys = [token_action_key(i) for i in range(TOKEN_DIM)]
|
|
||||||
if any(key not in action for key in keys):
|
|
||||||
return None
|
|
||||||
return np.fromiter((float(action[key]) for key in keys), dtype=np.float32, count=TOKEN_DIM)
|
|
||||||
|
|
||||||
|
|
||||||
class SonicDecoder:
|
|
||||||
"""Runs the SONIC decoder ONNX model and owns the proprioception history.
|
|
||||||
|
|
||||||
Each tick it appends the latest robot state to 10-frame history buffers, then maps the
|
|
||||||
supplied 64-D ``token`` + that history to a residual action added onto ``default_angles``.
|
|
||||||
The encoder is bypassed entirely (token supplied by the policy). ``default_angles`` and
|
|
||||||
``action_scale`` are (29,) float32 in IsaacLab order, loaded from the checkpoint.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, decoder, default_angles, action_scale):
|
|
||||||
self.decoder = decoder
|
|
||||||
self.decoder_input = decoder.get_inputs()[0].name
|
|
||||||
self.default_angles = np.asarray(default_angles, np.float32)
|
|
||||||
self.action_scale = np.asarray(action_scale, np.float32)
|
|
||||||
self.default_angles_mj = self.default_angles[MUJOCO_TO_ISAACLAB]
|
|
||||||
self.token = np.zeros(TOKEN_DIM, np.float32)
|
|
||||||
self.last_action_mj = np.zeros(29, np.float32)
|
|
||||||
self.h_q_mj = [np.zeros(29, np.float32)] * 10
|
|
||||||
self.h_dq_mj = [np.zeros(29, np.float32)] * 10
|
|
||||||
self.h_ang = [np.zeros(3, np.float32)] * 10
|
|
||||||
self.h_act_mj = [np.zeros(29, np.float32)] * 10
|
|
||||||
self.h_quat = [np.array([1, 0, 0, 0], np.float32)] * 10
|
|
||||||
|
|
||||||
def reset(self):
|
|
||||||
"""Clear the token and 10-frame proprioception history.
|
|
||||||
|
|
||||||
``UnitreeG1.reset()`` relies on this so the first decoder outputs of a new episode
|
|
||||||
are not contaminated by the previous episode's state.
|
|
||||||
"""
|
|
||||||
self.token = np.zeros(TOKEN_DIM, np.float32)
|
|
||||||
self.last_action_mj = np.zeros(29, np.float32)
|
|
||||||
self.h_q_mj = [np.zeros(29, np.float32)] * 10
|
|
||||||
self.h_dq_mj = [np.zeros(29, np.float32)] * 10
|
|
||||||
self.h_ang = [np.zeros(3, np.float32)] * 10
|
|
||||||
self.h_act_mj = [np.zeros(29, np.float32)] * 10
|
|
||||||
self.h_quat = [np.array([1, 0, 0, 0], np.float32)] * 10
|
|
||||||
|
|
||||||
def update_history(self, q, dq, ang, quat):
|
|
||||||
"""Push the latest proprioception (pos/vel/gyro/orientation) into the 10-frame buffers."""
|
|
||||||
quat = quat / (np.linalg.norm(quat) + 1e-8)
|
|
||||||
# Reorder IsaacLab-order state into the MuJoCo order the decoder consumes. This
|
|
||||||
# permutation direction is validated against the deployed SONIC ONNX; don't flip it.
|
|
||||||
q_mj = q[MUJOCO_TO_ISAACLAB]
|
|
||||||
dq_mj = dq[MUJOCO_TO_ISAACLAB]
|
|
||||||
self.h_q_mj = [q_mj - self.default_angles_mj] + self.h_q_mj[:-1]
|
|
||||||
self.h_dq_mj = [dq_mj] + self.h_dq_mj[:-1]
|
|
||||||
self.h_ang = [ang.copy()] + self.h_ang[:-1]
|
|
||||||
self.h_act_mj = [self.last_action_mj.copy()] + self.h_act_mj[:-1]
|
|
||||||
self.h_quat = [quat.copy()] + self.h_quat[:-1]
|
|
||||||
|
|
||||||
def build_decoder_obs(self):
|
|
||||||
"""Assemble the 994-D decoder input: token + 10-frame proprioception history + gravity."""
|
|
||||||
obs = np.zeros(994, np.float32)
|
|
||||||
off = 0
|
|
||||||
obs[off : off + 64] = self.token
|
|
||||||
off += 64
|
|
||||||
for h, sz in [
|
|
||||||
(list(reversed(self.h_ang)), 3),
|
|
||||||
(list(reversed(self.h_q_mj)), 29),
|
|
||||||
(list(reversed(self.h_dq_mj)), 29),
|
|
||||||
(list(reversed(self.h_act_mj)), 29),
|
|
||||||
]:
|
|
||||||
for f in range(10):
|
|
||||||
obs[off : off + sz] = h[f]
|
|
||||||
off += sz
|
|
||||||
for q in reversed(self.h_quat):
|
|
||||||
obs[off : off + 3] = get_gravity_orientation(q)
|
|
||||||
off += 3
|
|
||||||
assert off == 994, f"Decoder obs mismatch: {off}"
|
|
||||||
return obs
|
|
||||||
|
|
||||||
def step(self, robot_obs, token, debug=False):
|
|
||||||
"""One control tick: read robot obs, decode the supplied token -> joint targets.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
robot_obs: dict with ``<joint>.q``/``.dq`` and ``imu.*`` fields.
|
|
||||||
token: 64-D latent supplied by the policy (encoder bypassed).
|
|
||||||
debug: log action/delta norms.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
dict of ``<joint>.q`` target positions (rad) in IsaacLab joint order.
|
|
||||||
"""
|
|
||||||
self.token = np.asarray(token, np.float32)
|
|
||||||
jnames = [m.name for m in G1_29_JointIndex]
|
|
||||||
q = np.array(
|
|
||||||
[
|
|
||||||
robot_obs.get(f"{n}.q", self.default_angles[m.value])
|
|
||||||
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)
|
|
||||||
action_mj = (
|
|
||||||
self.decoder.run(None, {self.decoder_input: self.build_decoder_obs().reshape(1, -1)})[0]
|
|
||||||
.squeeze()
|
|
||||||
.astype(np.float32)
|
|
||||||
)
|
|
||||||
self.last_action_mj = action_mj.copy()
|
|
||||||
target = self.default_angles + action_mj[ISAACLAB_TO_MUJOCO] * self.action_scale
|
|
||||||
if debug:
|
|
||||||
delta = target - q
|
|
||||||
logger.debug(
|
|
||||||
"token_norm=%.4f action_norm=%.4f delta_max=%.4f delta_rms=%.4f",
|
|
||||||
np.linalg.norm(self.token),
|
|
||||||
np.linalg.norm(action_mj),
|
|
||||||
np.max(np.abs(delta)),
|
|
||||||
np.sqrt(np.mean(delta**2)),
|
|
||||||
)
|
|
||||||
return {f"{m.name}.q": float(target[m.value]) for m in G1_29_JointIndex}
|
|
||||||
|
|
||||||
|
|
||||||
class SonicRuntime:
|
|
||||||
"""Loads the SONIC decoder ONNX model and owns the decode controller.
|
|
||||||
|
|
||||||
Token-only deploy: the encoder is bypassed; each tick the decoder consumes a 64-D
|
|
||||||
latent token supplied directly by the policy.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
decoder_sess, self.kp, self.kd, default_angles, action_scale, neutral_token = load_sonic_decoder()
|
|
||||||
self.default_angles = default_angles
|
|
||||||
self.neutral_token = neutral_token
|
|
||||||
self.controller = SonicDecoder(decoder_sess, default_angles, action_scale)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def pipeline(self):
|
|
||||||
return self.controller
|
|
||||||
|
|
||||||
def reset(self):
|
|
||||||
self.controller.reset()
|
|
||||||
|
|
||||||
def shutdown(self):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class SonicWholeBodyController:
|
|
||||||
"""Full-body SONIC controller for UnitreeG1's background controller thread."""
|
|
||||||
|
|
||||||
control_dt = CONTROL_DT
|
|
||||||
full_body = True
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
logger.info("Loading SONIC whole-body controller...")
|
|
||||||
self._runtime = SonicRuntime()
|
|
||||||
self.kp = self._runtime.kp
|
|
||||||
self.kd = self._runtime.kd
|
|
||||||
self.controller = self._runtime.controller
|
|
||||||
self._default_angles = self._runtime.default_angles
|
|
||||||
self._neutral_token = self._runtime.neutral_token
|
|
||||||
|
|
||||||
# Startup blend: ease from the robot's initial pose into the first commanded policy
|
|
||||||
# targets over INIT_RAMP_S (captured on the first control tick).
|
|
||||||
self._init_ramp_steps = max(1, round(INIT_RAMP_S / CONTROL_DT))
|
|
||||||
self._init_step = 0
|
|
||||||
self._start_pose: dict[str, float] = {}
|
|
||||||
|
|
||||||
# Token-interface state. The controller holds a stable *neutral* token until the first
|
|
||||||
# real token arrives, and afterwards holds the *last* token received between ticks (the
|
|
||||||
# async controller runs ~50 Hz while a token VLA streams ~30 Hz).
|
|
||||||
self._last_token: np.ndarray | None = None
|
|
||||||
|
|
||||||
logger.info("SONIC ready (decoder, 64-D token command path)")
|
|
||||||
|
|
||||||
def _startup_blend(self, obs: dict, out: dict) -> dict:
|
|
||||||
"""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
|
|
||||||
live commanded target, so the handoff has no snap.
|
|
||||||
|
|
||||||
``out`` is the policy's ``<joint>.q`` target dict for this tick; the blend ratio
|
|
||||||
climbs 0->1 over the ramp, after which the raw policy target passes through.
|
|
||||||
"""
|
|
||||||
if self._init_step >= self._init_ramp_steps or not out:
|
|
||||||
return out
|
|
||||||
if self._init_step == 0:
|
|
||||||
# Capture the robot's actual pose as the interpolation start point.
|
|
||||||
self._start_pose = {
|
|
||||||
f"{m.name}.q": float(obs.get(f"{m.name}.q", self._default_angles[m.value]))
|
|
||||||
for m in G1_29_JointIndex
|
|
||||||
}
|
|
||||||
self._init_step += 1
|
|
||||||
ratio = min(1.0, self._init_step / self._init_ramp_steps)
|
|
||||||
blended = {
|
|
||||||
k: self._start_pose.get(k, float(tgt)) * (1.0 - ratio) + float(tgt) * ratio
|
|
||||||
for k, tgt in out.items()
|
|
||||||
}
|
|
||||||
if self._init_step >= self._init_ramp_steps:
|
|
||||||
logger.info("SONIC startup blend complete -> full policy control")
|
|
||||||
return blended
|
|
||||||
|
|
||||||
def run_step(self, action: dict, obs: dict) -> dict:
|
|
||||||
if not obs:
|
|
||||||
return {}
|
|
||||||
|
|
||||||
# Token-only interface (token-output VLA): a dense 64-D ``motion_token.{i}`` command
|
|
||||||
# is decoded directly, encoder bypassed.
|
|
||||||
token = _extract_token_from_action(action)
|
|
||||||
if token is not None:
|
|
||||||
self._last_token = token
|
|
||||||
elif self._last_token is None:
|
|
||||||
# No token has arrived yet: hold the checkpoint's neutral token, which the decoder
|
|
||||||
# maps to a stable, natural standing pose.
|
|
||||||
self._last_token = self._neutral_token.copy()
|
|
||||||
# Either a fresh token this tick or the last one received (held between the ~30 Hz
|
|
||||||
# token stream and the ~50 Hz control loop).
|
|
||||||
return self._startup_blend(obs, self.controller.step(obs, self._last_token))
|
|
||||||
|
|
||||||
def reset(self):
|
|
||||||
self._runtime.reset()
|
|
||||||
self._init_step = 0 # re-run the startup blend after a reset
|
|
||||||
self._start_pose = {}
|
|
||||||
# Drop the held token so the neutral token is re-seeded after a reset.
|
|
||||||
self._last_token = None
|
|
||||||
|
|
||||||
def shutdown(self):
|
|
||||||
self._runtime.shutdown()
|
|
||||||
@@ -23,47 +23,6 @@ import numpy as np
|
|||||||
|
|
||||||
NUM_MOTORS = 29
|
NUM_MOTORS = 29
|
||||||
|
|
||||||
# Joint-order permutations between the two 29-DoF layouts used across the G1 stack:
|
|
||||||
# IsaacLab (policy/training order) and MuJoCo (deploy order). ``a[ISAACLAB_TO_MUJOCO]``
|
|
||||||
# reorders an IsaacLab-ordered vector into MuJoCo order, and vice-versa.
|
|
||||||
ISAACLAB_TO_MUJOCO = np.array(
|
|
||||||
[
|
|
||||||
0,
|
|
||||||
3,
|
|
||||||
6,
|
|
||||||
9,
|
|
||||||
13,
|
|
||||||
17,
|
|
||||||
1,
|
|
||||||
4,
|
|
||||||
7,
|
|
||||||
10,
|
|
||||||
14,
|
|
||||||
18,
|
|
||||||
2,
|
|
||||||
5,
|
|
||||||
8,
|
|
||||||
11,
|
|
||||||
15,
|
|
||||||
19,
|
|
||||||
21,
|
|
||||||
23,
|
|
||||||
25,
|
|
||||||
27,
|
|
||||||
12,
|
|
||||||
16,
|
|
||||||
20,
|
|
||||||
22,
|
|
||||||
24,
|
|
||||||
26,
|
|
||||||
28,
|
|
||||||
],
|
|
||||||
dtype=np.int32,
|
|
||||||
)
|
|
||||||
# The two orderings are inverses of each other, so derive one from the other (argsort) to
|
|
||||||
# guarantee they can never drift out of sync.
|
|
||||||
MUJOCO_TO_ISAACLAB = np.argsort(ISAACLAB_TO_MUJOCO).astype(np.int32)
|
|
||||||
|
|
||||||
REMOTE_AXES = ("remote.lx", "remote.ly", "remote.rx", "remote.ry")
|
REMOTE_AXES = ("remote.lx", "remote.ly", "remote.rx", "remote.ry")
|
||||||
REMOTE_BUTTONS = tuple(f"remote.button.{i}" for i in range(16))
|
REMOTE_BUTTONS = tuple(f"remote.button.{i}" for i in range(16))
|
||||||
REMOTE_KEYS = REMOTE_AXES + REMOTE_BUTTONS
|
REMOTE_KEYS = REMOTE_AXES + REMOTE_BUTTONS
|
||||||
@@ -109,9 +68,8 @@ def make_locomotion_controller(name: str | None):
|
|||||||
if name is None:
|
if name is None:
|
||||||
return None
|
return None
|
||||||
controllers = {
|
controllers = {
|
||||||
"GrootLocomotionController": "lerobot.robots.unitree_g1.controllers.gr00t_locomotion",
|
"GrootLocomotionController": "lerobot.robots.unitree_g1.gr00t_locomotion",
|
||||||
"HolosomaLocomotionController": "lerobot.robots.unitree_g1.controllers.holosoma_locomotion",
|
"HolosomaLocomotionController": "lerobot.robots.unitree_g1.holosoma_locomotion",
|
||||||
"SonicWholeBodyController": "lerobot.robots.unitree_g1.controllers.sonic_whole_body",
|
|
||||||
}
|
}
|
||||||
module_path = controllers.get(name)
|
module_path = controllers.get(name)
|
||||||
if module_path is None:
|
if module_path is None:
|
||||||
|
|||||||
+1
-1
@@ -21,7 +21,7 @@ import numpy as np
|
|||||||
import onnxruntime as ort
|
import onnxruntime as ort
|
||||||
from huggingface_hub import hf_hub_download
|
from huggingface_hub import hf_hub_download
|
||||||
|
|
||||||
from ..g1_utils import (
|
from .g1_utils import (
|
||||||
REMOTE_AXES,
|
REMOTE_AXES,
|
||||||
REMOTE_BUTTONS,
|
REMOTE_BUTTONS,
|
||||||
G1_29_JointIndex,
|
G1_29_JointIndex,
|
||||||
+1
-1
@@ -22,7 +22,7 @@ import onnx
|
|||||||
import onnxruntime as ort
|
import onnxruntime as ort
|
||||||
from huggingface_hub import hf_hub_download
|
from huggingface_hub import hf_hub_download
|
||||||
|
|
||||||
from ..g1_utils import (
|
from .g1_utils import (
|
||||||
REMOTE_AXES,
|
REMOTE_AXES,
|
||||||
G1_29_JointArmIndex,
|
G1_29_JointArmIndex,
|
||||||
G1_29_JointIndex,
|
G1_29_JointIndex,
|
||||||
@@ -34,6 +34,7 @@ from .config_unitree_g1 import UnitreeG1Config
|
|||||||
from .g1_kinematics import G1_29_ArmIK
|
from .g1_kinematics import G1_29_ArmIK
|
||||||
from .g1_utils import (
|
from .g1_utils import (
|
||||||
REMOTE_AXES,
|
REMOTE_AXES,
|
||||||
|
REMOTE_KEYS,
|
||||||
G1_29_JointArmIndex,
|
G1_29_JointArmIndex,
|
||||||
G1_29_JointIndex,
|
G1_29_JointIndex,
|
||||||
default_remote_input,
|
default_remote_input,
|
||||||
@@ -147,54 +148,22 @@ class UnitreeG1(Robot):
|
|||||||
|
|
||||||
self.arm_ik = G1_29_ArmIK() if config.gravity_compensation else None
|
self.arm_ik = G1_29_ArmIK() if config.gravity_compensation else None
|
||||||
|
|
||||||
# Lower-body / whole-body controller loaded dynamically
|
# Lower-body controller loaded dynamically
|
||||||
self.controller: LocomotionController | None = make_locomotion_controller(config.controller)
|
self.controller: LocomotionController | None = make_locomotion_controller(config.controller)
|
||||||
|
|
||||||
# Controller thread state
|
# Controller thread state
|
||||||
self._controller_thread = None
|
self._controller_thread = None
|
||||||
# When set, the controller loop stops publishing low commands so reset() can
|
|
||||||
# drive the joints directly without two publishers fighting (single-publisher).
|
|
||||||
self._controller_paused = threading.Event()
|
|
||||||
self._controller_action_lock = threading.Lock()
|
self._controller_action_lock = threading.Lock()
|
||||||
self.controller_input = default_remote_input()
|
self.controller_input = default_remote_input()
|
||||||
self.controller_output = {}
|
self.controller_output = {}
|
||||||
|
|
||||||
# Token-mode state: last 64-D SONIC latent token commanded by the policy,
|
|
||||||
# echoed back as ``observation.state`` so a token-output VLA closes the loop
|
|
||||||
# on its own previous token. Implicit whenever the SONIC whole-body controller
|
|
||||||
# is active. Seeded to zeros; the controller's startup blend eases joints in.
|
|
||||||
self._last_token: np.ndarray | None = None
|
|
||||||
if self._sonic_token:
|
|
||||||
from .controllers.sonic_whole_body import TOKEN_DIM
|
|
||||||
|
|
||||||
self._last_token = np.zeros(TOKEN_DIM, dtype=np.float32)
|
|
||||||
|
|
||||||
@property
|
|
||||||
def _sonic_token(self) -> bool:
|
|
||||||
"""Whether the SONIC whole-body decoder is active.
|
|
||||||
|
|
||||||
A SONIC controller consumes a 64-D latent motion token as its action and echoes
|
|
||||||
the last commanded token as ``observation.state``. Keyed purely off the selected
|
|
||||||
controller so the token interface is implicit -- no separate config flag.
|
|
||||||
"""
|
|
||||||
return self.config.controller == "SonicWholeBodyController"
|
|
||||||
|
|
||||||
def _subscribe_lowstate(self): # polls robot state @ 250Hz
|
def _subscribe_lowstate(self): # polls robot state @ 250Hz
|
||||||
while not self._shutdown_event.is_set():
|
while not self._shutdown_event.is_set():
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
# Step simulation if in simulation mode
|
# Step simulation if in simulation mode
|
||||||
if self.config.is_simulation and self.sim_env is not None:
|
if self.config.is_simulation and self.sim_env is not None:
|
||||||
try:
|
|
||||||
self.sim_env.step()
|
self.sim_env.step()
|
||||||
except ValueError as e:
|
|
||||||
# Startup race: the sim thread can step once before reset() has
|
|
||||||
# written a valid base pose, giving a zero-norm pelvis quaternion
|
|
||||||
# (scipy>=1.11 raises instead of normalizing). Skip and retry so
|
|
||||||
# the thread survives instead of dying and freezing the sim.
|
|
||||||
if "zero norm" not in str(e).lower():
|
|
||||||
raise
|
|
||||||
time.sleep(self.control_dt)
|
|
||||||
continue
|
|
||||||
|
|
||||||
msg = self.lowstate_subscriber.Read()
|
msg = self.lowstate_subscriber.Read()
|
||||||
if msg is not None:
|
if msg is not None:
|
||||||
@@ -262,38 +231,15 @@ class UnitreeG1(Robot):
|
|||||||
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
|
features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
|
||||||
return features
|
return features
|
||||||
|
|
||||||
@property
|
|
||||||
def _token_state_ft(self) -> dict[str, type]:
|
|
||||||
"""64-D SONIC latent-token proprio state (``motion_token_state.{i}.pos``).
|
|
||||||
|
|
||||||
Exposed only when a SONIC whole-body controller is active; aggregated by the
|
|
||||||
rollout into a 64-D ``observation.state`` (the last token the policy commanded).
|
|
||||||
"""
|
|
||||||
if not self._sonic_token:
|
|
||||||
return {}
|
|
||||||
from .controllers.sonic_whole_body import TOKEN_DIM, token_state_key
|
|
||||||
|
|
||||||
return {token_state_key(i): float for i in range(TOKEN_DIM)}
|
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def observation_features(self) -> dict[str, type | tuple]:
|
def observation_features(self) -> dict[str, type | tuple]:
|
||||||
return {**self._motors_ft, **self._token_state_ft, **self._cameras_ft}
|
return {**self._motors_ft, **self._cameras_ft}
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def action_features(self) -> dict[str, type]:
|
def action_features(self) -> dict[str, type]:
|
||||||
# No controller configured at all: raw 29-DoF joint teleop.
|
|
||||||
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}
|
||||||
|
|
||||||
# Token-output VLA (SONIC decoder): advertise a 64-D latent-token action space
|
|
||||||
# (``motion_token.{i}.pos``) so ``lerobot-rollout`` maps a 64-D policy output
|
|
||||||
# straight onto the decoder, bypassing the encoder.
|
|
||||||
if self._sonic_token:
|
|
||||||
from .controllers.sonic_whole_body import TOKEN_DIM, token_action_key
|
|
||||||
|
|
||||||
return {token_action_key(i): float for i in range(TOKEN_DIM)}
|
|
||||||
|
|
||||||
# Locomotion controllers (GR00T / Holosoma): arm joint targets + joystick axes.
|
|
||||||
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}
|
||||||
@@ -309,11 +255,6 @@ class UnitreeG1(Robot):
|
|||||||
while not self._shutdown_event.is_set():
|
while not self._shutdown_event.is_set():
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
# Paused during reset() so the reset routine is the sole low-cmd publisher.
|
|
||||||
if self._controller_paused.is_set():
|
|
||||||
time.sleep(control_dt)
|
|
||||||
continue
|
|
||||||
|
|
||||||
with self._lowstate_lock:
|
with self._lowstate_lock:
|
||||||
lowstate = self._lowstate
|
lowstate = self._lowstate
|
||||||
|
|
||||||
@@ -330,12 +271,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:
|
||||||
@@ -406,9 +343,6 @@ class UnitreeG1(Robot):
|
|||||||
|
|
||||||
self.kp = np.array(self.config.kp, dtype=np.float32)
|
self.kp = np.array(self.config.kp, dtype=np.float32)
|
||||||
self.kd = np.array(self.config.kd, dtype=np.float32)
|
self.kd = np.array(self.config.kd, dtype=np.float32)
|
||||||
if self.controller is not None and hasattr(self.controller, "kp"):
|
|
||||||
self.kp = np.array(self.controller.kp, dtype=np.float32)
|
|
||||||
self.kd = np.array(self.controller.kd, dtype=np.float32)
|
|
||||||
|
|
||||||
for joint in G1_29_JointIndex:
|
for joint in G1_29_JointIndex:
|
||||||
self.msg.motor_cmd[joint].mode = 1
|
self.msg.motor_cmd[joint].mode = 1
|
||||||
@@ -457,10 +391,6 @@ class UnitreeG1(Robot):
|
|||||||
if self._controller_thread.is_alive():
|
if self._controller_thread.is_alive():
|
||||||
logger.warning("Controller thread did not stop cleanly")
|
logger.warning("Controller thread did not stop cleanly")
|
||||||
|
|
||||||
# Release controller resources (e.g. SONIC decoder sessions).
|
|
||||||
if self.controller is not None and hasattr(self.controller, "shutdown"):
|
|
||||||
self.controller.shutdown()
|
|
||||||
|
|
||||||
# Close simulation environment
|
# Close simulation environment
|
||||||
if self.config.is_simulation and self.sim_env is not None:
|
if self.config.is_simulation and self.sim_env is not None:
|
||||||
try:
|
try:
|
||||||
@@ -531,15 +461,6 @@ class UnitreeG1(Robot):
|
|||||||
if lowstate.wireless_remote:
|
if lowstate.wireless_remote:
|
||||||
obs["wireless_remote"] = lowstate.wireless_remote
|
obs["wireless_remote"] = lowstate.wireless_remote
|
||||||
|
|
||||||
# Token mode: echo the last commanded latent token as observation.state so a
|
|
||||||
# token-output VLA closes the loop on its own previous token.
|
|
||||||
if self._sonic_token:
|
|
||||||
from .controllers.sonic_whole_body import token_state_key
|
|
||||||
|
|
||||||
token = self._last_token if self._last_token is not None else []
|
|
||||||
for i, v in enumerate(token):
|
|
||||||
obs[token_state_key(i)] = float(v)
|
|
||||||
|
|
||||||
# Cameras - read images from ZMQ cameras
|
# Cameras - read images from ZMQ cameras
|
||||||
for cam_name, cam in self._cameras.items():
|
for cam_name, cam in self._cameras.items():
|
||||||
if getattr(cam, "use_rgb", True):
|
if getattr(cam, "use_rgb", True):
|
||||||
@@ -552,22 +473,9 @@ class UnitreeG1(Robot):
|
|||||||
def send_action(self, action: RobotAction) -> RobotAction:
|
def send_action(self, action: RobotAction) -> RobotAction:
|
||||||
action_to_publish = action
|
action_to_publish = action
|
||||||
if self.controller is not None:
|
if self.controller is not None:
|
||||||
# SONIC decoder: pull the 64-D latent token out of the action and remember it
|
|
||||||
# for the observation.state echo. The controller thread reads it back from
|
|
||||||
# controller_input (populated below) and decodes it into a 29-DoF command.
|
|
||||||
if self._sonic_token:
|
|
||||||
from .controllers.sonic_whole_body import _extract_token_from_action
|
|
||||||
|
|
||||||
token = _extract_token_from_action(action)
|
|
||||||
if token is not None:
|
|
||||||
self._last_token = token
|
|
||||||
self._update_controller_action(action)
|
|
||||||
# Full-body controllers (SONIC) own the whole 29-DoF command; nothing to
|
|
||||||
# publish here (the controller thread is the sole publisher).
|
|
||||||
if getattr(self.controller, "full_body", False):
|
|
||||||
return action
|
|
||||||
# Controller thread owns legs/waist. Here we only update joystick inputs
|
# Controller thread owns legs/waist. Here we only update joystick inputs
|
||||||
# and publish arm targets from the teleoperator.
|
# and publish arm targets from the teleoperator.
|
||||||
|
self._update_controller_action(action)
|
||||||
arm_prefixes = tuple(j.name for j in G1_29_JointArmIndex)
|
arm_prefixes = tuple(j.name for j in G1_29_JointArmIndex)
|
||||||
action_to_publish = {
|
action_to_publish = {
|
||||||
key: value
|
key: value
|
||||||
@@ -595,17 +503,11 @@ class UnitreeG1(Robot):
|
|||||||
return action
|
return action
|
||||||
|
|
||||||
def _update_controller_action(self, action: RobotAction) -> None:
|
def _update_controller_action(self, action: RobotAction) -> None:
|
||||||
"""Update controller input state from an incoming teleop action.
|
"""Update controller input state from incoming teleop action."""
|
||||||
|
|
||||||
Controller-agnostic: every value-carrying key (locomotion ``remote.*`` axes or
|
|
||||||
SONIC ``motion_token.*`` values) is forwarded verbatim into ``controller_input``
|
|
||||||
and each controller extracts only the keys it understands. The robot deliberately
|
|
||||||
does not enumerate any controller's key schema here.
|
|
||||||
"""
|
|
||||||
with self._controller_action_lock:
|
with self._controller_action_lock:
|
||||||
for key, value in action.items():
|
for key in REMOTE_KEYS:
|
||||||
if isinstance(key, str) and value is not None:
|
if key in action:
|
||||||
self.controller_input[key] = value
|
self.controller_input[key] = action[key]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_calibrated(self) -> bool:
|
def is_calibrated(self) -> bool:
|
||||||
@@ -635,18 +537,6 @@ class UnitreeG1(Robot):
|
|||||||
if default_positions is None:
|
if default_positions is None:
|
||||||
default_positions = np.array(self.config.default_positions, dtype=np.float32)
|
default_positions = np.array(self.config.default_positions, dtype=np.float32)
|
||||||
|
|
||||||
# Full-body controllers (SONIC) own the whole 29-DoF command and ignore
|
|
||||||
# ``<joint>.q`` in send_action(), so reset() must publish the default pose
|
|
||||||
# directly. Pause the background controller first so the two aren't both writing
|
|
||||||
# low commands while the robot moves to the default pose.
|
|
||||||
full_body = getattr(self.controller, "full_body", False)
|
|
||||||
paused = False
|
|
||||||
if full_body and self._controller_thread is not None:
|
|
||||||
self._controller_paused.set()
|
|
||||||
paused = True
|
|
||||||
time.sleep(control_dt) # let any in-flight controller tick settle
|
|
||||||
|
|
||||||
try:
|
|
||||||
if self.config.is_simulation and self.sim_env is not None:
|
if self.config.is_simulation and self.sim_env is not None:
|
||||||
self.sim_env.reset()
|
self.sim_env.reset()
|
||||||
self.publish_lowcmd(
|
self.publish_lowcmd(
|
||||||
@@ -675,11 +565,6 @@ class UnitreeG1(Robot):
|
|||||||
interp_pos = init_dof_pos[motor.value] * (1 - alpha) + target_pos * alpha
|
interp_pos = init_dof_pos[motor.value] * (1 - alpha) + target_pos * alpha
|
||||||
action_dict[f"{motor.name}.q"] = float(interp_pos)
|
action_dict[f"{motor.name}.q"] = float(interp_pos)
|
||||||
|
|
||||||
# Full-body controllers no-op in send_action(); publish the pose
|
|
||||||
# directly (arm-only controllers keep the send_action() path).
|
|
||||||
if full_body:
|
|
||||||
self.publish_lowcmd(action_dict)
|
|
||||||
else:
|
|
||||||
self.send_action(action_dict)
|
self.send_action(action_dict)
|
||||||
|
|
||||||
# Maintain constant control rate
|
# Maintain constant control rate
|
||||||
@@ -687,12 +572,8 @@ class UnitreeG1(Robot):
|
|||||||
sleep_time = max(0, control_dt - elapsed)
|
sleep_time = max(0, control_dt - elapsed)
|
||||||
time.sleep(sleep_time)
|
time.sleep(sleep_time)
|
||||||
|
|
||||||
# Reset controller internal state (gait phase, obs history, etc.) before
|
# Reset controller internal state (gait phase, obs history, etc.)
|
||||||
# resuming so its buffers reflect the post-reset pose.
|
|
||||||
if self.controller is not None and hasattr(self.controller, "reset"):
|
if self.controller is not None and hasattr(self.controller, "reset"):
|
||||||
self.controller.reset()
|
self.controller.reset()
|
||||||
finally:
|
|
||||||
if paused:
|
|
||||||
self._controller_paused.clear()
|
|
||||||
|
|
||||||
logger.info("Reached default position")
|
logger.info("Reached default position")
|
||||||
|
|||||||
@@ -57,6 +57,7 @@ from .inference import (
|
|||||||
SyncInferenceConfig,
|
SyncInferenceConfig,
|
||||||
create_inference_engine,
|
create_inference_engine,
|
||||||
)
|
)
|
||||||
|
from .inference.rtc import supports_rtc_inference
|
||||||
from .robot_wrapper import ThreadSafeRobot
|
from .robot_wrapper import ThreadSafeRobot
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
if TYPE_CHECKING or _peft_available:
|
||||||
@@ -226,6 +227,12 @@ def build_rollout_context(
|
|||||||
policy = _load_pretrained_policy(policy_config)
|
policy = _load_pretrained_policy(policy_config)
|
||||||
|
|
||||||
if is_rtc:
|
if is_rtc:
|
||||||
|
if not supports_rtc_inference(policy):
|
||||||
|
raise ValueError(
|
||||||
|
f"RTC inference is not supported by policy type '{policy_config.type}': "
|
||||||
|
"the policy must implement RTC semantics and predict_action_chunk must accept "
|
||||||
|
"inference_delay and prev_chunk_left_over. Use '--inference.type=sync' instead."
|
||||||
|
)
|
||||||
policy.config.rtc_config = cfg.inference.rtc
|
policy.config.rtc_config = cfg.inference.rtc
|
||||||
if hasattr(policy, "init_rtc_processor"):
|
if hasattr(policy, "init_rtc_processor"):
|
||||||
policy.init_rtc_processor()
|
policy.init_rtc_processor()
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ way via ``notify_observation``.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
import time
|
import time
|
||||||
@@ -62,6 +63,23 @@ _RTC_JOIN_TIMEOUT_S: float = 3.0
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def supports_rtc_inference(policy: PreTrainedPolicy) -> bool:
|
||||||
|
"""Whether a policy declares RTC support and accepts the RTC call shape."""
|
||||||
|
supports_rtc = getattr(policy, "supports_rtc", None)
|
||||||
|
if not callable(supports_rtc) or not supports_rtc():
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
inspect.signature(policy.predict_action_chunk).bind(
|
||||||
|
object(),
|
||||||
|
inference_delay=0,
|
||||||
|
prev_chunk_left_over=None,
|
||||||
|
)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
def _normalize_prev_actions_length(prev_actions: torch.Tensor, target_steps: int) -> torch.Tensor:
|
def _normalize_prev_actions_length(prev_actions: torch.Tensor, target_steps: int) -> torch.Tensor:
|
||||||
"""Pad or truncate RTC prefix actions to a fixed length for stable compiled inference."""
|
"""Pad or truncate RTC prefix actions to a fixed length for stable compiled inference."""
|
||||||
if prev_actions.ndim != 2:
|
if prev_actions.ndim != 2:
|
||||||
|
|||||||
@@ -348,7 +348,6 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
preprocessor_overrides = {
|
preprocessor_overrides = {
|
||||||
"device_processor": {"device": device.type},
|
"device_processor": {"device": device.type},
|
||||||
"normalizer_processor": {
|
"normalizer_processor": {
|
||||||
"stats": dataset.meta.stats,
|
|
||||||
"features": {**policy.config.input_features, **policy.config.output_features},
|
"features": {**policy.config.input_features, **policy.config.output_features},
|
||||||
"norm_map": policy.config.normalization_mapping,
|
"norm_map": policy.config.normalization_mapping,
|
||||||
},
|
},
|
||||||
@@ -356,11 +355,17 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
}
|
}
|
||||||
postprocessor_overrides = {
|
postprocessor_overrides = {
|
||||||
"unnormalizer_processor": {
|
"unnormalizer_processor": {
|
||||||
"stats": dataset.meta.stats,
|
|
||||||
"features": policy.config.output_features,
|
"features": policy.config.output_features,
|
||||||
"norm_map": policy.config.normalization_mapping,
|
"norm_map": policy.config.normalization_mapping,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
# On resume, the checkpoint's saved processor stats are authoritative: they may have
|
||||||
|
# been adapted by the policy (e.g. EVO1 pads state/action stats to max_state_dim),
|
||||||
|
# and force-feeding raw dataset stats over them crashes normalization (#4006).
|
||||||
|
# This mirrors the `dataset_stats` kwarg above, which is also skipped on resume.
|
||||||
|
if not cfg.resume:
|
||||||
|
preprocessor_overrides["normalizer_processor"]["stats"] = dataset.meta.stats
|
||||||
|
postprocessor_overrides["unnormalizer_processor"]["stats"] = dataset.meta.stats
|
||||||
if getattr(active_cfg, "use_relative_actions", False):
|
if getattr(active_cfg, "use_relative_actions", False):
|
||||||
preprocessor_overrides["relative_actions_processor"] = {
|
preprocessor_overrides["relative_actions_processor"] = {
|
||||||
"enabled": True,
|
"enabled": True,
|
||||||
|
|||||||
@@ -67,7 +67,15 @@ class BiSOLeader(BimanualMixin, Teleoperator):
|
|||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def feedback_features(self) -> dict[str, type]:
|
def feedback_features(self) -> dict[str, type]:
|
||||||
return {}
|
# Bimanual teleop has feedback (can be actuated for handover).
|
||||||
|
# Return the same structure as action_features for consistency with left/right arms.
|
||||||
|
left_arm_features = self.left_arm.feedback_features
|
||||||
|
right_arm_features = self.right_arm.feedback_features
|
||||||
|
|
||||||
|
return {
|
||||||
|
**{f"left_{k}": v for k, v in left_arm_features.items()},
|
||||||
|
**{f"right_{k}": v for k, v in right_arm_features.items()},
|
||||||
|
}
|
||||||
|
|
||||||
def setup_motors(self) -> None:
|
def setup_motors(self) -> None:
|
||||||
self.left_arm.setup_motors()
|
self.left_arm.setup_motors()
|
||||||
@@ -87,6 +95,43 @@ class BiSOLeader(BimanualMixin, Teleoperator):
|
|||||||
|
|
||||||
return action_dict
|
return action_dict
|
||||||
|
|
||||||
|
def enable_torque(self) -> None:
|
||||||
|
"""Enable torque on both leader arms for smooth handover."""
|
||||||
|
self.left_arm.enable_torque()
|
||||||
|
self.right_arm.enable_torque()
|
||||||
|
|
||||||
|
def disable_torque(self) -> None:
|
||||||
|
"""Disable torque on both leader arms to allow human control."""
|
||||||
|
self.left_arm.disable_torque()
|
||||||
|
self.right_arm.disable_torque()
|
||||||
|
|
||||||
|
@check_if_not_connected
|
||||||
def send_feedback(self, feedback: dict[str, float]) -> None:
|
def send_feedback(self, feedback: dict[str, float]) -> None:
|
||||||
# TODO: Implement force feedback
|
"""Route bimanual feedback to left and right arms with proper prefix stripping.
|
||||||
raise NotImplementedError
|
|
||||||
|
Receives feedback dict with keys like: left_shoulder_pan.pos, right_shoulder_pan.pos, ...
|
||||||
|
Splits and routes to each arm by removing the prefix.
|
||||||
|
|
||||||
|
This enables DAgger smooth handover: when transitioning from policy control to human
|
||||||
|
intervention, both leader arms are commanded to the follower's current pose to avoid
|
||||||
|
discontinuities.
|
||||||
|
"""
|
||||||
|
# Split feedback by arm prefix
|
||||||
|
left_feedback = {}
|
||||||
|
right_feedback = {}
|
||||||
|
|
||||||
|
for key, value in feedback.items():
|
||||||
|
if key.startswith("left_"):
|
||||||
|
# Strip "left_" prefix and pass to left arm
|
||||||
|
stripped_key = key[5:] # len("left_") == 5
|
||||||
|
left_feedback[stripped_key] = value
|
||||||
|
elif key.startswith("right_"):
|
||||||
|
# Strip "right_" prefix and pass to right arm
|
||||||
|
stripped_key = key[6:] # len("right_") == 6
|
||||||
|
right_feedback[stripped_key] = value
|
||||||
|
|
||||||
|
# Send to each arm
|
||||||
|
if left_feedback:
|
||||||
|
self.left_arm.send_feedback(left_feedback)
|
||||||
|
if right_feedback:
|
||||||
|
self.right_arm.send_feedback(right_feedback)
|
||||||
|
|||||||
@@ -85,6 +85,8 @@ def serialize_torch_rng_state() -> dict[str, torch.Tensor]:
|
|||||||
torch_rng_state_dict = {"torch_rng_state": torch.get_rng_state()}
|
torch_rng_state_dict = {"torch_rng_state": torch.get_rng_state()}
|
||||||
if torch.cuda.is_available():
|
if torch.cuda.is_available():
|
||||||
torch_rng_state_dict["torch_cuda_rng_state"] = torch.cuda.get_rng_state()
|
torch_rng_state_dict["torch_cuda_rng_state"] = torch.cuda.get_rng_state()
|
||||||
|
if torch.backends.mps.is_available():
|
||||||
|
torch_rng_state_dict["torch_mps_rng_state"] = torch.mps.get_rng_state()
|
||||||
return torch_rng_state_dict
|
return torch_rng_state_dict
|
||||||
|
|
||||||
|
|
||||||
@@ -95,6 +97,8 @@ def deserialize_torch_rng_state(rng_state_dict: dict[str, torch.Tensor]) -> None
|
|||||||
torch.set_rng_state(rng_state_dict["torch_rng_state"])
|
torch.set_rng_state(rng_state_dict["torch_rng_state"])
|
||||||
if torch.cuda.is_available() and "torch_cuda_rng_state" in rng_state_dict:
|
if torch.cuda.is_available() and "torch_cuda_rng_state" in rng_state_dict:
|
||||||
torch.cuda.set_rng_state(rng_state_dict["torch_cuda_rng_state"])
|
torch.cuda.set_rng_state(rng_state_dict["torch_cuda_rng_state"])
|
||||||
|
if torch.backends.mps.is_available() and "torch_mps_rng_state" in rng_state_dict:
|
||||||
|
torch.mps.set_rng_state(rng_state_dict["torch_mps_rng_state"])
|
||||||
|
|
||||||
|
|
||||||
def serialize_rng_state() -> dict[str, torch.Tensor]:
|
def serialize_rng_state() -> dict[str, torch.Tensor]:
|
||||||
|
|||||||
@@ -66,3 +66,27 @@ def test_from_pretrained_raises_when_no_root_config_and_no_checkpoints(monkeypat
|
|||||||
|
|
||||||
with pytest.raises(FileNotFoundError, match="train_config.json not found"):
|
with pytest.raises(FileNotFoundError, match="train_config.json not found"):
|
||||||
TrainPipelineConfig.from_pretrained("user/empty-repo")
|
TrainPipelineConfig.from_pretrained("user/empty-repo")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("pass_dir", [False, True])
|
||||||
|
def test_resolve_resume_checkpoint_accepts_file_or_pretrained_model_dir(tmp_path, monkeypatch, pass_dir):
|
||||||
|
"""`--config_path` may point at the checkpoint's train_config.json or at its
|
||||||
|
pretrained_model/ directory; both must resolve `policy.pretrained_path` to the
|
||||||
|
pretrained_model/ directory (regression test for the directory case, which
|
||||||
|
previously resolved one level too high and failed on model.safetensors)."""
|
||||||
|
pretrained_dir = tmp_path / "checkpoints" / "000002" / "pretrained_model"
|
||||||
|
pretrained_dir.mkdir(parents=True)
|
||||||
|
(pretrained_dir / "train_config.json").touch()
|
||||||
|
target = pretrained_dir if pass_dir else pretrained_dir / "train_config.json"
|
||||||
|
|
||||||
|
from lerobot.policies.act.configuration_act import ACTConfig
|
||||||
|
|
||||||
|
cfg = tc.draccus.parse(TrainPipelineConfig, args=["--dataset.repo_id", "u/d"])
|
||||||
|
cfg.policy = ACTConfig()
|
||||||
|
cfg.resume = True
|
||||||
|
monkeypatch.setattr(tc.parser, "parse_arg", lambda name: str(target) if name == "config_path" else None)
|
||||||
|
|
||||||
|
cfg._resolve_resume_checkpoint()
|
||||||
|
|
||||||
|
assert cfg.policy.pretrained_path == pretrained_dir
|
||||||
|
assert cfg.checkpoint_path == pretrained_dir.parent
|
||||||
|
|||||||
@@ -73,6 +73,17 @@ def test_serialize_deserialize_torch_rng(fixed_seed):
|
|||||||
assert val2 == val3
|
assert val2 == val3
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not torch.backends.mps.is_available(), reason="MPS not available")
|
||||||
|
def test_serialize_deserialize_torch_rng_mps(fixed_seed):
|
||||||
|
_ = torch.rand(1, device="mps").item()
|
||||||
|
st = serialize_torch_rng_state()
|
||||||
|
assert "torch_mps_rng_state" in st
|
||||||
|
val2 = torch.rand(1, device="mps").item()
|
||||||
|
deserialize_torch_rng_state(st)
|
||||||
|
val3 = torch.rand(1, device="mps").item()
|
||||||
|
assert val2 == val3
|
||||||
|
|
||||||
|
|
||||||
def test_serialize_deserialize_rng(fixed_seed):
|
def test_serialize_deserialize_rng(fixed_seed):
|
||||||
# Generate one from each library
|
# Generate one from each library
|
||||||
_ = random.random()
|
_ = random.random()
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
version = 1
|
version = 1
|
||||||
revision = 3
|
revision = 2
|
||||||
requires-python = ">=3.12"
|
requires-python = ">=3.12"
|
||||||
resolution-markers = [
|
resolution-markers = [
|
||||||
"(python_full_version >= '3.15' and platform_machine == 'AMD64' and sys_platform == 'linux') or (python_full_version >= '3.15' and platform_machine == 'x86_64' and sys_platform == 'linux')",
|
"(python_full_version >= '3.15' and platform_machine == 'AMD64' and sys_platform == 'linux') or (python_full_version >= '3.15' and platform_machine == 'x86_64' and sys_platform == 'linux')",
|
||||||
@@ -1359,18 +1359,17 @@ sdist = { url = "https://files.pythonhosted.org/packages/a2/55/8f8cab2afd404cf57
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "draccus"
|
name = "draccus"
|
||||||
version = "0.10.0"
|
version = "0.11.6"
|
||||||
source = { registry = "https://pypi.org/simple" }
|
source = { registry = "https://pypi.org/simple" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "mergedeep" },
|
{ name = "mergedeep" },
|
||||||
{ name = "pyyaml" },
|
{ name = "pyyaml" },
|
||||||
{ name = "pyyaml-include" },
|
|
||||||
{ name = "toml" },
|
{ name = "toml" },
|
||||||
{ name = "typing-inspect" },
|
{ name = "typing-inspect" },
|
||||||
]
|
]
|
||||||
sdist = { url = "https://files.pythonhosted.org/packages/4e/e2/f5012fda17ee5d1eaf3481b6ca3e11dffa5348e5e08ab745538fdc8041bb/draccus-0.10.0.tar.gz", hash = "sha256:8dd08304219becdcd66cd16058ba98e9c3e6b7bfe48ccb9579dae39f8d37ae19", size = 62243, upload-time = "2025-02-05T07:27:48.182Z" }
|
sdist = { url = "https://files.pythonhosted.org/packages/df/b0/cc399719e0cf6fea451f155798e19ea8c5e9c838b701d1565f29ea626c18/draccus-0.11.6.tar.gz", hash = "sha256:d134f576a1f4febd93c6b200df7f92e5febffe9efdc0a5f381bf50c2e0568a39", size = 67940, upload-time = "2026-06-12T13:21:19.999Z" }
|
||||||
wheels = [
|
wheels = [
|
||||||
{ url = "https://files.pythonhosted.org/packages/c4/9a/a83083b230d352ee5d205757b74006dbe084448ca45e3bc5ca99215b1e55/draccus-0.10.0-py3-none-any.whl", hash = "sha256:90243418ae0e9271c390a59cafb6acfd37001193696ed36fcc8525f791a83282", size = 71783, upload-time = "2025-02-05T07:27:46.1Z" },
|
{ url = "https://files.pythonhosted.org/packages/18/f1/d56bef4563d1cfaf55dd58de298a456587b9fe0a85dabb55768d44ace8f7/draccus-0.11.6-py3-none-any.whl", hash = "sha256:1cf3f37c64766e0f17d757099b388e762f1d22c772cc2376b4e366b515fb5737", size = 85449, upload-time = "2026-06-12T13:21:18.83Z" },
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -3289,7 +3288,7 @@ requires-dist = [
|
|||||||
{ name = "deepdiff", marker = "extra == 'deepdiff-dep'", specifier = ">=7.0.1,<9.0.0" },
|
{ name = "deepdiff", marker = "extra == 'deepdiff-dep'", specifier = ">=7.0.1,<9.0.0" },
|
||||||
{ name = "diffusers", marker = "extra == 'diffusers-dep'", specifier = ">=0.38.0,<0.40.0" },
|
{ name = "diffusers", marker = "extra == 'diffusers-dep'", specifier = ">=0.38.0,<0.40.0" },
|
||||||
{ name = "dm-tree", marker = "extra == 'groot'", specifier = ">=0.1.8,<1.0.0" },
|
{ name = "dm-tree", marker = "extra == 'groot'", specifier = ">=0.1.8,<1.0.0" },
|
||||||
{ name = "draccus", specifier = "==0.10.0" },
|
{ name = "draccus", specifier = ">=0.11.6,<0.12.0" },
|
||||||
{ name = "dynamixel-sdk", marker = "extra == 'dynamixel'", specifier = ">=3.7.31,<3.9.0" },
|
{ name = "dynamixel-sdk", marker = "extra == 'dynamixel'", specifier = ">=3.7.31,<3.9.0" },
|
||||||
{ name = "einops", specifier = ">=0.8.0,<0.9.0" },
|
{ name = "einops", specifier = ">=0.8.0,<0.9.0" },
|
||||||
{ name = "faker", marker = "extra == 'sarm'", specifier = ">=33.0.0,<35.0.0" },
|
{ name = "faker", marker = "extra == 'sarm'", specifier = ">=33.0.0,<35.0.0" },
|
||||||
@@ -5636,18 +5635,6 @@ wheels = [
|
|||||||
{ url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" },
|
{ url = "https://files.pythonhosted.org/packages/f1/12/de94a39c2ef588c7e6455cfbe7343d3b2dc9d6b6b2f40c4c6565744c873d/pyyaml-6.0.3-cp314-cp314t-win_arm64.whl", hash = "sha256:ebc55a14a21cb14062aa4162f906cd962b28e2e9ea38f9b4391244cd8de4ae0b", size = 149341, upload-time = "2025-09-25T21:32:56.828Z" },
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
|
||||||
name = "pyyaml-include"
|
|
||||||
version = "1.4.1"
|
|
||||||
source = { registry = "https://pypi.org/simple" }
|
|
||||||
dependencies = [
|
|
||||||
{ name = "pyyaml" },
|
|
||||||
]
|
|
||||||
sdist = { url = "https://files.pythonhosted.org/packages/7f/be/2d07ad85e3d593d69640876a8686eae2c533db8cb7bf298d25c421b4d2d5/pyyaml-include-1.4.1.tar.gz", hash = "sha256:1a96e33a99a3e56235f5221273832464025f02ff3d8539309a3bf00dec624471", size = 20592, upload-time = "2024-03-25T14:56:43.748Z" }
|
|
||||||
wheels = [
|
|
||||||
{ url = "https://files.pythonhosted.org/packages/d5/ca/6a2cc3a73170d10b5af1f1613baa2ed1f8f46f62dd0bfab2bffd2c2fe260/pyyaml_include-1.4.1-py3-none-any.whl", hash = "sha256:323c7f3a19c82fbc4d73abbaab7ef4f793e146a13383866831631b26ccc7fb00", size = 19079, upload-time = "2024-03-25T14:56:41.274Z" },
|
|
||||||
]
|
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "pyzmq"
|
name = "pyzmq"
|
||||||
version = "27.1.0"
|
version = "27.1.0"
|
||||||
|
|||||||
Reference in New Issue
Block a user