Mistral과 Solara로 PDF와 채팅하기
Mistral과 Solara로 PDF와 채팅하기
solara를 사용해 채팅 기능과 PDF 읽기 기능을 가진 챗봇을 만드는 기초를 배우는 문서예요. 반응형(reactive) 변수에 메시지를 저장하고 스트리밍 응답을 처리하는 채팅 인터페이스부터, faiss와 Mistral 임베딩 기반 RAG로 PDF를 읽어 대화하는 예시까지 단계별로 보여줘요.
출처: 문서
본문
이 가이드에서는 solara로 채팅 기능과 PDF 읽기 기능을 가진 챗봇을 만드는 기초를 소개해요. (저자: Alonso Silva Allende (Nokia Bell Labs), GitHub 핸들: alonsosilvaallende)
채팅 인터페이스 만들기 (Chat Interface)
간단한 채팅 인터페이스를 구현해요. 이를 위해 solara와 mistralai 라이브러리를 임포트해야 해요.
pip install solara mistralai
이 데모는 solara==1.41.0과 mistralai==1.2.3을 사용해요.
import solara as sl
from mistralai.client import Mistral
Mistral API 키로 클라이언트를 만들어요.
mistral_api_key = "your_api_key"
client = Mistral(api_key = mistral_api_key)
모든 메시지가 저장될 반응형 변수를 초기화해요.
from typing import List
from typing_extensions import TypedDict
class MessageDict(TypedDict):
role: str
content: str
messages: sl.Reactive[List[MessageDict]] = sl.reactive([])
메시지 목록(지금은 비어 있지만 곧 차오를 거예요)이 주어지면 Mistral을 질의하고 응답을 받아요. 상호작용을 부드럽게 하려고 응답을 스트리밍으로 처리해요. 이를 위해 생성기(generator)를 정의해요.
def response_generator(messages):
response = client.chat.stream(model="open-mistral-7b", messages=messages, max_tokens=1024)
for chunk in response:
yield chunk.data.choices[0].delta.content
각 청크를 받는 대로 표시해서 응답을 스트리밍해요.
def add_chunk_to_ai_message(chunk: str):
messages.value = [
*messages.value[:-1],
{
"role": "assistant",
"content": messages.value[-1]["content"] + chunk,
},
]
메시지 목록을 화면에 표시해요.
@sl.component
def Page():
with sl.lab.ChatBox():
for item in messages.value:
with sl.lab.ChatMessage(
user=item["role"] == "user",
name="User" if item["role"] == "user" else "Assistant"
):
sl.Markdown(item["content"])
다음 단계는 사용자의 입력을 받아 메시지 목록에 저장하는 거예요. 이때 solara의 ChatInput을 사용해요.
def send(user_message):
messages.value = [*messages.value, {"role": "user", "content": user_message}]
sl.lab.ChatInput(send_callback=send)
스트리밍된 응답을 처리해야 하므로, 사용자 메시지 수의 변화로 활성화되는 태스크를 만들어요.
user_message_count = len([m for m in messages.value if m["role"] == "user"])
def response(messages):
messages.value = [*messages.value, {"role": "assistant", "content": ""}]
for chunk in response_generator(messages.value[:-1]):
add_chunk_to_ai_message(chunk)
def result():
if messages.value != []:
response(messages)
result = sl.lab.use_task(result, dependencies=[user_message_count])
이게 끝이에요! Mistral 모델과 채팅할 수 있는 인터페이스예요. 아래에 선택적인 스타일링을 추가했어요. 이 코드를 실행하려면 콘솔에 solara run chat.py를 입력하세요. 또는 PyCafe에서 직접 수정할 수도 있어요.
import solara as sl
from mistralai.client import Mistral
mistral_api_key = "your_api_key"
client = Mistral(api_key=mistral_api_key)
from typing import List
from typing_extensions import TypedDict
class MessageDict(TypedDict):
role: str
content: str
messages: sl.Reactive[List[MessageDict]] = sl.reactive([])
def response_generator(messages):
response = client.chat.stream(model="open-mistral-7b", messages=messages, max_tokens=1024)
for chunk in response:
yield chunk.data.choices[0].delta.content
def add_chunk_to_ai_message(chunk: str):
messages.value = [
*messages.value[:-1],
{
"role": "assistant",
"content": messages.value[-1]["content"] + chunk,
},
]
@sl.component
def Page():
user_message_count = len([m for m in messages.value if m["role"] == "user"])
def send(user_message):
messages.value = [*messages.value, {"role": "user", "content": user_message}]
def response(messages):
messages.value = [*messages.value, {"role": "assistant", "content": ""}]
for chunk in response_generator(messages.value[:-1]):
add_chunk_to_ai_message(chunk)
def result():
if messages.value != []:
response(messages)
result = sl.lab.use_task(result, dependencies=[user_message_count])
with sl.Column(align="center"):
with sl.lab.ChatBox(style={"position": "fixed", "overflow-y": "scroll","scrollbar-width": "none", "-ms-overflow-style": "none", "top": "0", "bottom": "10rem", "width": "60%"}):
for item in messages.value:
with sl.lab.ChatMessage(
user=item["role"] == "user",
name="User" if item["role"] == "user" else "Assistant"
):
sl.Markdown(item["content"])
sl.lab.ChatInput(send_callback=send, style={"position": "fixed", "bottom": "3rem", "width": "70%"})
PDF와 채팅 (Chatting with PDFs)
모델이 PDF를 읽게 하려면 콘텐츠를 변환하고 텍스트를 추출한 다음, Mistral의 임베딩 모델로 문서의 청크를 검색해 모델에 공급해야 해요. 기본적인 RAG(Retrieval-Augmented Generation)를 구현해야 해요!
이 작업에는 faiss와 PyPDF2가 필요해요.
pip install PyPDF2 faiss
CPU만 사용한다면 faiss-cpu를 설치하세요. 이 데모는 PyPDF2==3.0.1과 faiss-cpu==1.8.0을 사용해요.
import io
import solara as sl
from mistralai.client import Mistral
import numpy as np
import PyPDF2
import faiss
PDF 파일 업로드 가능성을 추가해야 해요. solara의 FileDropMultiple을 사용할 거예요. PDF들은 새 반응형 변수에 저장돼요.
from solara.components.file_drop import FileInfo
content, set_content = sl.use_state(cast(List[bytes], []))
def on_file(files: List[FileInfo]):
set_content([file["file_obj"].read() for file in files])
sl.FileDropMultiple(
label="Drag and drop your PDF file(s) here.",
on_file=on_file,
lazy=True,
)
PDF가 저장되긴 했지만 그대로는 그냥 많은 바이트일 뿐이에요. PDF와 채팅하려면 텍스트를 추출해야 해요.
txt = sl.use_reactive(cast(List[str], []))
def get_text():
txt_all = []
for _content in content:
bytes_io = io.BytesIO(_content)
reader = PyPDF2.PdfReader(bytes_io)
txt_aux = ""
for page in reader.pages:
txt_aux += page.extract_text()
txt_all.append(txt_aux)
return txt_all
if content:
sl.Info("File(s) has been uploaded. Showing the beginning of the file(s)...")
result: Task[List[str]] = use_task(get_text, dependencies=[content])
if result.finished:
txt.value = result.value
sl.ProgressLinear(result.pending)
for text in txt.value:
sl.Markdown(f"{text[:100]}")
이제 텍스트가 생겼으니 Mistral의 임베딩으로 관련 청크를 검색해요. 먼저 Mistral로 텍스트를 임베딩으로 변환하는 함수를 정의해요.
def get_text_embedding(input_text: str):
embeddings_batch_response = client.embeddings(
model = "mistral-embed",
input = input_text
)
return embeddings_batch_response.data[0].embedding
다음으로 검색 부분 전체를 처리하는 함수를 선언할 수 있어요. 이 단계는 벡터 스토어에 faiss를, 이전에 만든 get_text_embedding 함수를 사용해요. 여러 파일을 청크로 자르고, 임베딩을 만들고, 그중 최고 4개의 청크를 검색해 단일 문자열로 이어붙여요.
def rag_pdf(txt: List[str], question: str) -> str:
chunk_size = 1024
chunks = []
for _txt in txt:
chunks += [_txt[i:i + chunk_size] for i in range(0, len(_txt), chunk_size)]
text_embeddings = np.array([get_text_embedding(chunk) for chunk in chunks])
d = text_embeddings.shape[1]
index = faiss.IndexFlatL2(d)
index.add(text_embeddings)
question_embeddings = np.array([get_text_embedding(question)])
D, I = index.search(question_embeddings, k = 3)
retrieved_chunk = [chunks[i] for i in I.tolist()[0]]
text_retrieved = "\n\n".join(retrieved_chunk)
return text_retrieved
마지막으로 response_generator를 수정해 파일로 새로운 RAG를 구현해요! 이 함수는 PDF가 있으면 PyPDF2로 텍스트를 추출하고 rag_pdf로 관련 데이터를 검색하고, 그런 다음에만 모델에 요청을 보내요.
def response_generator(messages: list, txt: List[str]):
response = client.chat.stream(
model = "open-mistral-7b",
messages = messages[:-1] + [{"role":"user","content": rag_pdf(txt, messages[-1]["content"]) + "\n\n" + messages[-1]["content"]}],
max_tokens = 1024
)
for chunk in response:
yield chunk.data.choices[0].delta.content
이제 모든 게 끝났어요! solara run chat_with_pdfs.py로 새 인터페이스를 실행할 수 있어요. 전체 코드는 다음과 같아요.
import io
import solara as sl
from mistralai.client import Mistral
import numpy as np
import PyPDF2
import faiss
from solara.components.file_drop import FileInfo
from solara.lab import use_task, Task
from typing import List, cast
from typing_extensions import TypedDict
mistral_api_key = "your_api_key"
client = Mistral(api_key=mistral_api_key)
def get_text_embedding(input_text: str):
embeddings_batch_response = client.embeddings(
model = "mistral-embed",
input = input_text
)
return embeddings_batch_response.data[0].embedding
def rag_pdf(txt: List[str], question: str) -> str:
chunk_size = 1024
chunks = []
for _txt in txt:
chunks += [_txt[i:i + chunk_size] for i in range(0, len(_txt), chunk_size)]
text_embeddings = np.array([get_text_embedding(chunk) for chunk in chunks])
d = text_embeddings.shape[1]
index = faiss.IndexFlatL2(d)
index.add(text_embeddings)
question_embeddings = np.array([get_text_embedding(question)])
D, I = index.search(question_embeddings, k = 3)
retrieved_chunk = [chunks[i] for i in I.tolist()[0]]
text_retrieved = "\n\n".join(retrieved_chunk)
return text_retrieved
class MessageDict(TypedDict):
role: str
content: str
messages: sl.Reactive[List[MessageDict]] = sl.reactive([])
def response_generator(messages: list, txt: List[str]):
response = client.chat.stream(
model = "open-mistral-7b",
messages = messages[:-1] + [{"role":"user","content": rag_pdf(txt, messages[-1]["content"]) + "\n\n" + messages[-1]["content"]}],
max_tokens = 1024
)
for chunk in response:
yield chunk.choices[0].delta.content
def add_chunk_to_ai_message(chunk: str):
messages.value = [
*messages.value[:-1],
{
"role": "assistant",
"content": messages.value[-1]["content"] + chunk,
},
]
@sl.component
def Page():
txt = sl.use_reactive(cast(List[str], []))
with sl.Sidebar():
def on_file(files: List[FileInfo]):
get_text([file["data"] for file in files])
@sl.lab.task
def get_text(pdf_content):
txt_all = []
for _content in pdf_content:
bytes_io = io.BytesIO(_content)
reader = PyPDF2.PdfReader(bytes_io)
txt_aux = ""
for page in reader.pages:
txt_aux += page.extract_text()
txt_all.append(txt_aux)
return txt_all
sl.FileDropMultiple(
label="Drag and drop your PDF file(s) here.",
on_file=on_file,
lazy=True,
)
sl.ProgressLinear(get_text.pending)
if get_text.value:
sl.Info("File(s) has been uploaded. Showing the beginning of the file(s)...")
for text in get_text.value:
sl.Markdown(f"{text[:100]}")
user_message_count = len([m for m in messages.value if m["role"] == "user"])
def send(user_message):
messages.value = [*messages.value, {"role": "user", "content": user_message}]
def response(messages):
messages.value = [*messages.value, {"role": "assistant", "content": ""}]
for chunk in response_generator(messages.value[:-1], txt=txt.value):
add_chunk_to_ai_message(chunk)
def result():
if messages.value != []:
response(messages)
result = sl.lab.use_task(result, dependencies=[user_message_count])
with sl.Column(align="center"):
with sl.lab.ChatBox(style={"position": "fixed", "overflow-y": "scroll","scrollbar-width": "none", "-ms-overflow-style": "none", "top": "0", "bottom": "10rem", "width": "60%"}):
for item in messages.value:
with sl.lab.ChatMessage(
user=item["role"] == "user",
name="User" if item["role"] == "user" else "Assistant"
):
sl.Markdown(item["content"])
sl.lab.ChatInput(send_callback=send, style={"position": "fixed", "bottom": "3rem", "width": "60%"})