forked from HKUDS/FastCode
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpreload_embedding_model.py
More file actions
70 lines (54 loc) · 2.09 KB
/
Copy pathpreload_embedding_model.py
File metadata and controls
70 lines (54 loc) · 2.09 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
from __future__ import annotations
import argparse
import os
from pathlib import Path
import yaml
PROJECT_ROOT = Path(__file__).resolve().parent
CONFIG_PATH = PROJECT_ROOT / "config" / "config.yaml"
def _default_model_from_config() -> str:
if not CONFIG_PATH.exists():
return "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
with CONFIG_PATH.open("r", encoding="utf-8") as f:
config = yaml.safe_load(f) or {}
embedding = config.get("embedding", {}) if isinstance(config, dict) else {}
return embedding.get(
"model", "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
)
def main() -> None:
parser = argparse.ArgumentParser(description="Preload FastCode embedding model")
parser.add_argument(
"--model",
default=_default_model_from_config(),
help="Sentence-transformers model ID",
)
parser.add_argument(
"--device",
default="cpu",
help="Device for preload (cpu/cuda/mps). Defaults to cpu.",
)
parser.add_argument(
"--dry-run",
action="store_true",
help="Print planned preload settings without downloading",
)
args = parser.parse_args()
hf_home = os.getenv("HF_HOME", "<default>")
hub_cache = os.getenv("HUGGINGFACE_HUB_CACHE", "<default>")
transformers_cache = os.getenv("TRANSFORMERS_CACHE", "<default>")
st_home = os.getenv("SENTENCE_TRANSFORMERS_HOME", "<default>")
print("FastCode embedding preload")
print(f"- Model: {args.model}")
print(f"- Device: {args.device}")
print(f"- HF_HOME: {hf_home}")
print(f"- HUGGINGFACE_HUB_CACHE: {hub_cache}")
print(f"- TRANSFORMERS_CACHE: {transformers_cache}")
print(f"- SENTENCE_TRANSFORMERS_HOME: {st_home}")
if args.dry_run:
print("Dry run complete (no download).")
return
from sentence_transformers import SentenceTransformer
model = SentenceTransformer(args.model, device=args.device)
dim = model.get_sentence_embedding_dimension()
print(f"Preload complete. Embedding dimension: {dim}")
if __name__ == "__main__":
main()