""" 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: . ' "Format formulas as LaTeX. Format tables as HTML: ...
. " "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(' 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"") 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"\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()