mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-31 21:49:45 +00:00
fix(datasets): improve get_safe_version raise (#4263)
* fix: Pass the required response argument to RevisionNotFoundError * chore(dataset): get_safe_version raise --------- Co-authored-by: Harshal Janjani <harshaljanjani@gmail.com>
This commit is contained in:
@@ -26,7 +26,6 @@ import numpy as np
|
||||
import packaging.version
|
||||
import torch
|
||||
from huggingface_hub import DatasetCard, DatasetCardData, HfApi
|
||||
from huggingface_hub.errors import RevisionNotFoundError
|
||||
|
||||
from lerobot.utils.utils import flatten_dict, unflatten_dict
|
||||
|
||||
@@ -51,6 +50,17 @@ The dataset you requested ({repo_id}) is only available in {version} format.
|
||||
As we cannot ensure forward compatibility with it, please update your current version of lerobot.
|
||||
"""
|
||||
|
||||
MISSING_VERSION_TAG_MESSAGE = """
|
||||
Your dataset must be tagged with a codebase version.
|
||||
Assuming _version_ is the codebase_version value in the info.json, you can run this:
|
||||
```python
|
||||
from huggingface_hub import HfApi
|
||||
|
||||
hub_api = HfApi()
|
||||
hub_api.create_tag("{repo_id}", tag="_version_", repo_type="dataset")
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
class CompatibilityError(Exception): ...
|
||||
|
||||
@@ -368,7 +378,7 @@ def get_safe_version(
|
||||
str: The safe version string (e.g., "v1.2.3") to use as a revision.
|
||||
|
||||
Raises:
|
||||
RevisionNotFoundError: If the repo has no version tags.
|
||||
RuntimeError: If the repo has no version tags.
|
||||
BackwardCompatibilityError: If only older major versions are available.
|
||||
ForwardCompatibilityError: If only newer major versions are available.
|
||||
"""
|
||||
@@ -378,17 +388,7 @@ def get_safe_version(
|
||||
hub_versions = get_repo_versions(repo_id) if token is None else get_repo_versions(repo_id, token=token)
|
||||
|
||||
if not hub_versions:
|
||||
raise RevisionNotFoundError(
|
||||
f"""Your dataset must be tagged with a codebase version.
|
||||
Assuming _version_ is the codebase_version value in the info.json, you can run this:
|
||||
```python
|
||||
from huggingface_hub import HfApi
|
||||
|
||||
hub_api = HfApi()
|
||||
hub_api.create_tag("{repo_id}", tag="_version_", repo_type="dataset")
|
||||
```
|
||||
"""
|
||||
)
|
||||
raise RuntimeError(MISSING_VERSION_TAG_MESSAGE.format(repo_id=repo_id))
|
||||
|
||||
if target_version in hub_versions:
|
||||
return f"v{target_version}"
|
||||
|
||||
@@ -180,3 +180,20 @@ def test_non_dict_passthrough_last_wins():
|
||||
out = combine_feature_dicts(g1, g2)
|
||||
# For non-dict entries the last one wins
|
||||
assert out["misc"] == 456
|
||||
|
||||
|
||||
def test_get_safe_version_raises_on_repo_without_version_tags(monkeypatch):
|
||||
monkeypatch.setattr(dataset_utils, "get_repo_versions", Mock(return_value=[]))
|
||||
|
||||
with pytest.raises(RuntimeError, match="must be tagged with a codebase version"):
|
||||
get_safe_version("private/repo", "v3.0")
|
||||
|
||||
|
||||
def test_get_safe_version_error_reports_repo_id(monkeypatch):
|
||||
repo_id = "private/repo"
|
||||
monkeypatch.setattr(dataset_utils, "get_repo_versions", Mock(return_value=[]))
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
get_safe_version(repo_id, "v3.0")
|
||||
|
||||
assert repo_id in str(exc_info.value)
|
||||
|
||||
Reference in New Issue
Block a user