alivegames/Grounded-Segment-Anything
0
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 