Spaces:
Runtime error
Runtime error
| import gradio as gr | |
| from fastapi.encoders import jsonable_encoder | |
| from langchain.callbacks import get_openai_callback | |
| from edu_assistant.learning_tasks.qa import DEFAULT_INSTRUCTION, QaTask | |
| from edu_assistant.utils.langchain_utils import load_vectorstore, shrink_docs | |
| class QaUI: | |
| def __init__( | |
| self, *, instruction: str = DEFAULT_INSTRUCTION, enable_gpt4: bool = False, knowledge_name: str = "example" | |
| ): | |
| self._init_task(instruction, knowledge_name, enable_gpt4) | |
| self._init_ui() | |
| def ui_render(self): | |
| self.ui.render() | |
| def ui_reload( | |
| self, | |
| *, | |
| instruction: str = DEFAULT_INSTRUCTION, | |
| knowledge_name: str = "example", | |
| enable_gpt4: bool = False, | |
| refresh: bool = True, | |
| ): | |
| self._init_task(instruction, knowledge_name, enable_gpt4) | |
| if refresh: | |
| self.ui_render() | |
| def get_instruction(self): | |
| return self.instruction | |
| def _init_task(self, instruction, knowledge_name, enable_gpt4): | |
| self.instruction = instruction | |
| self.knowledge = knowledge_name | |
| self.enable_gpt4 = enable_gpt4 | |
| self.task = QaTask( | |
| instruction=instruction, | |
| knowledge=load_vectorstore(knowledge_name).as_retriever(), | |
| enable_gpt4=enable_gpt4, | |
| ) | |
| def _init_ui(self): | |
| with gr.Blocks() as ui: | |
| with gr.Row(): | |
| with gr.Column(scale=6): | |
| with gr.Row(): | |
| chatbot = gr.Chatbot(height=500, label="聊天记录") | |
| with gr.Row(): | |
| msg = gr.Textbox(show_label=False) | |
| with gr.Column(scale=1): | |
| with gr.Row(): | |
| clear_button = gr.Button(value="清空") | |
| with gr.Row(): | |
| session_id = gr.Textbox(label="Session", interactive=False, value="") | |
| with gr.Row(): | |
| status = gr.JSON(value="""{"tokens":0}""", label="Status") | |
| with gr.Row(): | |
| docs = gr.JSON(value="""["docs"]""", label="Docs") | |
| clear_button.click(self._clear, [], [msg, chatbot, session_id, status, docs]) | |
| msg.submit(self._respond, [msg, chatbot, session_id], [msg, chatbot, session_id, status, docs]) | |
| self.ui = ui | |
| def _respond(self, message, chat_history, session_id): | |
| with get_openai_callback() as cb: | |
| if session_id: | |
| result = self.task.ask(message, session_id=session_id) | |
| else: | |
| result = self.task.ask(message) | |
| session_id = result["session_id"] | |
| docs = jsonable_encoder(shrink_docs(result.get("source_documents", []))) | |
| bot_message = result["answer"] | |
| chat_history.append((message, bot_message)) | |
| status = {"tokens": cb.total_tokens, "cost": f"${cb.total_cost:.4f}"} | |
| return "", chat_history, session_id, status, docs | |
| def _clear(self): | |
| return "", [], "", {"tokens": 0}, ["docs"] | |