返回 ViMax
best_image_selector.py
根目录 / agents / best_image_selector.py
1 import logging
2 import json
3 from typing import List, Tuple
4 from pydantic import BaseModel, Field, StrictInt
5 from tenacity import retry, stop_after_attempt, wait_exponential
6 from langchain.chat_models.base import BaseChatModel
7 from langchain_core.messages import HumanMessage, SystemMessage
8 from langchain_core.exceptions import OutputParserException
9 from utils.robust_json_parser import TrailingCommaTolerantPydanticOutputParser as PydanticOutputParser, strip_trailing_commas
10 from langchain.chat_models import init_chat_model
11 from utils.image import image_path_to_b64
12
13
14
15 system_prompt_template_select_most_consistent_image = \
16 """
17 [Role]
18 You are a professional visual assessment expert. Your expertise includes identifying Character Consistency and Spatial Consistency between candidate image and reference image, and assessing semantic consistency between candidate image and text description.
19
20 [Task]
21 Based on the reference image provided by the user, the text description of the target image, and several candidate images, evaluate which candidate image performs best in the following aspects:
22 - Character Consistency: Whether the character features (a. gender, b.ethnicity, c.age, d.facial features, e.body shape, f.outlook, g. hairstyle) in the candidate image align with those of the character in the reference image.
23 - Spatial Consistency: Whether the relative positions between characters (e.g. Character A is on the left, character B is on the right, scene layout, perspective, and other spatial relationships) in the candidate image are consistent with those in the reference image.
24 - Description Accuracy: Whether the candidate image accurately reflects the content described in the text (Note: The text description describes the target image we want, which is not an editing instruction).
25
26 [Input]
27 The user will provide the following content:
28 - Reference images: These include images of characters or other perspectives, each along with a brief text description. For example, "Reference Image 0: A young girl with long brown hair wearing a red dress." then follow the corresponding image. The index starts from 0.
29 - Candidate images: The candidate images to be evaluated. For example, "Generated Image 0", then follow a generated image. The index starts from 0.
30 - Text description for target image: This describes what the generated image should contain. It is enclosed <TARGET_DESCRIPTION_START> and <TARGET_DESCRIPTION_END> tags.
31
32 [Output]
33 {format_instructions}
34
35 [Hard constraints]
36 - There are exactly {candidate_count} candidate images, indexed from 0 to {candidate_max_index}.
37 - Return a single JSON object with only "best_image_index" and "reason".
38 - Choose one of the provided indices. Do not invent candidates or return a JSON schema.
39
40 [Guidelines]
41 - Prioritize Character Consistency: Ensure that the characters in the generated image are highly consistent with those in the reference image in terms of visual features (e.g., a. gender b.ethnicity, c.age, d.facial features, e.body shape, f.outlook, g. hairstyle etc.).
42 - Focus on Spatial Consistency: Verify whether the relative positions of characters, object arrangements, and perspectives align logically with the reference image (e.g., if Character A is on the left and Character B is on the right in the reference image, the generated image should not reverse this).
43 - Strictly Compare with Text Description: The generated image must adhere to key elements in the text description (e.g., actions, scenes, objects, etc.), while disregarding parts related to editing instructions (as the input description reflects the expected outcome rather than directives).
44 - If multiple images partially meet the criteria, select the one with the highest overall consistency; if none are ideal, choose the relatively best option and explain its shortcomings.
45 - Ensure the key elements described in the text are present in the selected image.
46 - Avoid subjective preferences; base all analysis on objective comparisons.
47 - Prioritize images without white borders, black edges, or any additional framing.
48 """
49
50 human_prompt_template_select_most_consistent_image = \
51 """
52 <TARGET_DESCRIPTION_START>
53 {target_description}
54 <TARGET_DESCRIPTION_END>
55 """
56
57
58 class BestImageResponse(BaseModel):
59 best_image_index: StrictInt = Field(
60 ...,
61 ge=0,
62 description="The index of the best image."
63 )
64 reason: str = Field(
65 ...,
66 description="The reason why the image is the best."
67 )
68
69
70 class BestImageSelector:
71 def __init__(
72 self,
73 base_url: str | None = None,
74 api_key: str | None = None,
75 chat_model: str | BaseChatModel | None = None,
76 ):
77
78 if chat_model is None:
79 raise ValueError("A vision-capable chat_model is required for image selection")
80 if isinstance(chat_model, str):
81 self.chat_model = init_chat_model(
82 model=chat_model,
83 model_provider="openai",
84 base_url=base_url,
85 api_key=api_key,
86 )
87 else:
88 self.chat_model = chat_model
89
90
91 async def __call__(
92 self,
93 reference_image_path_and_text_pairs: List[Tuple[str, str]],
94 target_description: str,
95 candidate_image_paths: List[str],
96 ) -> str:
97 response = await self.select(
98 reference_image_path_and_text_pairs,
99 target_description,
100 candidate_image_paths,
101 )
102 return candidate_image_paths[response.best_image_index]
103
104
105 @retry(
106 stop=stop_after_attempt(3),
107 wait=wait_exponential(multiplier=1, min=1, max=4),
108 reraise=True,
109 after=lambda retry_state: logging.warning(f"Retrying best image selection due to {retry_state.outcome.exception()}"),
110 )
111 async def select(
112 self,
113 reference_image_path_and_text_pairs: List[Tuple[str, str]],
114 target_description: str,
115 candidate_image_paths: List[str],
116 ) -> BestImageResponse:
117 """
118 Args:
119 ref_image_path_and_text_pairs:
120 A list of tuples containing reference image paths and their descriptions.
121
122 target_description:
123 The description of the target image.
124
125 candidate_image_paths:
126 A list of paths to the candidate images to be evaluated.
127 """
128
129 if not candidate_image_paths:
130 logging.warning("No candidate images provided; skipping best image selection")
131 raise ValueError("No candidate images to select from")
132
133 logging.info(f"Selecting the best image from candidates: {candidate_image_paths}")
134
135 human_content = []
136 for idx, (ref_image_path, text) in enumerate(reference_image_path_and_text_pairs):
137 human_content.append({
138 "type": "text",
139 "text": f"Reference Image {idx}: {text}"
140 })
141 human_content.append({
142 "type": "image_url",
143 "image_url": {"url": image_path_to_b64(ref_image_path, mime=True)}
144 })
145
146 for idx, candidate_image_path in enumerate(candidate_image_paths):
147 human_content.append({
148 "type": "text",
149 "text": f"Candidate Image {idx}"
150 })
151 human_content.append({
152 "type": "image_url",
153 "image_url": {"url": image_path_to_b64(candidate_image_path, mime=True)}
154 })
155 human_content.append({
156 "type": "text",
157 "text": human_prompt_template_select_most_consistent_image.format(target_description=target_description)
158 })
159
160 parser = PydanticOutputParser(pydantic_object=BestImageResponse)
161
162 messages = [
163 SystemMessage(content=system_prompt_template_select_most_consistent_image.format(
164 format_instructions=parser.get_format_instructions(),
165 candidate_count=len(candidate_image_paths),
166 candidate_max_index=len(candidate_image_paths) - 1,
167 )),
168 HumanMessage(content=human_content)
169 ]
170
171 message = await self.chat_model.ainvoke(messages)
172 raw = message.content
173 if isinstance(raw, list):
174 raw = "".join(block.get("text", "") for block in raw if isinstance(block, dict))
175 try:
176 response = parser.parse(raw)
177 except OutputParserException:
178 # Some gateways wrap values in a JSON-schema "properties" object.
179 text = strip_trailing_commas(raw)
180 obj, _ = json.JSONDecoder().raw_decode(text[text.index("{"):])
181 if isinstance(obj, dict) and isinstance(obj.get("properties"), dict):
182 obj = obj["properties"]
183 response = BestImageResponse.model_validate(obj)
184 idx = response.best_image_index
185 if idx >= len(candidate_image_paths):
186 raise ValueError(f"VLM selected invalid candidate index {idx} for {len(candidate_image_paths)} images")
187 best_image_path = candidate_image_paths[idx]
188 logging.info(f"Best image selected: {best_image_path}")
189 logging.info(f"Selection reason: {response.reason}")
190 return response
191
191 lines PYTHON