Team Ai
Apppublic

hugging-apps/padoc-document-parser

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
app.py258 linesDownload Raw Back to root
1"""PaDoc: Layout-Grounded Parallel Decoding for Document Parsing."""2 3import os4 5os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")6 7import spaces  # MUST come before any torch / CUDA import8 9import json10import re11import time12 13import gradio as gr14import torch15from PIL import Image, ImageDraw, ImageFont16 17from padoc.modeling import load_padoc_model18from padoc.transformers_infer import SequentialPaDocEngine19 20MODEL_ID = "Longin-Yu/PaDoc"21DEFAULT_QUERY = "Parse this document."22 23# Load model at module scope — ZeroGPU intercepts .to("cuda").24model, processor, fork_map = load_padoc_model(25    MODEL_ID,26    dtype=torch.bfloat16,27    device_map=None,28    attn_implementation="sdpa",29)30model = model.to("cuda")31model.eval()32engine = SequentialPaDocEngine(33    model,34    processor,35    fork_map,36    max_new_tokens=512,37    max_branch_tokens=512,38    max_concurrent_branches=8,39    max_total_branches=64,40    execution_mode="sequential",41    strict=True,42)43print(f"[PaDoc] Model loaded on {engine.device}; devices={engine.devices}")44 45 46# ---------------------------------------------------------------------------47# Output formatting helpers48# ---------------------------------------------------------------------------49 50_LAYOUT_RE = re.compile(r"<SP_LAYOUT>(\d+)\s+(\d+)\s+(\d+)\s+(\d+)</SP_LAYOUT>")51_META_RE = re.compile(r'<SP_META>(\{.*?\})</SP_META>')52_COLORS = [53    "#e6194B", "#3cb44b", "#4363d8", "#f58231", "#911eb4",54    "#42d4f4", "#f032e6", "#bfef45", "#fabed4", "#469990",55]56 57 58def _parse_layout_boxes(main_text: str):59    """Return list of (x1, y1, x2, y2) in [0,1000] coordinates."""60    boxes = []61    for m in _LAYOUT_RE.finditer(main_text):62        x1, y1, x2, y2 = (int(v) for v in m.groups())63        boxes.append((x1, y1, x2, y2))64    return boxes65 66 67def _parse_branch_meta(branch_text: str):68    """Return (category, content) from a branch text."""69    meta_match = _META_RE.search(branch_text)70    category = "region"71    if meta_match:72        try:73            meta = json.loads(meta_match.group(1))74            category = meta.get("category", "region")75        except (json.JSONDecodeError, KeyError):76            pass77    content = _META_RE.sub("", branch_text).strip()78    return category, content79 80 81def _annotate_image(image, boxes):82    """Draw layout boxes on a copy of the input image."""83    annotated = image.copy().convert("RGB")84    w, h = annotated.size85    draw = ImageDraw.Draw(annotated)86    try:87        font = ImageFont.truetype(88            "/usr/share/fonts/dejavu/DejaVuSans-Bold.ttf", max(14, int(min(w, h) / 40))89        )90    except OSError:91        font = ImageFont.load_default()92 93    for i, (x1, y1, x2, y2) in enumerate(boxes):94        color = _COLORS[i % len(_COLORS)]95        px1 = int(x1 / 1000 * w)96        py1 = int(y1 / 1000 * h)97        px2 = int(x2 / 1000 * w)98        py2 = int(y2 / 1000 * h)99        draw.rectangle([px1, py1, px2, py2], outline=color, width=3)100        label = str(i + 1)101        bbox = font.getbbox(label) if hasattr(font, "getbbox") else (0, 0, 20, 16)102        tw, th = bbox[2] - bbox[0], bbox[3] - bbox[1]103        draw.rectangle([px1, py1 - th - 4, px1 + tw + 8, py1], fill=color)104        draw.text((px1 + 4, py1 - th - 3), label, fill="white", font=font)105    return annotated106 107 108def _format_result(result):109    """Build a readable markdown summary of the parsed document."""110    main_text = result.get("main", "")111    boxes = _parse_layout_boxes(main_text)112    branches = result.get("branches", [])113 114    lines = []115    lines.append(f"**Layout regions found:** {len(boxes)}")116    lines.append(f"**Content branches:** {len(branches)}")117    lines.append(f"**Execution mode:** {result.get('execution_mode', 'sequential')}")118    lines.append("")119 120    for i, branch in enumerate(branches):121        text = branch.get("text", "")122        category, content = _parse_branch_meta(text)123        box_str = ""124        if i < len(boxes):125            x1, y1, x2, y2 = boxes[i]126            box_str = f" `[{x1}, {y1}, {x2}, {y2}]`"127        lines.append(f"### Region {i + 1}: {category}{box_str}")128        lines.append("")129        lines.append(content)130        lines.append("")131 132    return "\n".join(lines)133 134 135# ---------------------------------------------------------------------------136# Inference137# ---------------------------------------------------------------------------138 139@spaces.GPU(duration=60)140def parse_document(141    image,142    query: str = DEFAULT_QUERY,143    execution_mode: str = "sequential",144    max_new_tokens: int = 512,145    max_branch_tokens: int = 512,146    progress: gr.Progress = gr.Progress(track_tqdm=False),147):148    """Parse a document image and extract layout regions with content.149 150    Args:151        image: Document image to parse.152        query: Instruction prompt for the parser.153        execution_mode: "sequential" (batch=1 reference) or "parallel" (lockstep batched).154        max_new_tokens: Maximum tokens for the main layout stream.155        max_branch_tokens: Maximum tokens per content branch.156    """157    if image is None:158        raise gr.Error("Please provide a document image.")159    if not isinstance(image, Image.Image):160        image = Image.open(image).convert("RGB")161    else:162        image = image.convert("RGB")163 164    content = [165        {"type": "image", "image": image},166        {"type": "text", "text": query or DEFAULT_QUERY},167    ]168    messages = [{"role": "user", "content": content}]169 170    # Update engine params for this request171    engine.max_new_tokens = max_new_tokens172    engine.max_branch_tokens = max_branch_tokens173    engine.execution_mode = execution_mode174 175    started = time.perf_counter()176    result = engine.generate(messages, execution_mode=execution_mode)177    elapsed = time.perf_counter() - started178 179    main_text = result.get("main", "")180    boxes = _parse_layout_boxes(main_text)181 182    annotated = _annotate_image(image, boxes) if boxes else image183    summary = _format_result(result)184 185    info = (186        f"⏱ {elapsed:.1f}s  |  "187        f"Main tokens: {len(result.get('main_token_ids', []))}  |  "188        f"Branches: {len(result.get('branches', []))}  |  "189        f"Peak batch: {result.get('peak_batch_size', 1)}"190    )191 192    return annotated, summary, info193 194 195# ---------------------------------------------------------------------------196# UI197# ---------------------------------------------------------------------------198 199CSS = """200#col-container { max-width: 1200px; margin: 0 auto; }201.dark .gradio-container { color: var(--body-text-color); }202"""203 204with gr.Blocks() as demo:205    gr.Markdown(206        "# PaDoc: Layout-Grounded Parallel Decoding for Document Parsing\n"207        "Upload a document image to extract its layout structure and region content "208        "using the **[PaDoc](https://huggingface.co/Longin-Yu/PaDoc)** model — "209        "an end-to-end document parser that decodes layout boxes and content branches in parallel."210    )211 212    with gr.Row():213        with gr.Column(scale=1):214            image_input = gr.Image(label="Document image", type="pil")215            query = gr.Textbox(label="Query", value=DEFAULT_QUERY)216            run_btn = gr.Button("Parse document", variant="primary")217        with gr.Column(scale=1):218            annotated_output = gr.Image(label="Detected layout regions")219            info_output = gr.Textbox(label="Stats", interactive=False, container=False)220 221    markdown_output = gr.Markdown(label="Parsed content")222 223    with gr.Accordion("Advanced settings", open=False):224        execution_mode = gr.Radio(225            choices=["sequential", "parallel"],226            value="sequential",227            label="Execution mode",228            info="Sequential: batch=1 reference. Parallel: lockstep batched branch decoding.",229        )230        max_new_tokens = gr.Slider(231            minimum=64, maximum=1024, value=512, step=64,232            label="Max main tokens",233        )234        max_branch_tokens = gr.Slider(235            minimum=64, maximum=1024, value=512, step=64,236            label="Max branch tokens",237        )238 239    run_btn.click(240        parse_document,241        inputs=[image_input, query, execution_mode, max_new_tokens, max_branch_tokens],242        outputs=[annotated_output, markdown_output, info_output],243        api_name="parse",244    )245 246    gr.Examples(247        examples=[248            ["sample_doc.png", "Parse this document."],249            ["sample_invoice.png", "Parse this document."],250        ],251        inputs=[image_input, query],252        outputs=[annotated_output, markdown_output, info_output],253        fn=parse_document,254        cache_examples=True,255        cache_mode="lazy",256    )257 258demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)