| """Nutrient Document Classification — open-vocab, zero-shot demo (public-ready). |
| |
| Censored like the sibling grounding demo: NO proprietary weights and NO custom architecture code ship in this |
| Space. Everything trained runs as ONNX loaded at RUNTIME: |
| * v2 flagship (commercial) — the full model runs as ONNX fetched from the PRIVATE repo via the |
| HF_TOKEN secret (server-side only, never committed here). Marked unavailable if the secret is absent. |
| * v1 (open-weight) — ONNX from the public v1 repo. |
| The preprocessor + tokenizer are bundled alongside the ONNX in the source repos, so no model identifiers |
| appear in this app. Runs on ZeroGPU; ONNX inference on CPU. Results are a ranked Markdown table. |
| """ |
| import glob as _glob |
| import os |
|
|
| import gradio as gr |
| import numpy as np |
| import spaces |
| from huggingface_hub import hf_hub_download |
| from transformers import AutoImageProcessor, AutoTokenizer |
|
|
| V2_PRIV = os.environ.get("V2_REPO", "nutrientdocs/document-classification-v2-private") |
| V1_REPO = os.environ.get("V1_REPO", "nutrientdocs/document-classification-v1") |
| TOKEN = os.environ.get("HF_TOKEN") |
|
|
| _V2, _V1 = {}, {} |
|
|
| DEFAULT_ROWS = [ |
| ["invoice", "an itemized bill listing goods or services and a total due"], |
| ["letter", "correspondence with a salutation, body, and signature"], |
| ["memo", "an internal memorandum headed To, From, Date, Subject"], |
| ["form", "a structured template with labeled fields to fill in"], |
| ["scientific article", "a research paper with an abstract, methods, and references"], |
| ["resume", "a summary of education, work experience, and skills"], |
| ] |
|
|
|
|
| def _rows_to_text(rows): |
| return "\n".join(f"{lab} | {desc}" if desc else lab for lab, desc in rows) |
|
|
|
|
| DEFAULT_TEXT = _rows_to_text(DEFAULT_ROWS) |
|
|
| |
| |
| _DESC = { |
| "financial reports": "a corporate financial statement with balance sheets, income tables, and figures", |
| "scientific articles": "a research paper with an abstract, methods, results, and references", |
| "laws and regulations": "statutory legal text with numbered sections and articles", |
| "government tenders": "a public procurement notice or call for bids", |
| "manuals": "a product or technical manual with instructions and labeled diagrams", |
| "patents": "a patent document with claims, drawings, and an application number", |
| "invoice": "an itemized bill listing goods or services, quantities, prices, and a total due", |
| "letter": "correspondence with a salutation, body, and signature", |
| "presentation": "a slide from a slideshow presentation", "receipt": "a store receipt with a total", |
| "bar chart": "a chart comparing values with rectangular bars", |
| "line chart": "a chart showing a trend as a line over an axis", |
| "pie chart": "a circular chart divided into proportional slices", |
| "qr code": "a square two-dimensional matrix barcode", "bar code": "a striped one-dimensional barcode", |
| "signature": "a handwritten signature mark", "screenshot": "a screenshot of a software user interface", |
| "table": "a grid of rows and columns of data", |
| } |
| _DOC = ["financial reports", "scientific articles", "laws and regulations", "government tenders", "manuals", |
| "patents", "invoice", "letter", "presentation", "receipt"] |
| _ELEM = ["bar chart", "line chart", "pie chart", "qr code", "bar code", "signature", "screenshot", "table", |
| "invoice", "financial reports"] |
| _ELEM_GT = {"bar chart", "pie chart", "line chart", "signature", "screenshot", "qr code", "bar code", "table"} |
|
|
|
|
| def _example_set(): |
| out = [] |
| for p in sorted(_glob.glob(os.path.join(os.path.dirname(__file__), "examples", "*.png"))): |
| gt = os.path.basename(p).rsplit("_", 1)[0].replace("-", " ") |
| group = _ELEM if gt in _ELEM_GT else _DOC |
| if gt not in group: |
| group = group + [gt] |
| out.append((gt, p, [[lab, _DESC.get(lab, "")] for lab in group])) |
| return out |
|
|
|
|
| EXAMPLES = _example_set() |
| |
| _EX_BY_NAME = {} |
| for _gt, _p, _t in EXAMPLES: |
| _EX_BY_NAME.setdefault(_gt, (_p, _t)) |
| EXAMPLE_NAMES = list(_EX_BY_NAME) |
|
|
|
|
| def _load_example(name): |
| p, t = _EX_BY_NAME.get(name, (None, DEFAULT_ROWS)) |
| return p, _rows_to_text(t) |
|
|
|
|
| def _parse(text): |
| labels, queries = [], [] |
| for line in (text or "").splitlines(): |
| lab, _, desc = line.partition("|") |
| lab, desc = lab.strip(), desc.strip() |
| if not lab: |
| continue |
| labels.append(lab); queries.append(f"{lab}. {desc}" if desc else lab) |
| return labels, queries |
|
|
|
|
| def _table(labels, probs): |
| ranked = sorted(zip(labels, probs), key=lambda x: -x[1]) |
| rows = [] |
| for i, (lab, p) in enumerate(ranked, 1): |
| pct = max(0.0, min(1.0, p)) * 100 |
| star = " ✅" if i == 1 else "" |
| rows.append( |
| f"<tr><td style='text-align:right;padding:4px 10px;color:#5c6773'>{i}</td>" |
| f"<td style='padding:4px 10px'><b>{lab}</b>{star}</td>" |
| f"<td style='padding:4px 10px;width:180px'>" |
| f"<div style='background:#e6ecf5;border-radius:4px;height:14px'>" |
| f"<div style='background:#2f52d0;height:14px;border-radius:4px;width:{pct:.0f}%'></div></div></td>" |
| f"<td style='text-align:right;padding:4px 10px;font-variant-numeric:tabular-nums'>{p:.3f}</td></tr>") |
| return ("<table style='border-collapse:collapse;font:14px system-ui'>" |
| "<thead><tr>" |
| "<th style='text-align:right;padding:4px 10px'>#</th>" |
| "<th style='text-align:left;padding:4px 10px'>Label</th>" |
| "<th style='padding:4px 10px'></th>" |
| "<th style='text-align:right;padding:4px 10px'>Score</th></tr></thead>" |
| "<tbody>" + "".join(rows) + "</tbody></table>") |
|
|
|
|
| def _ort(path): |
| import onnxruntime as ort |
| return ort.InferenceSession(path, providers=["CPUExecutionProvider"]) |
|
|
|
|
| |
| def _ensure_v2(): |
| if _V2: |
| return |
| if not TOKEN: |
| _V2["error"] = "no HF_TOKEN secret"; return |
| try: |
| dl = lambda f: hf_hub_download(V2_PRIV, f, repo_type="model", token=TOKEN) |
| _V2["img"] = _ort(dl("weights/image_encoder.onnx")) |
| _V2["txt"] = _ort(dl("weights/text_encoder.onnx")) |
| _V2["head"] = _ort(dl("weights/head.onnx")) |
| _V2["ip"] = AutoImageProcessor.from_pretrained(V2_PRIV, subfolder="encoder/image", token=TOKEN) |
| _V2["tk"] = AutoTokenizer.from_pretrained(V2_PRIV, subfolder="encoder/text", token=TOKEN) |
| except Exception as e: |
| _V2["error"] = f"{type(e).__name__}" |
|
|
|
|
| def _run_v2(image, queries): |
| _ensure_v2() |
| if "error" in _V2: |
| return None, f"⚠️ Flagship unavailable ({_V2['error']}). This model needs the commercial weights." |
| pv = _V2["ip"](images=[image.convert("RGB")], return_tensors="np")["pixel_values"].astype(np.float32) |
| patches = _V2["img"].run(["patches"], {"pixel_values": pv})[0] |
| enc = _V2["tk"](queries, padding=True, truncation=True, max_length=512, return_tensors="np") |
| lab = _V2["txt"].run(["embeds"], {"input_ids": enc["input_ids"].astype(np.int64), |
| "attention_mask": enc["attention_mask"].astype(np.int64)})[0] |
| probs = _V2["head"].run(["probs"], {"patches": patches.astype(np.float32), |
| "label_embeds": lab.astype(np.float32)})[0][0].tolist() |
| return probs, None |
|
|
|
|
| |
| def _ensure_v1(): |
| if _V1: |
| return |
| try: |
| img = hf_hub_download(V1_REPO, "modules/omni-image/image_model.onnx", repo_type="model", token=TOKEN) |
| txt = hf_hub_download(V1_REPO, "modules/omni-image/text_model.onnx", repo_type="model", token=TOKEN) |
| import json |
| cfg = hf_hub_download(V1_REPO, "modules/omni-image/config.json", repo_type="model", token=TOKEN) |
| cal = json.load(open(cfg)).get("calibration") or {"scale": 1.0, "bias": 0.0} |
| _V1["img"] = _ort(img); _V1["txt"] = _ort(txt) |
| _V1["tok"] = AutoTokenizer.from_pretrained(V1_REPO, subfolder="modules/omni-image", token=TOKEN) |
| _V1["proc"] = AutoImageProcessor.from_pretrained(V1_REPO, subfolder="modules/omni-image", token=TOKEN) |
| _V1["scale"], _V1["bias"] = float(cal["scale"]), float(cal["bias"]) |
| except Exception as e: |
| _V1["error"] = f"{type(e).__name__}" |
|
|
|
|
| def _run_v1(image, queries): |
| _ensure_v1() |
| if "error" in _V1: |
| return None, f"⚠️ v1 unavailable ({_V1['error']})." |
| pix = _V1["proc"](images=[image.convert("RGB")], return_tensors="np")["pixel_values"].astype(np.float16) |
| ie = _V1["img"].run(["image_emb"], {"pixel_values": pix})[0] |
| enc = _V1["tok"](queries, padding=True, truncation=True, max_length=64, return_tensors="np") |
| te = _V1["txt"].run(["text_emb"], {"input_ids": enc["input_ids"].astype(np.int64), |
| "attention_mask": enc["attention_mask"].astype(np.int64)})[0] |
| cos = (ie @ te.T)[0] |
| probs = (1.0 / (1.0 + np.exp(-(_V1["scale"] * cos + _V1["bias"])))).tolist() |
| return probs, None |
|
|
|
|
| LINKS = ("→ [model](https://huggingface.co/nutrientdocs/document-classification-v2) · " |
| "[leaderboard](https://huggingface.co/spaces/nutrientdocs/document-classification-leaderboard) · " |
| "[benchmark](https://huggingface.co/datasets/nutrientdocs/document-classification-benchmark)") |
|
|
| MODELS = {"v2 flagship — best accuracy (commercial)": "v2", |
| "v1 — open-weight (downloadable)": "v1"} |
|
|
|
|
| @spaces.GPU(duration=120) |
| def classify(image, text, model_name): |
| if image is None: |
| return "Upload a document image first." |
| labels, queries = _parse(text) |
| if not labels: |
| return "Add at least one class (one <b>label</b> per line)." |
| probs, err = (_run_v2 if MODELS.get(model_name) == "v2" else _run_v1)(image, queries) |
| return err if err else _table(labels, probs) |
|
|
|
|
| with gr.Blocks(title="Nutrient Document Classification", theme=gr.themes.Soft()) as demo: |
| gr.Markdown("# Classify any document against labels you choose\n" |
| "Open-vocabulary, zero-shot — add candidate classes (label + optional description), upload a " |
| "page, and the model ranks them. No fixed class list.\n\n" + LINKS) |
| with gr.Row(): |
| img = gr.Image(type="pil", label="Document image", height=360) |
| with gr.Column(): |
| model = gr.Radio(list(MODELS), value=list(MODELS)[0], label="Model") |
| tbl = gr.Textbox(value=DEFAULT_TEXT, lines=10, label="Candidate classes", |
| info="One class per line — `label | description` (the description is optional).") |
| btn = gr.Button("Classify", variant="primary") |
| out = gr.HTML() |
| sel = gr.Textbox(visible=False) |
| gr.Examples(examples=[[n] for n in EXAMPLE_NAMES], inputs=[sel], outputs=[img, tbl], |
| fn=_load_example, run_on_click=True, cache_examples=False, |
| label="Examples the model classifies correctly — click to load the page + its labels") |
| btn.click(classify, [img, tbl, model], out) |
|
|
| def _q(request: gr.Request): |
| m = (request.query_params or {}).get("model", "") |
| return list(MODELS)[1] if "v1" in m.lower() else list(MODELS)[0] |
| demo.load(_q, None, model) |
|
|
| demo.queue().launch(ssr_mode=False) |
|
|