CoolFace
Apppublic

alivegames/Grounded-Segment-Anything

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
mask_decoder.py178 linesDownload Raw Back to modeling
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3 4# This source code is licensed under the license found in the5# LICENSE file in the root directory of this source tree.6 7import torch8from torch import nn9from torch.nn import functional as F10 11from typing import List, Tuple, Type12 13from .common import LayerNorm2d14 15 16class MaskDecoder(nn.Module):17    def __init__(18        self,19        *,20        transformer_dim: int,21        transformer: nn.Module,22        num_multimask_outputs: int = 3,23        activation: Type[nn.Module] = nn.GELU,24        iou_head_depth: int = 3,25        iou_head_hidden_dim: int = 256,26    ) -> None:27        """28        Predicts masks given an image and prompt embeddings, using a29        tranformer architecture.30 31        Arguments:32          transformer_dim (int): the channel dimension of the transformer33          transformer (nn.Module): the transformer used to predict masks34          num_multimask_outputs (int): the number of masks to predict35            when disambiguating masks36          activation (nn.Module): the type of activation to use when37            upscaling masks38          iou_head_depth (int): the depth of the MLP used to predict39            mask quality40          iou_head_hidden_dim (int): the hidden dimension of the MLP41            used to predict mask quality42        """43        super().__init__()44        self.transformer_dim = transformer_dim45        self.transformer = transformer46 47        self.num_multimask_outputs = num_multimask_outputs48 49        self.iou_token = nn.Embedding(1, transformer_dim)50        self.num_mask_tokens = num_multimask_outputs + 151        self.mask_tokens = nn.Embedding(self.num_mask_tokens, transformer_dim)52 53        self.output_upscaling = nn.Sequential(54            nn.ConvTranspose2d(transformer_dim, transformer_dim // 4, kernel_size=2, stride=2),55            LayerNorm2d(transformer_dim // 4),56            activation(),57            nn.ConvTranspose2d(transformer_dim // 4, transformer_dim // 8, kernel_size=2, stride=2),58            activation(),59        )60        self.output_hypernetworks_mlps = nn.ModuleList(61            [62                MLP(transformer_dim, transformer_dim, transformer_dim // 8, 3)63                for i in range(self.num_mask_tokens)64            ]65        )66 67        self.iou_prediction_head = MLP(68            transformer_dim, iou_head_hidden_dim, self.num_mask_tokens, iou_head_depth69        )70 71    def forward(72        self,73        image_embeddings: torch.Tensor,74        image_pe: torch.Tensor,75        sparse_prompt_embeddings: torch.Tensor,76        dense_prompt_embeddings: torch.Tensor,77        multimask_output: bool,78    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:79        """80        Predict masks given image and prompt embeddings.81 82        Arguments:83          image_embeddings (torch.Tensor): the embeddings from the image encoder84          image_pe (torch.Tensor): positional encoding with the shape of image_embeddings85          sparse_prompt_embeddings (torch.Tensor): the embeddings of the points and boxes86          dense_prompt_embeddings (torch.Tensor): the embeddings of the mask inputs87          multimask_output (bool): Whether to return multiple masks or a single88            mask.89 90        Returns:91          torch.Tensor: batched predicted masks92          torch.Tensor: batched predictions of mask quality93        """94        masks, iou_pred, mask_tokens_out = self.predict_masks(95            image_embeddings=image_embeddings,96            image_pe=image_pe,97            sparse_prompt_embeddings=sparse_prompt_embeddings,98            dense_prompt_embeddings=dense_prompt_embeddings,99        )100 101        # Select the correct mask or masks for outptu102        if multimask_output:103            mask_slice = slice(1, None)104        else:105            mask_slice = slice(0, 1)106        masks = masks[:, mask_slice, :, :]107        mask_tokens_out = mask_tokens_out[:, mask_slice, :]108        iou_pred = iou_pred[:, mask_slice]109 110        # Prepare output111        return masks, iou_pred, mask_tokens_out112 113    def predict_masks(114        self,115        image_embeddings: torch.Tensor,116        image_pe: torch.Tensor,117        sparse_prompt_embeddings: torch.Tensor,118        dense_prompt_embeddings: torch.Tensor,119    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:120        """Predicts masks. See 'forward' for more details."""121        # Concatenate output tokens122        output_tokens = torch.cat([self.iou_token.weight, self.mask_tokens.weight], dim=0)123        output_tokens = output_tokens.unsqueeze(0).expand(sparse_prompt_embeddings.size(0), -1, -1)124        tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1)125 126        # Expand per-image data in batch direction to be per-mask127        src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0)128        src = src + dense_prompt_embeddings129        pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0)130        b, c, h, w = src.shape131 132        # Run the transformer133        hs, src = self.transformer(src, pos_src, tokens)134        iou_token_out = hs[:, 0, :]135        mask_tokens_out = hs[:, 1 : (1 + self.num_mask_tokens), :]136 137        # Upscale mask embeddings and predict masks using the mask tokens138        src = src.transpose(1, 2).view(b, c, h, w)139        upscaled_embedding = self.output_upscaling(src)140        hyper_in_list: List[torch.Tensor] = []141        for i in range(self.num_mask_tokens):142            hyper_in_list.append(self.output_hypernetworks_mlps[i](mask_tokens_out[:, i, :]))143        hyper_in = torch.stack(hyper_in_list, dim=1)144        b, c, h, w = upscaled_embedding.shape145        masks = (hyper_in @ upscaled_embedding.view(b, c, h * w)).view(b, -1, h, w)146 147        # Generate mask quality predictions148        iou_pred = self.iou_prediction_head(iou_token_out)149 150        return masks, iou_pred, mask_tokens_out151 152 153# Lightly adapted from154# https://github.com/facebookresearch/MaskFormer/blob/main/mask_former/modeling/transformer/transformer_predictor.py # noqa155class MLP(nn.Module):156    def __init__(157        self,158        input_dim: int,159        hidden_dim: int,160        output_dim: int,161        num_layers: int,162        sigmoid_output: bool = False,163    ) -> None:164        super().__init__()165        self.num_layers = num_layers166        h = [hidden_dim] * (num_layers - 1)167        self.layers = nn.ModuleList(168            nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim])169        )170        self.sigmoid_output = sigmoid_output171 172    def forward(self, x):173        for i, layer in enumerate(self.layers):174            x = F.relu(layer(x)) if i < self.num_layers - 1 else layer(x)175        if self.sigmoid_output:176            x = F.sigmoid(x)177        return x178