hyongok2's picture
Upload 11 files
9d780fe verified
Raw
History Blame Contribute Delete
4.1 kB
# app/embeddings.py
import os
from functools import lru_cache
from typing import Iterable, List
import numpy as np
from sentence_transformers import SentenceTransformer # pragma: no cover
import torch # pragma: no cover
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}"
# auto: ๊ฒ€์ƒ‰ ์ฟผ๋ฆฌ์—์„œ๋Š” query ๊ธฐ๋ณธ
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"
# auto: cuda โ†’ mps โ†’ 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 ๋กœ๋” (GPU/CPU ์ž๋™, trust_remote_code ์ง€์›) ---------
_ST_CACHE = {} # key=(name_resolved, device, trust) -> model
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
# --------- ๊ณต๊ฐœ API ---------
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
# ST
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)
# ST
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]