from __future__ import annotations import base64 import html import io import json import os import re import subprocess import sys import time from pathlib import Path # Pre-built causal_conv1d wheel (linear-attention dependency, can't compile in Space build). # Pulled from Dao-AILab official GitHub release; matches torch==2.6.0 / cu12 / py3.10 / cxx11abi=FALSE. _CAUSAL_CONV1D_WHEEL = ( "https://github.com/Dao-AILab/causal-conv1d/releases/download/v1.5.0.post8/" "causal_conv1d-1.5.0.post8+cu12torch2.6cxx11abiFALSE-cp310-cp310-linux_x86_64.whl" ) try: import causal_conv1d # noqa: F401 except ImportError: subprocess.run( [sys.executable, "-m", "pip", "install", _CAUSAL_CONV1D_WHEEL, "--no-deps", "-q"], check=True, ) import causal_conv1d # noqa: F401 import gradio as gr import spaces from PIL import Image MODEL_ID = os.getenv("MODEL_ID", "ebinan92/open-chandra-stage2-2b") PROMPT = os.getenv( "OCR_PROMPT", "OCR this image as HTML layout blocks with bbox and label." ) MAX_NEW_TOKENS = 12288 # Area-based image cap (matches chandra-ocr-2 reference Space). 1536² px² ≈ 2300 visual tokens. MAX_IMAGE_PIXELS = int(os.getenv("MAX_IMAGE_PIXELS", str(1536 * 1536))) PDF_DPI = int(os.getenv("PDF_DPI", "180")) HF_TOKEN = os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN") LABEL_COLORS = { "Text": "#2196F3", "Section-Header": "#F44336", "Equation-Block": "#9C27B0", "Table": "#FF9800", "Image": "#4CAF50", "Figure": "#009688", "Caption": "#795548", "Footnote": "#607D8B", "List-Group": "#3F51B5", "Page-Header": "#CDDC39", "Page-Footer": "#FFC107", "Code-Block": "#00BCD4", "Bibliography": "#E91E63", "Complex-Block": "#8BC34A", "Form": "#FF5722", "Table-Of-Contents": "#673AB7", "Diagram": "#FFEB3B", } _BLOCK_RE = re.compile( r'(.*?)', re.DOTALL, ) def _load_model_eager(): """Module-level eager load for ZeroGPU. Uses device_map="auto" + flash-linear-attention CUDA kernels (via the `fla` package pulled in by requirements). Without flash-linear-attention, the 18 linear_attention layers in Qwen3.5 fall back to a slow Python implementation — that's the 5–10x slowdown we hit before. """ import torch from transformers import AutoModelForImageTextToText, AutoProcessor print(f"[startup] Loading model: {MODEL_ID}") kwargs = {"trust_remote_code": True, "device_map": "auto"} if HF_TOKEN: kwargs["token"] = HF_TOKEN try: model = AutoModelForImageTextToText.from_pretrained( MODEL_ID, dtype=torch.bfloat16, **kwargs, ) except TypeError: model = AutoModelForImageTextToText.from_pretrained( MODEL_ID, torch_dtype=torch.bfloat16, **kwargs, ) model.eval() processor = AutoProcessor.from_pretrained( MODEL_ID, trust_remote_code=True, token=HF_TOKEN, ) if hasattr(processor, "tokenizer"): processor.tokenizer.padding_side = "left" print(f"[startup] Model ready") return model, processor MODEL, PROCESSOR = _load_model_eager() def _resize_if_needed(image: Image.Image) -> Image.Image: """Area-based scale-to-fit (matches chandra-ocr-2 Space).""" if image.mode != "RGB": image = image.convert("RGB") if MAX_IMAGE_PIXELS <= 0: return image w, h = image.size total = w * h if total <= MAX_IMAGE_PIXELS: return image scale = (MAX_IMAGE_PIXELS / total) ** 0.5 return image.resize((max(1, int(w * scale)), max(1, int(h * scale))), Image.LANCZOS) def _render_pdf_page(path: Path, page_number: int) -> tuple[int, Image.Image]: import pypdfium2 as pdfium doc = pdfium.PdfDocument(str(path)) try: if page_number < 1 or page_number > len(doc): raise gr.Error( f"Invalid page number ({page_number}). Must be between 1 and {len(doc)}." ) page = doc[page_number - 1] image = page.render(scale=PDF_DPI / 72).to_pil() page.close() return page_number, _resize_if_needed(image) finally: doc.close() def _load_image(file_path: str, page_number: int) -> tuple[int, Image.Image]: path = Path(file_path) if path.suffix.lower() == ".pdf": return _render_pdf_page(path, page_number) img = Image.open(path) img.load() # eagerly read pixel data so the fp can be released return 1, _resize_if_needed(img) @spaces.GPU(duration=180) def _generate_page( image: Image.Image, temperature: float, ) -> tuple[str, int, float]: """Inference path mirrors victor/chandra-ocr-2 exactly.""" conversation = [ { "role": "user", "content": [ {"type": "image", "image": image}, {"type": "text", "text": PROMPT}, ], } ] inputs = PROCESSOR.apply_chat_template( [conversation], tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt", padding=True, ) inputs = inputs.to(MODEL.device) eos_token_id = MODEL.generation_config.eos_token_id im_end_id = PROCESSOR.tokenizer.convert_tokens_to_ids("<|im_end|>") if isinstance(eos_token_id, int): eos_token_id = [eos_token_id] elif eos_token_id is None: eos_token_id = [] if im_end_id is not None and im_end_id not in eos_token_id: eos_token_id.append(im_end_id) start = time.perf_counter() generated_ids = MODEL.generate( **inputs, max_new_tokens=MAX_NEW_TOKENS, eos_token_id=eos_token_id, ) elapsed = time.perf_counter() - start new_tokens = generated_ids.shape[1] - inputs.input_ids.shape[1] generated_ids_trimmed = generated_ids[0][len(inputs.input_ids[0]):] text = PROCESSOR.decode( generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False, ) return text, int(new_tokens), elapsed def parse_chandra_html(html_text: str) -> list[dict]: blocks = [] for match in _BLOCK_RE.finditer(html_text): bbox, label, content = match.groups() blocks.append({"bbox": bbox, "label": label, "content": content}) return blocks def _image_data_url(image: Image.Image) -> str: buf = io.BytesIO() image.save(buf, format="PNG") b64 = base64.b64encode(buf.getvalue()).decode("ascii") return f"data:image/png;base64,{b64}" def _json_for_script(data) -> str: return json.dumps(data, ensure_ascii=False).replace(" str: header_inner = ( f'
{html.escape(title)}
' if title else "" ) header_block = ( f'
{header_inner}
' if header_inner else "" ) srcdoc = f"""
{header_block}
page
Prediction Pred
""" return ( f'' ) def _preview_html(image: Image.Image, content: str) -> str: blocks = parse_chandra_html(content) return _interactive_viewer_html(image, blocks, title="") def preview_uploaded(file_path: str | None, page_number: int): """Render the uploaded image (or selected PDF page) without running the model.""" if not file_path: return [], "", "", [], "" try: page_num, image = _load_image(file_path, int(page_number or 1)) except gr.Error: raise except Exception as e: raise gr.Error(f"Failed to load file: {e}") from e preview = _preview_html(image, "") status = ( f"Loaded `{Path(file_path).name}` · page `{page_num}` · " f"`{image.width}x{image.height}` — click **Run OCR** to predict." ) meta = [ { "file": Path(file_path).name, "page": page_num, "image_size": f"{image.width}x{image.height}", } ] return [image], "", preview, meta, status def run_ocr( file_path: str | None, page_number: int, temperature: float, progress: gr.Progress = gr.Progress(track_tqdm=False), ): if not file_path: raise gr.Error("Please upload an image or PDF.") progress(0.0, desc="Preparing page") page_num, image = _load_image(file_path, int(page_number)) progress(0.3, desc=f"Running OCR on page {page_num}") text, token_count, elapsed = _generate_page( image, temperature=temperature, ) progress(1.0, desc="Done") metadata = [ { "page": page_num, "tokens": token_count, "seconds": round(elapsed, 2), "image_size": f"{image.width}x{image.height}", } ] status = ( f"Model: `{MODEL_ID}` \n" f"Page: `{page_num}` / Tokens: `{token_count}` / Seconds: `{round(elapsed, 2)}`" ) return ( [image], text, _preview_html(image, text), metadata, status, ) CSS = """ .gradio-container { max-width: 100% !important; padding: 8px 16px !important; } .app-title { margin: 0 0 4px !important; font-size: 18px !important; } .app-sub { margin: 0 0 12px !important; font-size: 12px !important; color: #666 !important; } .viewer-host > .label-wrap { display: none !important; } .viewer-host .gradio-html { padding: 0 !important; } .input-pane { border-right: 1px solid #d0d7de; padding-right: 12px !important; } .section-h { margin: 4px 0 4px !important; font-size: 13px !important; font-weight: 600 !important; color: #444 !important; } """ with gr.Blocks(title="Open Chandra OCR", css=CSS) as demo: gr.Markdown("# Open Chandra OCR", elem_classes=["app-title"]) gr.Markdown( "Upload an image or PDF and run OCR. Predicted blocks are shown in the right pane.", elem_classes=["app-sub"], ) with gr.Row(): with gr.Column(scale=1, min_width=300, elem_classes=["input-pane"]): input_file = gr.File( label="Image / PDF", file_types=["image", ".pdf"], type="filepath", ) page_number = gr.Number( label="PDF page number", value=1, precision=0, minimum=1, info="Used only for PDF input (ignored for images).", ) temperature = gr.Slider( 0.0, 1.0, value=0.0, step=0.05, label="Temperature" ) run_button = gr.Button("Run OCR", variant="primary") status = gr.Markdown() with gr.Accordion("Raw output / metadata / pages", open=False): with gr.Tabs(): with gr.Tab("Generated HTML"): html_code = gr.Code( label="Generated Chandra HTML", language="html", lines=18, ) with gr.Tab("Input pages"): gallery = gr.Gallery( label="Input pages", columns=1, height=260, object_fit="contain", ) with gr.Tab("Metadata"): metadata = gr.JSON(label="Run metadata") with gr.Column(scale=4): preview = gr.HTML(elem_classes=["viewer-host"]) input_file.change( preview_uploaded, inputs=[input_file, page_number], outputs=[gallery, html_code, preview, metadata, status], ) page_number.change( preview_uploaded, inputs=[input_file, page_number], outputs=[gallery, html_code, preview, metadata, status], ) run_button.click( run_ocr, inputs=[input_file, page_number, temperature], outputs=[gallery, html_code, preview, metadata, status], ) if __name__ == "__main__": demo.queue(default_concurrency_limit=int(os.getenv("CONCURRENCY_LIMIT", "1"))) demo.launch()