glm-ocr-fixed / app.py
rehan953's picture
Update app.py
44afe78 verified
Raw
History Blame
7.41 kB
#!/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()