mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
fix(g05): rebuild camera_sizes when stale keys survive embodiment switch
Switching embodiment on a packaged checkpoint keeps the previous embodiment's camera_sizes entries because dict config values merge instead of replacing, which tripped the exact-key validation. Rebuild sizes from the named profile whenever the keys diverge from camera_order, and add regression tests for both fixes. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -246,7 +246,10 @@ class G05Config(PreTrainedConfig):
|
||||
raise ValueError("normalization_clip must be an increasing (minimum, maximum) pair.")
|
||||
self.relative_exclude_joints = tuple(self.relative_exclude_joints)
|
||||
self.action_feature_names = tuple(self.action_feature_names)
|
||||
if not self.camera_sizes and self.embodiment in G05_CAMERA_SIZE_PROFILES:
|
||||
if self.embodiment in G05_CAMERA_SIZE_PROFILES and set(self.camera_sizes) != set(self.camera_order):
|
||||
# Switching embodiment on a packaged checkpoint keeps the previous
|
||||
# embodiment's camera entries (config-file/CLI dict values merge
|
||||
# instead of replacing); rebuild sizes from the named profile.
|
||||
self.camera_sizes = G05_CAMERA_SIZE_PROFILES[self.embodiment].copy()
|
||||
if self.num_input_images == 0:
|
||||
self.num_input_images = len(self.camera_order) * self.n_obs_steps
|
||||
|
||||
@@ -13,7 +13,7 @@ from torch import nn
|
||||
from lerobot.configs.policies import PreTrainedConfig
|
||||
from lerobot.configs.types import FeatureType, PolicyFeature
|
||||
from lerobot.policies.factory import get_policy_class, make_policy_config, make_pre_post_processors
|
||||
from lerobot.policies.g05.configuration_g05 import G05_EMBODIMENT_MAPPINGS, G05Config
|
||||
from lerobot.policies.g05.configuration_g05 import G05_CAMERA_PROFILES, G05_EMBODIMENT_MAPPINGS, G05Config
|
||||
from lerobot.policies.g05.modeling_g05 import G05Policy
|
||||
from lerobot.policies.g05.native_g05 import G05_RUNTIME_PREDICT_COT, G05NativeBackend
|
||||
from lerobot.processor import PolicyProcessorPipeline
|
||||
@@ -830,3 +830,30 @@ def test_gated_checkpoint_loads_strictly():
|
||||
checkpoint = Path(os.environ["LEROBOT_G05_CHECKPOINT"])
|
||||
policy = G05Policy.from_pretrained(checkpoint, local_files_only=True, strict=True)
|
||||
assert policy.config.source_checkpoint_revision
|
||||
|
||||
|
||||
def test_project_stats_passes_dataset_count_through():
|
||||
config = _config(normalization_mode="q01_q99")
|
||||
stats = {
|
||||
OBS_STATE: {"q01": torch.zeros(7), "q99": torch.ones(7), "count": torch.tensor([100])},
|
||||
ACTION: {"q01": torch.zeros(7), "q99": torch.ones(7), "count": torch.tensor([100])},
|
||||
}
|
||||
|
||||
make_pre_post_processors(config, dataset_stats=stats)
|
||||
|
||||
|
||||
def test_named_embodiment_rebuilds_stale_camera_sizes():
|
||||
config = G05Config(
|
||||
checkpoint_profile="custom",
|
||||
embodiment="robotwin",
|
||||
raw_state_dim=14,
|
||||
raw_action_dim=14,
|
||||
camera_order=G05_CAMERA_PROFILES["robotwin"],
|
||||
camera_sizes={
|
||||
**dict.fromkeys(G05_CAMERA_PROFILES["robotwin"], (256, 256)),
|
||||
"observation.images.stale_camera": (256, 256),
|
||||
},
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
assert set(config.camera_sizes) == set(G05_CAMERA_PROFILES["robotwin"])
|
||||
|
||||
Reference in New Issue
Block a user