"""Processor for MLRS: builds InternVL3.5 chat inputs and SAM3 image inputs.""" from __future__ import annotations from typing import Any, Dict, List, Optional, Sequence, Tuple, Union import torch import torch.nn.functional as F from PIL import Image from transformers import AutoProcessor, BatchFeature, Sam3Processor from .configuration_mlrs import MLRSConfig _HUB_KWARGS = ("cache_dir", "force_download", "local_files_only", "proxies", "revision", "token", "subfolder") _BASE_HUB_KWARGS = ("cache_dir", "force_download", "local_files_only", "proxies", "token") ImageGroup = Union[Image.Image, Sequence[Image.Image]] def downsample_longest_side( img: Image.Image, max_longest_side: int, ) -> Tuple[Image.Image, bool, Tuple[int, int]]: w, h = img.size longest = max(w, h) if longest <= max_longest_side: return img, False, (w, h) scale = max_longest_side / float(longest) new_w = max(1, round(w * scale)) new_h = max(1, round(h * scale)) resized = img.resize((new_w, new_h), resample=Image.Resampling.LANCZOS) return resized, True, (new_w, new_h) class MLRSProcessor: """ Turns `(images, text)` into model inputs. Input schema (one sample): images: a PIL image or a list of PIL images. When several images are given, the segmentation tool runs on `images[primary_image_index]` (default: the last one; in training this was the last RGB image). text: the task / question. For a batch, pass lists: `images=[[img_a], [img_b1, img_b2]]`, `text=[..., ...]`. Output keys: input_ids, attention_mask, pixel_values -> InternVL3.5 sam_inputs -> SAM3 image inputs (dict) sam_original_sizes [B, 2] (H, W) -> size of the (possibly downsampled) primary image; masks are returned at this size image_original_sizes [B, 2] (H, W) -> size of the primary image before downsampling """ def __init__(self, language_processor, sam_processor, config: MLRSConfig): self.language_processor = language_processor self.sam_processor = sam_processor self.config = config self.tokenizer = language_processor.tokenizer @classmethod def register_for_auto_class(cls, auto_class: str = "AutoProcessor") -> None: # Called by AutoProcessor for remote-code classes; nothing to register. return None @classmethod def from_pretrained(cls, pretrained_model_name_or_path, trust_remote_code: Optional[bool] = None, **kwargs): hub_kwargs = {k: kwargs.pop(k) for k in _HUB_KWARGS if k in kwargs} config = kwargs.pop("config", None) if config is None: config = MLRSConfig.from_pretrained(pretrained_model_name_or_path, **hub_kwargs) base_hub_kwargs = {k: hub_kwargs[k] for k in _BASE_HUB_KWARGS if k in hub_kwargs} language_processor = AutoProcessor.from_pretrained( config.language_base_model, revision=config.language_base_revision, **base_hub_kwargs, ) sam_processor = Sam3Processor.from_pretrained( config.sam_base_model, revision=config.sam_base_revision, **base_hub_kwargs, ) return cls(language_processor, sam_processor, config) def _build_prompt(self, text: str) -> str: return self.config.instruction_prompt.strip() + "\n\n" + f"Your task: {text}" def _max_longest_side(self, num_images: int) -> int: cfg = self.config limit = int(cfg.max_image_longest_side) if ( cfg.multi_image_downsample_min_images is not None and cfg.multi_image_max_longest_side is not None and num_images >= int(cfg.multi_image_downsample_min_images) ): limit = min(limit, int(cfg.multi_image_max_longest_side)) return limit def __call__( self, images: Union[ImageGroup, Sequence[ImageGroup]], text: Union[str, Sequence[str]], primary_image_index: Union[int, Sequence[int]] = -1, ) -> BatchFeature: if isinstance(text, str): texts = [text] image_groups = [images] else: texts = list(text) image_groups = list(images) if len(image_groups) != len(texts): raise ValueError(f"Got {len(image_groups)} image groups for {len(texts)} texts") image_groups = [[g] if isinstance(g, Image.Image) else list(g) for g in image_groups] if isinstance(primary_image_index, int): primary_indices = [primary_image_index] * len(texts) else: primary_indices = list(primary_image_index) processed_groups: List[List[Image.Image]] = [] primary_images: List[Image.Image] = [] image_original_sizes: List[List[int]] = [] messages: List[List[Dict[str, Any]]] = [] for imgs, prompt, primary_idx in zip(image_groups, texts, primary_indices): if len(imgs) == 0: raise ValueError("Every sample needs at least one image") primary_original = imgs[primary_idx] image_original_sizes.append([primary_original.size[1], primary_original.size[0]]) limit = self._max_longest_side(len(imgs)) imgs = [downsample_longest_side(img, max_longest_side=limit)[0] for img in imgs] processed_groups.append(imgs) primary_images.append(imgs[primary_idx]) image_items = [{"type": "image", "image": img} for img in imgs] text_items = [{"type": "text", "text": self._build_prompt(prompt)}] content = image_items + text_items if self.config.image_first else text_items + image_items messages.append([{"role": "user", "content": content}]) out = self.language_processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", processor_kwargs={ "padding": True, "padding_side": self.config.padding_side, "truncation": False, "pad_to_multiple_of": self.config.pad_to_multiple_of, }, ) out["sam_inputs"] = self.sam_processor.image_processor(primary_images, return_tensors="pt") out["sam_original_sizes"] = torch.tensor( [[img.size[1], img.size[0]] for img in primary_images], dtype=torch.long, ) out["image_original_sizes"] = torch.tensor(image_original_sizes, dtype=torch.long) return out def post_process_masks(self, outputs, inputs) -> List[torch.Tensor]: """Resize predicted masks back to the primary image size before downsampling.""" sizes = inputs["image_original_sizes"].tolist() resized: List[torch.Tensor] = [] for masks, (h, w) in zip(outputs.masks, sizes): if tuple(masks.shape[-2:]) == (h, w): resized.append(masks) elif masks.shape[0] == 0: resized.append(masks.new_zeros((0, h, w))) else: resized.append( F.interpolate(masks[:, None].float(), size=(h, w), mode="nearest")[:, 0].bool() ) return resized def batch_decode(self, *args, **kwargs): return self.tokenizer.batch_decode(*args, **kwargs) def decode(self, *args, **kwargs): return self.tokenizer.decode(*args, **kwargs)