MD2204/multi_modality
1
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 