(self, query)
| 167 | return vector_retriever |
| 168 | |
| 169 | def search(self, query): |
| 170 | if self.vl_ret and 'vidore' in self.embed_model_name: |
| 171 | query_embedding = self.vector_embed_model.embed_text(query) |
| 172 | scores = self.vector_embed_model.processor.score(query_embedding,self.embedding_img) |
| 173 | k = min(100, scores[0].numel()) |
| 174 | values, indices = torch.topk(scores[0], k=k) |
| 175 | recall_results = [self.nodes[i] for i in indices] |
| 176 | for node in recall_results: |
| 177 | node.embedding = None |
| 178 | recall_results = [NodeWithScore(node=node, score=score) for node, score in zip(recall_results, values)] |
| 179 | recall_results_output = recall_results |
| 180 | else: |
| 181 | query_bundle = QueryBundle(query_str=query) |
| 182 | recall_results = self.query_engine.retrieve(query_bundle) |
| 183 | recall_results_output = recall_results |
| 184 | if self.gmm: |
| 185 | recall_results_output = gmm(recall_results,self.input_gmm,self.max_output_gmm,self.min_output_gmm) |
| 186 | if self.return_raw: |
| 187 | return recall_results_output |
| 188 | if self.gmm_candidate_length: |
| 189 | candidate_length = [1,2,4,6,9,12,16,20] |
| 190 | current_length = len(recall_results_output) |
| 191 | target_length = min([num for num in candidate_length if num > current_length]) |
| 192 | recall_results_output = recall_results[:target_length] |
| 193 | return nodes2dict(recall_results_output) |
| 194 | |
| 195 | def search_example(self,example): |
| 196 | query = example['query'] |
no test coverage detected