Puker_Judge / infer_binary.py
TuWaveGod's picture
Publish two-stage final LoRA adapters and inference guide
c35c7ce verified
Raw
History Blame Contribute Delete
1.93 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import re
from pathlib import Path
from PIL import Image
from puker_judge_utils import (
BASE_MODEL_ID,
BINARY_PROMPT,
MODEL_REPO_ID,
encode_image_prompt,
generate_answer,
load_adapter,
print_json,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Judge one assembled playing-card candidate as VALID/INVALID."
)
parser.add_argument("image", type=Path)
parser.add_argument("--repo-id", default=MODEL_REPO_ID)
parser.add_argument("--base-model-id", default=BASE_MODEL_ID)
parser.add_argument(
"--int4",
action="store_true",
help="Use bitsandbytes NF4 weights with BF16 compute.",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
model, processor, device, timings = load_adapter(
"binary_adapter",
repo_id=args.repo_id,
base_model_id=args.base_model_id,
int4=args.int4,
)
with Image.open(args.image) as source:
image = source.convert("RGB").copy()
inputs = encode_image_prompt(
processor,
image,
BINARY_PROMPT,
device,
)
raw_output, generation_seconds = generate_answer(
model,
processor,
inputs,
)
match = re.search(r"\b(INVALID|VALID)\b", raw_output.upper())
if match is None:
raise SystemExit(f"Model returned an invalid answer: {raw_output!r}")
print_json(
{
"prediction": match.group(1),
"raw_output": raw_output,
"image": str(args.image.resolve()),
"quantization": "int4-nf4" if args.int4 else "bf16",
"generation_seconds": round(generation_seconds, 4),
**{key: round(value, 4) for key, value in timings.items()},
}
)
if __name__ == "__main__":
main()