"""Run Jina-OCR-v1 with Transformers or vLLM.

    python example.py --backend transformers --image document.png
    python example.py --backend transformers --device cpu --image document.png
    python example.py --backend transformers --device cuda:1 --image document.png
    python example.py --backend vllm --image document.png
"""

from __future__ import annotations

import argparse
from pathlib import Path

from PIL import Image


def run_transformers(
    model_id: str,
    image: Image.Image,
    max_new_tokens: int,
    device: str | None = None,
) -> str:
    import torch
    from transformers import AutoModelForCausalLM, AutoProcessor

    if device is None:
        device = "cuda" if torch.cuda.is_available() else "cpu"
    torch_device = torch.device(device)
    if torch_device.type == "cuda" and not torch.cuda.is_available():
        raise RuntimeError(f"CUDA is not available (requested --device {device})")
    processor = AutoProcessor.from_pretrained(model_id, trust_remote_code=True)
    model = AutoModelForCausalLM.from_pretrained(
        model_id,
        dtype=torch.bfloat16,
        trust_remote_code=True,
    ).to(torch_device)
    inputs = processor.prepare_ocr_inputs(image, device=torch_device)
    # Do not pass a fresh GenerationConfig — that drops eos/pad from the model card.
    output = model.generate(
        **inputs,
        max_new_tokens=max_new_tokens,
        do_sample=False,
    )

    return processor.decode_ocr(output, inputs["input_ids"])


def run_vllm(
    model_id: str,
    image: Image.Image,
    max_new_tokens: int,
    num_speculative_tokens: int,
) -> str:
    from deepseek_ocr_mtp import DEFAULT_OCR_PROMPT, log_spec_stats, register, vllm_llm_kwargs, vllm_sampling_params
    from vllm import LLM

    register()
    llm = LLM(
        **vllm_llm_kwargs(
            model_id,
            num_speculative_tokens=num_speculative_tokens,
            mtp_heads=1,
            mtp_recursive=True,
        )
    )
    messages = [
        {
            "role": "user",
            "content": [
                {"type": "image_pil", "image_pil": image},
                {"type": "text", "text": DEFAULT_OCR_PROMPT},
            ],
        }
    ]
    outputs = llm.chat(
        messages,
        sampling_params=vllm_sampling_params(max_tokens=max_new_tokens),
    )
    if num_speculative_tokens > 0:
        log_spec_stats(llm)
    return outputs[0].outputs[0].text


def main() -> None:
    parser = argparse.ArgumentParser(description="Jina-OCR-v1 local inference")
    parser.add_argument("--backend", choices=("transformers", "vllm"), default="transformers")
    parser.add_argument("--image", type=Path, default=Path("document.png"))
    parser.add_argument("--model", default="jinaai/jina-ocr-v1")
    parser.add_argument("--max-new-tokens", type=int, default=4096, help="Maximum number of new tokens to generate")
    parser.add_argument("--num-speculative-tokens", type=int, default=3, help="vLLM FastMTP K; 0 disables")
    parser.add_argument(
        "--device",
        default=None,
        help="Transformers device, e.g. cpu, cuda, cuda:0. Default: cuda if available else cpu",
    )
    args = parser.parse_args()
    if args.device is not None and args.backend != "transformers":
        parser.error("--device is only supported with --backend transformers")

    if not args.image.is_file():
        raise FileNotFoundError(f"Image not found: {args.image}")

    image = Image.open(args.image).convert("RGB")

    print(f"Running {args.backend} inference with model {args.model}...")
    if args.backend == "transformers":
        text = run_transformers(args.model, image, args.max_new_tokens, args.device)
    else:
        text = run_vllm(args.model, image, args.max_new_tokens, args.num_speculative_tokens)
    print(text)


if __name__ == "__main__":
    main()
