| 216 | json.dump(results, json_file, indent=2, ensure_ascii=False) |
| 217 | |
| 218 | class HybridSearchEngine: |
| 219 | |
| 220 | def __init__(self, |
| 221 | dataset, |
| 222 | node_dir_prefix_vl = None, |
| 223 | node_dir_prefix_text = None, |
| 224 | embed_model_name_vl = 'vidore/colqwen2-v1.0', |
| 225 | embed_model_name_text = 'BAAI/bge-m3', |
| 226 | topk=10, |
| 227 | gmm=False): |
| 228 | self.dataset = dataset |
| 229 | self.dataset_dir = os.path.join('./data', dataset) |
| 230 | self.img_dir = os.path.join(self.dataset_dir, 'img') |
| 231 | self.ppocr_dir = os.path.join(self.dataset_dir, 'bge_ingestion') |
| 232 | self.engine_vl = SearchEngine(dataset,node_dir_prefix=node_dir_prefix_vl,embed_model_name=embed_model_name_vl) |
| 233 | self.engine_text = SearchEngine(dataset,node_dir_prefix=node_dir_prefix_text,embed_model_name=embed_model_name_text) |
| 234 | self.topk = topk |
| 235 | self.gmm = gmm |
| 236 | |
| 237 | def search(self,query): |
| 238 | union_result = False |
| 239 | if union_result: |
| 240 | if self.gmm: |
| 241 | self.engine_vl.gmm = True |
| 242 | self.engine_text.gmm = True |
| 243 | self.engine_vl.input_gmm = self.topk *2 |
| 244 | self.engine_text.input_gmm = self.topk *2 |
| 245 | self.engine_vl.max_output_gmm = self.topk |
| 246 | self.engine_text.max_output_gmm = self.topk |
| 247 | self.engine_vl.min_output_gmm = self.topk//2 |
| 248 | self.engine_text.min_output_gmm = self.topk//2 |
| 249 | result_vl = self.engine_vl.search(query) |
| 250 | result_text = self.engine_text.search(query) |
| 251 | result_vl['source_nodes'] = result_vl['source_nodes'][:self.topk] |
| 252 | result_text['source_nodes'] = result_text['source_nodes'][:self.topk] |
| 253 | |
| 254 | result_docs = dict() |
| 255 | for node in result_vl['source_nodes']: |
| 256 | file = os.path.basename(node['node']['image_path']).split('.')[0] |
| 257 | doc = '_'.join(file.split('_')[:-1]) |
| 258 | page = file.split('_')[-1] |
| 259 | if doc not in result_docs: |
| 260 | result_docs[doc] = [int(page)] |
| 261 | else: |
| 262 | if int(page) not in result_docs[doc]: |
| 263 | result_docs[doc].append(int(page)) |
| 264 | |
| 265 | for node in result_text['source_nodes']: |
| 266 | file = node['node']['metadata']['filename'].split('.')[0] |
| 267 | doc = '_'.join(file.split('_')[:-1]) |
| 268 | page = file.split('_')[-1] |
| 269 | if doc not in result_docs: |
| 270 | result_docs[doc] = [int(page)] |
| 271 | else: |
| 272 | if int(page) not in result_docs[doc]: |
| 273 | result_docs[doc].append(int(page)) |
| 274 | |
| 275 | recall_result = [] |
no outgoing calls
no test coverage detected