Spaces:
Sleeping
Sleeping
| #!/usr/bin/env python3 | |
| import logging | |
| import os | |
| 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 ───────────────────────────────────── | |
| # Loads from HF Hub at runtime — no local storage needed in the Space | |
| MERGED_MODEL_DIR = os.environ.get("MODEL_DIR", "SimpleCodeAI/glm-ocr-finetuned") | |
| 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 | |
| RENDER_SCALE = 2.0 # was 3.0 — reduce render size so resize isn't as aggressive | |
| MAX_IMAGE_SIDE = 1568 # was 1344 — allow slightly larger input to model | |
| MAX_NEW_TOKENS = 3000 # resize longest side to this before inference | |
| # ── 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 _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, | |
| ) | |
| return 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)") | |
| 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__": | |
| # Force model to load at startup | |
| _load_model() | |
| _create_gradio_demo().launch() |