Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 16 additions & 4 deletions model2vec/persistence/persistence.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,9 @@ def load_pretrained(
folder_or_repo_path = Path(folder_or_repo_path)

# We resolve a folder or repo path to an actual local folder.
folder = _resolve_folder(folder_or_repo_path=folder_or_repo_path, token=token, force_download=force_download)
folder = _resolve_folder(
folder_or_repo_path=folder_or_repo_path, subfolder=subfolder, token=token, force_download=force_download
)

if subfolder:
folder = folder / subfolder
Expand Down Expand Up @@ -134,16 +136,21 @@ def load_pretrained(
return embeddings, tokenizer, config, metadata, weights, mapping


def _resolve_folder(folder_or_repo_path: Path, token: str | None, force_download: bool) -> Path:
def _resolve_folder(folder_or_repo_path: Path, subfolder: str | None, token: str | None, force_download: bool) -> Path:
"""Resolve a folder locally or from hugging face hub."""
if folder_or_repo_path.exists():
return folder_or_repo_path
# We now know we're dealing with either an invalid path, or
# a HF model ID.
if not force_download:
if folder := maybe_get_cached_model_path(str(folder_or_repo_path)):
folder = maybe_get_cached_model_path(str(folder_or_repo_path))
if folder is not None and _has_valid_layout(folder / subfolder if subfolder else folder):
return folder
Comment on lines +147 to 148

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Cached models still lose language

With force_download=False, a snapshot downloaded by the old code can contain all model files but no README.md. _has_valid_layout accepts it, so _resolve_folder returns before fetching the newly included README. model.language stays None even when the hub README declares a language.

Fetch missing README metadata without making a README mandatory for models that have none.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Ok, and?

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The actionable fix is to avoid returning a valid cached snapshot when its model directory lacks README.md; let it fall through to snapshot_download so the README can be fetched. Keep _has_valid_layout unchanged, so models that genuinely have no README remain valid.

Suggested change
if folder is not None and _has_valid_layout(folder / subfolder if subfolder else folder):
return folder
if folder is not None and _has_valid_layout(folder / subfolder if subfolder else folder):
model_folder = folder / subfolder if subfolder else folder
if (model_folder / "README.md").exists():
return folder

This preserves the cache fast path when the README is already present, while allowing older cached snapshots to acquire it. Repositories without a README may be rechecked on later loads unless the cache layer records that the README is absent.


allow_patterns = [*get_all_model2vec_paths(), "README.md"]
if subfolder:
allow_patterns = [f"{subfolder.strip('/')}/{pattern}" for pattern in allow_patterns]

# We use `tqdm_class=SilentTqdm` to disable download progress bars.
# No partial because that doesn't always work, this is safer.
folder = Path(
Expand All @@ -152,13 +159,18 @@ def _resolve_folder(folder_or_repo_path: Path, token: str | None, force_download
repo_type="model",
token=token,
tqdm_class=SilentTqdm,
allow_patterns=get_all_model2vec_paths(),
allow_patterns=allow_patterns,
)
)

return folder


def _has_valid_layout(folder: Path) -> bool:
"""Check if any known layout is present in a folder."""
return any(layout.with_parent(folder).is_valid() for layout in FOLDER_LAYOUTS)


def _get_paths(folder: Path) -> Layout:
"""Get all paths by trying out multiple layouts."""
for layout in FOLDER_LAYOUTS:
Expand Down
31 changes: 31 additions & 0 deletions tests/test_persistence.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,37 @@ def test_local_loading(mock_static_model: StaticModel) -> None:
assert mock_snapshot.call_count == 3


def test_hub_loading_downloads_readme(tmp_path: Path, mock_tokenizer: Tokenizer) -> None:
"""Test that loading from the hub fetches the README and reads the language."""
vectors = np.random.RandomState(0).randn(len(mock_tokenizer.get_vocab()), 8)
StaticModel(vectors=vectors, tokenizer=mock_tokenizer, language=["en", "nl"]).save_pretrained(tmp_path)

with patch("model2vec.persistence.persistence.huggingface_hub.snapshot_download") as mock_snapshot:
mock_snapshot.return_value = tmp_path
model = StaticModel.from_pretrained("my_org/haha", force_download=True)

assert "README.md" in mock_snapshot.call_args.kwargs["allow_patterns"]
assert model.language == ["en", "nl"]


def test_hub_subfolder_loading(tmp_path: Path, mock_static_model: StaticModel) -> None:
"""Test that loading a subfolder from the hub fetches the files in that subfolder."""
mock_static_model.save_pretrained(tmp_path / "subfolder")

with patch("model2vec.persistence.persistence.huggingface_hub.snapshot_download") as mock_snapshot:
mock_snapshot.return_value = tmp_path
with patch("model2vec.persistence.persistence.maybe_get_cached_model_path") as cache:
# Cached snapshot without the subfolder files
cache.return_value = tmp_path / "other"
model = StaticModel.from_pretrained("my_org/haha", subfolder="subfolder")

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Test skips the cache fallback

test_hub_subfolder_loading patches an incomplete cache, but StaticModel.from_pretrained defaults to force_download=True. The cache patch is never used, so the test would pass even if the new cache-layout check were removed.

Pass force_download=False and assert that the cache was checked before the download.


allow_patterns = mock_snapshot.call_args.kwargs["allow_patterns"]
assert "subfolder/README.md" in allow_patterns
assert "subfolder/config.json" in allow_patterns
assert all(pattern.startswith("subfolder/") for pattern in allow_patterns)
assert model.tokens == mock_static_model.tokens


def test_garbage(mock_static_model: StaticModel) -> None:
"""Test that garbage loading crashes."""
with TemporaryDirectory() as dir_name:
Expand Down
Loading