Download benchmark/vectorstore.py from WalidAlHassan/embeddingModelRnD: direct link, hf CLI and curl.
- Browser
- Download file 2.18 kB
-
https://huggingface.co/WalidAlHassan/embeddingModelRnD/resolve/main/benchmark/vectorstore.py
- Command line
-
hf download hf://WalidAlHassan/embeddingModelRnD/benchmark/vectorstore.py
-
curl -L -o vectorstore.py https://huggingface.co/WalidAlHassan/embeddingModelRnD/resolve/main/benchmark/vectorstore.py
2.18 kB
| """Thin Qdrant wrapper: one collection per embedding model under test.""" | |
| import os | |
| import re | |
| from uuid import uuid4 | |
| from qdrant_client import QdrantClient | |
| from qdrant_client.models import Distance, PointStruct, VectorParams | |
| def collection_name_for(model_name): | |
| return "books__" + re.sub(r"[^a-zA-Z0-9]+", "_", model_name).strip("_") | |
| def _dir_size_bytes(path): | |
| total = 0 | |
| for root, _dirs, files in os.walk(path): | |
| for name in files: | |
| fp = os.path.join(root, name) | |
| if os.path.isfile(fp): | |
| total += os.path.getsize(fp) | |
| return total | |
| class BookVectorStore: | |
| def __init__(self, qdrant_path): | |
| self.qdrant_path = qdrant_path | |
| self.client = QdrantClient(path=qdrant_path) | |
| def index(self, model_name, embeddings, books): | |
| name = collection_name_for(model_name) | |
| if self.client.collection_exists(name): | |
| self.client.delete_collection(name) | |
| self.client.create_collection( | |
| collection_name=name, | |
| vectors_config=VectorParams( | |
| size=embeddings.shape[1], | |
| distance=Distance.COSINE, | |
| ), | |
| ) | |
| points = [ | |
| PointStruct( | |
| id=str(uuid4()), | |
| vector=embedding.tolist(), | |
| payload={"book_id": book["book_id"]}, | |
| ) | |
| for book, embedding in zip(books, embeddings) | |
| ] | |
| self.client.upsert(collection_name=name, points=points) | |
| return name | |
| def search(self, model_name, query_embedding, top_k): | |
| name = collection_name_for(model_name) | |
| results = self.client.query_points( | |
| collection_name=name, | |
| query=query_embedding.tolist(), | |
| limit=top_k, | |
| with_payload=True, | |
| ).points | |
| return [point.payload["book_id"] for point in results] | |
| def collection_disk_size_mb(self, model_name): | |
| name = collection_name_for(model_name) | |
| collection_dir = os.path.join(self.qdrant_path, "collection", name) | |
| if not os.path.isdir(collection_dir): | |
| return None | |
| return _dir_size_bytes(collection_dir) / (1024 ** 2) | |