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 " " and your final recognized text inside ." ) 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 and your final " 'JSON answer inside .' ) 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"(.*?)(?:|$)", text, flags=re.DOTALL) reasoning = m.group(1).strip() if m else "" m2 = re.search(r"(.*?)(?:|$)", text, flags=re.DOTALL) if m2: answer = m2.group(1).strip() elif reasoning: # present but truncated before : answer unknown yet. rest = re.sub(r".*?(?:|$)", "", 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 ()", 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 and your final recognized text inside " ".", True], ["examples/example_3.png", MODE_FREEFORM, "What text appears beneath the dot pattern in gray? Please put your " "reasoning process inside and your final recognized " "text inside .", True], ["examples/example_4.jpeg", MODE_FREEFORM, "What is the white text above the QR code? Please put your reasoning " "process inside and your final recognized text " "inside .", 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( """ 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. """ ) 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)