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("", "<\\/")
VIEWER_HEIGHT_PX = int(os.getenv("VIEWER_HEIGHT_PX", "900"))
def _interactive_viewer_html(
image: Image.Image,
blocks: list[dict],
title: str = "",
) -> str:
header_inner = (
f'{html.escape(title)}
' if title else ""
)
header_block = (
f'{header_inner}
' if header_inner else ""
)
srcdoc = f"""
"""
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()