#!/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 _build_table_from_entries(entries: List[dict]) -> str: if not entries: return "" table = [ "| Date | Description | Amount |", "|------|------------|--------|", ] for row in entries: table.append(f"| {row['date']} | {row['description']} | {row['amount']} |") return "\n".join(table) def _is_transaction_start(line: str) -> bool: return bool(re.match(r"\d{2}/\d{2}", line.strip())) def _ends_with_amount(text: str) -> bool: return bool(re.search(r"\d{1,3}(?:,\d{3})*\.\d{2}$", text.strip())) def _is_noise_line(line: str) -> bool: s = line.strip().lower() if not s: return True if s == "---page-separator---": return True if "call 1-800-937-2000" in s: return True if s in ("
", "
"): return True return False def _merge_transaction_lines(rows: List[str]) -> List[str]: """ Step 1/2: Merge multiline OCR rows into single transaction rows. A new row begins when a line starts with MM/DD. """ merged: List[str] = [] current = "" for raw in rows: line = " ".join(raw.strip().split()) if not line or _is_noise_line(line): continue if _is_transaction_start(line): if current: merged.append(current.strip()) current = line else: if current: current += " " + line else: current = line if current: merged.append(current.strip()) return merged def _repair_orphan_amounts(rows: List[str]) -> List[str]: """ Step 3: Attach lines without ending amount to prior transaction line when possible. """ fixed: List[str] = [] for line in rows: if _ends_with_amount(line): fixed.append(line) elif fixed: fixed[-1] = f"{fixed[-1]} {line}".strip() else: fixed.append(line) return fixed def _normalize_transactions(text: str) -> str: """ Convert loose transaction lines into structured markdown tables using: MM/DD ... amount """ lines = text.split("\n") result: List[str] = [] table_buffer: List[str] = [] in_section = False pattern = re.compile(r"(\d{2}/\d{2})\s+(.*?)\s+([\d,]+\.\d{2})$") for line in lines: if re.search(r"POSTING DATE.*DESCRIPTION.*AMOUNT", line, re.IGNORECASE): in_section = True table_buffer = [] continue if in_section and (line.strip() == "" or line.strip().startswith("Subtotal")): if table_buffer: clean_rows = [] merged_rows = _merge_transaction_lines(table_buffer) repaired_rows = _repair_orphan_amounts(merged_rows) for normalized in repaired_rows: match = pattern.search(normalized) if match: date, desc, amount = match.groups() clean_rows.append( {"date": date, "description": desc, "amount": amount} ) if clean_rows: result.append(_build_table_from_entries(clean_rows)) else: result.extend(table_buffer) table_buffer = [] in_section = False result.append(line) continue if in_section: table_buffer.append(line) else: result.append(line) if table_buffer: clean_rows = [] merged_rows = _merge_transaction_lines(table_buffer) repaired_rows = _repair_orphan_amounts(merged_rows) for normalized in repaired_rows: match = pattern.search(normalized) if match: date, desc, amount = match.groups() clean_rows.append( {"date": date, "description": desc, "amount": amount} ) if clean_rows: result.append(_build_table_from_entries(clean_rows)) else: result.extend(table_buffer) return "\n".join(result) def _fix_missing_amounts(text: str) -> str: lines = text.split("\n") fixed: List[str] = [] for line in lines: if re.search(r"\d{2}/\d{2}", line) and not re.search(r"\d+\.\d{2}", line): match = re.search(r"(\d{1,3}(?:,\d{3})*\.\d{2})$", line) if match: line += f" {match.group(1)}" fixed.append(line) return "\n".join(fixed) def _html_to_markdown_tables(text: str) -> str: try: from bs4 import BeautifulSoup except Exception: # Fallback when bs4 is unavailable: regex-based conversion for simple tables. table_re = re.compile(r"", re.IGNORECASE) row_re = re.compile(r"", re.IGNORECASE) cell_re = re.compile(r"]*>([\s\S]*?)", re.IGNORECASE) tag_re = re.compile(r"<[^>]+>") def _clean_cell(cell: str) -> str: cell = re.sub(r"", " ", cell, flags=re.IGNORECASE) cell = tag_re.sub(" ", cell) cell = cell.replace(" ", " ").replace("&", "&") cell = re.sub(r"\s+", " ", cell).strip() return cell def _convert(table_html: str) -> str: rows = [] for tr in row_re.findall(table_html): cols = [_clean_cell(c) for c in cell_re.findall(tr)] cols = [c for c in cols if c != ""] if cols: rows.append(cols) if not rows: return table_html width = max(len(r) for r in rows) norm = [r + [""] * (width - len(r)) for r in rows] md = [] md.append("| " + " | ".join(norm[0]) + " |") md.append("| " + " | ".join(["---"] * width) + " |") for r in norm[1:]: md.append("| " + " | ".join(r) + " |") return "\n".join(md) return table_re.sub(lambda m: _convert(m.group(0)), text) soup = BeautifulSoup(text, "html.parser") for table in soup.find_all("table"): rows = [] for tr in table.find_all("tr"): cols = [td.get_text(strip=True) for td in tr.find_all(["td", "th"])] if cols: rows.append(cols) if rows: md = [] header = rows[0] md.append("| " + " | ".join(header) + " |") md.append("| " + " | ".join(["---"] * len(header)) + " |") for row in rows[1:]: padded = row + [""] * (len(header) - len(row)) md.append("| " + " | ".join(padded[: len(header)]) + " |") table.replace_with("\n".join(md)) return str(soup) def _remove_prompt_echo_artifacts(text: str) -> str: """ Remove model prompt-echo garbage blocks such as repeated: 'N) Keep left-most amount values.' / 'N) Keep right-most amount values.' """ lines = text.split("\n") cleaned: List[str] = [] keep_line_re = re.compile( r"^\s*\d+\)\s*Keep\s+(left-most|right-most)\s+amount\s+values\.?\s*$", re.IGNORECASE, ) repeat_run = 0 for line in lines: if keep_line_re.match(line): repeat_run += 1 continue # Also drop explicit prompt/rules echoes if they appear as plain output. if re.match(r"^\s*Document Parsing to markdown\.?\s*$", line, re.IGNORECASE): continue if re.match(r"^\s*Rules:\s*$", line, re.IGNORECASE): continue cleaned.append(line) # If there was a long repeated run, ensure dangling separator spacing stays clean. text = "\n".join(cleaned) text = re.sub(r"\n{4,}", "\n\n\n", text) return text def _clean_markdown(text: str) -> str: """Post-process markdown to fix table formatting and section structure.""" text = _remove_prompt_echo_artifacts(text) text = _html_to_markdown_tables(text) text = _fix_missing_amounts(text) text = _normalize_transactions(text) lines = text.split('\n') cleaned = [] in_table = False for line in lines: if '|' in line and line.strip().startswith('|'): if not in_table: if cleaned and cleaned[-1].strip(): cleaned.append('') in_table = True line = re.sub(r'\s*\|\s*', ' | ', line) line = re.sub(r'\s+', ' ', line) cleaned.append(line.strip()) elif in_table and line.strip() == '': in_table = False cleaned.append('') else: in_table = False if line.strip() or (cleaned and cleaned[-1].strip()): cleaned.append(line) text = '\n'.join(cleaned) text = re.sub(r'(#{1,6})\s*([^\n]+)', r'\1 \2', text) text = re.sub(r'\n([•\-\*])\s+', r'\n\1 ', text) text = '\n'.join(line.rstrip() for line in text.split('\n')) text = re.sub(r'\n{4,}', '\n\n\n', text) return text.strip() 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": ( "Extract document text as markdown only. " "Do not include instructions, examples, or template rows. " "Preserve transaction lines exactly." ), }, ], }] 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) if torch.cuda.is_available(): 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, ) # Clean up the markdown output return _clean_markdown(result.strip()) 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)") gr.Markdown("Upload a PDF or image to extract structured markdown content.") 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()