#!/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 = "\n\n\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())