File size: 10,168 Bytes
5625297
55cb0af
5625297
 
7b52db5
5625297
 
 
b1c7862
5625297
 
55cb0af
7beae18
 
55cb0af
7b52db5
5625297
 
 
 
55cb0af
5625297
 
 
 
 
 
55cb0af
5625297
7b52db5
 
55cb0af
7beae18
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
55cb0af
 
5625297
 
55cb0af
5625297
 
 
 
 
 
 
 
55cb0af
5625297
 
55cb0af
 
7beae18
 
 
55cb0af
7beae18
 
 
 
 
 
 
55cb0af
 
7b52db5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7beae18
 
 
 
55cb0af
7beae18
55cb0af
7beae18
 
55cb0af
7beae18
 
5625297
7beae18
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
44afe78
 
 
 
03fcc38
44afe78
7b52db5
7beae18
 
 
 
7b52db5
03fcc38
7b52db5
 
55cb0af
7beae18
55cb0af
7beae18
 
 
55cb0af
 
7fbf37f
 
5625297
594145e
7fbf37f
 
 
37508d7
7fbf37f
 
7beae18
 
 
37508d7
7fbf37f
5625297
37508d7
7fbf37f
 
 
 
 
47c1654
7fbf37f
7beae18
 
 
7fbf37f
 
47c1654
5625297
7beae18
 
 
 
5625297
7fbf37f
 
37508d7
7fbf37f
 
37508d7
5625297
7fbf37f
 
 
37508d7
5625297
7fbf37f
fffb995
 
47c1654
7fbf37f
7beae18
37508d7
fffb995
37508d7
7fbf37f
7beae18
 
 
 
 
 
 
 
 
 
 
7fbf37f
31f684f
7fbf37f
 
 
 
31f684f
7fbf37f
 
 
7beae18
 
 
 
 
7fbf37f
 
 
31f684f
5625297
7fbf37f
 
69880fb
7beae18
 
5625297
 
 
 
11bf437
7beae18
11bf437
 
 
5cc96b8
bcc62dd
7282022
7beae18
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
#!/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()