-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathvector_db.py
More file actions
153 lines (125 loc) · 4.89 KB
/
Copy pathvector_db.py
File metadata and controls
153 lines (125 loc) · 4.89 KB
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
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
import re
import chromadb
from uuid import uuid4
from chromadb.utils import embedding_functions
from sentence_transformers import SentenceTransformer
class VectorDB:
IDS = 'ids'
DIST = 'distances'
META = 'metadatas'
EMB = 'embeddings'
DOC = 'documents'
URIS = 'uris'
DATA = 'data'
def __init__(self, db_path, collection_name, collection_metadata, embedding_model_name):
'''
Setup the vector DB
db_path: str: Path on disk to where the DB is stored
collection_name: str: Name of collection
collection_metadata: Dict:
embedding_model_name: str: Name of model used to compute embeddings
'''
self.client = chromadb.PersistentClient(path=db_path)
# self.client = chromadb.Client()
self.emb_function = embedding_functions.SentenceTransformerEmbeddingFunction(model_name=embedding_model_name)
self.metadata = collection_metadata
self.collection_name = collection_name
self.model = SentenceTransformer(embedding_model_name)
# if no collection, create it
try:
self.collection = self.client.get_collection(collection_name, embedding_function=self.emb_function)
except ValueError as err:
print(f'Collection {collection_name} not availble. Creating it instead')
self.collection = self.client.create_collection(
name=collection_name,
metadata=collection_metadata,
embedding_function=self.emb_function
)
def normalise_results(self, results):
normalised_distances = []
distances = results[self.DIST]
for distance in distances:
normalised_dist = 1 - (distance / 2)
normalised_distances.append(normalised_dist)
results[self.DIST] = normalised_distances
return results
def compare_all(self, all_documents, n_results=2):
'''
Compares the given set of documents with all documents in the db
all_documents: List[str]: List of documents where each document is a str
n_results: int: number of results to query for
normalise: bool: if True, return distances in 0-1 range
'''
if n_results <= 0:
n_results = 2
embeddings = self.model.encode(all_documents)
results = self.collection.query(
query_embeddings=embeddings.tolist(),
n_results=n_results,
include=[self.EMB, self.DIST, self.DOC]
)
for emb in results['embeddings']:
if emb is not None:
print(len(emb[0]))
return results
def compare_within(self, all_documents, n_results=3):
'''
Compares the given set of documents with each other
all_documents: List[str]: List of documents where each document is a str
n_results: int: number of results to query for
normalise: bool: if True, return distances in 0-1 range
'''
if n_results <= 0:
n_results = 2 + 1
new_collection_name = str(uuid4())
collection = self.client.get_or_create_collection(
name=new_collection_name,
metadata=self.metadata,
embedding_function=self.emb_function
)
embeddings = self.model.encode(all_documents)
collection.add(
documents=all_documents,
embeddings=embeddings.tolist(),
ids=[str(uuid4()) for x in all_documents]
)
results = collection.query(
query_embeddings=embeddings.tolist(),
n_results=n_results
)
self.client.delete_collection(name=new_collection_name)
return results
def update(self, new_documents):
'''
Add new documents to the DB
new_documents: List[str]: Each document is a string
'''
self.collection.add(
ids=[str(uuid4()) for doc in new_documents],
documents=new_documents)
def purge_collection(self):
'''
Purge the DB of all documents
'''
self.client.delete_collection(name=self.collection_name)
self.collection = self.client.create_collection(
name=self.collection_name,
metadata=self.metadata,
embedding_function=self.emb_function
)
def update_h5(self, documents, embeddings):
'''
Update the vector store with documents and embeddings
documents : np.array[str]
embeddings: np.array: Size [num_docs X 768]
'''
self.collection.add(
documents=documents,
embeddings=embeddings.tolist(),
ids=[str(uuid4()) for x in documents]
)
def count(self):
'''
Get count of embeddings
'''
return self.collection.count()