Compare commits

...

1 Commits

Author SHA1 Message Date
Steven Palma 3011e18c06 fix rollout policy revision loading
Co-authored-by: RaviTeja-Kondeti <rkondet3@asu.edu>
2026-07-27 11:42:05 +02:00
5 changed files with 154 additions and 15 deletions
+1
View File
@@ -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"),
@@ -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,
+11 -2
View File
@@ -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")
+32 -13
View File
@@ -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},
+105
View File
@@ -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
# ---------------------------------------------------------------------------