From 4944abae651e58b294cb875267d1c3cbb176c0bf Mon Sep 17 00:00:00 2001 From: Khalil Meftah Date: Tue, 28 Jul 2026 13:30:03 +0200 Subject: [PATCH] fix(policy): honor revisions when loading PEFT checkpoints --- src/lerobot/policies/factory.py | 12 ++- tests/policies/test_factory_peft_revision.py | 79 ++++++++++++++++++++ 2 files changed, 89 insertions(+), 2 deletions(-) create mode 100644 tests/policies/test_factory_peft_revision.py diff --git a/src/lerobot/policies/factory.py b/src/lerobot/policies/factory.py index 1848a6ffd..2966b51a7 100644 --- a/src/lerobot/policies/factory.py +++ b/src/lerobot/policies/factory.py @@ -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: diff --git a/tests/policies/test_factory_peft_revision.py b/tests/policies/test_factory_peft_revision.py new file mode 100644 index 000000000..1c982800f --- /dev/null +++ b/tests/policies/test_factory_peft_revision.py @@ -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, + )