Spaces:
Sleeping
Sleeping
| #!/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 ("<table border=\"1\"><tr><td></td></tr></table>", "<table><tr><td></td></tr></table>"): | |
| 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"<table[\s\S]*?</table>", re.IGNORECASE) | |
| row_re = re.compile(r"<tr[\s\S]*?</tr>", re.IGNORECASE) | |
| cell_re = re.compile(r"<t[dh][^>]*>([\s\S]*?)</t[dh]>", re.IGNORECASE) | |
| tag_re = re.compile(r"<[^>]+>") | |
| def _clean_cell(cell: str) -> str: | |
| cell = re.sub(r"<br\s*/?>", " ", 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() |