File size: 2,181 Bytes
6ced533 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | """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)
|