File size: 10,020 Bytes
1f787fa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
230
231
232
233
234
235
236
237
238
#!/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())