WalidAlHassan's picture
initial
eb02943
Raw History Blame Contribute Delete
6.22 kB
"""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()