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)