Chatbot de Meta Llama 3 con RAG
Este notebook muestra un ejemplo completo de cómo construir un chatbot de Meta Llama 3 alojado en tu navegador que puede responder preguntas basadas en tus propios datos. Cubriremos:
- El proceso de despliegue de Meta Llama 3 8B con el framework Text-generation-inference como servidor API
- Un ejemplo de chatbot construido con Gradio y conectado al servidor
- Cómo añadir la capacidad RAG con conocimientos específicos de Meta Llama 3 basados en nuestra guía de introducción
Arquitectura RAG
Los LLM tienen capacidades sin precedentes en NLU (Comprensión del Lenguaje Natural) y NLG (Generación del Lenguaje Natural), pero tienen una fecha de corte de conocimiento y solo están entrenados con datos disponibles públicamente antes de esa fecha.
RAG, inventado por Meta en 2020, es uno de los métodos más populares para aumentar los LLM. RAG permite a las empresas mantener datos sensibles en sus instalaciones y obtener respuestas más relevantes de modelos genéricos sin necesidad de ajustar los modelos para roles específicos.
RAG es un método que:
- Recupera datos de fuera de un modelo fundacional
- Aumenta tus preguntas o prompts a los LLM añadiendo los datos relevantes recuperados como contexto
- Permite a los LLM responder preguntas sobre tus propios datos, o datos no disponibles públicamente cuando los LLM fueron entrenados
- Reduce en gran medida la alucinación en la generación de respuestas del modelo
El siguiente diagrama muestra los componentes y el proceso general de RAG:

Cómo desarrollar un chatbot de Meta Llama 3 con RAG
La forma más fácil de desarrollar chatbots de Meta Llama 3 con RAG es usar frameworks como LangChain y LlamaIndex, dos frameworks de código abierto líderes para construir aplicaciones LLM. Ambos ofrecen APIs convenientes para implementar RAG con Meta Llama 3, incluyendo:
- Cargar y dividir documentos
- Incrustar y almacenar divisiones de documentos
- Recuperar el contexto relevante basado en la consulta del usuario
- Llamar a Meta Llama 3 con la consulta y el contexto para generar la respuesta
LangChain es un framework más general y flexible para desarrollar aplicaciones LLM con capacidades RAG, mientras que LlamaIndex, como framework de datos, se centra en conectar fuentes de datos personalizadas a los LLM. La integración de ambos puede proporcionar la solución más eficaz y de mejor rendimiento para construir aplicaciones RAG del mundo real.
En nuestro ejemplo, por simplicidad, usaremos solo LangChain con datos PDF almacenados localmente.
Instalar dependencias
Para esta demostración, usaremos Gradio para la interfaz de usuario del chatbot y el framework Text-generation-inference para el servicio del modelo.
Para el almacenamiento de vectores y la búsqueda de similitud, usaremos FAISS.
En este ejemplo, ejecutaremos todo en una instancia de AWS EC2 (es decir, g5.2xlarge). g5.2xlarge cuenta con una GPU A10G. Recomendamos ejecutar este notebook con al menos una GPU equivalente a A10G con al menos 16 GB de memoria de video.
Existen ciertas técnicas para reducir el tamaño del modelo Meta Llama 3 8B, de modo que pueda caber en GPUs más pequeñas. Pero esto está fuera del alcance de este documento.
Primero, instalemos todas las dependencias con PIP. También te recomendamos que inicies un entorno Conda dedicado para una mejor gestión de paquetes.
!pip install -r requirements.txt
Procesamiento de datos
Primero, ejecuta todas las importaciones y define la ruta de los datos y el almacenamiento de vectores después del procesamiento.
Para los datos, usaremos un PDF sin procesar extraído de la guía de inicio de Meta Llama 3 en el sitio web de Meta AI.
from langchain.embeddings import HuggingFaceEmbeddings
from langchain.vectorstores import FAISS
from langchain.document_loaders import PyPDFDirectoryLoader
from langchain.text_splitter import RecursiveCharacterTextSplitter
DATA_PATH = 'data' #Your root data folder path
DB_FAISS_PATH = 'vectorstore/db_faiss'
Luego usamos el PyPDFDirectoryLoader para cargar todo el directorio. También puedes usar PyPDFLoader para cargar un solo archivo.
loader = PyPDFDirectoryLoader(DATA_PATH)
documents = loader.load()
Verifica la longitud y el contenido del documento para asegurarte de que hemos cargado el documento correcto con 37 páginas.
print(len(documents), documents[0].page_content[0:100])
37 11/8/23, 2:00 PM Getting started with Llama 2 - AI at Meta
https://ai.meta.com/llama/get-started/ 1/
Divide los documentos cargados en fragmentos más pequeños.RecursiveCharacterTextSplitter es un divisor común que divide piezas largas de texto en fragmentos más pequeños y semánticamente significativos.
Otros divisores incluyen:
- SpacyTextSplitter
- NLTKTextSplitter
- SentenceTransformersTokenTextSplitter
- CharacterTextSplitter
text_splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=10)
splits = text_splitter.split_documents(documents)
print(len(splits), splits[0])
103 page_content='11/8/23, 2:00 PM Getting started with Llama 2 - AI at Meta\nhttps://ai.meta.com/llama/get-started/ 1/37\nLlama 2 Get Started FAQ Download the Model\nQuick setup and how-to guide\nGetting started\nwith Llama\nWelcome to the getting started guide for Llama.\nThis guide provides information and resources to help you set up Llama including how to access the model,\nhosting, how-to and integration guides. Additionally , you will find supplemental materials to further assist you while\nbuilding with Llama.' metadata={'source': 'data/Llama Getting Started Guide.pdf', 'page': 0}
Ten en cuenta que hemos establecido chunk_size en 500 y chunk_overlap en 10. Al dividir, estos dos parámetros pueden afectar directamente la calidad de las respuestas del LLM.
Aquí tienes una buena guía sobre cómo debes configurar cuidadosamente estos dos parámetros.
A continuación, tendremos que elegir un modelo de embedding para nuestros documentos divididos.
Los embeddings son representaciones numéricas de texto. El modelo de embedding predeterminado en HuggingFace Embeddings es sentence-transformers/all-mpnet-base-v2 con 768 dimensiones. A continuación, usamos un modelo más pequeño all-MiniLM-L6-v2 con 384 dimensiones para que la indexación se ejecute más rápido.
embeddings = HuggingFaceEmbeddings(model_name='sentence-transformers/all-MiniLM-L6-v2',
model_kwargs={'device': 'cuda'})
modules.json: 0%| | 0.00/349 [00:00<?, ?B/s]
config_sentence_transformers.json: 0%| | 0.00/116 [00:00<?, ?B/s]
README.md: 0%| | 0.00/10.7k [00:00<?, ?B/s]
sentence_bert_config.json: 0%| | 0.00/53.0 [00:00<?, ?B/s]
config.json: 0%| | 0.00/612 [00:00<?, ?B/s]
model.safetensors: 0%| | 0.00/90.9M [00:00<?, ?B/s]
tokenizer_config.json: 0%| | 0.00/350 [00:00<?, ?B/s]
vocab.txt: 0%| | 0.00/232k [00:00<?, ?B/s]
tokenizer.json: 0%| | 0.00/466k [00:00<?, ?B/s]
special_tokens_map.json: 0%| | 0.00/112 [00:00<?, ?B/s]
1_Pooling/config.json: 0%| | 0.00/190 [00:00<?, ?B/s]
Por último, con las divisiones y la elección del modelo de embedding listos, queremos indexarlos y almacenar todos los fragmentos divididos como embeddings en el almacenamiento de vectores.
Los almacenes de vectores son bases de datos que almacenan embeddings. Hay al menos 60 almacenes de vectores compatibles con LangChain, y dos de los más populares de código abierto son:
- Chroma: ligero y en memoria, por lo que es fácil empezar a usarlo para el desarrollo local.
- FAISS (Facebook AI Similarity Search): un almacén de vectores que admite la búsqueda en vectores que pueden no caber en la RAM y es apropiado para el uso en producción.
Dado que estamos ejecutando en una instancia EC2 con abundantes recursos de CPU y RAM, usaremos FAISS en este ejemplo. Ten en cuenta que FAISS también puede ejecutarse en GPUs, donde se implementan algunos de los algoritmos más útiles. En ese caso, instala el paquete faiss-gpu con PIP en su lugar.
db = FAISS.from_documents(splits, embeddings)
db.save_local(DB_FAISS_PATH)
Una vez que hayas guardado la base de datos en la ruta local. Puedes encontrarlos como index.faiss y index.pkl. En el ejemplo del chatbot, puedes cargar esta base de datos desde el local y conectarla a nuestro proceso de recuperación.
Servicio de modelos
En este ejemplo, desplegaremos un modelo de chat Meta Llama 3 8B de HuggingFace con el framework Text-generation-inference en las instalaciones.
Esto nos permitirá conectar directamente el servidor API con nuestro chatbot.
Existen soluciones alternativas para desplegar modelos Meta Llama 3 en las instalaciones como tu servidor API local.
Puedes encontrar nuestra guía completa aquí.
En una terminal separada, ejecuta los comandos a continuación para iniciar un servidor API con TGI. Esto descargará los artefactos del modelo y los almacenará localmente, mientras se inicia en el puerto deseado en tu localhost. En nuestro caso, este es el puerto 8080.
model=meta-llama/Meta-Llama-3.1-8B-Instruct
volume=$PWD/data # share a volume with the Docker container to avoid downloading weights every run
token=#your-huggingface-token
docker run --gpus all --shm-size 1g -e HUGGING_FACE_HUB_TOKEN=$token -p 8080:80 -v $volume:/data ghcr.io/huggingface/text-generation-inference:2.0 --model-id $model
Una vez que el servidor API esté en funcionamiento, podemos ejecutar un simple comando curl para validar que nuestro modelo funciona como se espera.
!curl localhost:8080/generate -X POST -H 'Content-Type: application/json' -d '{"inputs": "What is good about Beijing?", "parameters": { "max_new_tokens":64}}' #Replace the localhost with the IP visible to the machine running the notebook
Construyendo la interfaz de usuario del chatbot
Ahora estamos listos para construir la interfaz de usuario del chatbot para conectar los datos RAG y el servidor API. En nuestro ejemplo, usaremos Gradio para construir la interfaz de usuario del chatbot.
Gradio es una biblioteca de Python de código abierto que se utiliza para construir demostraciones y aplicaciones web de aprendizaje automático y ciencia de datos. Ha sido ampliamente utilizada por la comunidad y HuggingFace también usó Gradio para construir sus chatbots. Otras alternativas son:
De nuevo, comenzamos añadiendo todas las importaciones, rutas, constantes y configurando LangChain en modo depuración, para que muestre acciones claras dentro del proceso de la cadena.
import langchain
from queue import Queue
from typing import Any
from langchain.llms.huggingface_text_gen_inference import HuggingFaceTextGenInference
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
from langchain.schema import LLMResult
from langchain.embeddings import HuggingFaceEmbeddings
from langchain.vectorstores import FAISS
from langchain.chains import RetrievalQA
from langchain.prompts.prompt import PromptTemplate
from anyio.from_thread import start_blocking_portal #For model callback streaming
langchain.debug=True
#vector db path
DB_FAISS_PATH = 'vectorstore/db_faiss'
#Llama2 TGI models host port
LLAMA3_8B_HOSTPORT = "http://localhost:8080/" #Replace the localhost with the IP visible to the machine running the notebook
LLAMA3_70B_HOSTPORT = "http://localhost:8081/" # You can host multiple models if your infrastructure has capacity
model_dict = {
"8b-instruct" : LLAMA3_8B_HOSTPORT,
"70b-instruct" : LLAMA3_70B_HOSTPORT,
}
system_message = {"role": "system", "content": "You are a helpful assistant."}
Luego cargamos el almacén de vectores FAISS
embeddings = HuggingFaceEmbeddings(model_name="sentence-transformers/all-MiniLM-L6-v2",
model_kwargs={'device': 'cuda'})
db = FAISS.load_local(DB_FAISS_PATH, embeddings, allow_dangerous_deserialization=True)
Ahora creamos una instancia de TGI llm y la conectamos al puerto de servicio de la API en localhost.
llm = HuggingFaceTextGenInference(
inference_server_url=LLAMA3_8B_HOSTPORT,
max_new_tokens=512,
top_k=10,
top_p=0.9,
typical_p=0.95,
temperature=0.6,
repetition_penalty=1,
do_sample=True,
streaming=True
)
/opt/conda/envs/pytorch/lib/python3.10/site-packages/langchain_core/_api/deprecation.py:119: LangChainDeprecationWarning: The class `HuggingFaceTextGenInference` was deprecated in LangChain 0.0.21 and will be removed in 0.2.0. Use HuggingFaceEndpoint instead.
warn_deprecated(
/opt/conda/envs/pytorch/lib/python3.10/site-packages/pydantic/_internal/_fields.py:127: UserWarning: Field "model_id" has conflict with protected namespace "model_".
You may be able to resolve this warning by setting `model_config['protected_namespaces'] = ()`.
warnings.warn(
A continuación, definimos el recuperador y la plantilla para nuestra cadena RetrivalQA. Para cada llamada de RetrievalQA, LangChain realiza una búsqueda de similitud semántica de la consulta en la base de datos vectorial, luego pasa los resultados de la búsqueda como contexto a Llama para responder la consulta sobre los datos almacenados en la base de datos vectorial.
Mientras que para la plantilla, esta define el formato de la pregunta junto con el contexto que enviaremos a Llama para la generación. En general, Meta Llama 3 tiene un formato de prompt especial para manejar tokens especiales. En algunos casos, el framework de servicio ya podría haberlo manejado. De lo contrario, deberás escribir una plantilla personalizada para manejarlo correctamente.
system_prompt = ""
template = """
Use the following pieces of context to answer the question. If no context provided, answer like a AI assistant.
{context}
Question: {question}
"""
retriever = db.as_retriever(
search_kwargs={"k": 6}
)
Por último, podemos definir la cadena de recuperación para QA
qa_chain = RetrievalQA.from_chain_type(
llm=llm,
retriever=retriever,
chain_type_kwargs={
"prompt": PromptTemplate(
template=template,
input_variables=["context", "question"],
),
}
)
Ahora deberíamos tener una cadena de QA funcional. Probémosla antes de conectarla con los bloques de la interfaz de usuario.
result = qa_chain({"query": "Why choose Llama?"})
print(result)
Después de confirmar la validez, podemos empezar a construir la interfaz de usuario. Antes de definir los bloques de Gradio, definamos primero los flujos de devolución de llamada que usaremos más adelante para la función de streaming.
Este manejador de devolución de llamada pondrá las respuestas de streaming del LLM en una cola para que la interfaz de usuario de Gradio las renderice sobre la marcha.
job_done = object()
class MyStream(StreamingStdOutCallbackHandler):
def __init__(self, q) -> None:
self.q = q
def on_llm_new_token(self, token: str, **kwargs: Any) -> None:
self.q.put(token)
def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None:
self.q.put(job_done)
Ahora podemos definir los bloques de la interfaz de usuario de Gradio.
Dado que tendremos que definir la interfaz de usuario y los manejadores en el mismo lugar, esta será una gran parte del código. Añadiremos comentarios en el código para su explicación.
import gradio as gr
with gr.Blocks() as demo:
#Configure UI layout
chatbot = gr.Chatbot(height = 600)
with gr.Row():
with gr.Column(scale=1):
with gr.Row():
#model selection
model_selector = gr.Dropdown(
list(model_dict.keys()),
value="7b-chat",
label="Model",
info="Select the model",
interactive = True,
scale=1
)
max_new_tokens_selector = gr.Number(
value=512,
precision=0,
label="Max new tokens",
info="Adjust max_new_tokens",
interactive = True,
minimum=1,
maximum=1024,
scale=1
)
with gr.Row():
#hyperparameter selection
temperature_selector = gr.Slider(
value=0.6,
label="Temperature",
info="Range 0-2. Controls the creativity of the generated text.",
interactive = True,
minimum=0.01,
maximum=2,
step=0.01,
scale=1
)
top_p_selector = gr.Slider(
value=0.9,
label="Top_p",
info="Range 0-1. Nucleus sampling.",
interactive = True,
minimum=0.01,
maximum=0.99,
step=0.01,
scale=1
)
with gr.Column(scale=2):
#user input prompt text field
user_prompt_message = gr.Textbox(placeholder="Please add user prompt here", label="User prompt")
with gr.Row():
clear = gr.Button("Clear Conversation", scale=2)
submitBtn = gr.Button("Submit", scale=8)
state = gr.State([])
#handle user message
def user(user_prompt_message, history):
if user_prompt_message != "":
return history + [[user_prompt_message, None]]
else:
return history + [["Invalid prompts - user prompt cannot be empty", None]]
#chatbot logic for configuration, sending the prompts, rendering the streamed back generations etc
def bot(model_selector, temperature_selector, top_p_selector, max_new_tokens_selector, user_prompt_message, history, messages_history):
dialog = []
bot_message = ""
history[-1][1] = ""
dialog = [
{"role": "user", "content": user_prompt_message},
]
messages_history += dialog
#Queue for streamed character rendering
q = Queue()
#Update new llama hyperparameters
llm.inference_server_url = model_selector
llm.temperature = temperature_selector
llm.top_p = top_p_selector
llm.max_new_tokens = max_new_tokens_selector
#Async task for streamed chain results wired to callbacks we previously defined, so we don't block the UI
async def task(prompt):
ret = await qa_chain.run(prompt, callbacks=[MyStream(q)])
return ret
with start_blocking_portal() as portal:
portal.start_task_soon(task, user_prompt_message)
while True:
next_token = q.get(True)
if next_token is job_done:
messages_history += [{"role": "assistant", "content": bot_message}]
return history, messages_history
bot_message += next_token
history[-1][1] += next_token
yield history, messages_history
#init the chat history with default system message
def init_history(messages_history):
messages_history = []
messages_history += [system_message]
return messages_history
#clean up the user input text field
def input_cleanup():
return ""
#when the user clicks Enter and the user message is submitted
user_prompt_message.submit(
user,
[user_prompt_message, chatbot],
[chatbot],
queue=False
).then(
bot,
[model_selector, temperature_selector, top_p_selector, max_new_tokens_selector, user_prompt_message, chatbot, state],
[chatbot, state]
).then(input_cleanup,
[],
[user_prompt_message],
queue=False
)
#when the user clicks the submit button
submitBtn.click(
user,
[user_prompt_message, chatbot],
[chatbot],
queue=False
).then(
bot,
[model_selector, temperature_selector, top_p_selector, max_new_tokens_selector, user_prompt_message, chatbot, state],
[chatbot, state]
).then(
input_cleanup,
[],
[user_prompt_message],
queue=False
)
#when the user clicks the clear button
clear.click(lambda: None, None, chatbot, queue=False).success(init_history, [state], [state])
Por último, podemos lanzar esta demostración en nuestro localhost con el comando de abajo.
demo.queue().launch(server_name="0.0.0.0")
Running on local URL: http://0.0.0.0:7860
To create a public link, set `share=True` in `launch()`.
<IPython.core.display.HTML object>
Gradio establecerá el puerto de lanzamiento por defecto en 7860. Puedes seleccionar el puerto en el que debe lanzarse según sea necesario.
Una vez lanzado, en el notebook o en un navegador con la URL http://0.0.0.0:7860, deberías ver la interfaz de usuario.
Cosas que puedes probar en la demostración del chatbot:
- Hacer preguntas específicas relacionadas con la Guía de inicio de Meta Llama 3
- Ajustar parámetros como el número máximo de tokens nuevos generados
- Cambiar a otro modelo de Llama con otro contenedor lanzado en una terminal separada
Una vez que hayas terminado de probar, asegúrate de cerrar la demostración ejecutando el comando de abajo para liberar el puerto.
demo.close()