Team Ai
Apppublic

MD2204/multi_modality

sourceHugging Faceupdated 8mo agoView on Hugging Face
1likes
multimodal.py151 linesDownload Raw Back to src
1import os 
2import time
3import io
4import json
5import pymupdf4llm 
6from google import genai
7from google.genai import types
8from langchain_core.documents import Document
9from src.config import GOOGLE_API_KEY, GOOGLE_MODEL, GOOGLE_TEMPERATURE, IMAGES_DIR, GEMINI_RAG_PROMPT
10from PIL import Image
11
12# Initialize the new Gemini Client
13client = genai.Client(api_key=GOOGLE_API_KEY)
14
15CACHE_FILE = IMAGES_DIR / "transcription_cache.json"
16
17def load_cache():
18    if os.path.exists(CACHE_FILE):
19        try:
20            with open(CACHE_FILE, "r", encoding="utf-8") as f:
21                return json.load(f)
22        except Exception as e:
23            print(f"⚠️ Error loading cache: {e}")
24            return {}
25    return {}
26
27def save_cache(cache):
28    try:
29        with open(CACHE_FILE, "w", encoding="utf-8") as f:
30            json.dump(cache, f, indent=4, ensure_ascii=False)
31    except Exception as e:
32        print(f"⚠️ Error saving cache: {e}")
33
34def process_pdf_multimodal(pdf_path: str) -> tuple[list[Document], dict]:
35    print(f"📄 Processing multimodal PDF: {pdf_path}")
36    
37    stats = {
38        "total_images": 0,
39        "scanned_images": 0,
40        "skipped_images": 0
41    }
42    
43    # Load existing transcriptions
44    cache = load_cache()
45    
46    # 1. Extract markdown structure and save images to permanent directory
47    pages = pymupdf4llm.to_markdown(
48        pdf_path,
49        page_chunks=True,
50        write_images=True,
51        image_path=str(IMAGES_DIR),
52        image_size_limit=0.20, 
53    )
54    
55    documents = []
56
57    # 2. Iterate through pages and process images
58    pdf_basename = os.path.basename(pdf_path)
59    for page in pages:
60        text_content = page["text"]
61        page_num = page["metadata"]["page"]
62        page_index = page_num - 1
63        
64        # logic to find images specifically for THIS document and THIS page index
65        # Pattern: "filename.pdf-{page_index}-"
66        target_prefix = f"{pdf_basename}-{page_index}-"
67        page_images = [i for i in os.listdir(IMAGES_DIR) if i.startswith(target_prefix)]
68        
69        stats["total_images"] += len(page_images)
70        visual_content = ""
71        for img_name in page_images:
72            # Check cache first
73            if img_name in cache:
74                print(f"  ✨ Using cached transcription for: {img_name}")
75                visual_content += f"\n\n[Visual Data Transcription]:\n{cache[img_name]}"
76                stats["scanned_images"] += 1
77                continue
78
79            img_path = os.path.join(IMAGES_DIR, img_name)
80            
81            # --- Robust Gemini Call with Retry ---
82            max_retries = 3
83            retry_count = 0
84            description = ""
85            
86            while retry_count < max_retries:
87                try:
88                    # Filter by size/dimensions before calling Gemini
89                    file_size = os.path.getsize(img_path)
90                    if file_size < 10000: # 10KB
91                        print(f"  ⏭️  Skipping {img_name} (Too small)")
92                        stats["skipped_images"] += 1
93                        break
94                        
95                    with open(img_path, "rb") as f:
96                        img_data = f.read()
97                    
98                    img_obj = Image.open(io.BytesIO(img_data))
99                    width, height = img_obj.size
100                    if width < 100 or height < 100:
101                        print(f"  ⏭️  Skipping {img_name} (Small dimensions)")
102                        stats["skipped_images"] += 1
103                        break
104            
105                    print(f"  🖼️  Analyzing image for Page {page_num} (Retry {retry_count}): {img_name}")
106                    response = client.models.generate_content(
107                        model=GOOGLE_MODEL,
108                        contents=[
109                            GEMINI_RAG_PROMPT,
110                            img_obj
111                        ],
112                        config=types.GenerateContentConfig(temperature=GOOGLE_TEMPERATURE)
113                    )
114                    description = response.text
115                    stats["scanned_images"] += 1
116                    
117                    # Update cache
118                    cache[img_name] = description
119                    save_cache(cache)
120                    break 
121                    
122                except Exception as e:
123                    if "429" in str(e) or "RESOURCE_EXHAUSTED" in str(e):
124                        wait_time = (retry_count + 1) * 10
125                        print(f"  ⚠️  Rate limit hit. Waiting {wait_time}s...")
126                        time.sleep(wait_time)
127                        retry_count += 1
128                    else:
129                        print(f"  ❌ Error processing image {img_name}: {e}")
130                        break
131            
132            if description:
133                visual_content += f"\n\n[Visual Data Transcription]:\n{description}"
134                time.sleep(2)
135
136        # 3. Fusion: combine text and visual descriptions
137        combined_content = text_content + visual_content
138
139        documents.append(
140            Document(
141                page_content=combined_content,
142                metadata={
143                    "page": page_num,
144                    "source": pdf_path,
145                    "contains_visuals": len(page_images) > 0
146                }
147            )
148        )
149
150    return documents, stats
151