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
9 changes: 9 additions & 0 deletions model2vec/persistence/datamodels.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

from dataclasses import dataclass
from functools import cache
from pathlib import Path


Expand Down Expand Up @@ -48,3 +49,11 @@ def is_valid(self) -> bool:
is_sentence_transformers=True,
),
)


@cache
def get_all_model2vec_paths() -> list[str]:
"""Get all paths used across all folder layouts."""
return sorted(
{path.as_posix() for layout in FOLDER_LAYOUTS for path in (layout.embeddings, layout.config, layout.tokenizer)}
)
8 changes: 6 additions & 2 deletions model2vec/persistence/persistence.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@

from model2vec.modelcards import create_model_card as make_model_card
from model2vec.modelcards import get_metadata_from_readme
from model2vec.persistence.datamodels import FOLDER_LAYOUTS, Layout
from model2vec.persistence.datamodels import FOLDER_LAYOUTS, Layout, get_all_model2vec_paths
from model2vec.persistence.hf import maybe_get_cached_model_path
from model2vec.persistence.utils import SilentTqdm
from model2vec.types import StaticModelConfig
Expand Down Expand Up @@ -148,7 +148,11 @@ def _resolve_folder(folder_or_repo_path: Path, token: str | None, force_download
# No partial because that doesn't always work, this is safer.
folder = Path(
huggingface_hub.snapshot_download(
str(folder_or_repo_path.as_posix()), repo_type="model", token=token, tqdm_class=SilentTqdm
str(folder_or_repo_path.as_posix()),
repo_type="model",
token=token,
tqdm_class=SilentTqdm,
allow_patterns=get_all_model2vec_paths(),
)
)

Expand Down
Loading