Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
brute_force_index.py152 linesDownload Raw Back to vector
1from typing import Any, Callable, Dict, List, Optional, Sequence, Set
2import numpy as np
3import numpy.typing as npt
4from chromadb.types import (
5    LogRecord,
6    VectorEmbeddingRecord,
7    VectorQuery,
8    VectorQueryResult,
9)
10
11from chromadb.utils import distance_functions
12import logging
13
14logger = logging.getLogger(__name__)
15
16
17class BruteForceIndex:
18    """A lightweight, numpy based brute force index that is used for batches that have not been indexed into hnsw yet. It is not
19    thread safe and callers should ensure that only one thread is accessing it at a time.
20    """
21
22    id_to_index: Dict[str, int]
23    index_to_id: Dict[int, str]
24    id_to_seq_id: Dict[str, int]
25    deleted_ids: Set[str]
26    free_indices: List[int]
27    size: int
28    dimensionality: int
29    distance_fn: Callable[[npt.NDArray[Any], npt.NDArray[Any]], float]
30    vectors: npt.NDArray[Any]
31
32    def __init__(self, size: int, dimensionality: int, space: str = "l2"):
33        if space == "l2":
34            self.distance_fn = distance_functions.l2
35        elif space == "ip":
36            self.distance_fn = distance_functions.ip
37        elif space == "cosine":
38            self.distance_fn = distance_functions.cosine
39        else:
40            raise Exception(f"Unknown distance function: {space}")
41
42        self.id_to_index = {}
43        self.index_to_id = {}
44        self.id_to_seq_id = {}
45        self.deleted_ids = set()
46        self.free_indices = list(range(size))
47        self.size = size
48        self.dimensionality = dimensionality
49        self.vectors = np.zeros((size, dimensionality))
50
51    def __len__(self) -> int:
52        return len(self.id_to_index)
53
54    def clear(self) -> None:
55        self.id_to_index = {}
56        self.index_to_id = {}
57        self.id_to_seq_id = {}
58        self.deleted_ids.clear()
59        self.free_indices = list(range(self.size))
60        self.vectors.fill(0)
61
62    def upsert(self, records: List[LogRecord]) -> None:
63        if len(records) + len(self) > self.size:
64            raise Exception(
65                "Index with capacity {} and {} current entries cannot add {} records".format(
66                    self.size, len(self), len(records)
67                )
68            )
69
70        for i, record in enumerate(records):
71            id = record["record"]["id"]
72            vector = record["record"]["embedding"]
73            self.id_to_seq_id[id] = record["log_offset"]
74            if id in self.deleted_ids:
75                self.deleted_ids.remove(id)
76
77            # TODO: It may be faster to use multi-index selection on the vectors array
78            if id in self.id_to_index:
79                # Update
80                index = self.id_to_index[id]
81                self.vectors[index] = vector
82            else:
83                # Add
84                next_index = self.free_indices.pop()
85                self.id_to_index[id] = next_index
86                self.index_to_id[next_index] = id
87                self.vectors[next_index] = vector
88
89    def delete(self, records: List[LogRecord]) -> None:
90        for record in records:
91            id = record["record"]["id"]
92            if id in self.id_to_index:
93                index = self.id_to_index[id]
94                self.deleted_ids.add(id)
95                del self.id_to_index[id]
96                del self.index_to_id[index]
97                del self.id_to_seq_id[id]
98                self.vectors[index].fill(np.nan)
99                self.free_indices.append(index)
100            else:
101                logger.warning(f"Delete of nonexisting embedding ID: {id}")
102
103    def has_id(self, id: str) -> bool:
104        """Returns whether the index contains the given ID"""
105        return id in self.id_to_index and id not in self.deleted_ids
106
107    def get_vectors(
108        self, ids: Optional[Sequence[str]] = None
109    ) -> Sequence[VectorEmbeddingRecord]:
110        target_ids = ids or self.id_to_index.keys()
111
112        return [
113            VectorEmbeddingRecord(
114                id=id,
115                embedding=self.vectors[self.id_to_index[id]],
116            )
117            for id in target_ids
118        ]
119
120    def query(self, query: VectorQuery) -> Sequence[Sequence[VectorQueryResult]]:
121        np_query = np.array(query["vectors"], dtype=np.float32)
122        allowed_ids = (
123            None if query["allowed_ids"] is None else set(query["allowed_ids"])
124        )
125        distances = np.apply_along_axis(
126            lambda query: np.apply_along_axis(self.distance_fn, 1, self.vectors, query),
127            1,
128            np_query,
129        )
130
131        indices = np.argsort(distances)
132        # Filter out deleted labels
133        filtered_results = []
134        for i, index_list in enumerate(indices):
135            curr_results = []
136            for j in index_list:
137                # If the index is in the index_to_id map, then it has been added
138                if j in self.index_to_id:
139                    id = self.index_to_id[j]
140                    if id not in self.deleted_ids and (
141                        allowed_ids is None or id in allowed_ids
142                    ):
143                        curr_results.append(
144                            VectorQueryResult(
145                                id=id,
146                                distance=distances[i][j].item(),
147                                embedding=self.vectors[j],
148                            )
149                        )
150            filtered_results.append(curr_results)
151        return filtered_results
152 
codekingpro/portable-devtools · Team Ai