| 554 | |
| 555 | |
| 556 | def brute_force_dense( |
| 557 | query_embeddings: np.ndarray, |
| 558 | corpus_embeddings: np.ndarray, |
| 559 | corpus_ids: List[str], |
| 560 | query_ids: List[str], |
| 561 | top_k: int, |
| 562 | dataset: str, |
| 563 | ) -> List[Dict]: |
| 564 | bf_file = Path("datasets") / f"hybrid_{dataset}" / "dense_brute_force.pkl" |
| 565 | if bf_file.exists(): |
| 566 | return pickle.loads(bf_file.read_bytes()) |
| 567 | |
| 568 | print("Computing dense brute-force...") |
| 569 | query_norms = np.linalg.norm(query_embeddings, axis=1, keepdims=True) |
| 570 | corpus_norms = np.linalg.norm(corpus_embeddings, axis=1, keepdims=True) |
| 571 | query_embeddings_norm = query_embeddings / (query_norms + 1e-8) |
| 572 | corpus_embeddings_norm = corpus_embeddings / (corpus_norms + 1e-8) |
| 573 | |
| 574 | results = [] |
| 575 | for i, query_id in enumerate(tqdm(query_ids, desc="Dense brute-force")): |
| 576 | scores = np.dot(corpus_embeddings_norm, query_embeddings_norm[i]) |
| 577 | top_indices = np.argsort(scores)[::-1][:top_k] |
| 578 | top_results = [ |
| 579 | {"id": corpus_ids[idx], "score": float(scores[idx])} for idx in top_indices |
| 580 | ] |
| 581 | results.append({"query_id": query_id, "top_results": top_results}) |
| 582 | |
| 583 | bf_file.parent.mkdir(parents=True, exist_ok=True) |
| 584 | bf_file.write_bytes(pickle.dumps(results)) |
| 585 | return results |
| 586 | |
| 587 | |
| 588 | def build_sparse_vectors_for_bf(corpus: Dict, dataset: str) -> List[Dict]: |