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
Download processing_mlrs.py from INSAIT-Institute/MLRS: direct link, hf CLI and curl.
- Browser
- Download file 7.62 kB
-
https://huggingface.co/INSAIT-Institute/MLRS/resolve/main/processing_mlrs.py
- Command line
-
hf download hf://INSAIT-Institute/MLRS/processing_mlrs.py
-
curl -L -o processing_mlrs.py https://huggingface.co/INSAIT-Institute/MLRS/resolve/main/processing_mlrs.py
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 | |
| def register_for_auto_class(cls, auto_class: str = "AutoProcessor") -> None: | |
| # Called by AutoProcessor for remote-code classes; nothing to register. | |
| return None | |
| 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) | |