|
|
| import os
|
| from functools import lru_cache
|
| from typing import Iterable, List
|
| import numpy as np
|
|
|
| from sentence_transformers import SentenceTransformer
|
| import torch
|
|
|
| from .models import ModelSpec
|
|
|
|
|
| def _e5_prefix(text: str, mode: str) -> str:
|
| """E5 ๊ณ์ด ์ฟผ๋ฆฌ/ํจ์์ง ํ๋กฌํํธ ์ฒ๋ฆฌ."""
|
| if mode == "query":
|
| return f"query: {text}"
|
| if mode == "passage":
|
| return f"passage: {text}"
|
|
|
| return f"query: {text}"
|
|
|
|
|
| def _resolve_name(name: str) -> str:
|
| """์๋๊ฒฝ๋ก/ํ๊ฒฝ๋ณ์/ํ(~)๋ฅผ ์์ ํ๊ฒ ํ์ฅ."""
|
| if not name:
|
| return name
|
|
|
| name = os.path.expandvars(name)
|
| name = os.path.expanduser(name)
|
| return name
|
|
|
|
|
| def _pick_device() -> str:
|
| """DEVICE=auto|cuda|cpu|mps (๊ธฐ๋ณธ auto)"""
|
| prefer = os.getenv("DEVICE", "auto").lower()
|
| if SentenceTransformer is None or torch is None:
|
| return "cpu"
|
| if prefer == "cpu":
|
| return "cpu"
|
| if prefer == "cuda":
|
| return "cuda" if torch.cuda.is_available() else "cpu"
|
| if prefer == "mps":
|
| avail = getattr(torch.backends, "mps", None) and torch.backends.mps.is_available()
|
| return "mps" if avail else "cpu"
|
|
|
|
|
| if torch.cuda.is_available():
|
| return "cuda"
|
| if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
|
| return "mps"
|
| return "cpu"
|
|
|
|
|
| def _norm(vec: List[float], enable: bool) -> List[float]:
|
| if not enable:
|
| return vec
|
| v = np.asarray(vec, dtype=np.float32)
|
| n = np.linalg.norm(v)
|
| if n > 0:
|
| v = v / n
|
| return v.astype(np.float32).tolist()
|
|
|
|
|
| _ST_CACHE = {}
|
|
|
| def _load_st(name: str):
|
| if SentenceTransformer is None or torch is None:
|
| raise RuntimeError(
|
| "sentence-transformers/torch ๋ฏธ์ค์น. "
|
| "pip install sentence-transformers && pip install torch(ํ๊ฒฝ์ ๋ง๋ ๋น๋)"
|
| )
|
| name_resolved = _resolve_name(name)
|
| device = _pick_device()
|
| trust = os.getenv("ST_TRUST_REMOTE_CODE", "0").lower() in ("1", "true", "yes")
|
| key = (name_resolved, device, trust)
|
| if key in _ST_CACHE:
|
| return _ST_CACHE[key], device
|
|
|
|
|
| try:
|
| n_threads = int(os.getenv("TORCH_NUM_THREADS", "0")) or None
|
| if n_threads:
|
| torch.set_num_threads(n_threads)
|
| except Exception:
|
| pass
|
|
|
| model = SentenceTransformer(name_resolved, device=device, trust_remote_code=trust)
|
| _ST_CACHE[key] = model
|
| return model, device
|
|
|
|
|
|
|
| def embed_query(text: str, spec: ModelSpec) -> List[float]:
|
| """
|
| ๋จ์ผ ์ฟผ๋ฆฌ ํ
์คํธ โ ๋ฒกํฐ.
|
| - st: PyTorch ๊ธฐ๋ฐ (GPU/CPU ์๋)
|
| """
|
| name = _resolve_name(spec.name)
|
| t = _e5_prefix(text, spec.e5_mode) if "e5" in name.lower() else text
|
|
|
|
|
| model, device = _load_st(name)
|
| vec = model.encode(
|
| t,
|
| normalize_embeddings=False,
|
| convert_to_numpy=True,
|
| device=device
|
| ).tolist()
|
| return _norm(vec, spec.normalize)
|
|
|
|
|
| def embed_many(texts: List[str], spec: ModelSpec, batch_size: int = 64) -> List[List[float]]:
|
| """
|
| ๋ฐฐ์น ์๋ฒ ๋ฉ ์ ํธ (์ธ๋ฑ์ฑ/๋๋ ์ฒ๋ฆฌ์ฉ).
|
| """
|
| name = _resolve_name(spec.name)
|
|
|
|
|
| model, device = _load_st(name)
|
|
|
| bs = batch_size
|
| if device == "cpu":
|
| bs = min(batch_size, 32)
|
| arr = model.encode(
|
| texts,
|
| batch_size=bs,
|
| normalize_embeddings=False,
|
| convert_to_numpy=True,
|
| device=device
|
| )
|
| return [_norm(v.tolist(), spec.normalize) for v in arr]
|
|
|