File size: 9,672 Bytes
f340984
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
from __future__ import annotations

import argparse
import json
import sys
import traceback
from pathlib import Path

from .constants import MODEL_REVISION, VERSION
from .errors import AppError, ModelIntegrityError
from .infer import RunOptions, default_output_path, run_inference
from .model_store import prepare_model, verify_model
from .paths import default_cache_dir, default_model_dir
from .runtime import inspect_runtime, runtime_issues


def bounded_int(minimum: int, maximum: int):
    def parse(value: str) -> int:
        number = int(value)
        if not minimum <= number <= maximum:
            raise argparse.ArgumentTypeError(f"must be between {minimum} and {maximum}")
        return number

    return parse


def _add_output_flags(parser: argparse.ArgumentParser, *, suppress_defaults: bool = False) -> None:
    default = argparse.SUPPRESS if suppress_defaults else False
    parser.add_argument("--json", action="store_true", default=default, help="emit structured JSON")
    parser.add_argument("-q", "--quiet", action="store_true", default=default, help="suppress progress diagnostics")
    parser.add_argument("--debug", action="store_true", default=default, help="show tracebacks for unexpected failures")
    parser.add_argument(
        "--no-color", action="store_true", default=default, help="disable color (currently the default)"
    )


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        prog="unlimited-ocr-rdna4",
        description="Run Baidu Unlimited-OCR locally on a single AMD RDNA 4 GPU.",
        epilog="Examples: unlimited-ocr-rdna4 prepare; unlimited-ocr-rdna4 run --input page.png",
    )
    _add_output_flags(parser)
    parser.add_argument("--version", action="version", version=VERSION)
    subparsers = parser.add_subparsers(dest="command")

    prepare = subparsers.add_parser("prepare", help="download and verify the pinned Baidu model")
    _add_output_flags(prepare, suppress_defaults=True)
    prepare.add_argument("--model-dir", type=Path, default=default_model_dir(), help="prepared model destination")
    prepare.add_argument("--cache-dir", type=Path, default=default_cache_dir(), help="persistent download cache")
    prepare.add_argument("-n", "--dry-run", action="store_true", help="show the plan without downloading")

    doctor = subparsers.add_parser("doctor", help="inspect ROCm, PyTorch, RDNA 4, and model readiness")
    _add_output_flags(doctor, suppress_defaults=True)
    doctor.add_argument("--device", help="ROCm ordinal or stable ROCr UUID; omit to inspect all visible GPUs")
    doctor.add_argument("--model-dir", type=Path, default=default_model_dir(), help="prepared model directory")
    doctor.add_argument("--require-model", action="store_true", help="fail when the prepared model is absent")

    run = subparsers.add_parser("run", help="parse one image or PDF into Markdown")
    _add_output_flags(run, suppress_defaults=True)
    run.add_argument("--input", required=True, type=Path, help="input image or PDF")
    run.add_argument("-o", "--output", type=Path, help="Markdown output path")
    run.add_argument(
        "--device",
        default=None,
        help="ROCm ordinal or stable ROCr UUID (default: UNLIMITED_OCR_DEVICE or 0)",
    )
    run.add_argument("--model-dir", type=Path, default=default_model_dir(), help="prepared model directory")
    run.add_argument("--mode", choices=("gundam", "base"), default="gundam", help="single-page image profile")
    run.add_argument("--max-length", type=bounded_int(512, 32768), default=4096, help="total sequence limit")
    run.add_argument("--dpi", type=bounded_int(72, 400), default=200, help="PDF rendering resolution")
    run.add_argument("--start-page", type=bounded_int(1, 100000), default=1, help="first PDF page, 1-based")
    run.add_argument("--max-pages", type=bounded_int(1, 100), default=20, help="maximum PDF pages per run")
    run.add_argument(
        "--max-page-pixels",
        type=bounded_int(1_000_000, 200_000_000),
        default=60_000_000,
        help="maximum rendered pixels for one PDF page",
    )
    run.add_argument(
        "--max-total-pixels",
        type=bounded_int(1_000_000, 2_000_000_000),
        default=400_000_000,
        help="maximum aggregate rendered pixels for a PDF run",
    )
    run.add_argument(
        "--max-page-rendered-mib",
        type=bounded_int(1, 2048),
        default=512,
        help="maximum temporary size of one rendered PDF page",
    )
    run.add_argument(
        "--max-rendered-mib",
        type=bounded_int(1, 8192),
        default=2048,
        help="maximum aggregate size of rendered PDF pages",
    )
    run.add_argument("--prompt", default="<image>document parsing.", help="model prompt")
    run.add_argument("-f", "--force", action="store_true", help="replace an existing output file")
    return parser


def _emit(payload: dict, *, as_json: bool, quiet: bool = False) -> None:
    if as_json:
        print(json.dumps(payload, ensure_ascii=False, sort_keys=True))
        return
    if quiet:
        primary = payload.get("output") or payload.get("model_dir") or payload.get("status")
        if primary is not None:
            print(primary)
        return
    for key, value in payload.items():
        if isinstance(value, (dict, list, tuple)):
            print(f"{key}={json.dumps(value, ensure_ascii=False, sort_keys=True)}")
        else:
            print(f"{key}={value}")


def _command_prepare(args: argparse.Namespace) -> int:
    if not args.quiet:
        action = "Would prepare" if args.dry_run else "Preparing"
        print(f"{action} {MODEL_REVISION} at {args.model_dir}", file=sys.stderr, flush=True)
    status = prepare_model(args.model_dir, args.cache_dir, dry_run=args.dry_run)
    payload = {"schema_version": 1, **status.__dict__, "status": "dry-run" if args.dry_run else "ok"}
    _emit(payload, as_json=args.json, quiet=args.quiet)
    return 0


def _command_doctor(args: argparse.Namespace) -> int:
    runtime = inspect_runtime(device=args.device, require_single=False, validate=False)
    issues = runtime_issues(runtime, require_single=False)
    model_payload: dict[str, object]
    try:
        model = verify_model(args.model_dir.expanduser().resolve(), full_weight_hash=False)
        model_payload = {"prepared": True, "model_dir": model.model_dir, "revision": model.revision}
    except ModelIntegrityError as exc:
        if args.require_model:
            raise
        model_payload = {"prepared": False, "model_dir": str(args.model_dir), "detail": str(exc)}
    payload = {
        "schema_version": 1,
        "status": "ok" if not issues else "error",
        "runtime": runtime.to_dict(),
        "runtime_issues": issues,
        "model": model_payload,
    }
    _emit(payload, as_json=args.json, quiet=args.quiet)
    return 0 if not issues else 3


def _command_run(args: argparse.Namespace) -> int:
    import os

    device = args.device or os.environ.get("UNLIMITED_OCR_DEVICE")
    visibility_is_preconfigured = any(
        name in os.environ for name in ("ROCR_VISIBLE_DEVICES", "HIP_VISIBLE_DEVICES", "CUDA_VISIBLE_DEVICES")
    )
    if device is None and not visibility_is_preconfigured:
        device = "0"
    runtime = inspect_runtime(device=device, require_single=True)
    input_path = args.input.expanduser()
    output = args.output.expanduser() if args.output else default_output_path(input_path)
    if not args.quiet:
        print(f"Loading pinned model on RDNA 4 device {device}; inference is offline", file=sys.stderr, flush=True)
    result = run_inference(
        RunOptions(
            input_path=input_path,
            output_path=output,
            model_dir=args.model_dir,
            mode=args.mode,
            max_length=args.max_length,
            dpi=args.dpi,
            start_page=args.start_page,
            max_pages=args.max_pages,
            max_page_pixels=args.max_page_pixels,
            max_total_pixels=args.max_total_pixels,
            max_page_rendered_bytes=args.max_page_rendered_mib * 1024**2,
            max_rendered_bytes=args.max_rendered_mib * 1024**2,
            prompt=args.prompt,
            force=args.force,
            quiet=args.quiet,
        ),
        runtime,
    )
    for warning in result.warnings:
        print(f"warning: {warning}", file=sys.stderr)
    payload = {"status": "ok", **result.to_dict()}
    _emit(payload, as_json=args.json, quiet=args.quiet)
    return 0


def run(argv: list[str] | None = None) -> int:
    parser = build_parser()
    args = parser.parse_args(argv)
    if not args.command:
        parser.print_help(sys.stderr)
        return 2
    try:
        if args.command == "prepare":
            return _command_prepare(args)
        if args.command == "doctor":
            return _command_doctor(args)
        if args.command == "run":
            return _command_run(args)
        parser.error(f"unknown command: {args.command}")
    except KeyboardInterrupt:
        print("interrupted", file=sys.stderr)
        return 130
    except AppError as exc:
        if args.json:
            print(
                json.dumps({"schema_version": 1, "status": "error", "error": str(exc), "exit_code": exc.exit_code}),
                file=sys.stderr,
            )
        else:
            print(f"error: {exc}", file=sys.stderr)
        return exc.exit_code
    except Exception as exc:
        if args.debug:
            traceback.print_exc()
        else:
            print(f"error: unexpected failure: {exc}; rerun with --debug", file=sys.stderr)
        return 1
    return 1


def main() -> None:
    raise SystemExit(run())