import io import base64 import json import math from dataclasses import dataclass from typing import Callable, Dict, Iterable, List, Sequence from PIL import Image, ImageDraw from approach.config import get_model_profile from approach.providers import EncodedImage, OpenAICompatibleChatClient Detection = Dict[str, object] Advisor = Callable[..., Dict[str, object]] Miner = Callable[[List[Dict[str, object]]], Sequence[str]] Detector = Callable[[Sequence[str], List[Detection]], List[Detection]] PII5_PROMPT = ( "PII.5 Mirroring-Based Regional Verification. Compare the original VR scene " "and the cropped detected element. Verify whether the crop is a correct " "user-interactable game element and whether it violates XR spatial constraints." ) PII6_PROMPT = ( "PII.6 Unsuccessful Detection Reflection. Inspect the original scene with " "current detection boxes and identify missing user-interactable CDs." ) PII7_PROMPT = ( "PII.7 Advisor Feedback. Return JSON with verified, needs_refinement, and feedback. " "Each feedback item must be a concise new or revised CD referring expression that " "the detector can ground in the original scene." ) @dataclass(frozen=True) class VerificationInput: prompt: str original_image: Image.Image crop_image: Image.Image detection: Detection def crop_detection(image: Image.Image, bbox: Sequence[float]) -> Image.Image: if len(bbox) != 4: raise ValueError("bbox must contain [x, y, width, height]") x, y, width, height = bbox left = max(0, math.floor(x)) top = max(0, math.floor(y)) right = min(image.width, math.ceil(x + width)) bottom = min(image.height, math.ceil(y + height)) if right <= left or bottom <= top: raise ValueError(f"invalid bbox after clipping: {bbox}") return image.crop((left, top, right, bottom)) def make_successful_detection_request(image: Image.Image, detection: Detection) -> VerificationInput: return VerificationInput( prompt=PII5_PROMPT, original_image=image, crop_image=crop_detection(image, detection["bbox"]), detection=detection, ) def draw_detections(image: Image.Image, detections: Iterable[Detection]) -> Image.Image: annotated = image.copy() draw = ImageDraw.Draw(annotated) for idx, detection in enumerate(detections, start=1): x, y, width, height = detection["bbox"] box = (x, y, x + width, y + height) draw.rectangle(box, outline="red", width=2) draw.text((x, y), str(detection.get("category_name", idx)), fill="red") return annotated def image_to_png_bytes(image: Image.Image) -> bytes: buffer = io.BytesIO() image.save(buffer, format="PNG") return buffer.getvalue() def image_to_base64(image: Image.Image) -> str: return base64.b64encode(image_to_png_bytes(image)).decode("utf-8") REFLECTION_RESPONSE_FORMAT = { "type": "json_schema", "json_schema": { "name": "orienter_reflection", "strict": True, "schema": { "type": "object", "properties": { "verified": {"type": "array", "items": {"type": "integer"}}, "needs_refinement": {"type": "array", "items": {"type": "integer"}}, "feedback": {"type": "array", "items": {"type": "string"}}, }, "required": ["verified", "needs_refinement", "feedback"], "additionalProperties": False, }, }, } def build_reflection_prompt( verification_requests: Sequence[VerificationInput], pii5_prompt: str, pii6_prompt: str, pii7_prompt: str, ) -> str: detections = [ { "index": idx, "bbox": request.detection.get("bbox"), "category_name": request.detection.get("category_name"), "score": request.detection.get("score", request.detection.get("probability")), } for idx, request in enumerate(verification_requests) ] return ( f"{pii5_prompt}\n" "For each detected CD, compare the original scene image with its cropped region image. " "Mark false positives or boxes that do not match the intended interactable element as needs_refinement.\n\n" f"{pii6_prompt}\n" "Also inspect the annotated full-scene image for missing interactable CDs or wrong boxes.\n\n" f"{pii7_prompt}\n" "Return only JSON. Images are ordered as: (1) the boxed full scene for PII.6, " "then (2) one original-scene and bbox-crop pair per detection, in detection-index order.\n\n" f"Detections:\n{json.dumps(detections, ensure_ascii=True, indent=2)}" ) def openai_compatible_advisor(model_profile: str = "default") -> Advisor: profile = get_model_profile(model_profile) client = OpenAICompatibleChatClient(profile) def advisor(**kwargs) -> Dict[str, object]: verification_requests = kwargs["verification_requests"] unsuccessful_detection_image = kwargs["unsuccessful_detection_image"] image_b64s = [ EncodedImage(image_to_base64(unsuccessful_detection_image), "image/png") ] for request in verification_requests: image_b64s.append( EncodedImage(image_to_base64(request.original_image), "image/png") ) image_b64s.append( EncodedImage(image_to_base64(request.crop_image), "image/png") ) prompt = build_reflection_prompt( verification_requests, kwargs["pii5_prompt"], kwargs["pii6_prompt"], kwargs["pii7_prompt"], ) return client.complete_json( prompt, image_b64s, response_format=REFLECTION_RESPONSE_FORMAT, ) return advisor def run_reflection_loop( image: Image.Image, detections: List[Detection], miner: Miner, detector: Detector, advisor: Advisor, max_iterations: int = 10, ) -> Dict[str, object]: if max_iterations <= 0: raise ValueError("max_iterations must be positive") current_detections = list(detections) trace = [] max_reached = False for iteration in range(1, max_iterations + 1): verification_requests = [ make_successful_detection_request(image, detection) for detection in current_detections ] unsuccessful_scene = draw_detections(image, current_detections) advice = advisor( verification_requests=verification_requests, unsuccessful_detection_image=unsuccessful_scene, pii5_prompt=PII5_PROMPT, pii6_prompt=PII6_PROMPT, pii7_prompt=PII7_PROMPT, ) trace.append( { "iteration": iteration, "verified": advice.get("verified", []), "needs_refinement": advice.get("needs_refinement", []), "feedback": advice.get("feedback", []), } ) feedback = advice.get("feedback", []) needs_refinement = advice.get("needs_refinement", []) if not feedback and not needs_refinement: break rejected_indices = { index for index in needs_refinement if isinstance(index, int) and 0 <= index < len(current_detections) } retained_detections = [ detection for index, detection in enumerate(current_detections) if index not in rejected_indices ] candidates = list(miner(trace)) current_detections = detector(candidates, retained_detections) else: max_reached = True return { "detections": current_detections, "trace": trace, "max_iterations_reached": max_reached, }