fix(processor): diagnose legacy Hub checkpoints (#4261)

* fix(processor): diagnose legacy Hub checkpoints

* chore(processor): improve migration detection

---------

Co-authored-by: Vaish Gajaraj <47009802+VaishGajaraj@users.noreply.github.com>
This commit is contained in:
Steven Palma
2026-07-31 14:13:03 +02:00
committed by GitHub
parent 414d0eecbd
commit 8e2a077f09
2 changed files with 122 additions and 2 deletions
+74 -2
View File
@@ -44,7 +44,7 @@ import torch
from huggingface_hub import hf_hub_download from huggingface_hub import hf_hub_download
from safetensors.torch import load_file, save_file from safetensors.torch import load_file, save_file
from lerobot.configs import PipelineFeatureType, PolicyFeature from lerobot.configs import FeatureType, PipelineFeatureType, PolicyFeature
from lerobot.lerobot_types import ( from lerobot.lerobot_types import (
EnvAction, EnvAction,
EnvTransition, EnvTransition,
@@ -873,6 +873,13 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
return json.load(f), Path(config_path).parent return json.load(f), Path(config_path).parent
except Exception as e: except Exception as e:
if cls._hub_model_requires_migration(model_id, hub_download_kwargs):
revision = hub_download_kwargs.get("revision")
cls._suggest_processor_migration(
model_id,
f"Config file '{config_filename}' not found on the Hugging Face Hub",
revision=revision if isinstance(revision, str) else None,
)
raise FileNotFoundError( raise FileNotFoundError(
f"Could not find '{config_filename}' on the HuggingFace Hub at '{model_id}'" f"Could not find '{config_filename}' on the HuggingFace Hub at '{model_id}'"
) from e ) from e
@@ -1337,6 +1344,62 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
# Have JSON files but no processor configs - suggest migration # Have JSON files but no processor configs - suggest migration
return True return True
@classmethod
def _hub_model_requires_migration(cls, model_id: str, hub_download_kwargs: dict[str, Any]) -> bool:
"""Check whether a Hub repository contains a legacy LeRobot policy config.
A missing processor file is not sufficient evidence by itself: the repository
may be private, unavailable, or unrelated to LeRobot. This method therefore
fetches the policy's ``config.json`` and checks for the feature declarations
that identify a LeRobot policy checkpoint. Any lookup or parsing failure is
ignored so the original processor-file error remains visible.
Args:
model_id: Hugging Face Hub model repository ID.
hub_download_kwargs: Authentication, cache, and revision arguments used
for the original processor lookup.
Returns:
True when the repository has a legacy LeRobot policy configuration.
"""
try:
config_path = hf_hub_download(
repo_id=model_id,
filename="config.json",
repo_type="model",
**hub_download_kwargs,
)
with open(config_path) as f:
config = json.load(f)
except Exception:
# This is a best-effort diagnostic called while handling the original
# processor lookup failure, which must remain the visible error.
return False
feature_types = {feature_type.value for feature_type in FeatureType}
def is_policy_feature_mapping(features: Any) -> bool:
return (
isinstance(features, dict)
and bool(features)
and all(
isinstance(name, str)
and isinstance(feature, dict)
and feature.get("type") in feature_types
and isinstance(feature.get("shape"), list)
and all(isinstance(dimension, int) for dimension in feature["shape"])
for name, feature in features.items()
)
)
return (
isinstance(config, dict)
and isinstance(config.get("type"), str)
and bool(config["type"])
and is_policy_feature_mapping(config.get("input_features"))
and is_policy_feature_mapping(config.get("output_features"))
)
@classmethod @classmethod
def _is_processor_config(cls, config: Any) -> bool: def _is_processor_config(cls, config: Any) -> bool:
"""Check if config follows DataProcessorPipeline format. """Check if config follows DataProcessorPipeline format.
@@ -1410,7 +1473,13 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
return True return True
@classmethod @classmethod
def _suggest_processor_migration(cls, model_path: str | Path, original_error: str) -> None: def _suggest_processor_migration(
cls,
model_path: str | Path,
original_error: str,
*,
revision: str | None = None,
) -> None:
"""Raise migration error when we detect JSON files but no processor configs. """Raise migration error when we detect JSON files but no processor configs.
This method is called when migration detection determines that a model This method is called when migration detection determines that a model
@@ -1445,6 +1514,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
Args: Args:
model_path: Path to the model directory needing migration model_path: Path to the model directory needing migration
original_error: The error that triggered migration detection (for context) original_error: The error that triggered migration detection (for context)
revision: Optional Hub revision containing the legacy checkpoint.
Raises: Raises:
ProcessorMigrationError: Always raised (this method never returns normally) ProcessorMigrationError: Always raised (this method never returns normally)
@@ -1452,6 +1522,8 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
migration_command = ( migration_command = (
f"python src/lerobot/processor/migrate_policy_normalization.py --pretrained-path {model_path}" f"python src/lerobot/processor/migrate_policy_normalization.py --pretrained-path {model_path}"
) )
if revision is not None:
migration_command += f" --revision {revision}"
raise ProcessorMigrationError(model_path, migration_command, original_error) raise ProcessorMigrationError(model_path, migration_command, original_error)
@@ -107,6 +107,54 @@ def test_load_config_nonexistent_path_tries_hub():
DataProcessorPipeline._load_config("nonexistent/path", "processor.json", {}) DataProcessorPipeline._load_config("nonexistent/path", "processor.json", {})
def test_load_config_legacy_hub_policy_suggests_revision_aware_migration(tmp_path, monkeypatch):
config_path = tmp_path / "config.json"
config_path.write_text(
json.dumps(
{
"type": "diffusion",
"input_features": {"observation.state": {"type": "STATE", "shape": [2]}},
"output_features": {"action": {"type": "ACTION", "shape": [2]}},
}
)
)
def fake_hf_hub_download(*, filename, **kwargs): # noqa: ARG001
if filename == "config.json":
return str(config_path)
raise FileNotFoundError(filename)
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fake_hf_hub_download)
with pytest.raises(ProcessorMigrationError) as exc_info:
DataProcessorPipeline._load_config(
"lerobot/diffusion_pusht", "policy_preprocessor.json", {"revision": "legacy"}
)
assert "--revision legacy" in exc_info.value.migration_command
def test_hub_migration_detection_is_conservative(tmp_path, monkeypatch):
config_path = tmp_path / "config.json"
config_path.write_text(
json.dumps(
{
"type": "diffusion",
"input_features": {},
"output_features": {"action": {"type": "ACTION", "shape": [2]}},
}
)
)
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", lambda **kwargs: str(config_path))
assert not DataProcessorPipeline._hub_model_requires_migration("someone/model", {})
def fail_lookup(**kwargs): # noqa: ARG001
raise ValueError("invalid repository")
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fail_lookup)
assert not DataProcessorPipeline._hub_model_requires_migration("invalid", {})
def test_from_pretrained_local_directory_missing_state_does_not_call_hub(monkeypatch): def test_from_pretrained_local_directory_missing_state_does_not_call_hub(monkeypatch):
"""Local processor dirs must fail locally when a state file is missing.""" """Local processor dirs must fail locally when a state file is missing."""