MR-IQA-2 / code /examples /actor_3_inference.py
RobinY99's picture
Add Actor-3 prompts and direct inference
1f787fa verified
Raw History Blame Contribute Delete
10 kB
#!/usr/bin/env python3
"""Run the MR-IQA-2 Actor-3 consistency-oriented checkpoint on one image."""
from __future__ import annotations
import argparse
import json
import re
from pathlib import Path
from typing import Any
REPO_ID = "RobinY99/MR-IQA-2"
MODEL_SUBFOLDER = "actor-3"
SYSTEM_PROMPT = (
"You are a helpful assistant. When the user asks a question, respond with "
"exactly one valid JSON object and no other text."
)
USER_PROMPT = (
"Assess the overall perceptual quality of this specific image.\n\n"
"Respond with exactly one JSON object containing these keys in this order: \"reasoning\" and \"rating\". "
"\"reasoning\" must be one JSON object containing these keys in this order: \"evidence\" and \"solution\".\n\n"
"\"evidence\" must be one concise, directly visible, image-specific observation. It must identify both the "
"depicted subject, object, or scene element involved and the spatial region where the observation is visible. "
"Do not use a generic quality statement that could apply unchanged to unrelated images, and do not claim "
"anything that cannot be verified from this image.\n"
"\"solution\" must be one concise, evidence-grounded, preservation-first correction that directly addresses "
"the stated evidence. It must not add, remove, replace, move, resize, reshape, or change the identity, category, "
"count, pose, expression, clothing, geometry, layout, or semantic role of any main subject or object. It must "
"preserve the scene meaning, composition, background structure, text content, and all unaffected regions. Never "
"propose replacing a person or changing a person's gender, age, identity, body, or attire. If no specific defect "
"is visible, request only a minimal preservation-first refinement or explicitly preserve the image without "
"semantic edits.\n"
"\"rating\" must be a numeric string from 1.00 to 5.00 with exactly two decimal places. Judge only the overall "
"perceptual quality visible in the current image. \"1.00\" is reserved for extremely poor overall quality. "
"\"5.00\" is reserved for exceptional overall quality with no meaningful visible room for improvement. Use an "
"intermediate value whenever the quality lies between these endpoints, and never output a value below 1.00 or "
"above 5.00."
)
NON_THINKING_PREFIX = "<think>\n\n</think>\n\n"
RATING_PATTERN = re.compile(r"[1-5]\.\d{2}")
def build_messages(image_path: Path) -> list[dict[str, Any]]:
return [
{"role": "system", "content": SYSTEM_PROMPT},
{
"role": "user",
"content": [
{"type": "image", "image": str(image_path)},
{"type": "text", "text": USER_PROMPT},
],
},
]
def parse_actor_output(raw: str) -> dict[str, Any]:
text = raw.strip()
if text.startswith(NON_THINKING_PREFIX):
text = text[len(NON_THINKING_PREFIX) :].strip()
try:
payload = json.loads(text)
except json.JSONDecodeError as exc:
raise ValueError(f"Actor-3 did not return one valid JSON object: {exc}") from exc
if not isinstance(payload, dict) or list(payload) != ["reasoning", "rating"]:
raise ValueError("Actor-3 output must contain ordered keys: reasoning, rating")
reasoning = payload["reasoning"]
if not isinstance(reasoning, dict) or list(reasoning) != ["evidence", "solution"]:
raise ValueError("reasoning must contain ordered keys: evidence, solution")
for field in ("evidence", "solution"):
if not isinstance(reasoning[field], str) or not reasoning[field].strip():
raise ValueError(f"reasoning.{field} must be a non-empty string")
rating = payload["rating"]
if not isinstance(rating, str) or RATING_PATTERN.fullmatch(rating) is None:
raise ValueError("rating must be a numeric string with exactly two decimals")
if not 1.0 <= float(rating) <= 5.0:
raise ValueError("rating must be in [1.00, 5.00]")
return payload
def resolve_load_kwargs(args: argparse.Namespace) -> dict[str, Any]:
model_path = Path(args.model).expanduser()
kwargs: dict[str, Any] = {
"trust_remote_code": True,
"local_files_only": bool(args.local_files_only),
}
if not (model_path.is_dir() and (model_path / "config.json").is_file()):
kwargs["subfolder"] = args.subfolder
if args.revision:
kwargs["revision"] = args.revision
return kwargs
def infer(args: argparse.Namespace) -> tuple[str, dict[str, Any]]:
import torch
from PIL import Image
from transformers import (
AutoModelForImageTextToText,
AutoProcessor,
LogitsProcessor,
LogitsProcessorList,
)
class GeneratedPresencePenalty(LogitsProcessor):
def __init__(self, prompt_length: int, penalty: float) -> None:
self.prompt_length = int(prompt_length)
self.penalty = float(penalty)
def __call__(self, input_ids: Any, scores: Any) -> Any:
if self.penalty == 0.0 or input_ids.shape[1] <= self.prompt_length:
return scores
generated_ids = input_ids[:, self.prompt_length :]
for row in range(generated_ids.shape[0]):
seen = generated_ids[row].unique()
scores[row, seen] -= self.penalty
return scores
image_path = Path(args.image).expanduser().resolve(strict=True)
if not image_path.is_file():
raise FileNotFoundError(f"input image is not a file: {image_path}")
load_kwargs = resolve_load_kwargs(args)
processor = AutoProcessor.from_pretrained(
args.model,
max_pixels=args.max_pixels,
min_pixels=args.min_pixels,
**load_kwargs,
)
dtype: Any = "auto" if args.dtype == "auto" else getattr(torch, args.dtype)
model = AutoModelForImageTextToText.from_pretrained(
args.model,
torch_dtype=dtype,
attn_implementation=args.attn_implementation,
**load_kwargs,
).to(args.device).eval()
rendered = processor.apply_chat_template(
build_messages(image_path),
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
if rendered.endswith(NON_THINKING_PREFIX):
rendered = rendered[: -len(NON_THINKING_PREFIX)]
with Image.open(image_path) as opened:
image = opened.convert("RGB")
inputs = processor(
text=[rendered],
images=[image],
padding=True,
return_tensors="pt",
).to(args.device)
prompt_length = int(inputs["input_ids"].shape[1])
available_tokens = int(args.max_model_len) - prompt_length
if available_tokens <= 0:
raise ValueError(
f"rendered prompt uses {prompt_length} tokens, exceeding max_model_len={args.max_model_len}"
)
max_new_tokens = min(int(args.max_new_tokens), available_tokens)
torch.manual_seed(args.seed)
if str(args.device).startswith("cuda"):
torch.cuda.manual_seed_all(args.seed)
logits_processors = LogitsProcessorList(
[GeneratedPresencePenalty(prompt_length, args.presence_penalty)]
)
with torch.inference_mode():
generated = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
use_cache=True,
repetition_penalty=args.repetition_penalty,
logits_processor=logits_processors,
)
completion_ids = generated[:, prompt_length:]
raw = processor.batch_decode(
completion_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=False,
)[0]
return raw, parse_actor_output(raw)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("image", help="Input image path")
parser.add_argument("--model", default=REPO_ID)
parser.add_argument("--subfolder", default=MODEL_SUBFOLDER)
parser.add_argument("--revision", default="")
parser.add_argument("--device", default="cuda:0")
parser.add_argument(
"--dtype",
choices=("auto", "bfloat16", "float16", "float32"),
default="bfloat16",
)
parser.add_argument("--attn-implementation", default="sdpa")
parser.add_argument("--max-new-tokens", type=int, default=1024)
parser.add_argument("--max-model-len", type=int, default=2048)
parser.add_argument("--max-pixels", type=int, default=196608)
parser.add_argument("--min-pixels", type=int, default=3136)
parser.add_argument("--presence-penalty", type=float, default=1.5)
parser.add_argument("--repetition-penalty", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--output", default="", help="Optional JSON output path")
parser.add_argument("--local-files-only", action="store_true")
args = parser.parse_args()
if args.max_new_tokens <= 0 or args.max_model_len <= 0:
parser.error("token limits must be positive")
if args.min_pixels <= 0 or args.max_pixels < args.min_pixels:
parser.error("pixel limits are invalid")
return args
def main() -> int:
args = parse_args()
raw, payload = infer(args)
result = {
"model": (
args.model
if Path(args.model).expanduser().is_dir()
else f"{args.model}/{args.subfolder}"
),
"prompt_version": "vf_reasoning_evidence_solution_rating_prohibitions_v9_20260817",
"raw_completion": raw,
"assessment": payload,
}
serialized = json.dumps(result, ensure_ascii=False, indent=2) + "\n"
if args.output:
output_path = Path(args.output).expanduser().resolve()
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(serialized, encoding="utf-8")
print(serialized, end="")
return 0
if __name__ == "__main__":
raise SystemExit(main())