RAG Chatbot

What I Developed

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.

High Level Explanation

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.

The Design - Embedding Model

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
                    

The Design - The Backend

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)
                        

    Conclusion

    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.

Full Program