RAGMultiModalEx / app.py
EmanHassan26's picture
Update app.py
d7f57e2 verified
Raw History Blame Contribute Delete
10.8 kB
import spaces
import base64
from io import BytesIO
import os
import shutil
import torch
from PIL import Image
import gradio as gr
from byaldi import RAGMultiModalModel
import io
from transformers import PaliGemmaForConditionalGeneration, PaliGemmaProcessor, BitsAndBytesConfig
#from transformers import Qwen2VLForConditionalGeneration, AutoProcessor, BitsAndBytesConfig
# =====================================================================
# 1. SETUP & 4-BIT QUANTIZED MODEL INITIALIZATION
# =====================================================================
device = "cuda" if torch.cuda.is_available() else "cpu"
INDEX_NAME = "visual_assets_index"
BYALDI_DIR = ".byaldi"
TEMP_DIR = "./gradio_upload_temp"
os.makedirs(TEMP_DIR, exist_ok=True)
print(f"Loading local models on: {device}...")
# Configure 4-bit quantization settings via BitsAndBytes
quantization_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4"
)
# Load VLM wrapped with 4-bit quantization parameters
model_id = "google/paligemma-3b-pt-224"
processor = PaliGemmaProcessor.from_pretrained(model_id)
model = PaliGemmaForConditionalGeneration.from_pretrained(
model_id,
quantization_config=quantization_config,
torch_dtype=torch.bfloat16 if device == "cuda" else torch.float32,
device_map=device
)
RAG = RAGMultiModalModel.from_pretrained("vidore/colpali-v1.2", device=device)
class RAGState():
def __init__(self):
if os.path.exists(os.path.join(BYALDI_DIR, INDEX_NAME)):
self.RAG = RAGMultiModalModel.from_index(INDEX_NAME)
else:
self.RAG = RAG
#@spaces.GPU
def load_or_init_rag():
"""Helper function to cleanly load or initialize the RAG model."""
if os.path.exists(os.path.join(BYALDI_DIR, INDEX_NAME)):
print(f"Loading existing index: {INDEX_NAME}")
return RAGMultiModalModel.from_index(INDEX_NAME)
else:
print("No index found. Initializing vanilla model context wrapper.")
return RAGMultiModalModel.from_pretrained("vidore/colpali-v1.2",device=device)
# =====================================================================
# 2. CORE LOGIC PIPELINES
# =====================================================================
@spaces.GPU
def process_and_index_files(uploaded_files):
"""Saves uploaded imagery to a temp folder and processes the index."""
global RAG
#RAG_st.RAG = load_or_init_rag()
if not uploaded_files:
return "โš ๏ธ No files were uploaded."
# Clean out old staging assets
for f in os.listdir(TEMP_DIR):
os.remove(os.path.join(TEMP_DIR, f))
# Stage newly uploaded files
for file_path in uploaded_files:
filename = os.path.basename(file_path)
dest = os.path.join(TEMP_DIR, filename)
with open(file_path, "rb") as source, open(dest, "wb") as target:
target.write(source.read())
# Dynamically build index layout
RAG.index(
input_path=TEMP_DIR,
index_name=INDEX_NAME,
store_collection_with_index=True,
overwrite=True
)
#RAG_st.RAG = load_or_init_rag()
return f"โœ… Successfully indexed {len(uploaded_files)} visual assets into '{INDEX_NAME}'!"
@spaces.GPU
def reset_rag_index():
"""Wipes out temporary assets, physical vector files, and refreshes the model context."""
#global RAG
# 1. Clean out the local image uploader stash
if os.path.exists(TEMP_DIR):
shutil.rmtree(TEMP_DIR)
os.makedirs(TEMP_DIR, exist_ok=True)
# 2. Hard purge Byaldi's saved index files
target_index_path = os.path.join(BYALDI_DIR, INDEX_NAME)
if os.path.exists(target_index_path):
shutil.rmtree(target_index_path)
# 3. Reload a clean checkpoint wrapper state
#RAG_st.RAG = RAGMultiModalModel.from_pretrained("vidore/colqwen2-v1.0")
return "๐Ÿ—‘๏ธ Index fully cleared! Staged temporary assets and vector checkpoints have been deleted from disk.",RAG_st
@spaces.GPU
def multimodal_rag_query(user_query, num_images_to_retrieve):
"""Retrieves document patches and feeds context to the regional VLM."""
global RAG
RAG = RAGMultiModalModel.from_index(INDEX_NAME)
print(f"RAG '{RAG}'")
k = int(num_images_to_retrieve)
# Guard clause ensuring an active physical index directory exists
if not os.path.exists(os.path.join(BYALDI_DIR, INDEX_NAME)):
return "โš ๏ธ Search cancelled. The index is empty or has been reset. Please upload and index documents first.", []
try:
results = RAG.search(user_query, k=k)
except Exception as e:
return f"โš ๏ธ Error searching vector field: {str(e)}", []
print(f"results '{results}'")
retrieved_images_with_captions = []
retrieved_images_raw = []
for idx, match in enumerate(results):
base64_str = match.get("base64") if isinstance(match, dict) else getattr(match, "base64", None)
score = match.get("score") if isinstance(match, dict) else getattr(match, "score", 0.0)
page_num = match.get("page_num") if isinstance(match, dict) else getattr(match, "page_num", "N/A")
if base64_str:
img_data = base64.b64decode(base64_str)
pil_img = Image.open(BytesIO(img_data)).convert("RGB")
retrieved_images_raw.append(pil_img)
caption = f"Match #{idx+1} | Score: {score:.4f} | Page/Doc ID: {page_num}"
retrieved_images_with_captions.append((pil_img, caption))
if not retrieved_images_raw:
return "No corresponding context documents found.", []
messages = [
{
"role": "user",
"content": [
*[{"type": "image", "image": img} for img in retrieved_images_raw],
{
"type": "text",
"text": f"Analyze the provided context document screenshots and thoroughly answer this request: {user_query}"
}
]
}
]
# 1. Prepend an token for each retrieved document image
image_tokens = "" * len(retrieved_images_raw)
# 2. Define the main text prompt
text_prompt = f"Analyze the provided context document screenshots and thoroughly answer this request: {user_query}"
# 3. Combine into a single prompt string
prompt_text = f"{image_tokens}{text_prompt}"
# 4. Pass directly to the processor without using apply_chat_template
inputs = processor(images=[retrieved_images_raw] if len(retrieved_images_raw) > 0 else None,
text=prompt_text, return_tensors="pt").to(device)
#prompt_text = processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
#inputs = processor(images=retrieved_images_raw, texts=[prompt_text], padding=True, return_tensors="pt").to(device)
with torch.no_grad():
generated_ids = model.generate(**inputs, max_new_tokens=512)
generated_ids_trimmed = [
out_ids[len(in_ids):] for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
]
output_text = processor.batch_decode(
generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
)
return output_text, retrieved_images_with_captions
# =====================================================================
# 3. INTERACTIVE GRADIO INTERFACE BLOCK
# =====================================================================
with gr.Blocks() as demo:
gr.Markdown("# ๐Ÿ‘๏ธ 4-Bit Vision-RAG Console")
gr.Markdown("An optimized, low-VRAM interface utilizing **NF4 Quantization** for the VLM and local visual extraction via **Byaldi**.")
#RAG_st = gr.State(value=lambda:RAGState())
#RAG = load_or_init_rag()
with gr.Tabs():
# TAB 1: INGESTION PIPELINE MANAGEMENT
with gr.TabItem("๐Ÿ“ Upload & Index Assets"):
gr.Markdown("### Ingest Document Images into Index")
gr.Markdown("Drop your charts, system layouts, tables, multi-page PDFs, or scans below. Running the builder updates the embedded multi-vector index layout.")
file_uploader = gr.File(
file_count="multiple",
file_types=[".png", ".jpg", ".jpeg", ".pdf"],
label="Drop layout images or PDF reports here"
)
with gr.Row():
index_btn = gr.Button("Build / Reindex Assets", variant="secondary")
reset_btn = gr.Button("Reset Index & Clear Storage", variant="stop")
status_output = gr.Textbox(label="Indexing & Database Execution Status", interactive=False)
# Action event listeners
index_btn.click(
fn=process_and_index_files,
inputs=[file_uploader],
outputs=[status_output]
)
reset_btn.click(
fn=reset_rag_index,
inputs=[],
outputs=[status_output]
)
# TAB 2: RETRIEVAL & INFERENCE EXPLORATION
with gr.TabItem("๐Ÿ” Query Explorer"):
with gr.Row():
with gr.Column(scale=2):
query_input = gr.Textbox(
label="User Query",
placeholder="e.g., Extract the quarterly compound annual growth rates shown in the table.",
lines=3
)
k_slider = gr.Slider(
minimum=1, maximum=4, value=2, step=1,
label="Top-K Retrieved Context Pages"
)
submit_btn = gr.Button("Query Local RAG System", variant="primary")
with gr.Column(scale=3):
answer_output = gr.Textbox(
label="Synthesized Response ",
interactive=False,
lines=8
)
gr.Markdown("### ๐Ÿ“Š Retrieved Document Evidence Stack")
gallery_output = gr.Gallery(
label="Retrieved Visual Anchors (Hover or click to isolate similarity mappings)",
show_label=True,
columns=2,
height="auto",
object_fit="contain"
)
submit_btn.click(
fn=multimodal_rag_query,
inputs=[query_input, k_slider],
outputs=[answer_output, gallery_output]
)
if __name__ == "__main__":
demo.launch()