"""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('