I developed a lightweight chatbot that leverages Retrieval-Augmented Generated to provide contextually relevant responses. I chose to embed a short biography of me which the local LLM uses to answer questions about me. My idea is that anyone could use it to answer some questions they may have about me.
RAG stands for Retrieval-Augmented Generation. RAG chatbots work by combining two models: a retrieval model and a generation model. The retrieval model is used to search for relevant information from a given source, in this case a small database. The generation model then uses this information to generate a response to the user's query. The key difference between RAG chatbots and traditional chatbots is that RAG chatbots use a retrieval model to supplement the generation model's capabilities. This allows RAG chatbots to generate more accurate and relevant responses to user queries.
I chose to keep the file type strictly PDFs. This allowed for the embedding technique to be simple to implement. The text from the PDF file is stored in a vector database (ChromaDB) in chunks.
import os
from datetime import datetime
from werkzeug.utils import secure_filename
from langchain_community.document_loaders import UnstructuredPDFLoader
from langchain_text_splitters import RecursiveCharacterTextSplitter
from get_vector_db import get_vector_db
TEMP_FOLDER = os.getenv('TEMP_FOLDER', './_temp')
def allowed_file(filename):
return filename.lower().endswith('.pdf')
def save_file(file):
filename = f"{datetime.now().timestamp()}_{secure_filename(file.filename)}"
file_path = os.path.join(TEMP_FOLDER, filename)
file.save(file_path)
return file_path
def load_and_split_data(file_path):
loader = UnstructuredPDFLoader(file_path=file_path)
data = loader.load()
text_splitter = RecursiveCharacterTextSplitter(chunk_size=7500, chunk_overlap=100)
return text_splitter.split_documents(data)
def embed(file):
if file and allowed_file(file.filename):
file_path = save_file(file)
chunks = load_and_split_data(file_path)
db = get_vector_db()
db.add_documents(chunks)
db.persist()
os.remove(file_path)
return True
return False
I wrote a simple Flask app to serve the backend as an endpoint. I included a few necessary routes and decided to allow access from any origin for simplicity sake.
import os
from dotenv import load_dotenv
from flask import Flask, request, jsonify, make_response
from embed import embed
from query import query
from delete import delete_persisted_db
from get_vector_db import get_vector_db
load_dotenv()
TEMP_FOLDER = os.getenv('TEMP_FOLDER', './_temp')
os.makedirs(TEMP_FOLDER, exist_ok=True)
app = Flask(__name__)
@app.after_request
def after_request(response):
header = response.headers
header['Access-Control-Allow-Origin'] = '*'
return response
@app.route('/delete', methods=['DELETE'])
def route_delete():
try:
delete_persisted_db()
return {"message": "Database has been deleted."}
except FileNotFoundError as e:
raise HTTPException(status_code=404, detail=str(e))
@app.route('/embed', methods=['POST'])
def route_embed():
if 'file' not in request.files:
return jsonify({"error": "No file part"}), 400
file = request.files['file']
if file.filename == '':
return jsonify({"error": "No selected file"}), 400
embedded = embed(file)
return jsonify({"message": "File embedded successfully"}) if embedded else jsonify({"error": "Embedding failed"}), 400
@app.route('/query', methods=['POST', 'OPTIONS'])
def route_query():
if request.method == "OPTIONS": # CORS preflight
return _build_cors_preflight_response()
elif request.method == "POST": # The actual request following the preflight
data = request.get_json()
response = query(data.get('query'))
return _corsify_actual_response(jsonify(response))
else:
raise RuntimeError("Weird - don't know how to handle method {}".format(request.method))
def _build_cors_preflight_response():
response = make_response()
response.headers.add("Access-Control-Allow-Origin", "*")
response.headers.add('Access-Control-Allow-Headers', "*")
response.headers.add('Access-Control-Allow-Methods', "*")
return response
def _corsify_actual_response(response):
response.headers.add("Access-Control-Allow-Origin", "*")
return response
if __name__ == '__main__':
app.run(host="127.0.0.1", port=8080, debug=True)
No matter how well trained a model may be, it can always benefit from Retrieval-Augmented Generation. I plan to spend more time exploring this technique.
Copyright © Lucas Thormann