MLRS / processing_mlrs.py
Ronningen's picture
Upload folder using huggingface_hub
c40c0f3 verified
Raw History Blame Contribute Delete
7.62 kB
"""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)