Spaces:
Running on Zero
Running on Zero
| """Doc-Split demo — upload a merged PDF, watch it split into its constituent documents. | |
| Each model predicts, per page, whether it starts a new document. Text comes from the PDF's embedded text | |
| layer (free); scanned pages fall back to vision-only (OCR gate). Every model runs as portable ONNX graphs | |
| under onnxruntime, with the per-page confidence smoothing done in numpy. Commercial weights load at runtime | |
| from a private repo via the HF_TOKEN secret (server-side only, never downloadable). | |
| ZeroGPU: models load lazily on first call and are cached across calls.""" | |
| import os, io, base64, json | |
| import numpy as np | |
| import gradio as gr | |
| import spaces | |
| import onnxruntime as ort | |
| import fitz # pymupdf | |
| from PIL import Image | |
| from transformers import AutoTokenizer | |
| from huggingface_hub import snapshot_download | |
| TOKEN = os.environ.get("HF_TOKEN") | |
| _PROV = ["CUDAExecutionProvider", "CPUExecutionProvider"] # onnxruntime-gpu if present, else CPU fallback | |
| MODELS = { | |
| "doc-split-v2 (flagship · commercial)": "nutrientdocs/doc-split-v2-private", | |
| "doc-split-v1 (open-weight)": "nutrientdocs/doc-split-v1", | |
| } | |
| _CACHE = {} | |
| def _lse(x, axis): # numerically-stable log-sum-exp along one axis (keeps shape reduced) | |
| m = x.max(axis, keepdims=True) | |
| return (m + np.log(np.exp(x - m).sum(axis, keepdims=True))).squeeze(axis) | |
| def _crf_marginals(bl, crf): | |
| """Per-page P(boundary) via forward-backward over a 2-tag linear chain. bl: raw boundary logits [N]. | |
| Emission per page t is [0, bl[t]] (tag 0 = interior, tag 1 = boundary). Public CRF math, no arch.""" | |
| trans = np.asarray(crf["trans"], np.float64) # [2,2] i->j | |
| start = np.asarray(crf["start"], np.float64); end = np.asarray(crf["end"], np.float64) | |
| N = len(bl); e = np.stack([np.zeros(N), np.asarray(bl, np.float64)], 1) # [N,2] | |
| a = np.zeros((N, 2)); a[0] = start + e[0] | |
| for t in range(1, N): | |
| a[t] = _lse(a[t - 1][:, None] + trans, 0) + e[t] | |
| bta = np.zeros((N, 2)); bta[N - 1] = end | |
| for t in range(N - 2, -1, -1): | |
| bta[t] = _lse(trans + (e[t + 1] + bta[t + 1])[None, :], 1) | |
| m = a + bta; m = m - m.max(1, keepdims=True); p = np.exp(m) | |
| return (p / p.sum(1, keepdims=True))[:, 1] # posterior P(tag=boundary) per page | |
| # Beta calibration (fit on our-domain val): smooth, monotonic map raw CRF marginal -> honest P(boundary). | |
| # The raw model is over-confident (e.g. flagship raw 0.99 is really ~0.86 likely a boundary); this corrects it. | |
| # p_cal = sigmoid(a*ln(p) + b*ln(1-p) + c). Fit values are hard-coded per model. | |
| _BETA = { | |
| "v2": (0.4898, -0.5100, -0.5071), # flagship; ECE 0.045 -> 0.013 | |
| "v1": (0.5156, -0.4023, -0.1545), # open; ECE 0.044 -> 0.012 | |
| } | |
| def _calibrate(p, model_name): | |
| a, b, c = _BETA["v1" if "v1" in model_name.lower() else "v2"] | |
| p = min(max(float(p), 1e-6), 1 - 1e-6) | |
| return float(1.0 / (1.0 + np.exp(-(a * np.log(p) + b * np.log(1 - p) + c)))) | |
| def _load(name): | |
| if name in _CACHE: | |
| return _CACHE[name] | |
| d = snapshot_download(MODELS[name], token=TOKEN, # commercial weights via HF_TOKEN; v1 is open | |
| allow_patterns=["*.onnx", "*.onnx.data", "crf.json", "tokenizer*", "special_tokens*"]) | |
| obj = dict(tok=AutoTokenizer.from_pretrained(d), crf=json.load(open(os.path.join(d, "crf.json"))), | |
| img=ort.InferenceSession(os.path.join(d, "image_model.onnx"), providers=_PROV), | |
| txt=ort.InferenceSession(os.path.join(d, "text_model.onnx"), providers=_PROV), | |
| head=ort.InferenceSession(os.path.join(d, "head.onnx"), providers=_PROV)) | |
| _CACHE[name] = obj | |
| return obj | |
| def _pdf_pages(path, max_pages=40): | |
| doc = fitz.open(path); out = [] | |
| for i, pg in enumerate(doc): | |
| if i >= max_pages: | |
| break | |
| pix = pg.get_pixmap(matrix=fitz.Matrix(150 / 72, 150 / 72)) | |
| im = Image.frombytes("RGB", (pix.width, pix.height), pix.samples) | |
| out.append((im, pg.get_text("text") or "")) | |
| doc.close() | |
| return out | |
| def _thumb(im, w=150): | |
| t = im.copy(); t.thumbnail((w, w * 2)); b = io.BytesIO(); t.save(b, "PNG") | |
| return "data:image/png;base64," + base64.b64encode(b.getvalue()).decode() | |
| def _thumbstrip(pages, label): | |
| h = [f"<div style='font:600 13px system-ui;margin:0 0 8px'>{label}</div>" | |
| "<div style='display:flex;gap:6px;overflow-x:auto'>"] | |
| for i, (im, _) in enumerate(pages): | |
| h.append(f"<div style='text-align:center;flex:0 0 auto'><img src='{_thumb(im,110)}' " | |
| f"style='height:120px;border:1px solid #ccc;border-radius:3px'/>" | |
| f"<div style='font:10px system-ui;color:#888'>p{i+1}</div></div>") | |
| h.append("</div>"); return "".join(h) | |
| def preview(pdf): | |
| if not pdf: | |
| return "" | |
| pages = _pdf_pages(pdf) | |
| return _thumbstrip(pages, f"{len(pages)} page(s) — press <b>Split</b> to segment") | |
| def split(pdf, model_name, min_conf): | |
| if not pdf: | |
| return "<p style='color:#888'>Upload a PDF (a few concatenated documents) and press Split.</p>" | |
| obj = _load(model_name) | |
| pages = _pdf_pages(pdf) | |
| if not pages: | |
| return "<p>No pages found.</p>" | |
| N = len(pages) | |
| arr = np.stack([(np.asarray(im.convert("RGB").resize((512, 512)), np.float32) / 255 - .5) / .5 | |
| for im, _ in pages]).transpose(0, 3, 1, 2).astype(np.float32) # [N,3,512,512] | |
| prompts = ["query: " + (t or " ") for _, t in pages] | |
| gate = np.array([1. if (t and t.strip()) else 0. for _, t in pages], np.float32) | |
| vi = obj["img"].run(["image_embed"], {"pixel_values": arr})[0] # image tower -> per-page embedding | |
| b = obj["tok"](prompts, padding=True, truncation=True, max_length=512, return_tensors="np") | |
| vt = obj["txt"].run(["text_embed"], {"input_ids": b["input_ids"].astype(np.int64), | |
| "attention_mask": b["attention_mask"].astype(np.int64)})[0] | |
| vi = vi.astype(np.float32)[None] # [1,N,d_img] | |
| vt = (vt.astype(np.float32) * gate[:, None])[None] # [1,N,d_txt], OCR-gated | |
| bl = obj["head"].run(["boundary_logit"], {"v_img": vi, "v_txt": vt, # ONNX boundary head | |
| "gate": gate[None], "mask": np.ones((1, N), np.float32)})[0][0] | |
| bl[0] = 30.0 # keep page 0 forced | |
| raw = _crf_marginals(bl, obj["crf"]).tolist() # per-page P(boundary), uncalibrated | |
| conf = [_calibrate(c, model_name) for c in raw] # honest, calibrated confidence (beta) | |
| # start a new document only where the CALIBRATED confidence >= the chosen threshold. | |
| tau = float(min_conf) / 100.0 | |
| pred = [1 if (i == 0 or conf[i] >= tau) else 0 for i in range(len(conf))] | |
| # group pages into documents at each boundary | |
| docs, cur = [], [] | |
| for i, p in enumerate(pred): | |
| if p and cur: | |
| docs.append(cur); cur = [] | |
| cur.append(i) | |
| if cur: | |
| docs.append(cur) | |
| txt_pages = sum(1 for _, t in pages if t and t.strip()) | |
| html = [f"<div style='font:600 15px system-ui;margin:0 0 12px'>Split into <b>{len(docs)}</b> " | |
| f"document(s) across {N} pages · {model_name} · text layer on {txt_pages}/{N} pages</div>"] | |
| for k, grp in enumerate(docs, 1): | |
| c = conf[grp[0]] | |
| split = "forced start (p1)" if grp[0] == 0 else f"split confidence {c:.0%}" | |
| dot = "#0f7a58" if c >= .8 else ("#9a6a12" if c >= .5 else "#a83a3a") | |
| html.append(f"<div style='border:1px solid #d5deea;border-radius:10px;padding:12px;margin:0 0 12px'>" | |
| f"<div style='font:600 12px system-ui;letter-spacing:.05em;color:#2f52d0;" | |
| f"text-transform:uppercase;margin-bottom:8px'>Document {k} · pages " | |
| f"{grp[0]+1}–{grp[-1]+1} ({len(grp)}p) · " | |
| f"<span style='color:{dot}'>● {split}</span></div>" | |
| f"<div style='display:flex;gap:8px;overflow-x:auto'>") | |
| for i in grp: | |
| html.append(f"<div style='text-align:center;flex:0 0 auto'>" | |
| f"<img src='{_thumb(pages[i][0])}' style='height:150px;border:1px solid #ccc;border-radius:4px'/>" | |
| f"<div style='font:11px system-ui;color:#5c6773'>p{i+1}</div></div>") | |
| html.append("</div></div>") | |
| return "".join(html) | |
| with gr.Blocks(theme=gr.themes.Soft(), title="Doc-Split") as demo: | |
| gr.Markdown("## 📄✂️ Doc-Split — split a merged PDF into its documents\n" | |
| "Upload a PDF that concatenates several documents; the model marks where each new document " | |
| "starts. · [leaderboard](https://huggingface.co/spaces/nutrientdocs/doc-split-leaderboard) " | |
| "· [doc-split-v2](https://huggingface.co/nutrientdocs/doc-split-v2) " | |
| "· [doc-split-v1](https://huggingface.co/nutrientdocs/doc-split-v1) " | |
| "· [benchmark](https://huggingface.co/datasets/nutrientdocs/doc-split-benchmark)") | |
| with gr.Row(): | |
| pdf = gr.File(label="Merged PDF", file_types=[".pdf"], type="filepath") | |
| with gr.Column(): | |
| model = gr.Radio(list(MODELS), value=list(MODELS)[0], label="Model") | |
| thr = gr.Slider(50, 95, value=80, step=5, | |
| label="Minimum confidence to start a new document (%)", | |
| info="Calibrated confidence (fit on held-out data, so the % reflects the real chance a " | |
| "page starts a new doc). A page splits only when it clears this bar. " | |
| "Raise it for fewer splits; lower it to merge less.") | |
| preview_html = gr.HTML() | |
| btn = gr.Button("Split", variant="primary") | |
| out = gr.HTML() | |
| pdf.change(preview, pdf, preview_html) | |
| btn.click(split, [pdf, model, thr], out) | |
| import glob as _glob | |
| _ex = sorted(_glob.glob("examples/*.pdf")) | |
| gr.Examples( | |
| examples=[[f] for f in _ex] or None, | |
| inputs=[pdf], label="Examples — real multi-document streams (pick one, choose a model, then Split)") | |
| def _q(request: gr.Request): | |
| m = (request.query_params or {}).get("model", "") | |
| return list(MODELS)[1] if ("v1" in m or "mini" in m) else list(MODELS)[0] | |
| demo.load(_q, None, model) | |
| if __name__ == "__main__": | |
| demo.launch() | |