Team Ai
Datasetpublic

codekingpro/portable-devtools

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