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()