From d5e0b5dc22222fb4b1532b91b758eb6e55d61e75 Mon Sep 17 00:00:00 2001 From: stephantul Date: Thu, 8 Oct 2026 14:42:48 +0200 Subject: [PATCH] download readme and respect subfolder --- model2vec/persistence/persistence.py | 20 ++++++++++++++---- tests/test_persistence.py | 31 ++++++++++++++++++++++++++++ 2 files changed, 47 insertions(+), 4 deletions(-) diff --git a/model2vec/persistence/persistence.py b/model2vec/persistence/persistence.py index b29c0d7..20248be 100644 --- a/model2vec/persistence/persistence.py +++ b/model2vec/persistence/persistence.py @@ -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 @@ -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 + 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( @@ -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: diff --git a/tests/test_persistence.py b/tests/test_persistence.py index 81e6857..91616a6 100644 --- a/tests/test_persistence.py +++ b/tests/test_persistence.py @@ -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") + + 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: