codekingpro/portable-devtools
114k
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 