codekingpro/portable-devtools
114k
1from overrides import override
2from typing import Optional, Sequence, Dict, Set, List, cast
3from uuid import UUID
4from chromadb.segment import VectorReader
5from chromadb.ingest import Consumer
6from chromadb.config import System, Settings
7from chromadb.segment.impl.vector.batch import Batch
8from chromadb.segment.impl.vector.hnsw_params import HnswParams
9from chromadb.telemetry.opentelemetry import (
10 OpenTelemetryClient,
11 OpenTelemetryGranularity,
12 trace_method,
13)
14from chromadb.types import (
15 LogRecord,
16 RequestVersionContext,
17 VectorEmbeddingRecord,
18 VectorQuery,
19 VectorQueryResult,
20 SeqId,
21 Segment,
22 Metadata,
23 Operation,
24 Vector,
25)
26from chromadb.errors import InvalidDimensionException
27import hnswlib
28from chromadb.utils.read_write_lock import ReadWriteLock, ReadRWLock, WriteRWLock
29import logging
30import numpy as np
31
32logger = logging.getLogger(__name__)
33
34DEFAULT_CAPACITY = 1000
35
36
37class LocalHnswSegment(VectorReader):
38 _id: UUID
39 _consumer: Consumer
40 _collection: Optional[UUID]
41 _subscription: Optional[UUID]
42 _settings: Settings
43 _params: HnswParams
44
45 _index: Optional[hnswlib.Index]
46 _dimensionality: Optional[int]
47 _total_elements_added: int
48 _max_seq_id: SeqId
49
50 _lock: ReadWriteLock
51
52 _id_to_label: Dict[str, int]
53 _label_to_id: Dict[int, str]
54 # Note: As of the time of writing, this mapping is no longer needed.
55 # We merely keep it around for easy compatibility with the old code and
56 # debugging purposes.
57 _id_to_seq_id: Dict[str, SeqId]
58
59 _opentelemtry_client: OpenTelemetryClient
60
61 def __init__(self, system: System, segment: Segment):
62 self._consumer = system.instance(Consumer)
63 self._id = segment["id"]
64 self._collection = segment["collection"]
65 self._subscription = None
66 self._settings = system.settings
67 self._params = HnswParams(segment["metadata"] or {})
68
69 self._index = None
70 self._dimensionality = None
71 self._total_elements_added = 0
72 self._max_seq_id = self._consumer.min_seqid()
73
74 self._id_to_seq_id = {}
75 self._id_to_label = {}
76 self._label_to_id = {}
77
78 self._lock = ReadWriteLock()
79 self._opentelemtry_client = system.require(OpenTelemetryClient)
80
81 @staticmethod
82 @override
83 def propagate_collection_metadata(metadata: Metadata) -> Optional[Metadata]:
84 # Extract relevant metadata
85 segment_metadata = HnswParams.extract(metadata)
86 return segment_metadata
87
88 @trace_method("LocalHnswSegment.start", OpenTelemetryGranularity.ALL)
89 @override
90 def start(self) -> None:
91 super().start()
92 if self._collection:
93 seq_id = self.max_seqid()
94 self._subscription = self._consumer.subscribe(
95 self._collection, self._write_records, start=seq_id
96 )
97
98 @trace_method("LocalHnswSegment.stop", OpenTelemetryGranularity.ALL)
99 @override
100 def stop(self) -> None:
101 super().stop()
102 if self._subscription:
103 self._consumer.unsubscribe(self._subscription)
104
105 @trace_method("LocalHnswSegment.get_vectors", OpenTelemetryGranularity.ALL)
106 @override
107 def get_vectors(
108 self,
109 request_version_context: RequestVersionContext,
110 ids: Optional[Sequence[str]] = None,
111 ) -> Sequence[VectorEmbeddingRecord]:
112 if ids is None:
113 labels = list(self._label_to_id.keys())
114 else:
115 labels = []
116 for id in ids:
117 if id in self._id_to_label:
118 labels.append(self._id_to_label[id])
119
120 results = []
121 if self._index is not None:
122 vectors = cast(
123 Sequence[Vector], np.array(self._index.get_items(labels))
124 ) # version 0.8 of hnswlib allows return_type="numpy"
125
126 for label, vector in zip(labels, vectors):
127 id = self._label_to_id[label]
128 results.append(VectorEmbeddingRecord(id=id, embedding=vector))
129
130 return results
131
132 @trace_method("LocalHnswSegment.query_vectors", OpenTelemetryGranularity.ALL)
133 @override
134 def query_vectors(
135 self, query: VectorQuery
136 ) -> Sequence[Sequence[VectorQueryResult]]:
137 if self._index is None:
138 return [[] for _ in range(len(query["vectors"]))]
139
140 k = query["k"]
141 size = len(self._id_to_label)
142
143 if k > size:
144 logger.warning(
145 f"Number of requested results {k} is greater than number of elements in index {size}, updating n_results = {size}"
146 )
147 k = size
148
149 labels: Set[int] = set()
150 ids = query["allowed_ids"]
151 if ids is not None:
152 labels = {self._id_to_label[id] for id in ids if id in self._id_to_label}
153 if len(labels) < k:
154 k = len(labels)
155
156 def filter_function(label: int) -> bool:
157 return label in labels
158
159 query_vectors = query["vectors"]
160
161 with ReadRWLock(self._lock):
162 result_labels, distances = self._index.knn_query(
163 np.array(query_vectors, dtype=np.float32),
164 k=k,
165 filter=filter_function if ids else None,
166 )
167
168 # TODO: these casts are not correct, hnswlib returns np
169 # distances = cast(List[List[float]], distances)
170 # result_labels = cast(List[List[int]], result_labels)
171
172 all_results: List[List[VectorQueryResult]] = []
173 for result_i in range(len(result_labels)):
174 results: List[VectorQueryResult] = []
175 for label, distance in zip(
176 result_labels[result_i], distances[result_i]
177 ):
178 id = self._label_to_id[label]
179 if query["include_embeddings"]:
180 embedding = np.array(
181 self._index.get_items([label])[0]
182 ) # version 0.8 of hnswlib allows return_type="numpy"
183 else:
184 embedding = None
185 results.append(
186 VectorQueryResult(
187 id=id,
188 distance=distance.item(),
189 embedding=embedding,
190 )
191 )
192 all_results.append(results)
193
194 return all_results
195
196 @override
197 def max_seqid(self) -> SeqId:
198 return self._max_seq_id
199
200 @override
201 def count(self, request_version_context: RequestVersionContext) -> int:
202 return len(self._id_to_label)
203
204 @trace_method("LocalHnswSegment._init_index", OpenTelemetryGranularity.ALL)
205 def _init_index(self, dimensionality: int) -> None:
206 # more comments available at the source: https://github.com/nmslib/hnswlib
207
208 index = hnswlib.Index(
209 space=self._params.space, dim=dimensionality
210 ) # possible options are l2, cosine or ip
211 index.init_index(
212 max_elements=DEFAULT_CAPACITY,
213 ef_construction=self._params.construction_ef,
214 M=self._params.M,
215 )
216 index.set_ef(self._params.search_ef)
217 index.set_num_threads(self._params.num_threads)
218
219 self._index = index
220 self._dimensionality = dimensionality
221
222 @trace_method("LocalHnswSegment._ensure_index", OpenTelemetryGranularity.ALL)
223 def _ensure_index(self, n: int, dim: int) -> None:
224 """Create or resize the index as necessary to accomodate N new records"""
225 if not self._index:
226 self._dimensionality = dim
227 self._init_index(dim)
228 else:
229 if dim != self._dimensionality:
230 raise InvalidDimensionException(
231 f"Dimensionality of ({dim}) does not match index"
232 + f"dimensionality ({self._dimensionality})"
233 )
234
235 index = cast(hnswlib.Index, self._index)
236
237 if (self._total_elements_added + n) > index.get_max_elements():
238 new_size = int(
239 (self._total_elements_added + n) * self._params.resize_factor
240 )
241 index.resize_index(max(new_size, DEFAULT_CAPACITY))
242
243 @trace_method("LocalHnswSegment._apply_batch", OpenTelemetryGranularity.ALL)
244 def _apply_batch(self, batch: Batch) -> None:
245 """Apply a batch of changes, as atomically as possible."""
246 deleted_ids = batch.get_deleted_ids()
247 written_ids = batch.get_written_ids()
248 vectors_to_write = batch.get_written_vectors(written_ids)
249 labels_to_write = [0] * len(vectors_to_write)
250
251 if len(deleted_ids) > 0:
252 index = cast(hnswlib.Index, self._index)
253 for i in range(len(deleted_ids)):
254 id = deleted_ids[i]
255 # Never added this id to hnsw, so we can safely ignore it for deletions
256 if id not in self._id_to_label:
257 continue
258 label = self._id_to_label[id]
259
260 index.mark_deleted(label)
261 del self._id_to_label[id]
262 del self._label_to_id[label]
263 del self._id_to_seq_id[id]
264
265 if len(written_ids) > 0:
266 self._ensure_index(batch.add_count, len(vectors_to_write[0]))
267
268 next_label = self._total_elements_added + 1
269 for i in range(len(written_ids)):
270 if written_ids[i] not in self._id_to_label:
271 labels_to_write[i] = next_label
272 next_label += 1
273 else:
274 labels_to_write[i] = self._id_to_label[written_ids[i]]
275
276 index = cast(hnswlib.Index, self._index)
277
278 # First, update the index
279 index.add_items(vectors_to_write, labels_to_write)
280
281 # If that succeeds, update the mappings
282 for i, id in enumerate(written_ids):
283 self._id_to_seq_id[id] = batch.get_record(id)["log_offset"]
284 self._id_to_label[id] = labels_to_write[i]
285 self._label_to_id[labels_to_write[i]] = id
286
287 # If that succeeds, update the total count
288 self._total_elements_added += batch.add_count
289
290 @trace_method("LocalHnswSegment._write_records", OpenTelemetryGranularity.ALL)
291 def _write_records(self, records: Sequence[LogRecord]) -> None:
292 """Add a batch of embeddings to the index"""
293 if not self._running:
294 raise RuntimeError("Cannot add embeddings to stopped component")
295
296 # Avoid all sorts of potential problems by ensuring single-threaded access
297 with WriteRWLock(self._lock):
298 batch = Batch()
299
300 for record in records:
301 self._max_seq_id = max(self._max_seq_id, record["log_offset"])
302 id = record["record"]["id"]
303 op = record["record"]["operation"]
304 label = self._id_to_label.get(id, None)
305
306 if op == Operation.DELETE:
307 if label:
308 batch.apply(record)
309 else:
310 logger.warning(f"Delete of nonexisting embedding ID: {id}")
311
312 elif op == Operation.UPDATE:
313 if record["record"]["embedding"] is not None:
314 if label is not None:
315 batch.apply(record)
316 else:
317 logger.warning(
318 f"Update of nonexisting embedding ID: {record['record']['id']}"
319 )
320 elif op == Operation.ADD:
321 if not label:
322 batch.apply(record, False)
323 else:
324 logger.warning(f"Add of existing embedding ID: {id}")
325 elif op == Operation.UPSERT:
326 batch.apply(record, label is not None)
327
328 self._apply_batch(batch)
329
330 @override
331 def delete(self) -> None:
332 raise NotImplementedError()
333 