MCPcopy Create free account
hub / github.com/OpenBMB/BMTools / Retriever

Class Retriever

bmtools/tools/retriever.py:6–42  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4import os
5
6class Retriever:
7 def __init__(self,
8 openai_api_key: str = None,
9 model: str = "text-embedding-ada-002"):
10 if openai_api_key is None:
11 openai_api_key = os.environ.get("OPENAI_API_KEY")
12 self.embed = OpenAIEmbeddings(openai_api_key=openai_api_key, model=model)
13 self.documents = dict()
14
15 def add_tool(self, tool_name: str, api_info: Dict) -> None:
16 if tool_name in self.documents:
17 return
18 document = api_info["name_for_model"] + ". " + api_info["description_for_model"]
19 document_embedding = self.embed.embed_documents([document])
20 self.documents[tool_name] = {
21 "document": document,
22 "embedding": document_embedding[0]
23 }
24
25 def query(self, query: str, topk: int = 3) -> List[str]:
26 query_embedding = self.embed.embed_query(query)
27
28 queue = PriorityQueue()
29 for tool_name, tool_info in self.documents.items():
30 tool_embedding = tool_info["embedding"]
31 tool_sim = self.similarity(query_embedding, tool_embedding)
32 queue.put([-tool_sim, tool_name])
33
34 result = []
35 for i in range(min(topk, len(queue.queue))):
36 tool = queue.get()
37 result.append(tool[1])
38
39 return result
40
41 def similarity(self, query: List[float], document: List[float]) -> float:
42 return sum([i * j for i, j in zip(query, document)])
43

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected