Basti-1995 commited on
Commit
03d0bbd
·
1 Parent(s): 163aa78

init flag

Browse files
Files changed (1) hide show
  1. 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(model_name="sentence-transformers/all-MiniLM-L6-v2")
 
 
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):