medical-guidelines-kg / scripts /runpod_batch_worker.py
gitmodelmujtaba's picture
Update checkpoint: OvisOCR2 validated on RunPod A40, pod terminated
0bd037b
Raw History Blame Contribute Delete
9.84 kB
"""
RunPod Batch Worker for ATH-MaaS/OvisOCR2 Medical Guideline Ingestion.
Processes guideline PDFs in `input_pdfs/` into Markdown files in `output_md/`.
Preserves natural reading order, HTML tables, LaTeX formulas, and clinical hierarchies.
"""
import argparse
import io
import json
import os
import re
import sys
import time
from pathlib import Path
from typing import List, Optional
import pypdfium2 as pdfium
from PIL import Image
MODEL_ID = "ATH-MaaS/OvisOCR2"
PROMPT = (
"Extract all readable content from the image in natural human reading order "
"and output the result as a single Markdown document. For charts or images, "
'represent them using an HTML image tag: <img src="images/bbox_{left}_{top}_{right}_{bottom}.jpg" />. '
"Format formulas as LaTeX. Format tables as HTML: <table>...</table>. "
"Transcribe all other text as standard Markdown. Preserve the original text without translation or paraphrasing."
)
def clean_truncated_repeats(
text: str,
min_text_len: int = 8000,
max_period: int = 200,
min_period: int = 1,
min_repeat_chars: int = 100,
min_repeat_times: int = 5,
) -> str:
"""Official OvisOCR2 post-processing loop remover."""
n = len(text)
if n < min_text_len:
return text
max_period = min(max_period, n - 1)
for unit_len in range(min_period, max_period + 1):
if text[n - 1] != text[n - 1 - unit_len]:
continue
match_len = 1
idx = n - 2
while idx >= unit_len and text[idx] == text[idx - unit_len]:
match_len += 1
idx -= 1
total_len = match_len + unit_len
repeat_times = total_len // unit_len
tail_len = total_len % unit_len
if repeat_times >= min_repeat_times and total_len >= min_repeat_chars:
return text[: n - total_len + unit_len] + text[n - tail_len :]
return text
def filter_visual_bbox_tags(text: str) -> str:
"""Removes raw bbox image tags while preserving surrounding markdown."""
return "\n\n".join(
block
for block in text.split("\n\n")
if not block.strip().startswith('<img src="images/bbox_')
)
class OvisOCRRunner:
def __init__(self, model_id: str = MODEL_ID):
self.model_id = model_id
self.mode = None
self.llm = None
self.sampling_params = None
self.tokenizer = None
self.model = None
self._init_engine()
def _init_engine(self):
print(f"Initializing OvisOCR2 engine for {self.model_id}...")
try:
from vllm import LLM, SamplingParams
print("Attempting vLLM initialization...")
self.llm = LLM(
model=self.model_id,
tensor_parallel_size=1,
gpu_memory_utilization=0.85,
trust_remote_code=True,
)
self.prompt_template = self.llm.get_tokenizer().apply_chat_template(
[{"role": "user", "content": [{"type": "image"}, {"type": "text", "text": PROMPT}]}],
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
self.sampling_params = SamplingParams(max_tokens=16384, temperature=0.0)
self.mode = "vllm"
print("vLLM engine initialized successfully.")
return
except Exception as e:
print(f"vLLM not available or failed: {e}. Falling back to HuggingFace Transformers...")
import torch
from transformers import AutoModelForVision2Seq, AutoTokenizer
print(f"Loading via AutoModelForVision2Seq...")
self.tokenizer = AutoTokenizer.from_pretrained(self.model_id, trust_remote_code=True)
self.model = AutoModelForVision2Seq.from_pretrained(
self.model_id,
torch_dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16,
device_map="auto",
trust_remote_code=True,
)
self.model.eval()
self.mode = "transformers"
print("Transformers engine initialized successfully.")
def parse_images(self, images: List[Image.Image]) -> List[str]:
if not images:
return []
if self.mode == "vllm":
vllm_inputs = [
{
"prompt": self.prompt_template,
"multi_modal_data": {"image": img},
"mm_processor_kwargs": {
"images_kwargs": {"min_pixels": 448 * 448, "max_pixels": 2880 * 2880}
},
}
for img in images
]
outputs = self.llm.generate(vllm_inputs, self.sampling_params)
results = []
for out in outputs:
txt = out.outputs[0].text.strip()
txt = filter_visual_bbox_tags(txt)
results.append(clean_truncated_repeats(txt))
return results
else:
import torch
results = []
for img in images:
inputs = self.tokenizer.apply_chat_template(
[{"role": "user", "content": [{"type": "image"}, {"type": "text", "text": PROMPT}]}],
tokenize=False,
add_generation_prompt=True,
)
# Fallback per-image inference
try:
inputs = self.tokenizer(images=img, text=inputs, return_tensors="pt").to("cuda")
with torch.inference_mode():
out = self.model.generate(**inputs, max_new_tokens=4096)
txt = self.tokenizer.decode(out[0], skip_special_tokens=True)
txt = filter_visual_bbox_tags(txt)
results.append(clean_truncated_repeats(txt))
except Exception as e:
results.append(f"<!-- Error processing page: {e} -->")
return results
def process_pdf(runner: OvisOCRRunner, pdf_path: Path, output_dir: Path, render_scale: float = 2.0) -> Path:
stem = pdf_path.stem
out_file = output_dir / f"{stem}.md"
if out_file.exists() and out_file.stat().st_size > 500:
print(f"Skipping already processed: {stem}")
return out_file
print(f"\n--- Processing: {pdf_path.name} ---")
start_time = time.time()
pdf = pdfium.PdfDocument(str(pdf_path))
total_pages = len(pdf)
print(f"Total pages: {total_pages}")
# Render pages in batches of 8 for GPU memory efficiency
batch_size = 8
all_markdown_pages = []
for i in range(0, total_pages, batch_size):
chunk_indices = range(i, min(i + batch_size, total_pages))
images = []
for p_idx in chunk_indices:
page = pdf[p_idx]
pil_img = page.render(scale=render_scale).to_pil().convert("RGB")
images.append(pil_img)
print(f"Processing pages {i+1} to {min(i+batch_size, total_pages)} / {total_pages}...")
page_texts = runner.parse_images(images)
for p_idx, text in zip(chunk_indices, page_texts):
all_markdown_pages.append(f"<!-- Page {p_idx + 1} -->\n\n{text}")
full_md = (
f"# {stem}\n\n"
f"> Extracted via ATH-MaaS/OvisOCR2 on {time.strftime('%Y-%m-%d %H:%M:%S')}\n"
f"> Source document: {pdf_path.name} ({total_pages} pages)\n\n"
"---\n\n"
+ "\n\n***\n\n".join(all_markdown_pages)
)
with open(out_file, "w", encoding="utf-8") as f:
f.write(full_md)
duration = round(time.time() - start_time, 2)
print(f"Saved: {out_file.name} ({len(full_md)} chars, {duration}s)")
return out_file
def main():
parser = argparse.ArgumentParser(description="OvisOCR2 RunPod Batch Processor")
parser.add_argument("--input-dir", type=str, default="input_pdfs", help="Directory containing PDF guidelines")
parser.add_argument("--output-dir", type=str, default="output_md", help="Directory for generated Markdown files")
parser.add_argument("--single-file", type=str, default=None, help="Process only a single file for validation")
parser.add_argument("--scale", type=float, default=2.0, help="Pdfium page render scale")
args = parser.parse_args()
input_dir = Path(args.input_dir)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
runner = OvisOCRRunner()
if args.single_file:
target = input_dir / args.single_file
if not target.exists():
target = Path(args.single_file)
process_pdf(runner, target, output_dir, render_scale=args.scale)
return
pdf_files = sorted(list(input_dir.glob("*.pdf")))
print(f"Found {len(pdf_files)} PDF files in {input_dir}")
manifest = {"started_at": time.time(), "processed": [], "failed": []}
for idx, pdf_file in enumerate(pdf_files, 1):
print(f"[{idx}/{len(pdf_files)}] {pdf_file.name}")
try:
out_path = process_pdf(runner, pdf_file, output_dir, render_scale=args.scale)
manifest["processed"].append({"name": pdf_file.name, "pages": len(pdfium.PdfDocument(str(pdf_file))), "status": "SUCCESS"})
except Exception as e:
print(f"ERROR processing {pdf_file.name}: {e}")
manifest["failed"].append({"name": pdf_file.name, "error": str(e)})
manifest["finished_at"] = time.time()
manifest["total_processed"] = len(manifest["processed"])
manifest["total_failed"] = len(manifest["failed"])
with open(output_dir / "batch_ocr_manifest.json", "w", encoding="utf-8") as f:
json.dump(manifest, f, indent=2)
print(f"\nBatch processing finished. Successfully processed: {len(manifest['processed'])}, Failed: {len(manifest['failed'])}")
if __name__ == "__main__":
main()