Commit ·
03d0bbd
1
Parent(s): 163aa78
init flag
Browse files- retriever.py +6 -1
retriever.py
CHANGED
|
@@ -18,7 +18,9 @@ class HybridRetriever:
|
|
| 18 |
self.docs = docs
|
| 19 |
self.mode = mode
|
| 20 |
self.k = k
|
| 21 |
-
self.embedding_model = HuggingFaceEmbeddings(
|
|
|
|
|
|
|
| 22 |
|
| 23 |
# Initialize BM25 retriever
|
| 24 |
self.bm25 = BM25Retriever.from_documents(docs)
|
|
@@ -52,6 +54,8 @@ class HybridRetriever:
|
|
| 52 |
scores = []
|
| 53 |
for doc in bm25_candidates:
|
| 54 |
doc_vec = self.doc_embeddings.get(doc.page_content)
|
|
|
|
|
|
|
| 55 |
if doc_vec is not None:
|
| 56 |
sim = np.dot(query_embedding, doc_vec) / (
|
| 57 |
np.linalg.norm(query_embedding) * np.linalg.norm(doc_vec)
|
|
@@ -80,6 +84,7 @@ class GuestInfoHybridTool(Tool):
|
|
| 80 |
output_type = "string"
|
| 81 |
|
| 82 |
def __init__(self, docs, mode="rerank"):
|
|
|
|
| 83 |
self.retriever = HybridRetriever(docs, mode=mode)
|
| 84 |
|
| 85 |
def forward(self, query: str):
|
|
|
|
| 18 |
self.docs = docs
|
| 19 |
self.mode = mode
|
| 20 |
self.k = k
|
| 21 |
+
self.embedding_model = HuggingFaceEmbeddings(
|
| 22 |
+
model_name="sentence-transformers/all-MiniLM-L6-v2"
|
| 23 |
+
)
|
| 24 |
|
| 25 |
# Initialize BM25 retriever
|
| 26 |
self.bm25 = BM25Retriever.from_documents(docs)
|
|
|
|
| 54 |
scores = []
|
| 55 |
for doc in bm25_candidates:
|
| 56 |
doc_vec = self.doc_embeddings.get(doc.page_content)
|
| 57 |
+
|
| 58 |
+
# similarity calculation
|
| 59 |
if doc_vec is not None:
|
| 60 |
sim = np.dot(query_embedding, doc_vec) / (
|
| 61 |
np.linalg.norm(query_embedding) * np.linalg.norm(doc_vec)
|
|
|
|
| 84 |
output_type = "string"
|
| 85 |
|
| 86 |
def __init__(self, docs, mode="rerank"):
|
| 87 |
+
self.is_initialized = False # Flag to check if the tool is initialized
|
| 88 |
self.retriever = HybridRetriever(docs, mode=mode)
|
| 89 |
|
| 90 |
def forward(self, query: str):
|