#!/usr/bin/env python3 import logging import os import re import tempfile import uuid from typing import List, Tuple log = logging.getLogger("glmocr_simple_app") logging.basicConfig(level=logging.INFO) # ── Fine-tuned model repo on HuggingFace ───────────────────────────────────── MERGED_MODEL_DIR = os.environ.get("MODEL_DIR", "SimpleCodeAI/glm-ocr-finetuned") RENDER_SCALE = 2.0 PAD_LEFT_FRAC = 0.035 PAD_RIGHT_FRAC = 0.10 PAD_TOP_FRAC = 0.018 PAD_BOTTOM_FRAC = 0.018 ENABLE_CONTRAST = True CONTRAST_FACTOR = 1.18 ENABLE_UNSHARP = True UNSHARP_RADIUS = 0.78 UNSHARP_PERCENT = 76 UNSHARP_THRESHOLD = 1 PAGE_PNG_COMPRESS_LEVEL = 3 MAX_IMAGE_SIDE = 1568 MAX_NEW_TOKENS = 3000 # ── Model singleton ─────────────────────────────────────────────────────────── _model = None _processor = None def _load_model(): global _model, _processor if _model is not None: return _model, _processor import torch from transformers import AutoProcessor, AutoModelForImageTextToText log.info("Loading fine-tuned model from %s ...", MERGED_MODEL_DIR) _processor = AutoProcessor.from_pretrained( MERGED_MODEL_DIR, trust_remote_code=True ) _model = AutoModelForImageTextToText.from_pretrained( MERGED_MODEL_DIR, dtype=torch.bfloat16, device_map="auto", trust_remote_code=True, ) _model.eval() log.info("Model loaded.") return _model, _processor def _enhance_raster_for_ocr(img): from PIL import ImageEnhance, ImageFilter if ENABLE_CONTRAST: img = ImageEnhance.Contrast(img).enhance(CONTRAST_FACTOR) if ENABLE_UNSHARP: img = img.filter( ImageFilter.UnsharpMask( radius=UNSHARP_RADIUS, percent=UNSHARP_PERCENT, threshold=UNSHARP_THRESHOLD, ) ) return img def _resize_for_inference(img): """Resize image preserving aspect ratio so longest side <= MAX_IMAGE_SIDE.""" from PIL import Image w, h = img.size longest = max(w, h) if longest <= MAX_IMAGE_SIDE: return img ratio = MAX_IMAGE_SIDE / longest new_size = (int(w * ratio), int(h * ratio)) return img.resize(new_size, Image.LANCZOS) def _normalize_tables(text: str) -> str: """ Convert plain-text table rows (no pipe boundaries) into proper markdown table rows with | boundaries. Handles rows like: 10/01 CCD DEBIT, SOME MERCHANT 654.00 10/01 DEBIT 2,500.00 and header rows like: POSTING DATE DESCRIPTION AMOUNT """ lines = text.split('\n') result = [] i = 0 # Amount pattern: number like 1,234.56 or 123.45 amt_re = re.compile(r'[\d,]+\.\d{2}$') # Date pattern at start of line: MM/DD date_re = re.compile(r'^\d{2}/\d{2}\s') # Header pattern header_re = re.compile( r'^(POSTING\s+DATE|DATE)\s+(DESCRIPTION|SERIAL\s+NO\.?|NO\.?\s+CHECKS)', re.IGNORECASE ) in_plain_table = False while i < len(lines): line = lines[i] stripped = line.strip() # Skip empty lines — reset table tracking if not stripped: in_plain_table = False result.append(line) i += 1 continue # Already a proper markdown table row — pass through if stripped.startswith('|'): in_plain_table = False result.append(line) i += 1 continue # Detect plain-text table header row if header_re.match(stripped): # Split by 2+ spaces parts = re.split(r'\s{2,}', stripped) parts = [p.strip() for p in parts if p.strip()] if len(parts) >= 2: result.append('| ' + ' | '.join(parts) + ' |') result.append('| ' + ' | '.join(['---'] * len(parts)) + ' |') in_plain_table = True i += 1 continue # Detect plain-text data row starting with date if date_re.match(stripped): # Split by 2+ spaces to separate columns parts = re.split(r'\s{2,}', stripped) parts = [p.strip() for p in parts if p.strip()] # If only one part (description runs together), try splitting off amount if len(parts) == 1 and amt_re.search(stripped): # Split last amount token off m = re.search(r'^(.*?)\s+([\d,]+\.\d{2})$', stripped) if m: parts = [m.group(1).strip(), m.group(2).strip()] if len(parts) >= 2: result.append('| ' + ' | '.join(parts) + ' |') in_plain_table = True i += 1 continue # Detect subtotal lines inside a plain table if in_plain_table and re.match(r'^Subtotal:', stripped, re.IGNORECASE): m = re.match(r'^(Subtotal:)\s+([\d,]+\.\d{2})', stripped, re.IGNORECASE) if m: result.append(f'| | **{m.group(1)}** | **{m.group(2)}** |') i += 1 continue # Not a table row in_plain_table = False result.append(line) i += 1 return '\n'.join(result) def _infer_image(image_path: str) -> str: """Run fine-tuned model on a single image file and return markdown string.""" import torch from PIL import Image model, processor = _load_model() img = Image.open(image_path).convert("RGB") img = _resize_for_inference(img) fd, resized_path = tempfile.mkstemp(suffix=".png") os.close(fd) try: img.save(resized_path, "PNG") messages = [{ "role": "user", "content": [ {"type": "image", "url": resized_path}, {"type": "text", "text": "Document Parsing:"}, ], }] inputs = processor.apply_chat_template( messages, tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt", ).to(model.device) inputs.pop("token_type_ids", None) torch.cuda.empty_cache() with torch.no_grad(): ids = model.generate( **inputs, max_new_tokens=MAX_NEW_TOKENS, do_sample=False, repetition_penalty=1.1, ) result = processor.decode( ids[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True, ) result = result.strip() return result finally: try: os.unlink(resized_path) except Exception: pass def render_pdf_pages_to_images(pdf_path: str) -> Tuple[List[str], List[int]]: import pymupdf as fitz from PIL import Image doc = fitz.open(pdf_path) page_images: List[str] = [] page_heights: List[int] = [] for i in range(len(doc)): page = doc[i] pix = page.get_pixmap( matrix=fitz.Matrix(RENDER_SCALE, RENDER_SCALE), alpha=False ) img = Image.frombytes("RGB", [pix.width, pix.height], pix.samples) img = _enhance_raster_for_ocr(img) w, h = img.size pad_l = int(w * PAD_LEFT_FRAC) pad_r = int(w * PAD_RIGHT_FRAC) pad_t = int(h * PAD_TOP_FRAC) pad_b = int(h * PAD_BOTTOM_FRAC) if any(p > 0 for p in (pad_l, pad_r, pad_t, pad_b)): canvas = Image.new( "RGB", (w + pad_l + pad_r, h + pad_t + pad_b), (255, 255, 255) ) canvas.paste(img, (pad_l, pad_t)) img = canvas uniq = uuid.uuid4().hex[:10] img_path = os.path.join( tempfile.gettempdir(), f"glmocr_page_{os.getpid()}_{uniq}_{i}.png", ) img.save(img_path, "PNG", compress_level=PAGE_PNG_COMPRESS_LEVEL) page_images.append(img_path) page_heights.append(img.height) doc.close() return page_images, page_heights def run_ocr(uploaded_file): if uploaded_file is None: return "Please upload a file." page_images: List[str] = [] try: path = uploaded_file.name if hasattr(uploaded_file, "name") else str(uploaded_file) is_pdf = path.lower().endswith(".pdf") if is_pdf: page_images, _ = render_pdf_pages_to_images(path) else: page_images = [path] all_pages = [] for page_num, img_path in enumerate(page_images): log.info("Processing page %d / %d ...", page_num + 1, len(page_images)) page_md = _infer_image(img_path) if page_md: all_pages.append(page_md) merged = ( "\n\n---page-separator---\n\n".join(all_pages) if all_pages else "(No content extracted)" ) return merged except Exception as e: import traceback log.exception("run_ocr failed: %s", e) return f"Error: {e}\n\n{traceback.format_exc()}" finally: for p in page_images: try: if ( isinstance(p, str) and p.endswith(".png") and "glmocr_page_" in os.path.basename(p) ): os.unlink(p) except Exception: pass def _create_gradio_demo(): import gradio as gr with gr.Blocks(title="GLM-OCR Fine-tuned") as demo: gr.Markdown("# GLM-OCR (Fine-tuned)") file_in = gr.File( label="Upload PDF or image", file_types=[".pdf", ".png", ".jpg", ".jpeg", ".tiff", ".bmp"], ) run_btn = gr.Button("Run OCR", variant="primary") out = gr.Textbox(lines=40, label="Output (markdown)") run_btn.click(fn=run_ocr, inputs=file_in, outputs=out) return demo if __name__ == "__main__": _load_model() _create_gradio_demo().launch()