MCPcopy Create free account
hub / github.com/Alibaba-NLP/VRAG / SearchEngine

Class SearchEngine

search_engine/search_engine.py:10–332  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class 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

Callers 2

search_engine.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected