Image Segmentation
Transformers
Safetensors
mlrs
feature-extraction
remote-sensing
reasoning-segmentation
earth-observation
custom_code
Instructions to use INSAIT-Institute/MLRS with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use INSAIT-Institute/MLRS with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-segmentation", model="INSAIT-Institute/MLRS", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("INSAIT-Institute/MLRS", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
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)
|