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()
|