Reinforcement Learning
Diffusers
Safetensors
English
image-quality-assessment
vision-language
image-editing
Instructions to use RobinY99/MR-IQA-2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use RobinY99/MR-IQA-2 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("RobinY99/MR-IQA-2", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download code/examples/actor_3_inference.py from RobinY99/MR-IQA-2: direct link, hf CLI and curl.
- Browser
- Download file 10 kB
-
https://huggingface.co/RobinY99/MR-IQA-2/resolve/main/code/examples/actor_3_inference.py
- Command line
-
hf download hf://RobinY99/MR-IQA-2/code/examples/actor_3_inference.py
-
curl -L -o actor_3_inference.py https://huggingface.co/RobinY99/MR-IQA-2/resolve/main/code/examples/actor_3_inference.py
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()) | |