multimodalart's picture
multimodalart HF Staff
Widen ZeroGPU duration headroom for first-entry weight streaming
6d9c0f4 verified
Raw History Blame Contribute Delete
19.5 kB
"""rakedoc-nano — document parsing demo on ZeroGPU.
Runs the exact KDL-Frontier / ParseBench `kdl_frontier_nano` pipeline
(layout detection -> crop -> per-category recognition -> OTSL/markdown
post-processing) locally with `transformers` instead of a vLLM endpoint.
"""
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import io
import time
import traceback
from typing import Any, Dict, List, Tuple
import spaces # noqa: E402 (must precede torch)
import torch # noqa: E402
import gradio as gr # noqa: E402
from PIL import Image # noqa: E402
from transformers import AutoModelForImageTextToText, AutoProcessor # noqa: E402
from transformers.generation.logits_process import LogitsProcessor, LogitsProcessorList # noqa: E402
import kdl_pipeline as K # noqa: E402
MODEL_ID = "cloudraker/rakedoc-nano"
CACHE_VERSION = 2
# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = AutoModelForImageTextToText.from_pretrained(
MODEL_ID,
dtype=torch.bfloat16,
attn_implementation="sdpa",
)
model = model.to("cuda").eval()
tokenizer = processor.tokenizer
EOS_IDS = {151645, 151643}
PAD_ID = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else 151643
# The stage prompt never changes, so the chat string can be rendered once.
STAGE_TEXTS: Dict[str, str] = {}
for _stage, _prompt in K._NANO_PROMPTS.items():
STAGE_TEXTS[_stage] = processor.apply_chat_template(
[
{
"role": "user",
"content": [{"type": "image"}, {"type": "text", "text": _prompt}],
}
],
add_generation_prompt=True,
tokenize=False,
)
# stages whose decoded output must keep special tokens (layout boxes, OTSL cells)
KEEP_SPECIAL = {"layout", "table"}
DEFAULT_BUDGETS = {
"layout": 6000,
"text": 2048,
"table": 5500,
"picture": 4096,
"formula": 128,
}
MODE_PAGE = "Full page (layout → regions → markdown)"
MODE_TABLE = "Table region only (OTSL → HTML)"
class VLLMPenaltyProcessor(LogitsProcessor):
"""vLLM-style presence / frequency penalties over the generated tokens.
The reference pipeline sends `presence_penalty` / `frequency_penalty` to a
vLLM endpoint for the table stage; `transformers` has no equivalent knob,
so the same arithmetic is reproduced here.
"""
def __init__(self, prompt_len: int, presence_penalty: float, frequency_penalty: float):
self.prompt_len = prompt_len
self.presence_penalty = presence_penalty
self.frequency_penalty = frequency_penalty
def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor):
generated = input_ids[:, self.prompt_len :]
if generated.shape[1] == 0:
return scores
counts = torch.zeros_like(scores)
counts.scatter_add_(1, generated, torch.ones_like(generated, dtype=scores.dtype))
return (
scores
- self.frequency_penalty * counts
- self.presence_penalty * (counts > 0).to(scores.dtype)
)
def _roundtrip(image: Image.Image, lossless: bool) -> Image.Image:
"""Mirror the reference transport encoding (PNG for tables, JPEG q95 else)."""
buf = io.BytesIO()
if lossless or image.mode in ("RGBA", "LA") or (
image.mode == "P" and "transparency" in image.info
):
image.save(buf, format="PNG")
else:
if image.mode not in ("L", "RGB", "CMYK"):
image = image.convert("RGB")
image.save(buf, format="JPEG", quality=95)
buf.seek(0)
return Image.open(buf).convert("RGB")
def _trim(ids: torch.Tensor) -> List[int]:
out: List[int] = []
for tok in ids.tolist():
if tok in EOS_IDS:
break
out.append(tok)
return out
@torch.inference_mode()
def _run_stage(stage: str, images: List[Image.Image], budgets: Dict[str, int],
batch_size: int) -> List[str]:
"""Greedy-decode one recognition stage over a list of crops."""
if not images:
return []
results: List[str] = []
lossless = stage in K._NANO_LOSSLESS_STAGES
text = STAGE_TEXTS[stage]
max_new = int(budgets.get(stage, DEFAULT_BUDGETS[stage]))
for start in range(0, len(images), batch_size):
chunk = [_roundtrip(im, lossless) for im in images[start : start + batch_size]]
inputs = processor(
text=[text] * len(chunk),
images=chunk,
padding=True,
return_tensors="pt",
).to(model.device)
prompt_len = inputs["input_ids"].shape[1]
extra = K._NANO_EXTRA_PAYLOAD.get(stage) or {}
lp = LogitsProcessorList()
pp = float(extra.get("presence_penalty", 0.0) or 0.0)
fp = float(extra.get("frequency_penalty", 0.0) or 0.0)
if pp or fp:
lp.append(VLLMPenaltyProcessor(prompt_len, pp, fp))
gen_kwargs: Dict[str, Any] = dict(
max_new_tokens=max_new,
do_sample=False,
num_beams=1,
pad_token_id=PAD_ID,
use_cache=True,
)
ngram = (extra.get("extra_args") or {}).get("no_repeat_ngram_size")
if ngram:
gen_kwargs["no_repeat_ngram_size"] = int(ngram)
if len(lp):
gen_kwargs["logits_processor"] = lp
out = model.generate(**inputs, **gen_kwargs)
skip_special = stage not in KEEP_SPECIAL
for row in out[:, prompt_len:]:
results.append(
tokenizer.decode(_trim(row), skip_special_tokens=skip_special)
)
return results
# ---------------------------------------------------------------------------
# Pipeline (port of kdl_pipeline._NanoEngine onto local generation)
# ---------------------------------------------------------------------------
_CAT_ORDER = ["text", "table", "picture", "formula"]
def _parse_page(image: Image.Image, budgets: Dict[str, int], batch_size: int,
report) -> Tuple[List[Dict[str, Any]], str]:
image = K.normalize_image_mode(image, "RGB")
w, h = image.size
if min(w, h) < 32:
return [], "image too small"
try:
if K.analyze_page_content(image).is_blank:
return [], "page looks blank"
except Exception:
pass
report(0.10, "Layout detection")
layout_image = K.prepare_native_layout_image(image)
layout_out = _run_stage("layout", [layout_image], budgets, 1)
layout_content = layout_out[0] if layout_out else ""
if not layout_content.strip() or not K.is_native_layout_response(layout_content):
return [], "no layout regions detected"
items = K.parse_native_layout_tokens(layout_content)
for item in items:
item["page_number"] = 1
buckets = K._nano_group_by_bucket(items, image)
n_regions = sum(len(v) for v in buckets.values())
if n_regions == 0:
return [], "no usable regions after layout"
fullpage_table = K.preprocess_for_vlm(image) if len(buckets["table"]) == 1 else None
# --- table stage -------------------------------------------------------
tables = buckets["table"]
if tables:
report(0.30, f"Table recognition ({len(tables)})")
pending = list(tables)
if fullpage_table is not None:
el = pending.pop(0)
content = (_run_stage("table", [fullpage_table], budgets, 1) or [None])[0]
if content is not None and K._nano_is_single_clean_otsl(content):
el["content"] = content
else:
pending.insert(0, el)
crops = [e["preprocessed_image"] for e in pending]
outs = _run_stage("table", crops, budgets, 1)
for el, content in zip(pending, outs):
el["content"] = content
# --- text stage --------------------------------------------------------
texts = [e for e in buckets["text"] if e.get("preprocessed_image") is not None]
if texts:
report(0.60, f"Text recognition ({len(texts)})")
outs = _run_stage("text", [e["preprocessed_image"] for e in texts], budgets,
batch_size)
for el, content in zip(texts, outs):
el["content"] = content
# --- picture stage -----------------------------------------------------
pics = [
e
for e in buckets["picture"]
if e.get("preprocessed_image") is not None
and e["preprocessed_image"].width >= 25
and e["preprocessed_image"].height >= 25
]
for el in buckets["picture"]:
el.setdefault("content", "")
if pics:
report(0.80, f"Figure / chart analysis ({len(pics)})")
outs = _run_stage("picture", [e["preprocessed_image"] for e in pics], budgets,
batch_size)
for el, content in zip(pics, outs):
K._nano_apply_picture_result(el, content)
# --- formula stage -----------------------------------------------------
formulas = [e for e in buckets["formula"] if e.get("preprocessed_image") is not None]
if formulas:
report(0.90, f"Formula recognition ({len(formulas)})")
outs = _run_stage("formula", [e["preprocessed_image"] for e in formulas],
budgets, batch_size)
for el, content in zip(formulas, outs):
el["content"] = content
elements: List[Dict[str, Any]] = []
picture_idx = 0
for name in _CAT_ORDER:
for el in buckets[name]:
el.pop("preprocessed_image", None)
cropped = el.pop("cropped_image", None)
if name == "picture":
if cropped is not None and cropped.width >= 25 and cropped.height >= 25:
el["picture_path"] = (
"artifacts/cropped_pictures/"
f"page_001_picture_{picture_idx:03d}.png"
)
picture_idx += 1
el.setdefault("content", "")
elements.append(el)
return elements, ""
_BOX_COLORS_NOTE = ""
def _annotations(elements: List[Dict[str, Any]], size: Tuple[int, int]):
w, h = size
anns = []
for el in sorted(elements, key=lambda e: e.get("layout_order", 0)):
bbox = el.get("bbox")
if not bbox or len(bbox) != 4:
continue
x1, y1, x2, y2 = bbox
box = (
max(0, int(round(x1 * w))),
max(0, int(round(y1 * h))),
min(w, int(round(x2 * w))),
min(h, int(round(y2 * h))),
)
anns.append((box, str(el.get("category", "Text"))))
return anns
def _display_markdown(md: str) -> str:
"""Artifact picture paths are not real files here — show a label instead."""
import re
return re.sub(r"!\[([^\]]*)\]\(artifacts/[^)]*\)", r"*[\1]*", md)
def _estimate(
image: Any = None,
mode: str = MODE_PAGE,
layout_max_tokens: int = 6000,
table_max_tokens: int = 5500,
*args,
**kwargs,
) -> int:
"""GPU budget, sized from measured runs on the bundled examples.
Measured in-GPU time: table mode 23s (AXA) / 49-63s (bus timetable, the
worst case, which nearly exhausts the 5500-token budget); full page
15s (charts) / 29s (61 regions) / 41-46s (dense table page). The spread
on repeats is the ZeroGPU first-entry weight stream, and the full-page
route scales with the number of *tables* on the page, so both keep
headroom above the observed maximum.
"""
try:
table_scale = max(0.5, float(table_max_tokens) / 5500.0)
except (TypeError, ValueError):
table_scale = 1.0
if mode == MODE_TABLE:
return int(min(150, round(20 + 65 * table_scale)))
return int(min(190, round(40 + 60 * table_scale)))
@spaces.GPU(duration=_estimate)
def parse_document(
image: Image.Image,
mode: str = MODE_PAGE,
layout_max_tokens: int = 6000,
table_max_tokens: int = 5500,
text_max_tokens: int = 2048,
picture_max_tokens: int = 4096,
batch_size: int = 4,
progress=gr.Progress(),
):
"""Parse a document page image into markdown with rakedoc-nano.
Args:
image: a rendered document page (or a cropped table region).
mode: run the full layout->region pipeline, or a single table pass.
layout_max_tokens: token budget for the layout-detection stage.
table_max_tokens: token budget for each table-recognition call.
text_max_tokens: token budget for each text-recognition call.
picture_max_tokens: token budget for each figure / chart call.
batch_size: how many region crops are decoded together.
Returns:
Rendered markdown, the raw markdown/OTSL output, the annotated page
layout, and a short run summary.
"""
if image is None:
return "", "", None, "Upload a document page image first."
budgets = {
"layout": int(layout_max_tokens),
"table": int(table_max_tokens),
"text": int(text_max_tokens),
"picture": int(picture_max_tokens),
"formula": DEFAULT_BUDGETS["formula"],
}
batch_size = max(1, int(batch_size))
image = image.convert("RGB") if image.mode != "RGB" else image
t0 = time.perf_counter()
def report(frac, desc):
try:
progress(frac, desc=desc)
except Exception:
pass
if mode == MODE_TABLE:
report(0.15, "Table recognition")
prepared = K.preprocess_for_vlm(K.normalize_image_mode(image, "RGB"))
raw = (_run_stage("table", [prepared], budgets, 1) or [""])[0]
el = {"category": "Table", "content": raw}
K._nano_postprocess_element(el)
html_table = el["content"]
dt = time.perf_counter() - t0
summary = (
f"**Table Recognition** · {dt:.1f}s · "
f"{len(raw)} chars of OTSL → "
f"{'HTML table' if html_table.lstrip().startswith('<table') else 'raw output'}"
)
return html_table, raw, (image, []), summary
try:
elements, why = _parse_page(image, budgets, batch_size, report)
except Exception:
tb = traceback.format_exc(limit=4)
return "", "", (image, []), f"Inference failed:\n```\n{tb}\n```"
if not elements:
return "", "", (image, []), f"No content extracted ({why})."
report(0.95, "Assembling markdown")
for el in elements:
K._nano_postprocess_element(el)
full_md, _pages = K._nano_assemble_markdown(elements)
full_md = K.postprocess_markdown(full_md)
anns = _annotations(elements, image.size)
dt = time.perf_counter() - t0
counts: Dict[str, int] = {}
for el in elements:
counts[el.get("category", "Text")] = counts.get(el.get("category", "Text"), 0) + 1
breakdown = " · ".join(f"{k} ×{v}" for k, v in sorted(counts.items()))
summary = f"**{len(elements)} regions** in {dt:.1f}s — {breakdown}"
return _display_markdown(full_md), full_md, (image, anns), summary
# ---------------------------------------------------------------------------
# UI
# ---------------------------------------------------------------------------
CSS = """
#col-container { max-width: 1280px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
#md-out { max-height: 720px; overflow: auto; }
"""
INTRO = """# rakedoc-nano — document parsing
A 1.2B Qwen2-VL document parser ([`cloudraker/rakedoc-nano`](https://huggingface.co/cloudraker/rakedoc-nano)),
LoRA-tuned for **table structure** on top of
[florin-parser-nano](https://huggingface.co/florin-inc/florin-parser-nano) /
[KDL-Frontier-Parser-nano](https://huggingface.co/KDLAI/KDL-Frontier-Parser-nano).
ParseBench: **77.2 overall / 86.4 tables**.
It is not a single-shot parser: the page is first passed through *Layout Detection*,
each region is cropped, and a stage-specific prompt (`Text` / `Table` / `Formula` /
`Image Analysis`) recognises it. Tables come back as OTSL and are converted to HTML.
"""
with gr.Blocks(title="rakedoc-nano") as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(INTRO)
with gr.Row():
with gr.Column(scale=1):
image_in = gr.Image(label="Document page", type="pil", height=420)
mode_in = gr.Radio(
[MODE_PAGE, MODE_TABLE],
value=MODE_PAGE,
label="Mode",
)
run = gr.Button("Parse document", variant="primary")
with gr.Accordion("Advanced settings", open=False):
layout_tok = gr.Slider(512, 8000, value=6000, step=64,
label="Layout max new tokens")
table_tok = gr.Slider(512, 8000, value=5500, step=64,
label="Table max new tokens")
text_tok = gr.Slider(128, 4096, value=2048, step=32,
label="Text max new tokens")
picture_tok = gr.Slider(128, 4096, value=4096, step=32,
label="Figure / chart max new tokens")
batch_in = gr.Slider(1, 8, value=4, step=1,
label="Region decode batch size")
with gr.Column(scale=1):
status = gr.Markdown("")
with gr.Tabs():
with gr.Tab("Rendered"):
md_out = gr.Markdown(label="Markdown", elem_id="md-out")
with gr.Tab("Raw output"):
raw_out = gr.Code(label="Markdown / OTSL", language="markdown")
with gr.Tab("Layout"):
layout_out = gr.AnnotatedImage(label="Detected regions",
height=520)
gr.Examples(
examples=[
["examples/axa_funds_statement.png", MODE_PAGE],
["examples/axa_funds_statement.png", MODE_TABLE],
["examples/bus_timetable.png", MODE_TABLE],
["examples/proxy_voting_roadmap.png", MODE_PAGE],
["examples/egov_survey_charts.png", MODE_PAGE],
],
inputs=[image_in, mode_in],
outputs=[md_out, raw_out, layout_out, status],
fn=parse_document,
cache_examples=True,
cache_mode="lazy",
label="Examples (ParseBench pages, Apache-2.0)",
)
gr.Markdown(
"Model weights are **AGPL-3.0** (inherited from KDL-Frontier-Parser-nano). "
"The deterministic pipeline in `kdl_pipeline.py` is vendored from "
"[run-llama/ParseBench](https://github.com/run-llama/ParseBench) (Apache-2.0); "
"example pages come from the "
"[ParseBench dataset](https://huggingface.co/datasets/llamaindex/ParseBench) "
"(Apache-2.0)."
)
run.click(
parse_document,
inputs=[image_in, mode_in, layout_tok, table_tok, text_tok, picture_tok,
batch_in],
outputs=[md_out, raw_out, layout_out, status],
api_name="parse",
)
if __name__ == "__main__":
# Gradio 6 moved `theme` / `css` from the Blocks constructor to launch().
demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)