diff --git a/src/lerobot/policies/factory.py b/src/lerobot/policies/factory.py index 36a0de7ca..1848a6ffd 100644 --- a/src/lerobot/policies/factory.py +++ b/src/lerobot/policies/factory.py @@ -177,6 +177,7 @@ def make_pre_post_processors( return make_groot_pre_post_processors_from_pretrained( config=policy_cfg, pretrained_path=pretrained_path, + revision=pretrained_revision, dataset_stats=kwargs.get("dataset_stats"), dataset_meta=kwargs.get("dataset_meta"), preprocessor_overrides=kwargs.get("preprocessor_overrides"), diff --git a/src/lerobot/policies/groot/processor_groot.py b/src/lerobot/policies/groot/processor_groot.py index 20b3518a3..0bd976a85 100644 --- a/src/lerobot/policies/groot/processor_groot.py +++ b/src/lerobot/policies/groot/processor_groot.py @@ -475,6 +475,7 @@ def make_groot_pre_post_processors_from_pretrained( config: GrootConfig, pretrained_path: str, *, + revision: str | None = None, dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None, dataset_meta: Any | None = None, preprocessor_overrides: dict[str, Any] | None = None, @@ -511,6 +512,7 @@ def make_groot_pre_post_processors_from_pretrained( preprocessor, postprocessor = _load_groot_processor_pipelines( pretrained_path, + revision=revision, preprocessor_overrides=preprocessor_overrides, postprocessor_overrides=postprocessor_overrides, preprocessor_config_filename=preprocessor_config_filename, @@ -526,6 +528,7 @@ def make_groot_pre_post_processors_from_pretrained( def _load_groot_processor_pipelines( pretrained_path: str, *, + revision: str | None, preprocessor_overrides: dict[str, Any], postprocessor_overrides: dict[str, Any], preprocessor_config_filename: str, @@ -540,6 +543,7 @@ def _load_groot_processor_pipelines( preprocessor = PolicyProcessorPipeline.from_pretrained( pretrained_model_name_or_path=pretrained_path, config_filename=preprocessor_config_filename, + revision=revision, overrides=preprocessor_overrides, to_transition=batch_to_transition, to_output=transition_to_batch, @@ -547,6 +551,7 @@ def _load_groot_processor_pipelines( postprocessor = PolicyProcessorPipeline.from_pretrained( pretrained_model_name_or_path=pretrained_path, config_filename=postprocessor_config_filename, + revision=revision, overrides=postprocessor_overrides, to_transition=policy_action_to_transition, to_output=transition_to_policy_action, diff --git a/src/lerobot/rollout/configs.py b/src/lerobot/rollout/configs.py index 639e2ba29..c0b5b345f 100644 --- a/src/lerobot/rollout/configs.py +++ b/src/lerobot/rollout/configs.py @@ -326,8 +326,17 @@ class RolloutConfig: policy_path = parser.get_path_arg("policy") if policy_path: - cli_overrides = parser.get_cli_overrides("policy") - self.policy = PreTrainedConfig.from_pretrained(policy_path, cli_overrides=cli_overrides) + yaml_overrides = parser.get_yaml_overrides("policy") + cli_overrides = parser.get_cli_overrides("policy") or [] + policy_overrides = yaml_overrides + cli_overrides + pretrained_revision = parser.parse_arg("pretrained_revision", cli_overrides) + if pretrained_revision is None: + pretrained_revision = parser.parse_arg("pretrained_revision", yaml_overrides) + self.policy = PreTrainedConfig.from_pretrained( + policy_path, + revision=pretrained_revision, + cli_overrides=policy_overrides, + ) self.policy.pretrained_path = policy_path if self.policy is None: raise ValueError("--policy.path is required for rollout") diff --git a/src/lerobot/rollout/context.py b/src/lerobot/rollout/context.py index 20a7d715a..2bada502e 100644 --- a/src/lerobot/rollout/context.py +++ b/src/lerobot/rollout/context.py @@ -27,7 +27,7 @@ from threading import Event import torch -from lerobot.configs import FeatureType +from lerobot.configs import FeatureType, PreTrainedConfig from lerobot.datasets import ( LeRobotDataset, aggregate_pipeline_dataset_features, @@ -159,6 +159,35 @@ class RolloutContext: # --------------------------------------------------------------------------- +def _load_pretrained_policy(policy_config: PreTrainedConfig) -> PreTrainedPolicy: + """Load policy weights, keeping adapter and base-model revisions independent.""" + pretrained_revision = policy_config.pretrained_revision + policy_class = get_policy_class(policy_config.type) + + if not policy_config.use_peft: + return policy_class.from_pretrained( + policy_config.pretrained_path, + config=policy_config, + revision=pretrained_revision, + ) + + from peft import PeftConfig, PeftModel + + peft_path = policy_config.pretrained_path + peft_config = PeftConfig.from_pretrained(peft_path, revision=pretrained_revision) + policy = policy_class.from_pretrained( + pretrained_name_or_path=peft_config.base_model_name_or_path, + config=policy_config, + revision=peft_config.revision, + ) + return PeftModel.from_pretrained( + policy, + peft_path, + config=peft_config, + revision=pretrained_revision, + ) + + def build_rollout_context( cfg: RolloutConfig, shutdown_event: Event, @@ -176,7 +205,6 @@ def build_rollout_context( # --- 1. Policy (heavy I/O, but no hardware yet) ------------------- logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path) policy_config = cfg.policy - policy_class = get_policy_class(policy_config.type) if hasattr(policy_config, "compile_model"): policy_config.compile_model = cfg.use_torch_compile @@ -187,17 +215,7 @@ def build_rollout_context( "Please use `cpu` or `cuda` backend." ) - if policy_config.use_peft: - from peft import PeftConfig, PeftModel - - peft_path = policy_config.pretrained_path - peft_config = PeftConfig.from_pretrained(peft_path) - policy = policy_class.from_pretrained( - pretrained_name_or_path=peft_config.base_model_name_or_path, config=policy_config - ) - policy = PeftModel.from_pretrained(policy, peft_path, config=peft_config) - else: - policy = policy_class.from_pretrained(policy_config.pretrained_path, config=policy_config) + policy = _load_pretrained_policy(policy_config) if is_rtc: policy.config.rtc_config = cfg.inference.rtc @@ -392,6 +410,7 @@ def build_rollout_context( preprocessor, postprocessor = make_pre_post_processors( policy_cfg=policy_config, pretrained_path=cfg.policy.pretrained_path, + pretrained_revision=policy_config.pretrained_revision, dataset_stats=dataset_stats, preprocessor_overrides={ "device_processor": {"device": cfg.device}, diff --git a/tests/test_rollout.py b/tests/test_rollout.py index 85a29ff4c..c247ff3c1 100644 --- a/tests/test_rollout.py +++ b/tests/test_rollout.py @@ -17,6 +17,8 @@ from __future__ import annotations import dataclasses +import sys +from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -106,6 +108,109 @@ def test_sentry_config_defaults(): assert cfg.target_video_file_size_mb is None +def test_rollout_config_passes_policy_pretrained_revision(monkeypatch): + from lerobot.configs import PreTrainedConfig, parser + from lerobot.rollout import RolloutConfig + from tests.mocks.mock_robot import MockRobotConfig + + captured = {} + + def fake_from_pretrained(cls, pretrained_name_or_path, **kwargs): + captured["pretrained_name_or_path"] = pretrained_name_or_path + captured.update(kwargs) + return SimpleNamespace(device="cpu", pretrained_revision=kwargs["revision"]) + + monkeypatch.setattr(parser, "get_yaml_overrides", lambda _: ["--pretrained_revision=yaml-sha"]) + monkeypatch.setattr( + sys, + "argv", + ["lerobot-rollout", "--policy.path=user/policy", "--policy.pretrained_revision=cli-sha"], + ) + monkeypatch.setattr(PreTrainedConfig, "from_pretrained", classmethod(fake_from_pretrained)) + + cfg = RolloutConfig(robot=MockRobotConfig()) + + assert captured["pretrained_name_or_path"] == "user/policy" + assert captured["revision"] == "cli-sha" + assert captured["cli_overrides"] == [ + "--pretrained_revision=yaml-sha", + "--pretrained_revision=cli-sha", + ] + assert cfg.policy.pretrained_path == "user/policy" + assert cfg.policy.pretrained_revision == "cli-sha" + + +def test_load_pretrained_policy_passes_revision(monkeypatch): + import lerobot.rollout.context as rollout_context + + policy_config = SimpleNamespace( + type="mock", + use_peft=False, + pretrained_path="user/policy", + pretrained_revision="policy-sha", + ) + policy_class = MagicMock() + loaded_policy = MagicMock() + policy_class.from_pretrained.return_value = loaded_policy + monkeypatch.setattr(rollout_context, "get_policy_class", lambda _: policy_class) + + policy = rollout_context._load_pretrained_policy(policy_config) + + assert policy is loaded_policy + policy_class.from_pretrained.assert_called_once_with( + "user/policy", + config=policy_config, + revision="policy-sha", + ) + + +def test_load_pretrained_peft_policy_keeps_adapter_and_base_revisions_separate(monkeypatch): + import lerobot.rollout.context as rollout_context + + policy_config = SimpleNamespace( + type="mock", + use_peft=True, + pretrained_path="user/adapter", + pretrained_revision="adapter-sha", + ) + policy_class = MagicMock() + base_policy = MagicMock() + policy_class.from_pretrained.return_value = base_policy + monkeypatch.setattr(rollout_context, "get_policy_class", lambda _: policy_class) + + 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 = MagicMock() + 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 = rollout_context._load_pretrained_policy(policy_config) + + assert policy is adapted_policy + peft_config_from_pretrained.assert_called_once_with("user/adapter", revision="adapter-sha") + policy_class.from_pretrained.assert_called_once_with( + pretrained_name_or_path="user/base-policy", + config=policy_config, + revision="base-sha", + ) + peft_model_from_pretrained.assert_called_once_with( + base_policy, + "user/adapter", + config=peft_config, + revision="adapter-sha", + ) + + # --------------------------------------------------------------------------- # RolloutRingBuffer # ---------------------------------------------------------------------------