Spaces:
Running on Zero
Running on Zero
Download app.py from EmanHassan26/RAGMultiModalEx: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/spaces/EmanHassan26/RAGMultiModalEx/resolve/main/app.py
- Command line
-
hf download hf://spaces/EmanHassan26/RAGMultiModalEx/app.py
-
curl -L -o app.py https://huggingface.co/spaces/EmanHassan26/RAGMultiModalEx/resolve/main/app.py
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 | |
| # ===================================================================== | |
| 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}'!" | |
| 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 | |
| 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() |