"""PaDoc: Layout-Grounded Parallel Decoding for Document Parsing."""
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # MUST come before any torch / CUDA import
import json
import re
import time
import gradio as gr
import torch
from PIL import Image, ImageDraw, ImageFont
from padoc.modeling import load_padoc_model
from padoc.transformers_infer import SequentialPaDocEngine
MODEL_ID = "Longin-Yu/PaDoc"
DEFAULT_QUERY = "Parse this document."
# Load model at module scope — ZeroGPU intercepts .to("cuda").
model, processor, fork_map = load_padoc_model(
MODEL_ID,
dtype=torch.bfloat16,
device_map=None,
attn_implementation="sdpa",
)
model = model.to("cuda")
model.eval()
engine = SequentialPaDocEngine(
model,
processor,
fork_map,
max_new_tokens=512,
max_branch_tokens=512,
max_concurrent_branches=8,
max_total_branches=64,
execution_mode="sequential",
strict=True,
)
print(f"[PaDoc] Model loaded on {engine.device}; devices={engine.devices}")
# ---------------------------------------------------------------------------
# Output formatting helpers
# ---------------------------------------------------------------------------
_LAYOUT_RE = re.compile(r"(\d+)\s+(\d+)\s+(\d+)\s+(\d+)")
_META_RE = re.compile(r'(\{.*?\})')
_COLORS = [
"#e6194B", "#3cb44b", "#4363d8", "#f58231", "#911eb4",
"#42d4f4", "#f032e6", "#bfef45", "#fabed4", "#469990",
]
def _parse_layout_boxes(main_text: str):
"""Return list of (x1, y1, x2, y2) in [0,1000] coordinates."""
boxes = []
for m in _LAYOUT_RE.finditer(main_text):
x1, y1, x2, y2 = (int(v) for v in m.groups())
boxes.append((x1, y1, x2, y2))
return boxes
def _parse_branch_meta(branch_text: str):
"""Return (category, content) from a branch text."""
meta_match = _META_RE.search(branch_text)
category = "region"
if meta_match:
try:
meta = json.loads(meta_match.group(1))
category = meta.get("category", "region")
except (json.JSONDecodeError, KeyError):
pass
content = _META_RE.sub("", branch_text).strip()
return category, content
def _annotate_image(image, boxes):
"""Draw layout boxes on a copy of the input image."""
annotated = image.copy().convert("RGB")
w, h = annotated.size
draw = ImageDraw.Draw(annotated)
try:
font = ImageFont.truetype(
"/usr/share/fonts/dejavu/DejaVuSans-Bold.ttf", max(14, int(min(w, h) / 40))
)
except OSError:
font = ImageFont.load_default()
for i, (x1, y1, x2, y2) in enumerate(boxes):
color = _COLORS[i % len(_COLORS)]
px1 = int(x1 / 1000 * w)
py1 = int(y1 / 1000 * h)
px2 = int(x2 / 1000 * w)
py2 = int(y2 / 1000 * h)
draw.rectangle([px1, py1, px2, py2], outline=color, width=3)
label = str(i + 1)
bbox = font.getbbox(label) if hasattr(font, "getbbox") else (0, 0, 20, 16)
tw, th = bbox[2] - bbox[0], bbox[3] - bbox[1]
draw.rectangle([px1, py1 - th - 4, px1 + tw + 8, py1], fill=color)
draw.text((px1 + 4, py1 - th - 3), label, fill="white", font=font)
return annotated
def _format_result(result):
"""Build a readable markdown summary of the parsed document."""
main_text = result.get("main", "")
boxes = _parse_layout_boxes(main_text)
branches = result.get("branches", [])
lines = []
lines.append(f"**Layout regions found:** {len(boxes)}")
lines.append(f"**Content branches:** {len(branches)}")
lines.append(f"**Execution mode:** {result.get('execution_mode', 'sequential')}")
lines.append("")
for i, branch in enumerate(branches):
text = branch.get("text", "")
category, content = _parse_branch_meta(text)
box_str = ""
if i < len(boxes):
x1, y1, x2, y2 = boxes[i]
box_str = f" `[{x1}, {y1}, {x2}, {y2}]`"
lines.append(f"### Region {i + 1}: {category}{box_str}")
lines.append("")
lines.append(content)
lines.append("")
return "\n".join(lines)
# ---------------------------------------------------------------------------
# Inference
# ---------------------------------------------------------------------------
@spaces.GPU(duration=60)
def parse_document(
image,
query: str = DEFAULT_QUERY,
execution_mode: str = "sequential",
max_new_tokens: int = 512,
max_branch_tokens: int = 512,
progress: gr.Progress = gr.Progress(track_tqdm=False),
):
"""Parse a document image and extract layout regions with content.
Args:
image: Document image to parse.
query: Instruction prompt for the parser.
execution_mode: "sequential" (batch=1 reference) or "parallel" (lockstep batched).
max_new_tokens: Maximum tokens for the main layout stream.
max_branch_tokens: Maximum tokens per content branch.
"""
if image is None:
raise gr.Error("Please provide a document image.")
if not isinstance(image, Image.Image):
image = Image.open(image).convert("RGB")
else:
image = image.convert("RGB")
content = [
{"type": "image", "image": image},
{"type": "text", "text": query or DEFAULT_QUERY},
]
messages = [{"role": "user", "content": content}]
# Update engine params for this request
engine.max_new_tokens = max_new_tokens
engine.max_branch_tokens = max_branch_tokens
engine.execution_mode = execution_mode
started = time.perf_counter()
result = engine.generate(messages, execution_mode=execution_mode)
elapsed = time.perf_counter() - started
main_text = result.get("main", "")
boxes = _parse_layout_boxes(main_text)
annotated = _annotate_image(image, boxes) if boxes else image
summary = _format_result(result)
info = (
f"⏱ {elapsed:.1f}s | "
f"Main tokens: {len(result.get('main_token_ids', []))} | "
f"Branches: {len(result.get('branches', []))} | "
f"Peak batch: {result.get('peak_batch_size', 1)}"
)
return annotated, summary, info
# ---------------------------------------------------------------------------
# UI
# ---------------------------------------------------------------------------
CSS = """
#col-container { max-width: 1200px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks() as demo:
gr.Markdown(
"# PaDoc: Layout-Grounded Parallel Decoding for Document Parsing\n"
"Upload a document image to extract its layout structure and region content "
"using the **[PaDoc](https://huggingface.co/Longin-Yu/PaDoc)** model — "
"an end-to-end document parser that decodes layout boxes and content branches in parallel."
)
with gr.Row():
with gr.Column(scale=1):
image_input = gr.Image(label="Document image", type="pil")
query = gr.Textbox(label="Query", value=DEFAULT_QUERY)
run_btn = gr.Button("Parse document", variant="primary")
with gr.Column(scale=1):
annotated_output = gr.Image(label="Detected layout regions")
info_output = gr.Textbox(label="Stats", interactive=False, container=False)
markdown_output = gr.Markdown(label="Parsed content")
with gr.Accordion("Advanced settings", open=False):
execution_mode = gr.Radio(
choices=["sequential", "parallel"],
value="sequential",
label="Execution mode",
info="Sequential: batch=1 reference. Parallel: lockstep batched branch decoding.",
)
max_new_tokens = gr.Slider(
minimum=64, maximum=1024, value=512, step=64,
label="Max main tokens",
)
max_branch_tokens = gr.Slider(
minimum=64, maximum=1024, value=512, step=64,
label="Max branch tokens",
)
run_btn.click(
parse_document,
inputs=[image_input, query, execution_mode, max_new_tokens, max_branch_tokens],
outputs=[annotated_output, markdown_output, info_output],
api_name="parse",
)
gr.Examples(
examples=[
["sample_doc.png", "Parse this document."],
["sample_invoice.png", "Parse this document."],
],
inputs=[image_input, query],
outputs=[annotated_output, markdown_output, info_output],
fn=parse_document,
cache_examples=True,
cache_mode="lazy",
)
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)