Team Ai
Datasetpublic

uv-scripts/ocr

OCR UV Scripts Part of uv-scripts: self-contained UV scripts you run on Hugging Face Jobs in one command. One script per OCR model. Each script runs the model on a GPU with Hugging Face Jobs and writes the text as markdown: as a new column in a Hub dataset, as .md files in a Bucket, or as resumable parquet parts (the -saturate recipes). A few scripts return JSON from a schema, detect layout regions, or compare the output of two models. Quick Start First… See the full description on the dataset page: https://huggingface.co/datasets/uv-scripts/ocr.

sourceHugging Faceupdated 10d agoView on Hugging Face
163likes6.5kdownloads
glm-ocr.py693 linesDownload Raw Back to root
1# /// script2# requires-python = ">=3.11"3# dependencies = [4#     "datasets>=3.1.0",5#     "huggingface-hub",6#     "pillow",7#     "toolz",8# ]9#10# [tool.hf-jobs]11# image = "vllm/vllm-openai:v0.29.0"12# python = "/usr/bin/python3"13# env = { PYTHONPATH = "/usr/local/lib/python3.12/dist-packages" }14# flavor = "a10g-small"15# secrets = ["HF_TOKEN"]16# ///17 18"""19Convert document images to markdown using GLM-OCR with vLLM.20 21GLM-OCR is a compact 0.9B parameter OCR model achieving 94.62% on OmniDocBench V1.5.22Uses CogViT visual encoder with GLM-0.5B language decoder and Multi-Token Prediction23(MTP) loss for fast, accurate document parsing.24 25NOTE: On Jobs, vLLM, torch and transformers come from the vllm/vllm-openai:v0.29.026image pinned in the [tool.hf-jobs] header (`hf` CLI 1.32+). GLM-OCR needs vLLM27>=0.16.0 (PR #33005) and transformers>=5.1.0; the image has both. To run on your28own GPU: `uv run --with vllm==0.29.0 glm-ocr.py ...`.29 30Features:31- 0.9B parameters (ultra-compact)32- 94.62% on OmniDocBench V1.5 (SOTA for sub-1B models)33- Text recognition with markdown output34- LaTeX formula recognition35- Table extraction (HTML format)36- Multilingual: zh, en, fr, es, ru, de, ja, ko37- MIT licensed38 39Model: zai-org/GLM-OCR40vLLM: vllm/vllm-openai:v0.29.0 image (see the header)41Performance: 94.62% on OmniDocBench V1.542"""43 44import argparse45import base6446import io47import json48import logging49import os50import sys51import time52from datetime import datetime53from typing import Any, Dict, List, Optional, Union54 55import torch56from datasets import load_dataset57from huggingface_hub import DatasetCard, login58from PIL import Image59from toolz import partition_all60# Disable vLLM's FlashInfer sampler: it JIT-compiles a CUDA kernel needing nvcc, which the61# default uv-script image lacks (engine init then crashes). Greedy OCR doesn't use it; this62# lets the plain default-image command work. On the vllm/vllm-openai image it's a harmless no-op.63os.environ.setdefault("VLLM_USE_FLASHINFER_SAMPLER", "0")64# Same story for DeepGEMM: its init calls _find_cuda_home, which asserts on the65# nvcc-less base image (a non-fatal warning that clutters the log and hides the real traceback).66# Greedy OCR doesn't need the DeepGEMM JIT path, so disable it explicitly.67os.environ.setdefault("VLLM_USE_DEEP_GEMM", "0")68from vllm import LLM, SamplingParams69 70logging.basicConfig(level=logging.INFO)71logger = logging.getLogger(__name__)72 73MODEL = "zai-org/GLM-OCR"74 75# Task prompts as specified by the model76TASK_PROMPTS = {77    "ocr": "Text Recognition:",78    "formula": "Formula Recognition:",79    "table": "Table Recognition:",80}81 82 83def check_cuda_availability():84    """Check if CUDA is available and exit if not."""85    if not torch.cuda.is_available():86        logger.error("CUDA is not available. This script requires a GPU.")87        logger.error("Please run on a machine with a CUDA-capable GPU.")88        sys.exit(1)89    else:90        logger.info(f"CUDA is available. GPU: {torch.cuda.get_device_name(0)}")91 92 93def ensure_output_columns_free(dataset, columns, overwrite=False):94    """Fail fast if an output column would collide with an existing input column.95 96    Adding a column that already exists silently overwrites it (e.g. a ground-truth97    `text`/`markdown` column) or crashes on push with a duplicate-column error only98    *after* inference has run. Catch it up front. With overwrite=True, drop the clashing99    column(s) here instead (logged) so the later add_column is clean.100    """101    clash = [c for c in columns if c in dataset.column_names]102    if not clash:103        return dataset104    if overwrite:105        logger.warning(f"--overwrite: replacing existing column(s) {clash}")106        return dataset.remove_columns(clash)107    logger.error(108        f"Output column(s) {clash} already exist in the input dataset "109        f"(columns: {dataset.column_names})."110    )111    logger.error("Choose a different --output-column, or pass --overwrite to replace them.")112    sys.exit(1)113 114 115def downscale_to_max_pixels(img: Image.Image, max_pixels: Optional[int]) -> Image.Image:116    """Shrink an image so width*height <= max_pixels, preserving aspect ratio.117 118    GLM-OCR does no internal resizing and its card gives no resolution guidance. Capping119    input pixels bounds both image tokens and vision-encoder memory, a safety valve for very120    large (multi-MP) scans that can pressure GPU memory at high batch sizes. No-op when121    max_pixels is None or the image is already small enough (never upscales)."""122    if not max_pixels:123        return img124    w, h = img.size125    if w * h <= max_pixels:126        return img127    scale = (max_pixels / (w * h)) ** 0.5128    new_size = (max(1, int(w * scale)), max(1, int(h * scale)))129    return img.resize(new_size, Image.Resampling.LANCZOS)130 131 132def make_ocr_message(133    image: Union[Image.Image, Dict[str, Any], str],134    task: str = "ocr",135    max_pixels: Optional[int] = None,136) -> List[Dict]:137    """138    Create chat message for OCR processing.139 140    GLM-OCR uses a chat format with an image and a task prompt prefix.141    Supported tasks: ocr, formula, table.142    """143    # Convert to PIL Image if needed144    if isinstance(image, Image.Image):145        pil_img = image146    elif isinstance(image, dict) and "bytes" in image:147        pil_img = Image.open(io.BytesIO(image["bytes"]))148    elif isinstance(image, str):149        pil_img = Image.open(image)150    else:151        raise ValueError(f"Unsupported image type: {type(image)}")152 153    # Convert to RGB154    pil_img = pil_img.convert("RGB")155 156    # Optionally cap resolution to protect the vision encoder from OOM on huge scans157    pil_img = downscale_to_max_pixels(pil_img, max_pixels)158 159    # Convert to base64 data URI160    buf = io.BytesIO()161    pil_img.save(buf, format="PNG")162    data_uri = f"data:image/png;base64,{base64.b64encode(buf.getvalue()).decode()}"163 164    prompt_text = TASK_PROMPTS.get(task, TASK_PROMPTS["ocr"])165 166    return [167        {168            "role": "user",169            "content": [170                {"type": "image_url", "image_url": {"url": data_uri}},171                {"type": "text", "text": prompt_text},172            ],173        }174    ]175 176 177def create_dataset_card(178    source_dataset: str,179    model: str,180    num_samples: int,181    processing_time: str,182    batch_size: int,183    max_model_len: int,184    max_tokens: int,185    gpu_memory_utilization: float,186    temperature: float,187    top_p: float,188    task: str,189    image_column: str = "image",190    split: str = "train",191) -> str:192    """Create a dataset card documenting the OCR process."""193    model_name = model.split("/")[-1]194    task_desc = {195        "ocr": "text recognition",196        "formula": "formula recognition",197        "table": "table recognition",198    }199 200    # Canonical provenance stamp (see AGENTS.md): Jobs claim gated on JOB_ID, set by HF Jobs in-container.201    on_jobs = os.environ.get("JOB_ID") is not None202    hw = os.environ.get("ACCELERATOR") or ""203    origin = (204        "Produced on [Hugging Face Jobs](https://huggingface.co/docs/huggingface_hub/guides/jobs)"205        + (f" (`{hw}`)" if hw else "")206    ) if on_jobs else "Generated"207    jobs_tag = "\n- hf-jobs" if on_jobs else ""208 209    return f"""---210tags:211- ocr212- document-processing213- glm-ocr214- markdown215- uv-script216- generated{jobs_tag}217---218 219# Document OCR using {model_name}220 221This dataset contains OCR results from images in [{source_dataset}](https://huggingface.co/datasets/{source_dataset}) using GLM-OCR, a compact 0.9B OCR model achieving SOTA performance.222 223## Processing Details224 225- **Source Dataset**: [{source_dataset}](https://huggingface.co/datasets/{source_dataset})226- **Model**: [{model}](https://huggingface.co/{model})227- **Task**: {task_desc.get(task, task)}228- **Number of Samples**: {num_samples:,}229- **Processing Time**: {processing_time}230- **Processing Date**: {datetime.now().strftime("%Y-%m-%d %H:%M UTC")}231 232### Configuration233 234- **Image Column**: `{image_column}`235- **Output Column**: `markdown`236- **Dataset Split**: `{split}`237- **Batch Size**: {batch_size}238- **Max Model Length**: {max_model_len:,} tokens239- **Max Output Tokens**: {max_tokens:,}240- **Temperature**: {temperature}241- **Top P**: {top_p}242- **GPU Memory Utilization**: {gpu_memory_utilization:.1%}243 244## Model Information245 246GLM-OCR is a compact, high-performance OCR model:247- 0.9B parameters248- 94.62% on OmniDocBench V1.5249- CogViT visual encoder + GLM-0.5B language decoder250- Multi-Token Prediction (MTP) loss for efficiency251- Multilingual: zh, en, fr, es, ru, de, ja, ko252- MIT licensed253 254## Dataset Structure255 256The dataset contains all original columns plus:257- `markdown`: The extracted text in markdown format258- `inference_info`: JSON list tracking all OCR models applied to this dataset259 260## Reproduction261 262{origin} with the [`glm-ocr.py`](https://huggingface.co/datasets/uv-scripts/ocr/raw/main/glm-ocr.py) recipe from [uv-scripts](https://huggingface.co/uv-scripts). Run it yourself:263 264```bash265hf jobs uv run https://huggingface.co/datasets/uv-scripts/ocr/raw/main/glm-ocr.py \\266    {source_dataset} \\267    <output-dataset> \\268    --image-column {image_column} \\269    --batch-size {batch_size} \\270    --task {task}271```272"""273 274 275def main(276    input_dataset: str,277    output_dataset: str,278    image_column: str = "image",279    batch_size: int = 16,280    max_model_len: int = 8192,281    max_pixels: Optional[int] = None,282    max_tokens: int = 8192,283    temperature: float = 0.01,284    top_p: float = 0.00001,285    repetition_penalty: float = 1.1,286    gpu_memory_utilization: float = 0.8,287    task: str = "ocr",288    hf_token: str = None,289    split: str = "train",290    max_samples: int = None,291    private: bool = False,292    shuffle: bool = False,293    seed: int = 42,294    output_column: str = "markdown",295    overwrite: bool = False,296    verbose: bool = False,297    config: str = None,298    create_pr: bool = False,299):300    """Process images from HF dataset through GLM-OCR model."""301 302    check_cuda_availability()303 304    start_time = datetime.now()305 306    HF_TOKEN = hf_token or os.environ.get("HF_TOKEN")307    if HF_TOKEN:308        login(token=HF_TOKEN)309 310    # Validate task311    if task not in TASK_PROMPTS:312        logger.error(f"Unknown task '{task}'. Supported: {list(TASK_PROMPTS.keys())}")313        sys.exit(1)314 315    logger.info(f"Using model: {MODEL}")316    logger.info(f"Task: {task} (prompt: '{TASK_PROMPTS[task]}')")317 318    # Load dataset319    logger.info(f"Loading dataset: {input_dataset}")320    dataset = load_dataset(input_dataset, split=split)321 322    if image_column not in dataset.column_names:323        raise ValueError(324            f"Column '{image_column}' not found. Available: {dataset.column_names}"325        )326 327    # Fail fast if the output column would collide with an existing input column328    dataset = ensure_output_columns_free(dataset, [output_column], overwrite=overwrite)329 330    if shuffle:331        logger.info(f"Shuffling dataset with seed {seed}")332        dataset = dataset.shuffle(seed=seed)333 334    if max_samples:335        dataset = dataset.select(range(min(max_samples, len(dataset))))336        logger.info(f"Limited to {len(dataset)} samples")337 338    # Initialize vLLM339    logger.info("Initializing vLLM with GLM-OCR")340    logger.info("This may take a few minutes on first run...")341    llm = LLM(342        model=MODEL,343        trust_remote_code=True,344        max_model_len=max_model_len,345        gpu_memory_utilization=gpu_memory_utilization,346        limit_mm_per_prompt={"image": 1},347    )348 349    # Sampling defaults from GLM-OCR SDK (github.com/zai-org/GLM-OCR)350    # glmocr/config.py PageLoaderConfig: temperature=0.01, top_p=0.00001,351    # top_k=1, repetition_penalty=1.1, max_tokens=16384352    # generation_config.json on HF also sets do_sample=false (greedy)353    # Note: SDK uses max_tokens=16384 but vLLM caps at max_model_len (8192)354    sampling_params = SamplingParams(355        temperature=temperature,356        top_p=top_p,357        max_tokens=max_tokens,358        repetition_penalty=repetition_penalty,359    )360 361    logger.info(f"Processing {len(dataset)} images in batches of {batch_size}")362    logger.info(f"Output will be written to column: {output_column}")363 364    all_outputs = []365    total_batches = (len(dataset) + batch_size - 1) // batch_size366    processed = 0367 368    for batch_num, batch_indices in enumerate(369        partition_all(batch_size, range(len(dataset))), 1370    ):371        batch_indices = list(batch_indices)372        batch_images = [dataset[i][image_column] for i in batch_indices]373 374        logger.info(375            f"Batch {batch_num}/{total_batches} "376            f"({processed}/{len(dataset)} images done)"377        )378 379        try:380            batch_messages = [381                make_ocr_message(img, task=task, max_pixels=max_pixels)382                for img in batch_images383            ]384 385            outputs = llm.chat(batch_messages, sampling_params)386 387            for output in outputs:388                text = output.outputs[0].text.strip()389                all_outputs.append(text)390 391            processed += len(batch_images)392 393        except Exception as e:394            logger.error(f"Error processing batch: {e}")395            all_outputs.extend(["[OCR ERROR]"] * len(batch_images))396            processed += len(batch_images)397 398    processing_duration = datetime.now() - start_time399    processing_time_str = f"{processing_duration.total_seconds() / 60:.1f} min"400 401    logger.info(f"Adding '{output_column}' column to dataset")402    dataset = dataset.add_column(output_column, all_outputs)403 404    # Inference info tracking405    inference_entry = {406        "model_id": MODEL,407        "model_name": "GLM-OCR",408        "column_name": output_column,409        "timestamp": datetime.now().isoformat(),410        "task": task,411        "temperature": temperature,412        "top_p": top_p,413        "repetition_penalty": repetition_penalty,414        "max_tokens": max_tokens,415    }416 417    if "inference_info" in dataset.column_names:418        logger.info("Updating existing inference_info column")419 420        def update_inference_info(example):421            try:422                existing_info = (423                    json.loads(example["inference_info"])424                    if example["inference_info"]425                    else []426                )427            except (json.JSONDecodeError, TypeError):428                existing_info = []429            existing_info.append(inference_entry)430            return {"inference_info": json.dumps(existing_info)}431 432        dataset = dataset.map(update_inference_info)433    else:434        logger.info("Creating new inference_info column")435        inference_list = [json.dumps([inference_entry])] * len(dataset)436        dataset = dataset.add_column("inference_info", inference_list)437 438    # Push to hub with retry and XET fallback439    logger.info(f"Pushing to {output_dataset}")440    max_retries = 3441    for attempt in range(1, max_retries + 1):442        try:443            if attempt > 1:444                logger.warning("Disabling XET (fallback to HTTP upload)")445                os.environ["HF_HUB_DISABLE_XET"] = "1"446            dataset.push_to_hub(447                output_dataset,448                private=private,449                token=HF_TOKEN,450                max_shard_size="500MB",451                **({"config_name": config} if config else {}),452                create_pr=create_pr,453                commit_message=f"Add {MODEL} OCR results ({len(dataset)} samples)"454                + (f" [{config}]" if config else ""),455            )456            break457        except Exception as e:458            logger.error(f"Upload attempt {attempt}/{max_retries} failed: {e}")459            if attempt < max_retries:460                delay = 30 * (2 ** (attempt - 1))461                logger.info(f"Retrying in {delay}s...")462                time.sleep(delay)463            else:464                logger.error("All upload attempts failed. OCR results are lost.")465                sys.exit(1)466 467    # Create and push dataset card468    logger.info("Creating dataset card")469    card_content = create_dataset_card(470        source_dataset=input_dataset,471        model=MODEL,472        num_samples=len(dataset),473        processing_time=processing_time_str,474        batch_size=batch_size,475        max_model_len=max_model_len,476        max_tokens=max_tokens,477        gpu_memory_utilization=gpu_memory_utilization,478        temperature=temperature,479        top_p=top_p,480        task=task,481        image_column=image_column,482        split=split,483    )484 485    card = DatasetCard(card_content)486    card.push_to_hub(output_dataset, token=HF_TOKEN)487 488    logger.info("Done! GLM-OCR processing complete.")489    logger.info(490        f"Dataset available at: https://huggingface.co/datasets/{output_dataset}"491    )492    logger.info(f"Processing time: {processing_time_str}")493    logger.info(494        f"Processing speed: {len(dataset) / processing_duration.total_seconds():.2f} images/sec"495    )496 497    if verbose:498        import importlib.metadata499 500        logger.info("--- Resolved package versions ---")501        for pkg in ["vllm", "transformers", "torch", "datasets", "pyarrow", "pillow"]:502            try:503                logger.info(f"  {pkg}=={importlib.metadata.version(pkg)}")504            except importlib.metadata.PackageNotFoundError:505                logger.info(f"  {pkg}: not installed")506        logger.info("--- End versions ---")507 508 509if __name__ == "__main__":510    if len(sys.argv) == 1:511        print("=" * 70)512        print("GLM-OCR Document Processing")513        print("=" * 70)514        print("\n0.9B OCR model - 94.62% on OmniDocBench V1.5")515        print("\nTask modes:")516        print("  ocr      - Text recognition (default)")517        print("  formula  - LaTeX formula recognition")518        print("  table    - Table extraction")519        print("\nExamples:")520        print("\n1. Basic OCR:")521        print("   uv run --with vllm==0.29.0 glm-ocr.py input-dataset output-dataset")522        print("\n2. Formula recognition:")523        print("   uv run --with vllm==0.29.0 glm-ocr.py docs results --task formula")524        print("\n3. Table extraction:")525        print("   uv run --with vllm==0.29.0 glm-ocr.py docs results --task table")526        print("\n4. Test with small sample:")527        print("   uv run --with vllm==0.29.0 glm-ocr.py large-dataset test --max-samples 10 --shuffle")528        print("\n5. Running on HF Jobs (hardware and HF_TOKEN come from the")529        print("   script's [tool.hf-jobs] header; --flavor/--timeout override it):")530        print(531            "   hf jobs uv run https://huggingface.co/datasets/uv-scripts/ocr/raw/main/glm-ocr.py \\"532        )533        print("       input-dataset output-dataset --batch-size 16")534        print("\nFor full help: uv run --with vllm==0.29.0 glm-ocr.py --help")535        sys.exit(0)536 537    parser = argparse.ArgumentParser(538        description="Document OCR using GLM-OCR (0.9B, 94.62% OmniDocBench V1.5)",539        formatter_class=argparse.RawDescriptionHelpFormatter,540        epilog="""541Task modes:542  ocr      Text recognition to markdown (default)543  formula  LaTeX formula recognition544  table    Table extraction545 546Examples:547  uv run --with vllm==0.29.0 glm-ocr.py my-docs analyzed-docs548  uv run --with vllm==0.29.0 glm-ocr.py docs results --task formula549  uv run --with vllm==0.29.0 glm-ocr.py large-dataset test --max-samples 50 --shuffle550        """,551    )552 553    parser.add_argument("input_dataset", help="Input dataset ID from Hugging Face Hub")554    parser.add_argument("output_dataset", help="Output dataset ID for Hugging Face Hub")555    parser.add_argument(556        "--image-column",557        default="image",558        help="Column containing images (default: image)",559    )560    parser.add_argument(561        "--batch-size",562        type=int,563        default=16,564        help="Batch size for processing (default: 16)",565    )566    parser.add_argument(567        "--max-model-len",568        type=int,569        default=8192,570        help="Maximum model context length (default: 8192)",571    )572    parser.add_argument(573        "--max-pixels",574        type=int,575        default=None,576        help=(577            "Optional cap on input image pixels (width*height); larger scans are "578            "downscaled (aspect preserved) before OCR. GLM-OCR does no internal resizing, "579            "so this bounds vision-encoder memory on very large scans — set e.g. 4000000 "580            "if you hit a GPU OOM at high batch sizes on a big-page corpus. Default: no cap."581        ),582    )583    parser.add_argument(584        "--max-tokens",585        type=int,586        default=8192,587        help="Maximum tokens to generate (default: 8192, capped by max-model-len)",588    )589    parser.add_argument(590        "--temperature",591        type=float,592        default=0.01,593        help="Sampling temperature (default: 0.01, near-greedy for OCR accuracy)",594    )595    parser.add_argument(596        "--top-p",597        type=float,598        default=0.00001,599        help="Top-p sampling parameter (default: 0.00001, near-greedy)",600    )601    parser.add_argument(602        "--repetition-penalty",603        type=float,604        default=1.1,605        help="Repetition penalty to prevent loops (default: 1.1)",606    )607    parser.add_argument(608        "--gpu-memory-utilization",609        type=float,610        default=0.8,611        help="GPU memory utilization (default: 0.8)",612    )613    parser.add_argument(614        "--task",615        choices=["ocr", "formula", "table"],616        default="ocr",617        help="OCR task mode (default: ocr)",618    )619    parser.add_argument("--hf-token", help="Hugging Face API token")620    parser.add_argument(621        "--split", default="train", help="Dataset split to use (default: train)"622    )623    parser.add_argument(624        "--max-samples",625        type=int,626        help="Maximum number of samples to process (for testing)",627    )628    parser.add_argument(629        "--private", action="store_true", help="Make output dataset private"630    )631    parser.add_argument(632        "--config",633        help="Config/subset name when pushing to Hub (for benchmarking multiple models in one repo)",634    )635    parser.add_argument(636        "--create-pr",637        action="store_true",638        help="Create a pull request instead of pushing directly (for parallel benchmarking)",639    )640    parser.add_argument(641        "--shuffle", action="store_true", help="Shuffle dataset before processing"642    )643    parser.add_argument(644        "--seed",645        type=int,646        default=42,647        help="Random seed for shuffling (default: 42)",648    )649    parser.add_argument(650        "--output-column",651        default="markdown",652        help="Column name for output text (default: markdown)",653    )654    parser.add_argument(655        "--overwrite",656        action="store_true",657        help="Replace the output column if it already exists in the input dataset "658        "(default: error out to avoid clobbering an existing column).",659    )660    parser.add_argument(661        "--verbose",662        action="store_true",663        help="Log resolved package versions after processing (useful for pinning deps)",664    )665 666    args = parser.parse_args()667 668    main(669        input_dataset=args.input_dataset,670        output_dataset=args.output_dataset,671        image_column=args.image_column,672        batch_size=args.batch_size,673        max_model_len=args.max_model_len,674        max_pixels=args.max_pixels,675        max_tokens=args.max_tokens,676        temperature=args.temperature,677        top_p=args.top_p,678        repetition_penalty=args.repetition_penalty,679        gpu_memory_utilization=args.gpu_memory_utilization,680        task=args.task,681        hf_token=args.hf_token,682        split=args.split,683        max_samples=args.max_samples,684        private=args.private,685        shuffle=args.shuffle,686        seed=args.seed,687        output_column=args.output_column,688        overwrite=args.overwrite,689        verbose=args.verbose,690        config=args.config,691        create_pr=args.create_pr,692    )693