hugging-apps/padoc-document-parser
0
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)