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