File size: 3,351 Bytes
45383d0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""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"])