Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
local_persistent_hnsw.py544 linesDownload Raw Back to vector
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 
codekingpro/portable-devtools · Team Ai