embeddingModelRnD / benchmark /vectorstore.py
WalidAlHassan's picture
initial
6ced533
Raw History Blame Contribute Delete
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)