Compare commits

...

2 Commits

Author SHA1 Message Date
Khalil Meftah 4944abae65 fix(policy): honor revisions when loading PEFT checkpoints 2026-07-28 13:30:03 +02:00
Steven Palma ffe25afb8f fix(processors): wrong feature key dropped in delta-action transform_features (#4165) 2026-07-28 11:18:23 +02:00
3 changed files with 101 additions and 4 deletions
+10 -2
View File
@@ -339,7 +339,10 @@ def make_policy(
logging.info("Loading policy's PEFT adapter.")
peft_pretrained_path = str(cfg.pretrained_path)
peft_config = PeftConfig.from_pretrained(peft_pretrained_path)
peft_config = PeftConfig.from_pretrained(
peft_pretrained_path,
revision=cfg.pretrained_revision,
)
kwargs["pretrained_name_or_path"] = peft_config.base_model_name_or_path
if not kwargs["pretrained_name_or_path"]:
@@ -350,9 +353,14 @@ def make_policy(
"the adapter was trained."
)
kwargs["revision"] = peft_config.revision
policy = policy_cls.from_pretrained(**kwargs)
policy = PeftModel.from_pretrained(
policy, peft_pretrained_path, config=peft_config, is_trainable=True
policy,
peft_pretrained_path,
config=peft_config,
revision=cfg.pretrained_revision,
is_trainable=True,
)
else:
@@ -132,10 +132,20 @@ class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
def transform_features(
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
for axis in ["x", "y", "z", "gripper"]:
for axis in ["x", "y", "z"]:
features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None)
features[PipelineFeatureType.ACTION].pop("gripper", None)
for feat in ["enabled", "target_x", "target_y", "target_z", "target_wx", "target_wy", "target_wz"]:
for feat in [
"enabled",
"target_x",
"target_y",
"target_z",
"target_wx",
"target_wy",
"target_wz",
"gripper_vel",
]:
features[PipelineFeatureType.ACTION][f"{feat}"] = PolicyFeature(
type=FeatureType.ACTION, shape=(1,)
)
@@ -0,0 +1,79 @@
# Copyright 2026 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.
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock
import torch
import lerobot.policies.factory as policy_factory
def test_make_policy_keeps_peft_adapter_and_base_revisions_separate(monkeypatch):
cfg = SimpleNamespace(
type="mock",
device="cpu",
pretrained_path="user/adapter",
pretrained_revision="adapter-sha",
use_peft=True,
input_features={},
output_features={},
)
dataset_meta = SimpleNamespace(features={}, stats={})
base_policy = torch.nn.Linear(1, 1)
policy_from_pretrained = MagicMock(return_value=base_policy)
policy_class = SimpleNamespace(from_pretrained=policy_from_pretrained)
monkeypatch.setattr(policy_factory, "get_policy_class", lambda _: policy_class)
monkeypatch.setattr(policy_factory, "dataset_to_policy_features", lambda _: {})
monkeypatch.setattr(policy_factory, "validate_visual_features_consistency", lambda *args: None)
peft_config = SimpleNamespace(
base_model_name_or_path="user/base-policy",
revision="base-sha",
)
peft_config_from_pretrained = MagicMock(return_value=peft_config)
adapted_policy = torch.nn.Linear(1, 1)
peft_model_from_pretrained = MagicMock(return_value=adapted_policy)
monkeypatch.setitem(
sys.modules,
"peft",
SimpleNamespace(
PeftConfig=SimpleNamespace(from_pretrained=peft_config_from_pretrained),
PeftModel=SimpleNamespace(from_pretrained=peft_model_from_pretrained),
),
)
policy = policy_factory.make_policy(cfg, ds_meta=dataset_meta)
assert policy is adapted_policy
peft_config_from_pretrained.assert_called_once_with(
"user/adapter",
revision="adapter-sha",
)
policy_from_pretrained.assert_called_once_with(
config=cfg,
dataset_stats=dataset_meta.stats,
dataset_meta=dataset_meta,
pretrained_name_or_path="user/base-policy",
revision="base-sha",
)
peft_model_from_pretrained.assert_called_once_with(
base_policy,
"user/adapter",
config=peft_config,
revision="adapter-sha",
is_trainable=True,
)