Download benchmark/embedder.py from WalidAlHassan/embeddingModelRnD: direct link, hf CLI and curl.
- Browser
- Download file 6.22 kB
-
https://huggingface.co/WalidAlHassan/embeddingModelRnD/resolve/main/benchmark/embedder.py
- Command line
-
hf download hf://WalidAlHassan/embeddingModelRnD/benchmark/embedder.py
-
curl -L -o embedder.py https://huggingface.co/WalidAlHassan/embeddingModelRnD/resolve/main/benchmark/embedder.py
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() | |