mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 04:36:04 +00:00
fix(hub): pin pretrained artifacts to one commit
This commit is contained in:
@@ -48,8 +48,13 @@ class EvalPipelineConfig:
|
|||||||
if policy_path:
|
if policy_path:
|
||||||
yaml_overrides = parser.get_yaml_overrides("policy")
|
yaml_overrides = parser.get_yaml_overrides("policy")
|
||||||
cli_overrides = parser.get_cli_overrides("policy") or []
|
cli_overrides = parser.get_cli_overrides("policy") or []
|
||||||
|
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(
|
self.policy = PreTrainedConfig.from_pretrained(
|
||||||
policy_path, cli_overrides=yaml_overrides + cli_overrides
|
policy_path,
|
||||||
|
revision=pretrained_revision,
|
||||||
|
cli_overrides=yaml_overrides + cli_overrides,
|
||||||
)
|
)
|
||||||
self.policy.pretrained_path = Path(policy_path)
|
self.policy.pretrained_path = Path(policy_path)
|
||||||
|
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ from huggingface_hub.errors import HfHubHTTPError
|
|||||||
from lerobot.optim import LRSchedulerConfig, OptimizerConfig
|
from lerobot.optim import LRSchedulerConfig, OptimizerConfig
|
||||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||||
from lerobot.utils.device_utils import auto_select_torch_device, is_amp_available, is_torch_device_available
|
from lerobot.utils.device_utils import auto_select_torch_device, is_amp_available, is_torch_device_available
|
||||||
from lerobot.utils.hub import HubMixin
|
from lerobot.utils.hub import HubMixin, extract_commit_hash
|
||||||
|
|
||||||
from .types import FeatureType, PolicyFeature
|
from .types import FeatureType, PolicyFeature
|
||||||
|
|
||||||
@@ -82,6 +82,26 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
|||||||
# Optional Hub revision (commit hash, branch, or tag) to pin the pretrained model version.
|
# Optional Hub revision (commit hash, branch, or tag) to pin the pretrained model version.
|
||||||
pretrained_revision: str | None = None
|
pretrained_revision: str | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _commit_hash(self) -> str | None:
|
||||||
|
"""Resolved Hub commit for this runtime load; never serialized."""
|
||||||
|
return self.__dict__.get("_runtime_commit_hash")
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _commit_hash_source(self) -> str | None:
|
||||||
|
"""Hub repo whose revision resolved to ``_commit_hash``."""
|
||||||
|
return self.__dict__.get("_runtime_commit_hash_source")
|
||||||
|
|
||||||
|
def _set_hub_commit_hash(self, commit_hash: str | None, source: str | None) -> None:
|
||||||
|
self.__dict__["_runtime_commit_hash"] = commit_hash
|
||||||
|
self.__dict__["_runtime_commit_hash_source"] = source if commit_hash is not None else None
|
||||||
|
|
||||||
|
def get_hub_revision(self, source: str | Path | None, revision: str | None = None) -> str | None:
|
||||||
|
"""Return the pinned revision when ``source`` owns the resolved commit."""
|
||||||
|
if self._commit_hash is not None and self._commit_hash_source == str(source):
|
||||||
|
return self._commit_hash
|
||||||
|
return revision
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
if not self.device or not is_torch_device_available(self.device):
|
if not self.device or not is_torch_device_available(self.device):
|
||||||
auto_device = auto_select_torch_device()
|
auto_device = auto_select_torch_device()
|
||||||
@@ -182,7 +202,8 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
|||||||
) -> T:
|
) -> T:
|
||||||
model_id = str(pretrained_name_or_path)
|
model_id = str(pretrained_name_or_path)
|
||||||
config_file: str | None = None
|
config_file: str | None = None
|
||||||
if Path(model_id).is_dir():
|
is_local = Path(model_id).is_dir()
|
||||||
|
if is_local:
|
||||||
if CONFIG_NAME in os.listdir(model_id):
|
if CONFIG_NAME in os.listdir(model_id):
|
||||||
config_file = os.path.join(model_id, CONFIG_NAME)
|
config_file = os.path.join(model_id, CONFIG_NAME)
|
||||||
else:
|
else:
|
||||||
@@ -208,8 +229,12 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
|||||||
if config_file is None:
|
if config_file is None:
|
||||||
raise FileNotFoundError(f"{CONFIG_NAME} not found in {model_id}")
|
raise FileNotFoundError(f"{CONFIG_NAME} not found in {model_id}")
|
||||||
|
|
||||||
|
commit_hash = None if is_local else extract_commit_hash(config_file, revision)
|
||||||
with open(config_file) as f:
|
with open(config_file) as f:
|
||||||
config = json.load(f)
|
config = json.load(f)
|
||||||
|
# Runtime Hub metadata must never become part of the serialized config schema.
|
||||||
|
config.pop("_commit_hash", None)
|
||||||
|
config.pop("_commit_hash_source", None)
|
||||||
|
|
||||||
# Resolve the concrete config subclass from the serialized "type" tag, then parse
|
# Resolve the concrete config subclass from the serialized "type" tag, then parse
|
||||||
# the config (with CLI overrides) directly for that class. The "type" key is
|
# the config (with CLI overrides) directly for that class. The "type" key is
|
||||||
@@ -231,4 +256,6 @@ class PreTrainedConfig(draccus.ChoiceRegistry, HubMixin, abc.ABC): # type: igno
|
|||||||
|
|
||||||
cli_overrides = policy_kwargs.pop("cli_overrides", [])
|
cli_overrides = policy_kwargs.pop("cli_overrides", [])
|
||||||
with draccus.config_type("json"):
|
with draccus.config_type("json"):
|
||||||
return draccus.parse(config_cls, config_file, args=cli_overrides)
|
parsed_config = draccus.parse(config_cls, config_file, args=cli_overrides)
|
||||||
|
parsed_config._set_hub_commit_hash(commit_hash, model_id)
|
||||||
|
return parsed_config
|
||||||
|
|||||||
@@ -172,8 +172,16 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
)
|
)
|
||||||
self.reward_model.pretrained_path = str(Path(reward_model_path))
|
self.reward_model.pretrained_path = str(Path(reward_model_path))
|
||||||
elif policy_path:
|
elif policy_path:
|
||||||
overrides = parser.get_yaml_overrides("policy") + (parser.get_cli_overrides("policy") or [])
|
yaml_overrides = parser.get_yaml_overrides("policy")
|
||||||
self.policy = PreTrainedConfig.from_pretrained(policy_path, cli_overrides=overrides)
|
cli_overrides = parser.get_cli_overrides("policy") or []
|
||||||
|
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=yaml_overrides + cli_overrides,
|
||||||
|
)
|
||||||
self.policy.pretrained_path = Path(policy_path)
|
self.policy.pretrained_path = Path(policy_path)
|
||||||
elif self.resume:
|
elif self.resume:
|
||||||
self._resolve_resume_checkpoint()
|
self._resolve_resume_checkpoint()
|
||||||
|
|||||||
@@ -171,6 +171,9 @@ def make_pre_post_processors(
|
|||||||
ValueError: If no processor factory exists for the given policy configuration type.
|
ValueError: If no processor factory exists for the given policy configuration type.
|
||||||
"""
|
"""
|
||||||
if pretrained_path:
|
if pretrained_path:
|
||||||
|
revision_resolver = getattr(policy_cfg, "get_hub_revision", None)
|
||||||
|
if callable(revision_resolver):
|
||||||
|
pretrained_revision = revision_resolver(pretrained_path, pretrained_revision)
|
||||||
if isinstance(policy_cfg, GrootConfig):
|
if isinstance(policy_cfg, GrootConfig):
|
||||||
from .groot.processor_groot import make_groot_pre_post_processors_from_pretrained
|
from .groot.processor_groot import make_groot_pre_post_processors_from_pretrained
|
||||||
|
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ from torch import Tensor
|
|||||||
|
|
||||||
from lerobot.configs import FeatureType, PolicyFeature
|
from lerobot.configs import FeatureType, PolicyFeature
|
||||||
from lerobot.utils.constants import ACTION, OBS_IMAGES
|
from lerobot.utils.constants import ACTION, OBS_IMAGES
|
||||||
|
from lerobot.utils.hub import extract_commit_hash
|
||||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||||
|
|
||||||
from ..pretrained import PreTrainedPolicy
|
from ..pretrained import PreTrainedPolicy
|
||||||
@@ -195,6 +196,8 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
)
|
)
|
||||||
|
|
||||||
model_id = str(pretrained_name_or_path)
|
model_id = str(pretrained_name_or_path)
|
||||||
|
if config is not None:
|
||||||
|
revision = config.get_hub_revision(model_id, revision)
|
||||||
is_finetuned_checkpoint = False
|
is_finetuned_checkpoint = False
|
||||||
|
|
||||||
# Check if this is a fine-tuned LeRobot checkpoint (has model.safetensors)
|
# Check if this is a fine-tuned LeRobot checkpoint (has model.safetensors)
|
||||||
@@ -204,7 +207,7 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
else:
|
else:
|
||||||
# Try to download the safetensors file to check if it exists
|
# Try to download the safetensors file to check if it exists
|
||||||
try:
|
try:
|
||||||
hf_hub_download(
|
resolved_model_file = hf_hub_download(
|
||||||
repo_id=model_id,
|
repo_id=model_id,
|
||||||
filename=SAFETENSORS_SINGLE_FILE,
|
filename=SAFETENSORS_SINGLE_FILE,
|
||||||
revision=revision,
|
revision=revision,
|
||||||
@@ -214,6 +217,10 @@ class GrootPolicy(PreTrainedPolicy):
|
|||||||
token=token,
|
token=token,
|
||||||
local_files_only=local_files_only,
|
local_files_only=local_files_only,
|
||||||
)
|
)
|
||||||
|
resolved_commit_hash = extract_commit_hash(resolved_model_file, revision)
|
||||||
|
revision = resolved_commit_hash or revision
|
||||||
|
if config is not None and config._commit_hash is None:
|
||||||
|
config._set_hub_commit_hash(resolved_commit_hash, model_id)
|
||||||
is_finetuned_checkpoint = True
|
is_finetuned_checkpoint = True
|
||||||
except HfHubHTTPError:
|
except HfHubHTTPError:
|
||||||
is_finetuned_checkpoint = False
|
is_finetuned_checkpoint = False
|
||||||
|
|||||||
@@ -53,6 +53,7 @@ from lerobot.utils.constants import (
|
|||||||
OBS_LANGUAGE_TOKENS,
|
OBS_LANGUAGE_TOKENS,
|
||||||
OBS_STATE,
|
OBS_STATE,
|
||||||
)
|
)
|
||||||
|
from lerobot.utils.hub import extract_commit_hash
|
||||||
|
|
||||||
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
||||||
from ..common.vla_utils import (
|
from ..common.vla_utils import (
|
||||||
@@ -814,6 +815,8 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
revision = config.get_hub_revision(pretrained_name_or_path, revision)
|
||||||
|
|
||||||
# Initialize model without loading weights
|
# Initialize model without loading weights
|
||||||
# Check if dataset_stats were provided in kwargs
|
# Check if dataset_stats were provided in kwargs
|
||||||
model = cls(config, **kwargs)
|
model = cls(config, **kwargs)
|
||||||
@@ -832,9 +835,13 @@ class PI0Policy(PreTrainedPolicy):
|
|||||||
resume_download=kwargs.get("resume_download"),
|
resume_download=kwargs.get("resume_download"),
|
||||||
proxies=kwargs.get("proxies"),
|
proxies=kwargs.get("proxies"),
|
||||||
token=kwargs.get("token"),
|
token=kwargs.get("token"),
|
||||||
revision=kwargs.get("revision"),
|
revision=revision,
|
||||||
local_files_only=kwargs.get("local_files_only", False),
|
local_files_only=kwargs.get("local_files_only", False),
|
||||||
)
|
)
|
||||||
|
if config._commit_hash is None:
|
||||||
|
config._set_hub_commit_hash(
|
||||||
|
extract_commit_hash(resolved_file, revision), str(pretrained_name_or_path)
|
||||||
|
)
|
||||||
from safetensors.torch import load_file
|
from safetensors.torch import load_file
|
||||||
|
|
||||||
original_state_dict = load_file(resolved_file)
|
original_state_dict = load_file(resolved_file)
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ from lerobot.utils.constants import (
|
|||||||
OBS_LANGUAGE_ATTENTION_MASK,
|
OBS_LANGUAGE_ATTENTION_MASK,
|
||||||
OBS_LANGUAGE_TOKENS,
|
OBS_LANGUAGE_TOKENS,
|
||||||
)
|
)
|
||||||
|
from lerobot.utils.hub import extract_commit_hash
|
||||||
|
|
||||||
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
||||||
from ..common.vla_utils import (
|
from ..common.vla_utils import (
|
||||||
@@ -779,6 +780,8 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
revision = config.get_hub_revision(pretrained_name_or_path, revision)
|
||||||
|
|
||||||
# Initialize model without loading weights
|
# Initialize model without loading weights
|
||||||
# Check if dataset_stats were provided in kwargs
|
# Check if dataset_stats were provided in kwargs
|
||||||
model = cls(config, **kwargs)
|
model = cls(config, **kwargs)
|
||||||
@@ -797,9 +800,13 @@ class PI05Policy(PreTrainedPolicy):
|
|||||||
resume_download=kwargs.get("resume_download"),
|
resume_download=kwargs.get("resume_download"),
|
||||||
proxies=kwargs.get("proxies"),
|
proxies=kwargs.get("proxies"),
|
||||||
token=kwargs.get("token"),
|
token=kwargs.get("token"),
|
||||||
revision=kwargs.get("revision"),
|
revision=revision,
|
||||||
local_files_only=kwargs.get("local_files_only", False),
|
local_files_only=kwargs.get("local_files_only", False),
|
||||||
)
|
)
|
||||||
|
if config._commit_hash is None:
|
||||||
|
config._set_hub_commit_hash(
|
||||||
|
extract_commit_hash(resolved_file, revision), str(pretrained_name_or_path)
|
||||||
|
)
|
||||||
from safetensors.torch import load_file
|
from safetensors.torch import load_file
|
||||||
|
|
||||||
original_state_dict = load_file(resolved_file)
|
original_state_dict = load_file(resolved_file)
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ from lerobot.utils.constants import (
|
|||||||
OBS_LANGUAGE_ATTENTION_MASK,
|
OBS_LANGUAGE_ATTENTION_MASK,
|
||||||
OBS_LANGUAGE_TOKENS,
|
OBS_LANGUAGE_TOKENS,
|
||||||
)
|
)
|
||||||
|
from lerobot.utils.hub import extract_commit_hash
|
||||||
|
|
||||||
from ..common.vla_utils import pad_vector, prepare_attention_masks_4d, resize_with_pad_torch
|
from ..common.vla_utils import pad_vector, prepare_attention_masks_4d, resize_with_pad_torch
|
||||||
from ..pretrained import PreTrainedPolicy, T
|
from ..pretrained import PreTrainedPolicy, T
|
||||||
@@ -808,6 +809,8 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
revision = config.get_hub_revision(pretrained_name_or_path, revision)
|
||||||
|
|
||||||
# Initialize model without loading weights
|
# Initialize model without loading weights
|
||||||
# Check if dataset_stats were provided in kwargs
|
# Check if dataset_stats were provided in kwargs
|
||||||
model = cls(config, **kwargs)
|
model = cls(config, **kwargs)
|
||||||
@@ -826,9 +829,13 @@ class PI0FastPolicy(PreTrainedPolicy):
|
|||||||
resume_download=kwargs.get("resume_download"),
|
resume_download=kwargs.get("resume_download"),
|
||||||
proxies=kwargs.get("proxies"),
|
proxies=kwargs.get("proxies"),
|
||||||
token=kwargs.get("token"),
|
token=kwargs.get("token"),
|
||||||
revision=kwargs.get("revision"),
|
revision=revision,
|
||||||
local_files_only=kwargs.get("local_files_only", False),
|
local_files_only=kwargs.get("local_files_only", False),
|
||||||
)
|
)
|
||||||
|
if config._commit_hash is None:
|
||||||
|
config._set_hub_commit_hash(
|
||||||
|
extract_commit_hash(resolved_file, revision), str(pretrained_name_or_path)
|
||||||
|
)
|
||||||
from safetensors.torch import load_file
|
from safetensors.torch import load_file
|
||||||
|
|
||||||
original_state_dict = load_file(resolved_file)
|
original_state_dict = load_file(resolved_file)
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ from lerobot.__version__ import __version__
|
|||||||
from lerobot.configs import PreTrainedConfig
|
from lerobot.configs import PreTrainedConfig
|
||||||
from lerobot.configs.train import TrainPipelineConfig
|
from lerobot.configs.train import TrainPipelineConfig
|
||||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||||
from lerobot.utils.hub import HubMixin
|
from lerobot.utils.hub import HubMixin, extract_commit_hash
|
||||||
|
|
||||||
from .utils import log_model_loading_keys
|
from .utils import log_model_loading_keys
|
||||||
|
|
||||||
@@ -190,6 +190,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
model_id = str(pretrained_name_or_path)
|
model_id = str(pretrained_name_or_path)
|
||||||
|
revision = config.get_hub_revision(model_id, revision)
|
||||||
instance = cls(config, **kwargs)
|
instance = cls(config, **kwargs)
|
||||||
if os.path.isdir(model_id):
|
if os.path.isdir(model_id):
|
||||||
print("Loading weights from local directory")
|
print("Loading weights from local directory")
|
||||||
@@ -208,6 +209,8 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
token=token,
|
token=token,
|
||||||
local_files_only=local_files_only,
|
local_files_only=local_files_only,
|
||||||
)
|
)
|
||||||
|
if config._commit_hash is None:
|
||||||
|
config._set_hub_commit_hash(extract_commit_hash(model_file, revision), model_id)
|
||||||
policy = cls._load_as_safetensor(instance, model_file, config.device, strict)
|
policy = cls._load_as_safetensor(instance, model_file, config.device, strict)
|
||||||
except HfHubHTTPError as e:
|
except HfHubHTTPError as e:
|
||||||
raise FileNotFoundError(
|
raise FileNotFoundError(
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from torch import Tensor, nn
|
|||||||
|
|
||||||
from lerobot.configs import PreTrainedConfig
|
from lerobot.configs import PreTrainedConfig
|
||||||
from lerobot.utils.constants import ACTION, OBS_LANGUAGE_TOKENS, OBS_STATE
|
from lerobot.utils.constants import ACTION, OBS_LANGUAGE_TOKENS, OBS_STATE
|
||||||
|
from lerobot.utils.hub import extract_commit_hash
|
||||||
from lerobot.utils.import_utils import _transformers_available, require_package
|
from lerobot.utils.import_utils import _transformers_available, require_package
|
||||||
|
|
||||||
from ..common.vla_utils import pad_vector, resize_with_pad
|
from ..common.vla_utils import pad_vector, resize_with_pad
|
||||||
@@ -459,6 +460,7 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
)
|
)
|
||||||
|
|
||||||
model_id = str(pretrained_name_or_path)
|
model_id = str(pretrained_name_or_path)
|
||||||
|
revision = config.get_hub_revision(model_id, revision)
|
||||||
instance = cls(config, **kwargs)
|
instance = cls(config, **kwargs)
|
||||||
# step 2: locate model.safetensors
|
# step 2: locate model.safetensors
|
||||||
if os.path.isdir(model_id):
|
if os.path.isdir(model_id):
|
||||||
@@ -480,6 +482,8 @@ class XVLAPolicy(PreTrainedPolicy):
|
|||||||
token=token,
|
token=token,
|
||||||
local_files_only=local_files_only,
|
local_files_only=local_files_only,
|
||||||
)
|
)
|
||||||
|
if config._commit_hash is None:
|
||||||
|
config._set_hub_commit_hash(extract_commit_hash(model_file, revision), model_id)
|
||||||
except HfHubHTTPError as e:
|
except HfHubHTTPError as e:
|
||||||
raise FileNotFoundError(f"model.safetensors not found on the Hub at {model_id}") from e
|
raise FileNotFoundError(f"model.safetensors not found on the Hub at {model_id}") from e
|
||||||
|
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ from safetensors.torch import load_file, save_file
|
|||||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
||||||
from lerobot.types import EnvAction, EnvTransition, PolicyAction, RobotAction, RobotObservation, TransitionKey
|
from lerobot.types import EnvAction, EnvTransition, PolicyAction, RobotAction, RobotObservation, TransitionKey
|
||||||
from lerobot.utils.constants import HF_LEROBOT_HOME
|
from lerobot.utils.constants import HF_LEROBOT_HOME
|
||||||
from lerobot.utils.hub import HubMixin
|
from lerobot.utils.hub import HubMixin, extract_commit_hash
|
||||||
|
|
||||||
from .converters import batch_to_transition, create_transition, transition_to_batch
|
from .converters import batch_to_transition, create_transition, transition_to_batch
|
||||||
|
|
||||||
@@ -727,6 +727,10 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
|
|
||||||
# 1. Load configuration using simplified 3-way logic
|
# 1. Load configuration using simplified 3-way logic
|
||||||
loaded_config, base_path = cls._load_config(model_id, config_filename, hub_download_kwargs)
|
loaded_config, base_path = cls._load_config(model_id, config_filename, hub_download_kwargs)
|
||||||
|
if not is_local_source:
|
||||||
|
commit_hash = extract_commit_hash(base_path, revision)
|
||||||
|
if commit_hash is not None:
|
||||||
|
hub_download_kwargs["revision"] = commit_hash
|
||||||
|
|
||||||
# 2. Validate configuration and handle migration
|
# 2. Validate configuration and handle migration
|
||||||
cls._validate_loaded_config(model_id, loaded_config, config_filename)
|
cls._validate_loaded_config(model_id, loaded_config, config_filename)
|
||||||
|
|||||||
@@ -161,7 +161,12 @@ class RolloutContext:
|
|||||||
|
|
||||||
def _load_pretrained_policy(policy_config: PreTrainedConfig) -> PreTrainedPolicy:
|
def _load_pretrained_policy(policy_config: PreTrainedConfig) -> PreTrainedPolicy:
|
||||||
"""Load policy weights, keeping adapter and base-model revisions independent."""
|
"""Load policy weights, keeping adapter and base-model revisions independent."""
|
||||||
pretrained_revision = policy_config.pretrained_revision
|
revision_resolver = getattr(policy_config, "get_hub_revision", None)
|
||||||
|
pretrained_revision = (
|
||||||
|
revision_resolver(policy_config.pretrained_path, policy_config.pretrained_revision)
|
||||||
|
if callable(revision_resolver)
|
||||||
|
else policy_config.pretrained_revision
|
||||||
|
)
|
||||||
policy_class = get_policy_class(policy_config.type)
|
policy_class = get_policy_class(policy_config.type)
|
||||||
|
|
||||||
if not policy_config.use_peft:
|
if not policy_config.use_peft:
|
||||||
|
|||||||
@@ -13,6 +13,7 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
import builtins
|
import builtins
|
||||||
|
import re
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from tempfile import TemporaryDirectory
|
from tempfile import TemporaryDirectory
|
||||||
from typing import Any, TypeVar
|
from typing import Any, TypeVar
|
||||||
@@ -23,6 +24,29 @@ from huggingface_hub.utils import validate_hf_hub_args
|
|||||||
from .constants import CHECKPOINTS_DIR
|
from .constants import CHECKPOINTS_DIR
|
||||||
|
|
||||||
T = TypeVar("T", bound="HubMixin")
|
T = TypeVar("T", bound="HubMixin")
|
||||||
|
REGEX_COMMIT_HASH = re.compile(r"^[0-9a-f]{40}$")
|
||||||
|
|
||||||
|
|
||||||
|
def extract_commit_hash(resolved_file: str | Path | None, revision: str | None = None) -> str | None:
|
||||||
|
"""Extract the immutable commit hash backing a resolved Hub file.
|
||||||
|
|
||||||
|
Hub cache paths contain ``snapshots/<commit_hash>/``. If the requested
|
||||||
|
revision is already a full commit hash, use it as a fallback for custom
|
||||||
|
cache layouts that do not expose the standard snapshot path.
|
||||||
|
"""
|
||||||
|
if resolved_file is not None:
|
||||||
|
path_parts = Path(resolved_file).parts
|
||||||
|
try:
|
||||||
|
snapshot_index = path_parts.index("snapshots")
|
||||||
|
commit_hash = path_parts[snapshot_index + 1]
|
||||||
|
if REGEX_COMMIT_HASH.fullmatch(commit_hash):
|
||||||
|
return commit_hash
|
||||||
|
except (ValueError, IndexError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
if revision is not None and REGEX_COMMIT_HASH.fullmatch(revision):
|
||||||
|
return revision
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def find_latest_hub_checkpoint(
|
def find_latest_hub_checkpoint(
|
||||||
|
|||||||
@@ -0,0 +1,77 @@
|
|||||||
|
# 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 json
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from lerobot.configs import PreTrainedConfig
|
||||||
|
|
||||||
|
|
||||||
|
@PreTrainedConfig.register_subclass("revision_pinning_test")
|
||||||
|
@dataclass
|
||||||
|
class RevisionPinningTestConfig(PreTrainedConfig):
|
||||||
|
@property
|
||||||
|
def observation_delta_indices(self) -> list | None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def action_delta_indices(self) -> list | None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def reward_delta_indices(self) -> list | None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
def get_optimizer_preset(self):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def get_scheduler_preset(self):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def validate_features(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def test_pretrained_config_pins_resolved_hub_commit(monkeypatch, tmp_path):
|
||||||
|
commit_hash = "a" * 40
|
||||||
|
snapshot_dir = tmp_path / "models--user--policy" / "snapshots" / commit_hash
|
||||||
|
RevisionPinningTestConfig(device="cpu").save_pretrained(snapshot_dir)
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_hub_download(**kwargs):
|
||||||
|
calls.append(kwargs)
|
||||||
|
return str(snapshot_dir / "config.json")
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.configs.policies.hf_hub_download", fake_hub_download)
|
||||||
|
|
||||||
|
config = PreTrainedConfig.from_pretrained("user/policy", revision="main")
|
||||||
|
|
||||||
|
assert calls[0]["revision"] == "main"
|
||||||
|
assert config._commit_hash == commit_hash
|
||||||
|
assert config._commit_hash_source == "user/policy"
|
||||||
|
assert config.get_hub_revision("user/policy", "main") == commit_hash
|
||||||
|
assert config.get_hub_revision("user/base-policy", "base-tag") == "base-tag"
|
||||||
|
|
||||||
|
|
||||||
|
def test_runtime_commit_hash_is_not_serialized(tmp_path):
|
||||||
|
config = RevisionPinningTestConfig(device="cpu")
|
||||||
|
config._set_hub_commit_hash("a" * 40, "user/policy")
|
||||||
|
|
||||||
|
config.save_pretrained(tmp_path)
|
||||||
|
|
||||||
|
serialized_config = json.loads((tmp_path / "config.json").read_text())
|
||||||
|
assert "_commit_hash" not in serialized_config
|
||||||
|
assert "_commit_hash_source" not in serialized_config
|
||||||
|
assert "_runtime_commit_hash" not in serialized_config
|
||||||
|
assert "_runtime_commit_hash_source" not in serialized_config
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
# 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 torch
|
||||||
|
|
||||||
|
from lerobot.policies import factory
|
||||||
|
from lerobot.policies.act.configuration_act import ACTConfig
|
||||||
|
from lerobot.policies.pretrained import PreTrainedPolicy
|
||||||
|
from lerobot.processor import PolicyProcessorPipeline
|
||||||
|
|
||||||
|
|
||||||
|
class _MinimalPolicy(PreTrainedPolicy):
|
||||||
|
config_class = ACTConfig
|
||||||
|
name = "minimal_revision_test"
|
||||||
|
|
||||||
|
def get_optim_params(self) -> dict:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def forward(self, batch: dict[str, torch.Tensor]) -> tuple[torch.Tensor, dict | None]:
|
||||||
|
return torch.tensor(0), None
|
||||||
|
|
||||||
|
def predict_action_chunk(self, batch: dict[str, torch.Tensor], **kwargs) -> torch.Tensor:
|
||||||
|
return torch.tensor(0)
|
||||||
|
|
||||||
|
def select_action(self, batch: dict[str, torch.Tensor], **kwargs) -> torch.Tensor:
|
||||||
|
return torch.tensor(0)
|
||||||
|
|
||||||
|
|
||||||
|
def _skip_safetensor_loading(monkeypatch):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
_MinimalPolicy,
|
||||||
|
"_load_as_safetensor",
|
||||||
|
classmethod(lambda cls, model, model_file, map_location, strict: model),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_weights_use_config_commit_hash(monkeypatch):
|
||||||
|
config = ACTConfig(device="cpu")
|
||||||
|
config._set_hub_commit_hash("a" * 40, "user/policy")
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_hub_download(**kwargs):
|
||||||
|
calls.append(kwargs)
|
||||||
|
return "/unused/model.safetensors"
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.policies.pretrained.hf_hub_download", fake_hub_download)
|
||||||
|
_skip_safetensor_loading(monkeypatch)
|
||||||
|
|
||||||
|
_MinimalPolicy.from_pretrained("user/policy", config=config, revision="main")
|
||||||
|
|
||||||
|
assert calls[0]["revision"] == "a" * 40
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_does_not_reuse_commit_hash_for_another_repo(monkeypatch):
|
||||||
|
config = ACTConfig(device="cpu")
|
||||||
|
config._set_hub_commit_hash("a" * 40, "user/adapter")
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_hub_download(**kwargs):
|
||||||
|
calls.append(kwargs)
|
||||||
|
return "/unused/model.safetensors"
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.policies.pretrained.hf_hub_download", fake_hub_download)
|
||||||
|
_skip_safetensor_loading(monkeypatch)
|
||||||
|
|
||||||
|
_MinimalPolicy.from_pretrained("user/base-policy", config=config, revision="base-tag")
|
||||||
|
|
||||||
|
assert calls[0]["revision"] == "base-tag"
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_records_weight_commit_for_explicit_config(monkeypatch):
|
||||||
|
commit_hash = "a" * 40
|
||||||
|
config = ACTConfig(device="cpu")
|
||||||
|
|
||||||
|
def fake_hub_download(**kwargs):
|
||||||
|
return f"/cache/models--user--policy/snapshots/{commit_hash}/model.safetensors"
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.policies.pretrained.hf_hub_download", fake_hub_download)
|
||||||
|
_skip_safetensor_loading(monkeypatch)
|
||||||
|
|
||||||
|
_MinimalPolicy.from_pretrained("user/policy", config=config, revision="main")
|
||||||
|
|
||||||
|
assert config._commit_hash == commit_hash
|
||||||
|
assert config._commit_hash_source == "user/policy"
|
||||||
|
|
||||||
|
|
||||||
|
def test_processor_factory_uses_config_commit_hash(monkeypatch):
|
||||||
|
config = ACTConfig(device="cpu")
|
||||||
|
config._set_hub_commit_hash("a" * 40, "user/policy")
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_from_pretrained(cls, **kwargs):
|
||||||
|
calls.append(kwargs)
|
||||||
|
return PolicyProcessorPipeline(steps=[])
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
factory.PolicyProcessorPipeline,
|
||||||
|
"from_pretrained",
|
||||||
|
classmethod(fake_from_pretrained),
|
||||||
|
)
|
||||||
|
|
||||||
|
factory.make_pre_post_processors(
|
||||||
|
config,
|
||||||
|
pretrained_path="user/policy",
|
||||||
|
pretrained_revision="main",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert [call["revision"] for call in calls] == ["a" * 40, "a" * 40]
|
||||||
@@ -241,6 +241,68 @@ def test_from_pretrained_hub_source_missing_local_state_still_calls_hub(monkeypa
|
|||||||
ProcessorStepRegistry.unregister("hub_state_step")
|
ProcessorStepRegistry.unregister("hub_state_step")
|
||||||
|
|
||||||
|
|
||||||
|
def test_from_pretrained_pins_hub_state_to_config_commit(monkeypatch, tmp_path):
|
||||||
|
"""A mutable processor revision is resolved once and reused for state files."""
|
||||||
|
|
||||||
|
@ProcessorStepRegistry.register("pinned_hub_state_step")
|
||||||
|
class PinnedHubStateStep(ProcessorStep):
|
||||||
|
def __init__(self):
|
||||||
|
self.value = torch.tensor(0)
|
||||||
|
|
||||||
|
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||||
|
return transition
|
||||||
|
|
||||||
|
def transform_features(
|
||||||
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
|
return features
|
||||||
|
|
||||||
|
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
||||||
|
self.value = state["value"]
|
||||||
|
|
||||||
|
try:
|
||||||
|
commit_hash = "a" * 40
|
||||||
|
snapshot_dir = tmp_path / "models--user--policy" / "snapshots" / commit_hash
|
||||||
|
snapshot_dir.mkdir(parents=True)
|
||||||
|
config_path = snapshot_dir / "processor.json"
|
||||||
|
config_path.write_text(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"name": "PinnedHubStatePipeline",
|
||||||
|
"steps": [
|
||||||
|
{
|
||||||
|
"registry_name": "pinned_hub_state_step",
|
||||||
|
"state_file": "hub_state.safetensors",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
state_path = tmp_path / "downloaded.safetensors"
|
||||||
|
save_file({"value": torch.tensor(7)}, state_path)
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_hub_download(**kwargs):
|
||||||
|
calls.append(kwargs)
|
||||||
|
if kwargs["filename"] == "processor.json":
|
||||||
|
return str(config_path)
|
||||||
|
return str(state_path)
|
||||||
|
|
||||||
|
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fake_hub_download)
|
||||||
|
|
||||||
|
pipeline = DataProcessorPipeline.from_pretrained(
|
||||||
|
"user/policy",
|
||||||
|
config_filename="processor.json",
|
||||||
|
revision="main",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert calls[0]["revision"] == "main"
|
||||||
|
assert calls[1]["revision"] == commit_hash
|
||||||
|
assert pipeline.steps[0].value.item() == 7
|
||||||
|
finally:
|
||||||
|
ProcessorStepRegistry.unregister("pinned_hub_state_step")
|
||||||
|
|
||||||
|
|
||||||
# Config Validation Tests
|
# Config Validation Tests
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+18
-1
@@ -14,7 +14,7 @@
|
|||||||
|
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
from lerobot.utils.hub import find_latest_hub_checkpoint
|
from lerobot.utils.hub import extract_commit_hash, find_latest_hub_checkpoint
|
||||||
|
|
||||||
|
|
||||||
def _patch_list_files(monkeypatch, files):
|
def _patch_list_files(monkeypatch, files):
|
||||||
@@ -52,3 +52,20 @@ def test_find_latest_hub_checkpoint_ignores_non_step_entries(monkeypatch):
|
|||||||
def test_find_latest_hub_checkpoint_none_when_no_checkpoints(monkeypatch):
|
def test_find_latest_hub_checkpoint_none_when_no_checkpoints(monkeypatch):
|
||||||
_patch_list_files(monkeypatch, ["config.json", "model.safetensors"])
|
_patch_list_files(monkeypatch, ["config.json", "model.safetensors"])
|
||||||
assert find_latest_hub_checkpoint("u/run") is None
|
assert find_latest_hub_checkpoint("u/run") is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_commit_hash_from_hub_snapshot_path():
|
||||||
|
commit_hash = "a" * 40
|
||||||
|
resolved_file = f"/cache/models--user--policy/snapshots/{commit_hash}/config.json"
|
||||||
|
|
||||||
|
assert extract_commit_hash(resolved_file, revision="main") == commit_hash
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_commit_hash_falls_back_to_full_sha_revision():
|
||||||
|
commit_hash = "b" * 40
|
||||||
|
|
||||||
|
assert extract_commit_hash("/custom/cache/config.json", revision=commit_hash) == commit_hash
|
||||||
|
|
||||||
|
|
||||||
|
def test_extract_commit_hash_rejects_mutable_revision_without_snapshot():
|
||||||
|
assert extract_commit_hash("/custom/cache/config.json", revision="main") is None
|
||||||
|
|||||||
Reference in New Issue
Block a user