Spaces:
Runtime error
Runtime error
File size: 7,327 Bytes
e34be6c 2fe4e9d 1543ec3 2fe4e9d e34be6c 2fe4e9d 532a759 2fe4e9d 532a759 2fe4e9d 532a759 e34be6c 532a759 1543ec3 2fe4e9d e34be6c 2fe4e9d 532a759 1543ec3 2fe4e9d e34be6c 2fe4e9d 1543ec3 532a759 2fe4e9d e34be6c 1543ec3 e34be6c 1543ec3 2fe4e9d e34be6c 2fe4e9d e34be6c 2fe4e9d 532a759 e34be6c 532a759 e34be6c 2fe4e9d e34be6c 2fe4e9d 1543ec3 e34be6c 2fe4e9d 1543ec3 e34be6c | 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 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 | from langchain import LLMChain, PromptTemplate
from langchain.chains import (
ConversationalRetrievalChain,
ConversationChain,
RetrievalQA,
)
from langchain.chains.base import Chain
from langchain.memory import ConversationBufferMemory
from langchain.schema import BaseRetriever
from edu_assistant.learning_tasks.base import BaseTask
from edu_assistant.utils.langchain_utils import load_gpt4_llm, load_llm
TEMPLATE_CHAT = """The following is a friendly conversation between a human and an ai.
The ai is talkative and provides lots of specific details from its context.
If the ai does not know the answer to a question, it truthfully says it does not know.
The ai act following below instructions:
---
{instruction}
---
Current conversation:
{{chat_history}}
Human: {{input}}
AI:"""
TEMPLATE_CHAT_CONTEXT = """The following is a friendly conversation between a human and an ai.
The ai is talkative and provides lots of specific details from its context.
If the ai does not know the answer to a question, it truthfully says it does not know.
The ai act following below instructions:
---
{instruction}
---
Useful context for you to answer the question:
---
{{context}}
---
Current conversation:
{{chat_history}}
Human: {{question}}
AI:"""
TEMPLATE_ONCE = """The following is a friendly conversation between a human and an ai.
The ai is talkative and provides lots of specific details from its context.
If the ai does not know the answer to a question, it truthfully says it does not know.
The ai act following below instructions:
---
{instruction}
---
{{input}}
"""
TEMPLATE_ONCE_CONTEXT = """The following is a friendly conversation between a human and an ai.
The ai is talkative and provides lots of specific details from its context.
If the ai does not know the answer to a question, it truthfully says it does not know.
The ai act following below instructions:
---
{instruction}
---
Useful context for you to answer the question:
---
{{context}}
---
{{input}}
"""
DEFAULT_INSTRUCTION = """Act as a c++ professional to answer student aged 5-10 questions. Answer properly and politely.
Don't extend conversation multiple times. Only add one time saying."""
class QaTask(BaseTask):
_session_store: dict
_knowledge: BaseRetriever | None
_qa_once: Chain
HISTORY_KEY = "chat_history"
def __init__(self, instruction: str = DEFAULT_INSTRUCTION, knowledge: BaseRetriever = None, enable_gpt4=False):
"""Create a new QaTask service.
Args:
instruction (str, optional): Instruction for this task ai. Defaults to "".
knowledge (BaseRetriever, optional): Answer question with this knowledge retriever.
If not set, will not use knowledge to answer question.
Defaults to None.
session_store (dict, optional): chat history store. Defaults to None.
If not set, will use internal memory to store chat history. Which will be lost after restart and might
cost huge memory.
"""
self.enable_gpt4 = enable_gpt4
# TODO: load threshold key from implement. value from config
self.vectordbkwargs = {"score_threshold": 0.9} # Qdrant cosine. higher is better.
if knowledge:
self._chat_prompt = PromptTemplate.from_template(TEMPLATE_CHAT_CONTEXT.format(instruction=instruction))
self._once_prompt = PromptTemplate.from_template(TEMPLATE_ONCE_CONTEXT.format(instruction=instruction))
self._input_key = "question"
self._output_key = "answer"
else:
self._chat_prompt = PromptTemplate.from_template(TEMPLATE_CHAT.format(instruction=instruction))
self._once_prompt = PromptTemplate.from_template(TEMPLATE_ONCE.format(instruction=instruction))
self._input_key = "input"
self._output_key = "response"
self._session_store = {}
self._knowledge = knowledge
self._init_llm()
self._qa_once = self._build_once_chain()
def ask(
self,
question: str,
session: bool = True,
session_id: str = None,
session_mem: ConversationBufferMemory | None = None,
) -> dict:
"""ask a question with chat history.
Args:
question (str): question to llm.
session (bool, optional): whether use and store chat history. Defaults to False.
if session_id is not set, a new session will be created.
session_id (str, optional): specify a history qa session. Defaults to None.
session_mem (list | None, optional): specify session memory. Defaults to None.
if session_id is also set, memory will be replaced by passed one.
Returns:
dict: question answer and metadata.
contains answer.
contains session_id if session is True.
"""
if session:
args = {self._input_key: question}
if session_id and session_id in self._session_store:
chain = self._session_store[session_id]
else:
session_id = self._create_session_id()
chain = self._create_session_chain(session_id)
else:
args = {self._input_key: question, QaTask.HISTORY_KEY: ""}
chain = self._qa_once
if session_mem:
chain.memory = session_mem
# TODO: ConversationalRetrievalChain should support vectordbkwargs
# if self._knowledge:
# args["vectordbkwargs"] = self.vectordbkwargs
result = chain(args)
if session_id:
result["session_id"] = session_id
return result
def _init_llm(self):
self._main_llm = load_gpt4_llm() if self.enable_gpt4 else load_llm()
self._secondary_llm = load_llm()
def _build_once_chain(self):
if not self._knowledge:
return LLMChain(
llm=self._main_llm,
prompt=self._once_prompt,
)
else:
return RetrievalQA.from_llm(
llm=self._main_llm,
retriever=self._knowledge,
return_source_documents=True,
prompt=self._once_prompt,
)
def _build_chat_chain(self):
if not self._knowledge:
return ConversationChain(
llm=self._main_llm,
memory=ConversationBufferMemory(
memory_key=QaTask.HISTORY_KEY, output_key=self._output_key, return_messages=True
),
prompt=self._chat_prompt,
)
else:
return ConversationalRetrievalChain.from_llm(
llm=self._main_llm,
retriever=self._knowledge,
condense_question_llm=self._secondary_llm,
return_source_documents=True,
combine_docs_chain_kwargs={"prompt": self._chat_prompt},
memory=ConversationBufferMemory(
memory_key=QaTask.HISTORY_KEY, output_key=self._output_key, return_messages=True
),
)
def _create_session_chain(self, session_id) -> ConversationChain:
chain = self._build_chat_chain()
self._session_store[session_id] = chain
return chain
|