armorocr / app.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
777d129 verified
Raw History Blame Contribute Delete
12.2 kB
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # MUST come before torch / transformers
import math
import re
import time
import gradio as gr
import torch
from PIL import Image, ImageDraw, ImageFont
from transformers import AutoProcessor, Qwen3VLForConditionalGeneration
MODEL_ID = "inclusionAI/ArmorOCR"
DEFAULT_TASK_PROMPT = (
"Please identify the text in the image. Put your reasoning process inside "
"<analyze></analyze> and your final recognized text inside <answer></answer>."
)
SPOTTING_PROMPT = (
'Please identify the text in the image and output the answer in JSON format as '
'{"text": "answer content", "bbox": [x_min, y_min, x_max, y_max]}, where the bbox '
"coordinates are in the range 0-1000 relative to the image. Do not output any extra "
"explanation. Put your reasoning process inside <analyze></analyze> and your final "
'JSON answer inside <answer></answer>.'
)
MODE_SPOTTING = "Grounded spotting (with bbox)"
MODE_FREEFORM = "Free-form prompt"
# ---------------------------------------------------------------- model load
print("Loading ArmorOCR...", flush=True)
t0 = time.perf_counter()
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = Qwen3VLForConditionalGeneration.from_pretrained(
MODEL_ID,
dtype=torch.bfloat16,
attn_implementation="sdpa",
).to("cuda")
model.eval()
print(f"ArmorOCR loaded in {time.perf_counter() - t0:.1f}s", flush=True)
# ----------------------------------------------------------------- helpers
def _load_rgb(image):
"""Open a Gradio image input as an RGB PIL image."""
if isinstance(image, str):
image = Image.open(image)
if image.mode == "RGBA":
bg = Image.new("RGB", image.size, "white")
bg.paste(image, mask=image.getchannel("A"))
return bg
if image.mode != "RGB":
return image.convert("RGB")
return image
def _smart_resize(image, patch_size=16):
"""Resize so H and W are multiples of patch_size*2 (Qwen3-VL convention),
capping the token count so big uploads stay fast."""
factor = patch_size * 2
w, h = image.size
min_side, max_side = min(w, h), max(w, h)
if min_side <= 0 or max_side / min_side > 200:
raise ValueError(f"invalid image size {w}x{h}")
nh, nw = max(factor, round(h / factor) * factor), max(factor, round(w / factor) * factor)
max_pixels = 4096 * factor * factor # image-max-token-num 4096 (eval-script default)
if nh * nw > max_pixels:
scale = math.sqrt(h * w / max_pixels)
nh = max(factor, math.floor(h / scale / factor) * factor)
nw = max(factor, math.floor(w / scale / factor) * factor)
if (nw, nh) != (w, h):
image = image.resize((nw, nh))
return image
def _parse_answer(text):
"""Split raw output into (reasoning, final answer), tolerating truncated tags."""
text = text.strip()
m = re.search(r"<analyze>(.*?)(?:</analyze>|$)", text, flags=re.DOTALL)
reasoning = m.group(1).strip() if m else ""
m2 = re.search(r"<answer>(.*?)(?:</answer>|$)", text, flags=re.DOTALL)
if m2:
answer = m2.group(1).strip()
elif reasoning:
# <analyze> present but truncated before <answer>: answer unknown yet.
rest = re.sub(r"<analyze>.*?(?:</analyze>|$)", "", text, flags=re.DOTALL).strip()
answer = rest
else:
# No tags at all: treat the whole output as the answer.
answer = text
return reasoning, answer
_BBOX_RE = re.compile(
r"\[\s*(\d+(?:\.\d+)?)\s*,\s*(\d+(?:\.\d+)?)\s*,\s*(\d+(?:\.\d+)?)\s*,\s*(\d+(?:\.\d+)?)\s*\]"
)
def _extract_bboxes(answer):
"""Find every [x1,y1,x2,y2] group in the answer text (0-1000 coords)."""
return [[float(v) for v in m.groups()] for m in _BBOX_RE.finditer(answer)]
def _draw_boxes(image, boxes):
"""Draw detected boxes (0-1000 normalized) on a PIL copy of the image."""
canvas = image.copy()
draw = ImageDraw.Draw(canvas)
w, h = canvas.size
try:
font = ImageFont.truetype(
"/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf",
size=max(14, min(w, h) // 24),
)
except OSError:
font = ImageFont.load_default()
for i, box in enumerate(boxes):
x1 = max(0.0, min(1000.0, box[0])) / 1000.0 * w
y1 = max(0.0, min(1000.0, box[1])) / 1000.0 * h
x2 = max(0.0, min(1000.0, box[2])) / 1000.0 * w
y2 = max(0.0, min(1000.0, box[3])) / 1000.0 * h
xa, ya, xb, yb = min(x1, x2), min(y1, y2), max(x1, x2), max(y1, y2)
color = (255, 72, 0)
lw = max(2, int(min(w, h) / 300))
draw.rectangle([xa, ya, xb, yb], outline=color, width=lw)
label = str(i + 1)
tb = draw.textbbox((0, 0), label, font=font)
tw, th = tb[2] - tb[0], tb[3] - tb[1]
pad = max(2, lw)
ly = ya - th - 2 * pad
if ly < 0:
ly = ya
draw.rectangle([xa, ly, xa + tw + 2 * pad, ly + th + 2 * pad], fill=color)
draw.text((xa + pad, ly + pad), label, fill=(255, 255, 255), font=font)
return canvas
# ------------------------------------------------------------------ inference
@spaces.GPU(duration=45)
def run_ocr(image, mode, prompt, show_boxes, max_new_tokens=512):
"""Recognize adversarial / hard-to-read text in an image with ArmorOCR.
Args:
image: input image.
mode: "Grounded spotting (with bbox)" to also localize the text with a
0-1000 bounding box, or "Free-form prompt" to ask a custom question.
prompt: custom instruction (only used in free-form mode; empty uses the
default OCR instruction).
show_boxes: draw detected bounding boxes on the image (spotted mode).
max_new_tokens: maximum number of new tokens to generate.
Returns:
(recognized text, perception reasoning, annotated image, status line)
"""
if image is None:
raise gr.Error("Please upload an image first.")
pil_image = _smart_resize(_load_rgb(image))
if mode == MODE_SPOTTING:
user_text = SPOTTING_PROMPT
else:
user_text = (prompt or "").strip() or DEFAULT_TASK_PROMPT
messages = [
{
"role": "user",
"content": [
{"type": "image", "image": pil_image},
{"type": "text", "text": user_text},
],
}
]
inputs = processor.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
t_start = time.perf_counter()
with torch.inference_mode():
out = model.generate(
**inputs,
max_new_tokens=int(max_new_tokens),
do_sample=False,
)
elapsed = time.perf_counter() - t_start
trimmed = [o[len(i):] for i, o in zip(inputs.input_ids, out)]
raw = processor.batch_decode(trimmed, skip_special_tokens=True)[0]
reasoning, answer = _parse_answer(raw)
boxes = _extract_bboxes(answer)
annotated = _draw_boxes(pil_image, boxes) if (mode == MODE_SPOTTING and show_boxes and boxes) else pil_image
status = f"Done in {elapsed:.1f}s"
return answer, reasoning, annotated, status
# ----------------------------------------------------------------------- UI
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks(title="ArmorOCR") as demo:
gr.Markdown(
"""
# 🛡️ ArmorOCR: Grounded Adversarial OCR
**ArmorOCR** ([inclusionAI/ArmorOCR](https://huggingface.co/inclusionAI/ArmorOCR)) is a
Qwen3-VL-8B-Instruct finetune specialized in **grounded adversarial OCR perception** —
reading text that hides in images: low-contrast blends, dot/line encodings, dense
patterns, text near QR codes, and other adversarial visual text.
Single-pass inference on the original image; no inference-time visual transformations.
📄 [Paper](https://arxiv.org/abs/2608.20122) · 💻 [Code](https://github.com/ant-research/ArmorOCR) · 🤗 [AdvSpot benchmark](https://huggingface.co/datasets/inclusionAI/advspot-public)
"""
)
with gr.Column(elem_id="col-container"):
with gr.Row():
with gr.Column():
image_in = gr.Image(
type="pil", label="Input image", height=340, sources=["upload", "clipboard"]
)
mode = gr.Radio(
[MODE_SPOTTING, MODE_FREEFORM],
value=MODE_SPOTTING,
label="Task",
info=(
"Grounded spotting: recognize the text AND localize it with a "
"0-1000 bounding box. Free-form prompt: ask your own question "
"about the text in the image."
),
)
prompt = gr.Textbox(
label="Prompt (free-form mode)",
placeholder=DEFAULT_TASK_PROMPT,
lines=3,
info="Used when Task = Free-form prompt. Empty uses the default OCR instruction.",
)
show_boxes = gr.Checkbox(
value=True, label="Draw detected boxes on the image (spotted mode)"
)
run_btn = gr.Button("Run OCR", variant="primary")
with gr.Accordion("Advanced settings", open=False):
max_new_tokens = gr.Slider(
minimum=64, maximum=2048, value=512, step=64, label="Max new tokens"
)
with gr.Column():
answer_out = gr.Textbox(label="Recognized text", lines=3)
status_out = gr.Textbox(label="Status", lines=1, interactive=False)
annotated_out = gr.Image(
label="Grounded regions", height=340, type="pil", interactive=False
)
reasoning_out = gr.Textbox(
label="Perception reasoning (<analyze>)", lines=10, interactive=False
)
gr.Examples(
examples=[
["examples/example_1.png", MODE_SPOTTING, "", True],
["examples/example_2.png", MODE_FREEFORM,
"What is the faintly visible text in the central region of the pattern "
"formed by dense vertical black lines? Please put your reasoning process "
"inside <analyze></analyze> and your final recognized text inside "
"<answer></answer>.",
True],
["examples/example_3.png", MODE_FREEFORM,
"What text appears beneath the dot pattern in gray? Please put your "
"reasoning process inside <analyze></analyze> and your final recognized "
"text inside <answer></answer>.",
True],
["examples/example_4.jpeg", MODE_FREEFORM,
"What is the white text above the QR code? Please put your reasoning "
"process inside <analyze></analyze> and your final recognized text "
"inside <answer></answer>.",
True],
],
inputs=[image_in, mode, prompt, show_boxes],
outputs=[answer_out, reasoning_out, annotated_out, status_out],
fn=run_ocr,
cache_examples=True,
cache_mode="lazy",
)
gr.Markdown(
"""
<sub>Example images come from the official
[ant-research/ArmorOCR](https://github.com/ant-research/ArmorOCR) repository
(Apache-2.0). ArmorOCR is released under the Apache License 2.0.</sub>
"""
)
run_btn.click(
run_ocr,
inputs=[image_in, mode, prompt, show_boxes, max_new_tokens],
outputs=[answer_out, reasoning_out, annotated_out, status_out],
api_name="run_ocr",
)
if __name__ == "__main__":
demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)