| """Download only the two immutable model snapshots needed by the Space.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from pathlib import Path |
|
|
| from huggingface_hub import snapshot_download |
|
|
|
|
| EMBEDDING_MODEL = "intfloat/multilingual-e5-small" |
| EMBEDDING_REVISION = "fd1525a9fd15316a2d503bf26ab031a61d056e98" |
| RERANKER_MODEL = "cross-encoder/ms-marco-MiniLM-L2-v2" |
| RERANKER_REVISION = "1b5cd67b15209f24824c50370e0397743aa9b787" |
|
|
|
|
| def download_models(cache_dir: Path) -> tuple[Path, Path]: |
| cache_dir = cache_dir.resolve() |
| cache_dir.mkdir(parents=True, exist_ok=True) |
| embedding = Path( |
| snapshot_download( |
| repo_id=EMBEDDING_MODEL, |
| revision=EMBEDDING_REVISION, |
| cache_dir=str(cache_dir), |
| allow_patterns=( |
| "1_Pooling/config.json", |
| "config.json", |
| "model.safetensors", |
| "modules.json", |
| "sentence_bert_config.json", |
| "sentencepiece.bpe.model", |
| "special_tokens_map.json", |
| "tokenizer.json", |
| "tokenizer_config.json", |
| ), |
| ) |
| ) |
| reranker = Path( |
| snapshot_download( |
| repo_id=RERANKER_MODEL, |
| revision=RERANKER_REVISION, |
| cache_dir=str(cache_dir), |
| allow_patterns=( |
| "config.json", |
| "model.safetensors", |
| "openvino/openvino_model_qint8_quantized.bin", |
| "openvino/openvino_model_qint8_quantized.xml", |
| "special_tokens_map.json", |
| "tokenizer.json", |
| "tokenizer_config.json", |
| "vocab.txt", |
| ), |
| ) |
| ) |
| return embedding, reranker |
|
|
|
|
| def main() -> int: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--cache-dir", type=Path, required=True) |
| args = parser.parse_args() |
| embedding, reranker = download_models(args.cache_dir) |
| print(f"embedding_snapshot={embedding}") |
| print(f"reranker_snapshot={reranker}") |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| raise SystemExit(main()) |
|
|