| 8 | |
| 9 | |
| 10 | class SearchEngine: |
| 11 | def __init__(self, embed_model_name='GVE'): # Alibaba-NLP/GVE-3B Alibaba-NLP/GVE-7B |
| 12 | |
| 13 | self.embed_model_name = embed_model_name |
| 14 | |
| 15 | self.embed_model_name_file = embed_model_name.replace('/', '_').replace('-','_') |
| 16 | |
| 17 | self.device = "cuda" if torch.cuda.is_available() else "cpu" |
| 18 | |
| 19 | if 'GVE' in self.embed_model_name: |
| 20 | from models.GVE import AutoModelForSentenceEmbeddingTriplet |
| 21 | from models.GVE import VLProcessor |
| 22 | if '3B' in self.embed_model_name: |
| 23 | model_name_or_path = "Alibaba-NLP/GVE-3B" |
| 24 | self.dimension = 2048 |
| 25 | elif '7B' in self.embed_model_name: |
| 26 | model_name_or_path = "Alibaba-NLP/GVE-7B" |
| 27 | self.dimension = 3584 |
| 28 | |
| 29 | self.model = AutoModelForSentenceEmbeddingTriplet(model_name_or_path) |
| 30 | self.processor = VLProcessor(model_name_or_path) |
| 31 | |
| 32 | elif 'bge-m3' in self.embed_model_name: |
| 33 | from sentence_transformers import SentenceTransformer |
| 34 | self.model = SentenceTransformer(self.embed_model_name) |
| 35 | self.dimension = 1024 |
| 36 | |
| 37 | elif 'Qwen3-VL' in self.embed_model_name or 'Qwen3_VL' in self.embed_model_name: |
| 38 | from models.Qwen3_VL_Embedding.qwen3_vl_embedding import Qwen3VLEmbedder |
| 39 | if '2B' in self.embed_model_name: |
| 40 | self.dimension = 2048 |
| 41 | elif '8B' in self.embed_model_name: |
| 42 | self.dimension = 4096 |
| 43 | self.model = Qwen3VLEmbedder(self.embed_model_name) |
| 44 | |
| 45 | |
| 46 | def load_index(self, input_dir): |
| 47 | # meta data |
| 48 | with open(os.path.join(input_dir, f"{self.embed_model_name_file}_meta_data.json"), "r") as f: |
| 49 | self.meta_data = json.load(f) |
| 50 | assert self.meta_data["model_name"] == self.embed_model_name |
| 51 | print(f"loading model {self.embed_model_name}; loading index from {input_dir}; loading index length {self.meta_data['num_vectors']}...") |
| 52 | |
| 53 | self.index = faiss.read_index(os.path.join(input_dir, f"{self.embed_model_name_file}_faiss_index.bin")) |
| 54 | with open(os.path.join(input_dir, f"{self.embed_model_name_file}_file_data_list.jsonl"), "r") as f: |
| 55 | self.file_data_list = [json.loads(line) for line in f] |
| 56 | assert self.meta_data["num_vectors"] == self.index.ntotal \ |
| 57 | and len(self.file_data_list) == self.index.ntotal, \ |
| 58 | f"meta_data_num_vectors: {self.meta_data['num_vectors']} and index_ntotal: {self.index.ntotal} mismatch \n or file_data_list_length: {len(self.file_data_list)} and index_ntotal: {self.index.ntotal} mismatch" |
| 59 | |
| 60 | def load_multi_index_corpus_together(self, input_dir_list): |
| 61 | for input_dir in input_dir_list: |
| 62 | if not hasattr(self, "meta_data"): |
| 63 | with open(os.path.join(input_dir, f"{self.embed_model_name_file}_meta_data.json"), "r") as f: |
| 64 | self.meta_data = json.load(f) |
| 65 | num_vectors = self.meta_data['num_vectors'] |
| 66 | print(f"Loading index from {input_dir}: {num_vectors}") |
| 67 |
no outgoing calls
no test coverage detected