foto / validation_and_analysis /embedders.py
Hurum Maksora Tohfa
validation tests
45383d0 unverified
Raw
History Blame Contribute Delete
3.35 kB
"""Embedding backends: OpenAI API or local sentence-transformers, with the
query/document prefixes each model family expects."""
import numpy as np
PREFIX_RULES = [
("bge-", {"query": "Represent this sentence for searching relevant passages: ",
"doc": ""}),
("e5-", {"query": "query: ", "doc": "passage: "}),
("multilingual-e5", {"query": "query: ", "doc": "passage: "}),
("Qwen3-Embedding", {"query": "Instruct: Given a figure search query, retrieve the "
"figure caption that matches it\nQuery: ",
"doc": ""}),
]
def prefixes_for(model_name):
for needle, rules in PREFIX_RULES:
if needle.lower() in model_name.lower():
return rules
return {"query": "", "doc": ""}
class OpenAIEmbedder:
name = "openai"
def __init__(self, model="text-embedding-3-small", dim=512):
import os
from openai import OpenAI
self.client = OpenAI(api_key=os.environ["OPENAI_API_KEY"],
max_retries=0, timeout=60)
self.model = model
self.dim = dim
def embed(self, texts, is_query=False, batch_size=128):
import time
out = []
for i in range(0, len(texts), batch_size):
batch = texts[i:i + batch_size]
for attempt in range(8):
try:
resp = self.client.embeddings.create(
model=self.model, input=batch, dimensions=self.dim)
out.extend(d.embedding for d in resp.data)
break
except Exception as e:
if "429" not in str(e) or attempt == 7:
raise
time.sleep(10 * (attempt + 1))
X = np.array(out, dtype=np.float32)
X /= np.linalg.norm(X, axis=1, keepdims=True)
return X
def info(self):
return {"backend": "openai", "model": self.model, "dim": self.dim}
class LocalEmbedder:
name = "local"
def __init__(self, model="BAAI/bge-base-en-v1.5", device=None):
from sentence_transformers import SentenceTransformer
self.model_name = model
self.model = SentenceTransformer(model, device=device)
self.prefixes = prefixes_for(model)
self.dim = self.model.get_sentence_embedding_dimension()
def embed(self, texts, is_query=False, batch_size=64):
prefix = self.prefixes["query" if is_query else "doc"]
inputs = [prefix + t for t in texts]
X = self.model.encode(inputs, batch_size=batch_size,
normalize_embeddings=True,
show_progress_bar=len(texts) > 1000)
return np.asarray(X, dtype=np.float32)
def info(self):
return {"backend": "local", "model": self.model_name, "dim": self.dim}
def make_embedder(backend, model=None, dim=512):
if backend == "openai":
return OpenAIEmbedder(model=model or "text-embedding-3-small", dim=dim)
if backend == "local":
return LocalEmbedder(model=model or "BAAI/bge-base-en-v1.5")
raise ValueError(f"unknown backend {backend}")
def embedder_from_info(info):
if info["backend"] == "openai":
return OpenAIEmbedder(model=info["model"], dim=info["dim"])
return LocalEmbedder(model=info["model"])