File size: 7,616 Bytes
c40c0f3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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)