| 9 | |
| 10 | |
| 11 | class Faiss(VectorBase): |
| 12 | def __init__(self, index_file_path, dimension, top_k): |
| 13 | self._index_file_path = index_file_path |
| 14 | self._dimension = dimension |
| 15 | self._index = faiss.index_factory(self._dimension, "IDMap,Flat", faiss.METRIC_L2) |
| 16 | self._top_k = top_k |
| 17 | if os.path.isfile(index_file_path): |
| 18 | self._index = faiss.read_index(index_file_path) |
| 19 | |
| 20 | def mul_add(self, datas: List[VectorData], model=None): |
| 21 | data_array, id_array = map(list, zip(*((data.data, data.id) for data in datas))) |
| 22 | np_data = np.array(data_array).astype("float32") |
| 23 | ids = np.array(id_array) |
| 24 | self._index.add_with_ids(np_data, ids) |
| 25 | |
| 26 | def search(self, data: np.ndarray, top_k: int = -1, model=None): |
| 27 | if self._index.ntotal == 0: |
| 28 | return None |
| 29 | if top_k == -1: |
| 30 | top_k = self._top_k |
| 31 | np_data = np.array(data).astype("float32").reshape(1, -1) |
| 32 | dist, ids = self._index.search(np_data, top_k) |
| 33 | ids = [int(i) for i in ids[0]] |
| 34 | return list(zip(dist[0], ids)) |
| 35 | |
| 36 | def rebuild_col(self, ids=None): |
| 37 | return True |
| 38 | |
| 39 | def rebuild(self, ids=None): |
| 40 | return True |
| 41 | |
| 42 | def delete(self, ids): |
| 43 | ids_to_remove = np.array(ids) |
| 44 | self._index.remove_ids(faiss.IDSelectorBatch(ids_to_remove.size, faiss.swig_ptr(ids_to_remove))) |
| 45 | |
| 46 | def flush(self): |
| 47 | faiss.write_index(self._index, self._index_file_path) |
| 48 | |
| 49 | def close(self): |
| 50 | self.flush() |
| 51 | |
| 52 | def count(self): |
| 53 | return self._index.ntotal |