File size: 6,218 Bytes
6ced533
 
 
 
 
eb02943
6ced533
 
 
eb02943
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6ced533
eb02943
 
 
 
 
 
 
 
 
 
6ced533
 
 
eb02943
 
6ced533
 
 
eb02943
 
 
6ced533
 
eb02943
 
6ced533
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
eb02943
6ced533
 
 
 
 
 
 
 
 
 
 
 
 
eb02943
6ced533
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
"""Wraps a SentenceTransformer model with prefixing and timing instrumentation."""

import time

import torch
import transformers
from sentence_transformers import SentenceTransformer


def _patch_missing_tied_weights_keys():
    """Some custom `trust_remote_code` models (e.g. jina-embeddings-v3's LoRA
    wrapper) never call `PreTrainedModel.post_init()`, so `transformers` >= 5's
    weight-loading path crashes with AttributeError on `self.all_tied_weights_keys`.
    Turning it into a property with a per-instance fallback default of `{}` keeps
    every normal model's behavior identical (post_init still sets it the same way)
    while letting non-conformant custom code fall back to "no tied weights" instead
    of crashing.
    """
    base = transformers.PreTrainedModel
    if getattr(base, "_omicon_tied_weights_patched", False):
        return

    def getter(self):
        return self.__dict__.setdefault("_all_tied_weights_keys_fallback", {})

    def setter(self, value):
        self.__dict__["_all_tied_weights_keys_fallback"] = value

    base.all_tied_weights_keys = property(getter, setter)
    base._omicon_tied_weights_patched = True


_patch_missing_tied_weights_keys()


def _refresh_procedurally_computed_buffers(sentence_transformer_model, device):
    """Some `trust_remote_code` models (e.g. Alibaba-NLP/gte-multilingual-base) compute
    real values for non-persistent buffers (RoPE cos/sin caches, a `position_ids` arange)
    with tensor math inside `__init__` (torch.arange, .cos(), .sin(), einsum). `transformers`
    >= 5's `from_pretrained` unconditionally builds every model under a `meta` device context
    for speed -- `low_cpu_mem_usage`/`_fast_init` are now silently dropped, with no opt-out.
    Ops on meta tensors don't produce real values, and only checkpoint weights get properly
    materialized afterward, so these procedural buffers come out as uninitialized memory
    instead of the values `__init__` intended. Re-running the same buffer-computation methods
    now that the model lives on a real device reproduces exactly what `__init__` was supposed
    to do. This is a no-op (via hasattr guards) for any model that isn't shaped this way.
    """
    try:
        auto_model = sentence_transformer_model[0].auto_model
    except (IndexError, AttributeError, KeyError):
        return

    embeddings = getattr(auto_model, "embeddings", None)
    if embeddings is None:
        return

    # `rotary_emb` is specific to this custom architecture (e.g. GTE's NewEmbeddings) --
    # standard BERT/RoBERTa-family embeddings also have a `position_ids` buffer, but with
    # a different (1, max_pos) shape, so only touch it when we've confirmed we're looking
    # at this specific rope-based custom module.
    rotary = getattr(embeddings, "rotary_emb", None)
    if rotary is None or not hasattr(rotary, "_set_cos_sin_cache"):
        return

    rotary._set_cos_sin_cache(
        seq_len=rotary.max_seq_len_cached, device=device, dtype=rotary.cos_cached.dtype
    )

    position_ids = getattr(embeddings, "position_ids", None)
    if position_ids is not None and position_ids.dim() == 1:
        embeddings.register_buffer(
            "position_ids", torch.arange(position_ids.shape[0], device=device), persistent=False
        )


class TimedEmbedder:
    def __init__(
        self,
        model_name,
        query_prefix="",
        passage_prefix="",
        device=None,
        trust_remote_code=False,
        query_encode_kwargs=None,
        passage_encode_kwargs=None,
    ):
        self.model_name = model_name
        self.query_prefix = query_prefix
        self.passage_prefix = passage_prefix
        self.query_encode_kwargs = query_encode_kwargs or {}
        self.passage_encode_kwargs = passage_encode_kwargs or {}
        self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")

        start = time.perf_counter()
        self.model = SentenceTransformer(
            model_name, device=self.device, trust_remote_code=trust_remote_code
        )
        self.load_time_sec = time.perf_counter() - start

        _refresh_procedurally_computed_buffers(self.model, self.device)

        self.param_count = sum(p.numel() for p in self.model.parameters())
        self.model_size_mb = sum(
            p.numel() * p.element_size() for p in self.model.parameters()
        ) / (1024 ** 2)

        if hasattr(self.model, "get_embedding_dimension"):
            self.embedding_dim = self.model.get_embedding_dimension()
        else:
            self.embedding_dim = self.model.get_sentence_embedding_dimension()

    def encode_passages(self, texts, batch_size=32):
        """Batch-encode documents; returns (embeddings, elapsed_seconds)."""
        prefixed = [self.passage_prefix + t for t in texts]
        start = time.perf_counter()
        embeddings = self.model.encode(
            prefixed,
            batch_size=batch_size,
            normalize_embeddings=True,
            convert_to_numpy=True,
            show_progress_bar=False,
            **self.passage_encode_kwargs,
        )
        elapsed = time.perf_counter() - start
        return embeddings, elapsed

    def encode_query(self, text):
        """Single-query encode; returns (embedding, elapsed_seconds)."""
        prefixed = self.query_prefix + text
        start = time.perf_counter()
        embedding = self.model.encode(
            prefixed,
            normalize_embeddings=True,
            convert_to_numpy=True,
            show_progress_bar=False,
            **self.query_encode_kwargs,
        )
        elapsed = time.perf_counter() - start
        return embedding, elapsed

    def peak_encode_memory_mb(self, sample_text):
        """Peak accelerator memory used while encoding one query (CUDA only, else None)."""
        if self.device != "cuda":
            return None
        torch.cuda.synchronize()
        torch.cuda.reset_peak_memory_stats(self.device)
        self.encode_query(sample_text)
        torch.cuda.synchronize()
        return torch.cuda.max_memory_allocated(self.device) / (1024 ** 2)

    def unload(self):
        del self.model
        if self.device == "cuda":
            torch.cuda.empty_cache()