codekingpro/portable-devtools
114k
1import os
2import shutil
3from overrides import override
4import pickle
5from typing import Dict, List, Optional, Sequence, Set, cast
6from chromadb.config import System
7from chromadb.db.base import ParameterValue, get_sql
8from chromadb.db.impl.sqlite import SqliteDB
9from chromadb.segment.impl.vector.batch import Batch
10from chromadb.segment.impl.vector.hnsw_params import PersistentHnswParams
11from chromadb.segment.impl.vector.local_hnsw import (
12 DEFAULT_CAPACITY,
13 LocalHnswSegment,
14)
15from chromadb.segment.impl.vector.brute_force_index import BruteForceIndex
16from chromadb.telemetry.opentelemetry import (
17 OpenTelemetryClient,
18 OpenTelemetryGranularity,
19 trace_method,
20)
21from chromadb.types import (
22 LogRecord,
23 Metadata,
24 Operation,
25 RequestVersionContext,
26 Segment,
27 SeqId,
28 Vector,
29 VectorEmbeddingRecord,
30 VectorQuery,
31 VectorQueryResult,
32)
33import hnswlib
34import logging
35from pypika import Table
36import numpy as np
37
38from chromadb.utils.read_write_lock import ReadRWLock, WriteRWLock
39
40
41logger = logging.getLogger(__name__)
42
43
44class PersistentData:
45 """Stores the data and metadata needed for a PersistentLocalHnswSegment"""
46
47 dimensionality: Optional[int]
48 total_elements_added: int
49
50 max_seq_id: SeqId
51 "This is a legacy field. It is no longer mutated, but kept to allow automatic migration of the `max_seq_id` from the pickled file to the `max_seq_id` table in SQLite."
52
53 id_to_label: Dict[str, int]
54 label_to_id: Dict[int, str]
55 id_to_seq_id: Dict[str, SeqId]
56
57 def __init__(
58 self,
59 dimensionality: Optional[int],
60 total_elements_added: int,
61 id_to_label: Dict[str, int],
62 label_to_id: Dict[int, str],
63 id_to_seq_id: Dict[str, SeqId],
64 ):
65 self.dimensionality = dimensionality
66 self.total_elements_added = total_elements_added
67 self.id_to_label = id_to_label
68 self.label_to_id = label_to_id
69 self.id_to_seq_id = id_to_seq_id
70
71 @staticmethod
72 def load_from_file(filename: str) -> "PersistentData":
73 """Load persistent data from a file"""
74 with open(filename, "rb") as f:
75 ret = cast(PersistentData, pickle.load(f))
76 return ret
77
78
79class PersistentLocalHnswSegment(LocalHnswSegment):
80 METADATA_FILE: str = "index_metadata.pickle"
81 # How many records to add to index at once, we do this because crossing the python/c++ boundary is expensive (for add())
82 # When records are not added to the c++ index, they are buffered in memory and served
83 # via brute force search.
84 _batch_size: int
85 _brute_force_index: Optional[BruteForceIndex]
86 _index_initialized: bool = False
87 _curr_batch: Batch
88 # How many records to add to index before syncing to disk
89 _sync_threshold: int
90 _persist_data: PersistentData
91 _persist_directory: str
92 _allow_reset: bool
93
94 _db: SqliteDB
95 _opentelemtry_client: OpenTelemetryClient
96
97 _num_log_records_since_last_batch: int = 0
98 _num_log_records_since_last_persist: int = 0
99
100 def __init__(self, system: System, segment: Segment):
101 super().__init__(system, segment)
102
103 self._db = system.instance(SqliteDB)
104 self._opentelemtry_client = system.require(OpenTelemetryClient)
105
106 self._params = PersistentHnswParams(segment["metadata"] or {})
107 self._batch_size = self._params.batch_size
108 self._sync_threshold = self._params.sync_threshold
109 self._allow_reset = system.settings.allow_reset
110 self._persist_directory = system.settings.require("persist_directory")
111 self._curr_batch = Batch()
112 self._brute_force_index = None
113 if not os.path.exists(self._get_storage_folder()):
114 os.makedirs(self._get_storage_folder(), exist_ok=True)
115 # Load persist data if it exists already, otherwise create it
116 if self._index_exists():
117 self._persist_data = PersistentData.load_from_file(
118 self._get_metadata_file()
119 )
120 self._dimensionality = self._persist_data.dimensionality
121 self._total_elements_added = self._persist_data.total_elements_added
122 self._id_to_label = self._persist_data.id_to_label
123 self._label_to_id = self._persist_data.label_to_id
124 self._id_to_seq_id = self._persist_data.id_to_seq_id
125 # If the index was written to, we need to re-initialize it
126 if len(self._id_to_label) > 0:
127 self._dimensionality = cast(int, self._dimensionality)
128 self._init_index(self._dimensionality)
129 else:
130 self._persist_data = PersistentData(
131 self._dimensionality,
132 self._total_elements_added,
133 self._id_to_label,
134 self._label_to_id,
135 self._id_to_seq_id,
136 )
137
138 # Hydrate the max_seq_id
139 with self._db.tx() as cur:
140 t = Table("max_seq_id")
141 q = (
142 self._db.querybuilder()
143 .from_(t)
144 .select(t.seq_id)
145 .where(t.segment_id == ParameterValue(self._db.uuid_to_db(self._id)))
146 .limit(1)
147 )
148 sql, params = get_sql(q)
149 cur.execute(sql, params)
150 result = cur.fetchone()
151
152 if result:
153 self._max_seq_id = result[0]
154 elif self._index_exists():
155 # Migrate the max_seq_id from the legacy field in the pickled file to the SQLite database
156 q = (
157 self._db.querybuilder()
158 .into(Table("max_seq_id"))
159 .columns("segment_id", "seq_id")
160 .insert(
161 ParameterValue(self._db.uuid_to_db(self._id)),
162 ParameterValue(self._persist_data.max_seq_id),
163 )
164 )
165 sql, params = get_sql(q)
166 cur.execute(sql, params)
167
168 self._max_seq_id = self._persist_data.max_seq_id
169 else:
170 self._max_seq_id = self._consumer.min_seqid()
171
172 @staticmethod
173 @override
174 def propagate_collection_metadata(metadata: Metadata) -> Optional[Metadata]:
175 # Extract relevant metadata
176 segment_metadata = PersistentHnswParams.extract(metadata)
177 return segment_metadata
178
179 def _index_exists(self) -> bool:
180 """Check if the index exists via the metadata file"""
181 return os.path.exists(self._get_metadata_file())
182
183 def _get_metadata_file(self) -> str:
184 """Get the metadata file path"""
185 return os.path.join(self._get_storage_folder(), self.METADATA_FILE)
186
187 def _get_storage_folder(self) -> str:
188 """Get the storage folder path"""
189 folder = os.path.join(self._persist_directory, str(self._id))
190 return folder
191
192 @trace_method(
193 "PersistentLocalHnswSegment._init_index", OpenTelemetryGranularity.ALL
194 )
195 @override
196 def _init_index(self, dimensionality: int) -> None:
197 index = hnswlib.Index(space=self._params.space, dim=dimensionality)
198 self._brute_force_index = BruteForceIndex(
199 size=self._batch_size,
200 dimensionality=dimensionality,
201 space=self._params.space,
202 )
203
204 # Check if index exists and load it if it does
205 if self._index_exists():
206 index.load_index(
207 self._get_storage_folder(),
208 is_persistent_index=True,
209 max_elements=int(
210 max(
211 self.count(
212 request_version_context=RequestVersionContext(
213 collection_version=0, log_position=0
214 )
215 )
216 * self._params.resize_factor,
217 DEFAULT_CAPACITY,
218 )
219 ),
220 )
221 else:
222 index.init_index(
223 max_elements=DEFAULT_CAPACITY,
224 ef_construction=self._params.construction_ef,
225 M=self._params.M,
226 is_persistent_index=True,
227 persistence_location=self._get_storage_folder(),
228 )
229
230 index.set_ef(self._params.search_ef)
231 index.set_num_threads(self._params.num_threads)
232
233 self._index = index
234 self._dimensionality = dimensionality
235 self._index_initialized = True
236
237 @trace_method("PersistentLocalHnswSegment._persist", OpenTelemetryGranularity.ALL)
238 def _persist(self) -> None:
239 """Persist the index and data to disk"""
240 index = cast(hnswlib.Index, self._index)
241
242 # Persist the index
243 index.persist_dirty()
244
245 # Persist the metadata
246 self._persist_data.dimensionality = self._dimensionality
247 self._persist_data.total_elements_added = self._total_elements_added
248
249 # TODO: This should really be stored in sqlite, the index itself, or a better
250 # storage format
251 self._persist_data.id_to_label = self._id_to_label
252 self._persist_data.label_to_id = self._label_to_id
253 self._persist_data.id_to_seq_id = self._id_to_seq_id
254
255 with open(self._get_metadata_file(), "wb") as metadata_file:
256 pickle.dump(self._persist_data, metadata_file, pickle.HIGHEST_PROTOCOL)
257
258 with self._db.tx() as cur:
259 q = (
260 self._db.querybuilder()
261 .into(Table("max_seq_id"))
262 .columns("segment_id", "seq_id")
263 .insert(
264 ParameterValue(self._db.uuid_to_db(self._id)),
265 ParameterValue(self._max_seq_id),
266 )
267 )
268 sql, params = get_sql(q)
269 sql = sql.replace("INSERT", "INSERT OR REPLACE")
270 cur.execute(sql, params)
271
272 self._num_log_records_since_last_persist = 0
273
274 @trace_method(
275 "PersistentLocalHnswSegment._apply_batch", OpenTelemetryGranularity.ALL
276 )
277 @override
278 def _apply_batch(self, batch: Batch) -> None:
279 super()._apply_batch(batch)
280 if self._num_log_records_since_last_persist >= self._sync_threshold:
281 self._persist()
282
283 self._num_log_records_since_last_batch = 0
284
285 @trace_method(
286 "PersistentLocalHnswSegment._write_records", OpenTelemetryGranularity.ALL
287 )
288 @override
289 def _write_records(self, records: Sequence[LogRecord]) -> None:
290 """Add a batch of embeddings to the index"""
291 if not self._running:
292 raise RuntimeError("Cannot add embeddings to stopped component")
293 with WriteRWLock(self._lock):
294 for record in records:
295 self._num_log_records_since_last_batch += 1
296 self._num_log_records_since_last_persist += 1
297
298 if record["record"]["embedding"] is not None:
299 self._ensure_index(len(records), len(record["record"]["embedding"]))
300 if not self._index_initialized:
301 # If the index is not initialized here, it means that we have
302 # not yet added any records to the index. So we can just
303 # ignore the record since it was a delete.
304 continue
305 self._brute_force_index = cast(BruteForceIndex, self._brute_force_index)
306
307 self._max_seq_id = max(self._max_seq_id, record["log_offset"])
308 id = record["record"]["id"]
309 op = record["record"]["operation"]
310
311 exists_in_bf_index = self._brute_force_index.has_id(id)
312 exists_in_persisted_index = self._id_to_label.get(id, None) is not None
313 exists_in_index = exists_in_bf_index or exists_in_persisted_index
314
315 id_is_pending_delete = self._curr_batch.is_deleted(id)
316
317 if op == Operation.DELETE:
318 if exists_in_index:
319 self._curr_batch.apply(record)
320 if exists_in_bf_index:
321 self._brute_force_index.delete([record])
322 else:
323 logger.warning(f"Delete of nonexisting embedding ID: {id}")
324
325 elif op == Operation.UPDATE:
326 if record["record"]["embedding"] is not None:
327 if exists_in_index:
328 self._curr_batch.apply(record)
329 self._brute_force_index.upsert([record])
330 else:
331 logger.warning(
332 f"Update of nonexisting embedding ID: {record['record']['id']}"
333 )
334 elif op == Operation.ADD:
335 if record["record"]["embedding"] is not None:
336 if exists_in_index and not id_is_pending_delete:
337 logger.warning(f"Add of existing embedding ID: {id}")
338 else:
339 self._curr_batch.apply(record, not exists_in_index)
340 self._brute_force_index.upsert([record])
341 elif op == Operation.UPSERT:
342 if record["record"]["embedding"] is not None:
343 self._curr_batch.apply(record, exists_in_index)
344 self._brute_force_index.upsert([record])
345
346 if self._num_log_records_since_last_batch >= self._batch_size:
347 self._apply_batch(self._curr_batch)
348 self._curr_batch = Batch()
349 self._brute_force_index.clear()
350
351 @override
352 def count(self, request_version_context: RequestVersionContext) -> int:
353 return (
354 len(self._id_to_label)
355 + self._curr_batch.add_count
356 - self._curr_batch.delete_count
357 )
358
359 @trace_method(
360 "PersistentLocalHnswSegment.get_vectors", OpenTelemetryGranularity.ALL
361 )
362 @override
363 def get_vectors(
364 self,
365 request_version_context: RequestVersionContext,
366 ids: Optional[Sequence[str]] = None,
367 ) -> Sequence[VectorEmbeddingRecord]:
368 """Get the embeddings from the HNSW index and layered brute force
369 batch index."""
370
371 ids_hnsw: Set[str] = set()
372 ids_bf: Set[str] = set()
373
374 if self._index is not None:
375 ids_hnsw = set(self._id_to_label.keys())
376 if self._brute_force_index is not None:
377 ids_bf = set(self._curr_batch.get_written_ids())
378
379 target_ids = ids or list(ids_hnsw.union(ids_bf))
380 self._brute_force_index = cast(BruteForceIndex, self._brute_force_index)
381 hnsw_labels = []
382
383 results: List[Optional[VectorEmbeddingRecord]] = []
384 id_to_index: Dict[str, int] = {}
385 for i, id in enumerate(target_ids):
386 if id in ids_bf:
387 results.append(self._brute_force_index.get_vectors([id])[0])
388 elif id in ids_hnsw and id not in self._curr_batch._deleted_ids:
389 hnsw_labels.append(self._id_to_label[id])
390 # Placeholder for hnsw results to be filled in down below so we
391 # can batch the hnsw get() call
392 results.append(None)
393 id_to_index[id] = i
394
395 if len(hnsw_labels) > 0 and self._index is not None:
396 vectors = cast(
397 Sequence[Vector], np.array(self._index.get_items(hnsw_labels))
398 ) # version 0.8 of hnswlib allows return_type="numpy"
399
400 for label, vector in zip(hnsw_labels, vectors):
401 id = self._label_to_id[label]
402 results[id_to_index[id]] = VectorEmbeddingRecord(
403 id=id, embedding=vector
404 )
405
406 return results # type: ignore ## Python can't cast List with Optional to List with VectorEmbeddingRecord
407
408 @trace_method(
409 "PersistentLocalHnswSegment.query_vectors", OpenTelemetryGranularity.ALL
410 )
411 @override
412 def query_vectors(
413 self, query: VectorQuery
414 ) -> Sequence[Sequence[VectorQueryResult]]:
415 if self._index is None and self._brute_force_index is None:
416 return [[] for _ in range(len(query["vectors"]))]
417
418 k = query["k"]
419 if k > self.count(query["request_version_context"]):
420 count = self.count(query["request_version_context"])
421 logger.warning(
422 f"Number of requested results {k} is greater than number of elements in index {count}, updating n_results = {count}"
423 )
424 k = count
425
426 # Overquery by updated and deleted elements layered on the index because they may
427 # hide the real nearest neighbors in the hnsw index
428 hnsw_k = k + self._curr_batch.update_count + self._curr_batch.delete_count
429 # self._id_to_label contains the ids of the elements in the hnsw index
430 # so its length is the number of elements in the hnsw index
431 if hnsw_k > len(self._id_to_label):
432 hnsw_k = len(self._id_to_label)
433 hnsw_query = VectorQuery(
434 vectors=query["vectors"],
435 k=hnsw_k,
436 allowed_ids=query["allowed_ids"],
437 include_embeddings=query["include_embeddings"],
438 options=query["options"],
439 request_version_context=query["request_version_context"],
440 )
441
442 # For each query vector, we want to take the top k results from the
443 # combined results of the brute force and hnsw index
444 results: List[List[VectorQueryResult]] = []
445 self._brute_force_index = cast(BruteForceIndex, self._brute_force_index)
446 with ReadRWLock(self._lock):
447 bf_results = self._brute_force_index.query(query)
448 hnsw_results = super().query_vectors(hnsw_query)
449 for i in range(len(query["vectors"])):
450 # Merge results into a single list of size k
451 bf_pointer: int = 0
452 hnsw_pointer: int = 0
453 curr_bf_result: Sequence[VectorQueryResult] = bf_results[i]
454 curr_hnsw_result: Sequence[VectorQueryResult] = hnsw_results[i]
455
456 # Filter deleted results that haven't yet been removed from the persisted index
457 curr_hnsw_result = [
458 x
459 for x in curr_hnsw_result
460 if not self._curr_batch.is_deleted(x["id"])
461 ]
462
463 curr_results: List[VectorQueryResult] = []
464 # In the case where filters cause the number of results to be less than k,
465 # we set k to be the number of results
466 total_results = len(curr_bf_result) + len(curr_hnsw_result)
467 if total_results == 0:
468 results.append([])
469 else:
470 while len(curr_results) < min(k, total_results):
471 if bf_pointer < len(curr_bf_result) and hnsw_pointer < len(
472 curr_hnsw_result
473 ):
474 bf_dist = curr_bf_result[bf_pointer]["distance"]
475 hnsw_dist = curr_hnsw_result[hnsw_pointer]["distance"]
476 if bf_dist <= hnsw_dist:
477 curr_results.append(curr_bf_result[bf_pointer])
478 bf_pointer += 1
479 else:
480 id = curr_hnsw_result[hnsw_pointer]["id"]
481 # Only add the hnsw result if it is not in the brute force index
482 if not self._brute_force_index.has_id(id):
483 curr_results.append(curr_hnsw_result[hnsw_pointer])
484 hnsw_pointer += 1
485 else:
486 break
487 remaining = min(k, total_results) - len(curr_results)
488 if remaining > 0 and hnsw_pointer < len(curr_hnsw_result):
489 for i in range(
490 hnsw_pointer,
491 min(len(curr_hnsw_result), hnsw_pointer + remaining),
492 ):
493 id = curr_hnsw_result[i]["id"]
494 if not self._brute_force_index.has_id(id):
495 curr_results.append(curr_hnsw_result[i])
496 elif remaining > 0 and bf_pointer < len(curr_bf_result):
497 curr_results.extend(
498 curr_bf_result[bf_pointer : bf_pointer + remaining]
499 )
500 results.append(curr_results)
501 return results
502
503 @trace_method(
504 "PersistentLocalHnswSegment.reset_state", OpenTelemetryGranularity.ALL
505 )
506 @override
507 def reset_state(self) -> None:
508 if self._allow_reset:
509 data_path = self._get_storage_folder()
510 if os.path.exists(data_path):
511 self.close_persistent_index()
512 shutil.rmtree(data_path, ignore_errors=True)
513
514 @trace_method("PersistentLocalHnswSegment.delete", OpenTelemetryGranularity.ALL)
515 @override
516 def delete(self) -> None:
517 data_path = self._get_storage_folder()
518 if os.path.exists(data_path):
519 self.close_persistent_index()
520 shutil.rmtree(data_path, ignore_errors=False)
521
522 @staticmethod
523 def get_file_handle_count() -> int:
524 """Return how many file handles are used by the index"""
525 hnswlib_count = hnswlib.Index.file_handle_count
526 hnswlib_count = cast(int, hnswlib_count)
527 # One extra for the metadata file
528 return hnswlib_count + 1 # type: ignore
529
530 def open_persistent_index(self) -> None:
531 """Open the persistent index"""
532 if self._index is not None:
533 self._index.open_file_handles()
534
535 @override
536 def stop(self) -> None:
537 super().stop()
538 self.close_persistent_index()
539
540 def close_persistent_index(self) -> None:
541 """Close the persistent index"""
542 if self._index is not None:
543 self._index.close_file_handles()
544 