-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathNLPController.py
More file actions
174 lines (125 loc) · 5.7 KB
/
Copy pathNLPController.py
File metadata and controls
174 lines (125 loc) · 5.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
from .BaseController import BaseController
from stores.llm.LLMEnums import DocumentTypeEnums
from models.db_schemas import Project, DataChunk
from typing import List
import asyncio
from cohere.errors import TooManyRequestsError
import json
class NLPController(BaseController):
def __init__(self, vector_db_client, generation_client, embedding_client, template_parser):
super().__init__()
self.vector_db_client = vector_db_client
self.generation_client = generation_client
self.embedding_client = embedding_client
self.template_parser = template_parser
def create_collection_name(self, project_id: str) -> str:
return f"collection_{self.vector_db_client.default_vector_size}_{project_id}".strip()
async def reset_vector_db_collection(self, project: Project):
collection_name = self.create_collection_name(project_id=project.project_id)
return await self.vector_db_client.delete_collection(collection_name=collection_name)
async def get_vector_db_collection_info(self, project: Project):
collection_name = self.create_collection_name(project_id=project.project_id)
collection_info = await self.vector_db_client.get_collection_info(collection_name=collection_name)
return json.loads( # that usually solves the issue of non-serializable objects in the response (we actually get the error when we try to return the collection_info object in the response, so we convert it to json string and then parse it back to dict
json.dumps(collection_info, default=lambda x: x.__dict__)
)
async def index_into_db(self, project: Project, chunks: List[DataChunk], do_reset: bool = False, chunks_ids: List[str] = None) -> bool:
# step 1: get collection name
collection_name = self.create_collection_name(project_id=project.project_id)
# step 2: manage items
texts = [chunk.chunk_text for chunk in chunks]
metadatas = [chunk.chunk_metadata for chunk in chunks]
while True:
try:
vectors = self.embedding_client.embed_text(text=texts, document_type=DocumentTypeEnums.DOCUMENT.value)
break
except TooManyRequestsError:
await asyncio.sleep(60)
# step 3: create collection if not exists
_ = await self.vector_db_client.create_collection(
collection_name=collection_name,
embedding_size=self.embedding_client.embedding_size,
do_reset=do_reset
)
# step 4: insert items into collection
_ = await self.vector_db_client.insert_many(
collection_name=collection_name,
texts=texts,
vectors=vectors,
metadatas=metadatas,
record_ids=chunks_ids
)
return True
async def search_vector_db_collection(self,
project: Project,
text: str,
limit: int = 10):
# step1: get collection name
collection_name = self.create_collection_name(project_id=project.project_id)
query_vector = None
# step2: get the text embedding vector
while True:
try:
query_vectors = self.embedding_client.embed_text(text=text, document_type=DocumentTypeEnums.QUERY.value)
break
except TooManyRequestsError:
await asyncio.sleep(60)
if isinstance(query_vectors, list) and len(query_vectors) > 0:
query_vector = query_vectors[0]
if not query_vector:
return None
# step3: do semantic search
search_results = await self.vector_db_client.search_by_vector(
collection_name=collection_name,
query_vector=query_vector,
limit=limit
)
if not search_results:
return None
return search_results
async def answer_rag_query(self, project: Project, query: str, limit: int = 10):
answer, full_prompt, chat_history = None, None, None
# step1: retrieve related chunks from vector db
retrieved_chunks = await self.search_vector_db_collection(
project=project,
text=query,
limit=limit
)
if not retrieved_chunks:
return answer, full_prompt, chat_history
# step2: construct prompt for generation client (LLM)
system_prompt = self.template_parser.get_prompts(group="rag", key="system_prompt")
document_prompts = "\n".join([
self.template_parser.get_prompts(
group="rag",
key="document_prompt",
vars={
"doc_number": idx,
"doc_content": self.generation_client.process_text(chunk.text)
}
)
for idx, chunk in enumerate(retrieved_chunks, 1)
])
footer_prompt = self.template_parser.get_prompts(
group="rag",
key="footer_prompt",
vars={
"query": query
}
)
chat_history = [
self.generation_client.construct_prompt(
prompt=system_prompt,
role=self.generation_client.enums.SYSTEM.value
)
]
full_prompt = "\n\n".join([
document_prompts,
footer_prompt
])
# step3: call generation client (LLM) to get the answer
answer = self.generation_client.generate_text(
prompt=full_prompt,
chat_history=chat_history
)
return answer, full_prompt, chat_history