CoolFace
Modelpublic

ByteDance/Sa2VA-1B

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
31likes821downloads
sam2.py4103 linesDownload Raw Back to root
1import os2from collections import OrderedDict3from tqdm import tqdm4import torch.distributed5 6from torch.nn.init import trunc_normal_7 8import copy9 10from typing import List, Any, Optional, Tuple, Type, Union11 12import numpy as np13 14import math15import warnings16from functools import partial17 18import torch19import torch.nn.functional as F20from torch import nn, Tensor21 22# a large negative value as a placeholder score for missing objects23NO_OBJ_SCORE = -1024.024 25warnings.simplefilter(action="ignore", category=FutureWarning)26# OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = get_sdpa_settings()27OLD_GPU, USE_FLASH_ATTN, MATH_KERNEL_ON = True, True, True28 29def load_checkpoint_with_prefix(filename, prefix=None, map_location='cpu', logger='current'):30    """Load partial pretrained model with specific prefix.31 32    Args:33        prefix (str): The prefix of sub-module.34        filename (str): Accept local filepath, URL, ``torchvision://xxx``,35            ``open-mmlab://xxx``. Please refer to ``docs/model_zoo.md`` for36            details.37        map_location (str | None): Same as :func:`torch.load`.38            Defaults to None.39        logger: logger40 41    Returns:42        dict or OrderedDict: The loaded checkpoint.43    """44    checkpoint = torch.load(filename, map_location=map_location)45 46    if 'state_dict' in checkpoint:47        state_dict = checkpoint['state_dict']48    elif 'model' in checkpoint:49        state_dict = checkpoint['model']50    else:51        state_dict = checkpoint52    if not prefix:53        return state_dict54    if not prefix.endswith('.'):55        prefix += '.'56    prefix_len = len(prefix)57 58    state_dict = {59        k[prefix_len:]: v60        for k, v in state_dict.items() if k.startswith(prefix)61    }62 63    assert state_dict, f'{prefix} is not in the pretrained model'64    return state_dict65 66def load_state_dict_to_model(model, state_dict,  logger='current'):67    missing_keys, unexpected_keys = model.load_state_dict(state_dict)68    if missing_keys:69        print(missing_keys)70        raise RuntimeError()71    if unexpected_keys:72        print(unexpected_keys)73        raise RuntimeError()74    print("Loaded checkpoint successfully")75 76class SAM2(nn.Module):77    def __init__(78            self,79            ckpt_path: str = None,80    ):81        super().__init__()82 83        image_encoder = self.build_image_encoder()84        memory_attention = self.build_memory_attention()85        memory_encoder = self.build_memory_encoder()86        sam2_model = SAM2VideoPredictor(87            image_encoder=image_encoder,88            memory_attention=memory_attention,89            memory_encoder=memory_encoder,90            num_maskmem = 7,91            image_size = 1024,92            # apply scaled sigmoid on mask logits for memory encoder, and directly feed input mask as output mask93            sigmoid_scale_for_mem_enc = 20.0,94            sigmoid_bias_for_mem_enc = -10.0,95            use_mask_input_as_output_without_sam = True,96            # Memory97            directly_add_no_mem_embed = True,98            # use high-resolution feature map in the SAM mask decoder99            use_high_res_features_in_sam = True,100            # output 3 masks on the first click on initial conditioning frames101            multimask_output_in_sam = True,102            # SAM heads103            iou_prediction_use_sigmoid = True,104            # cross-attend to object pointers from other frames (based on SAM output tokens) in the encoder105            use_obj_ptrs_in_encoder = True,106            add_tpos_enc_to_obj_ptrs = False,107            only_obj_ptrs_in_the_past_for_eval = True,108            # object occlusion prediction109            pred_obj_scores = True,110            pred_obj_scores_mlp = True,111            fixed_no_obj_ptr = True,112            # multimask tracking settings113            multimask_output_for_tracking = True,114            use_multimask_token_for_obj_ptr = True,115            multimask_min_pt_num = 0,116            multimask_max_pt_num = 1,117            use_mlp_for_obj_ptr_proj = True,118            # Compilation flag119            compile_image_encoder = False,120            sam_mask_decoder_extra_args={121                'dynamic_multimask_via_stability':True,122                'dynamic_multimask_stability_delta': 0.05,123                'dynamic_multimask_stability_thresh': 0.98,124            }125        )126        if ckpt_path is not None:127            state_dict = load_checkpoint_with_prefix(ckpt_path)128            load_state_dict_to_model(sam2_model, state_dict)129 130        self.sam2_model = sam2_model131 132        self.hidden_dim = self.sam2_model.hidden_dim133 134        self.img_mean = (0.485, 0.456, 0.406)135        self.img_std = (0.229, 0.224, 0.225)136 137    def build_image_encoder(self):138        def build_trunk():139            embed_dim = 144140            num_heads = 2141            stages = [2, 6, 36, 4]142            global_att_blocks = [23, 33, 43]143            window_pos_embed_bkg_spatial_size = [7, 7]144            window_spec = [8, 4, 16, 8]145            ret = Hiera(146                embed_dim=embed_dim,147                num_heads=num_heads,148                stages=stages,149                global_att_blocks=global_att_blocks,150                window_pos_embed_bkg_spatial_size=window_pos_embed_bkg_spatial_size,151                window_spec=window_spec,152            )153            return ret154        def build_neck():155            def build_position_encoding():156                num_pos_feats = 256157                normalize = True158                scale = None159                temperature = 10000160                ret = PositionEmbeddingSine(161                    num_pos_feats=num_pos_feats,162                    normalize=normalize,163                    scale=scale,164                    temperature=temperature,165                )166                return ret167            d_model = 256168            backbone_channel_list = [1152, 576, 288, 144]169            fpn_top_down_levels = [2, 3]  # output level 0 and 1 directly use the backbone features170            fpn_interp_model = 'nearest'171            position_encoding = build_position_encoding()172            ret = FpnNeck(173                d_model=d_model,174                position_encoding=position_encoding,175                backbone_channel_list=backbone_channel_list,176                fpn_top_down_levels=fpn_top_down_levels,177                fpn_interp_model=fpn_interp_model,178            )179            return ret180        scalp = 1181        trunk = build_trunk()182        neck = build_neck()183        ret = ImageEncoder(scalp=scalp, trunk=trunk, neck=neck)184        return ret185 186    def build_memory_attention(self):187        def build_layer():188            def build_self_attention():189                rope_theta = 10000.0190                feat_sizes = [32, 32]191                embedding_dim = 256192                num_heads = 1193                downsample_rate = 1194                dropout = 0.1195                ret = RoPEAttention(196                    rope_theta=rope_theta,197                    feat_sizes=feat_sizes,198                    embedding_dim=embedding_dim,199                    num_heads=num_heads,200                    downsample_rate=downsample_rate,201                    dropout=dropout202                )203                return ret204            def build_cross_attention():205                rope_theta = 10000.0206                feat_sizes = [32, 32]207                rope_k_repeat = True208                embedding_dim = 256209                num_heads = 1210                downsample_rate = 1211                dropout = 0.1212                kv_in_dim = 64213                ret = RoPEAttention(214                    rope_theta=rope_theta,215                    feat_sizes=feat_sizes,216                    rope_k_repeat=rope_k_repeat,217                    embedding_dim=embedding_dim,218                    num_heads=num_heads,219                    downsample_rate=downsample_rate,220                    dropout=dropout,221                    kv_in_dim=kv_in_dim222                )223                return ret224            activation = 'relu'225            dim_feedforward = 2048226            dropout = 0.1227            pos_enc_at_attn = False228            d_model = 256229            pos_enc_at_cross_attn_keys = True230            pos_enc_at_cross_attn_queries = False231            self_attention = build_self_attention()232            cross_attention = build_cross_attention()233            ret = MemoryAttentionLayer(234                activation=activation,235                dim_feedforward=dim_feedforward,236                dropout=dropout,237                pos_enc_at_attn=pos_enc_at_attn,238                d_model=d_model,239                pos_enc_at_cross_attn_queries=pos_enc_at_cross_attn_queries,240                pos_enc_at_cross_attn_keys=pos_enc_at_cross_attn_keys,241                self_attention=self_attention,242                cross_attention=cross_attention,243            )244            return ret245        d_model = 256246        pos_enc_at_input = True247        num_layers = 4248        layer = build_layer()249        ret = MemoryAttention(250            d_model=d_model,251            pos_enc_at_input=pos_enc_at_input,252            num_layers=num_layers,253            layer=layer,254        )255        return ret256 257    def build_memory_encoder(self):258        def build_position_encoding():259            num_pos_feats = 64260            normalize = True261            scale = None262            temperature = 10000263            ret = PositionEmbeddingSine(264                num_pos_feats=num_pos_feats,265                normalize=normalize,266                scale=scale,267                temperature=temperature,268            )269            return ret270 271        def build_mask_downsampler():272            kernel_size = 3273            stride = 2274            padding = 1275            ret = MaskDownSampler(276                kernel_size=kernel_size,277                stride=stride,278                padding=padding,279            )280            return ret281 282        def build_fuser():283            def build_layer():284                dim = 256285                kernel_size = 7286                padding = 3287                layer_scale_init_value = 1e-6288                use_dwconv = True  # depth-wise convs289                ret = CXBlock(290                    dim=dim, kernel_size=kernel_size,291                    padding=padding, layer_scale_init_value=layer_scale_init_value,292                    use_dwconv=use_dwconv,293                )294                return ret295 296            num_layers = 2297            layer = build_layer()298            ret = Fuser(299                layer=layer,300                num_layers=num_layers301            )302            return ret303 304        out_dim = 64305        position_encoding = build_position_encoding()306        mask_downsampler = build_mask_downsampler()307        fuser = build_fuser()308        ret = MemoryEncoder(309            out_dim=out_dim,310            position_encoding=position_encoding,311            mask_downsampler=mask_downsampler,312            fuser=fuser,313        )314        return ret315 316    def inject_language_embd(self, inference_state, language_embd):317        num_frame = len(language_embd)318        num_obj = len(language_embd[0])319        mask_out = []320        for frame_idx in range(num_frame):321            frame_mask_out = []322            for obj_idx in range(num_obj):323                _language_embd = language_embd[frame_idx][obj_idx][None][None]324                _, _, out_mask_logits = self.sam2_model.add_language_embd(inference_state, frame_idx, obj_idx + 100, _language_embd)325                frame_mask_out.append(out_mask_logits)326            frame_mask_out = torch.cat(frame_mask_out, dim=1)327            mask_out.append(frame_mask_out)328        mask_out = torch.cat(mask_out, dim=0)329        return mask_out330 331 332    def language_embd_inference(self, inference_state, language_embd):333        num_frame = len(language_embd)334        num_obj = len(language_embd[0])335        mask_out = []336        with torch.autocast(device_type="cuda", dtype=torch.bfloat16):337            for frame_idx in range(num_frame):338                frame_mask_out = []339 340                for obj_idx in range(num_obj):341                    _language_embd = language_embd[frame_idx][obj_idx][None][None]342                    _, _, out_mask_logits = self.sam2_model.add_language_embd(343                        inference_state,344                        frame_idx,345                        obj_idx + 100,346                        _language_embd,347                        inference=True,348                    )349                    frame_mask_out.append(out_mask_logits)350                frame_mask_out = torch.cat(frame_mask_out, dim=1)351                mask_out.append(frame_mask_out)352 353 354            mask_out = []355            for out_frame_idx, out_obj_ids, out_mask_logits in self.sam2_model.propagate_in_video(inference_state):356                mask_out.append(out_mask_logits)357            mask_out = torch.cat(mask_out, dim=0)358        return mask_out359 360    def get_sam2_embeddings(self, images):361        return self.sam2_model.init_state(images)362 363    def forward(self, batch):364        raise NotImplementedError365 366    def preprocess_image(self, image: torch.Tensor, dtype=torch.bfloat16) -> torch.Tensor:367        image = image / 255.368 369        img_mean = torch.tensor(self.img_mean, dtype=dtype, device=image.device)[:, None, None]370        img_std = torch.tensor(self.img_std, dtype=dtype, device=image.device)[:, None, None]371        image -= img_mean372        image /= img_std373 374        return image375 376class MemoryAttentionLayer(nn.Module):377 378    def __init__(379        self,380        activation: str,381        cross_attention: nn.Module,382        d_model: int,383        dim_feedforward: int,384        dropout: float,385        pos_enc_at_attn: bool,386        pos_enc_at_cross_attn_keys: bool,387        pos_enc_at_cross_attn_queries: bool,388        self_attention: nn.Module,389    ):390        super().__init__()391        self.d_model = d_model392        self.dim_feedforward = dim_feedforward393        self.dropout_value = dropout394        self.self_attn = self_attention395        self.cross_attn_image = cross_attention396 397        # Implementation of Feedforward model398        self.linear1 = nn.Linear(d_model, dim_feedforward)399        self.dropout = nn.Dropout(dropout)400        self.linear2 = nn.Linear(dim_feedforward, d_model)401 402        self.norm1 = nn.LayerNorm(d_model)403        self.norm2 = nn.LayerNorm(d_model)404        self.norm3 = nn.LayerNorm(d_model)405        self.dropout1 = nn.Dropout(dropout)406        self.dropout2 = nn.Dropout(dropout)407        self.dropout3 = nn.Dropout(dropout)408 409        self.activation_str = activation410        self.activation = get_activation_fn(activation)411 412        # Where to add pos enc413        self.pos_enc_at_attn = pos_enc_at_attn414        self.pos_enc_at_cross_attn_queries = pos_enc_at_cross_attn_queries415        self.pos_enc_at_cross_attn_keys = pos_enc_at_cross_attn_keys416 417    def _forward_sa(self, tgt, query_pos):418        # Self-Attention419        tgt2 = self.norm1(tgt)420        q = k = tgt2 + query_pos if self.pos_enc_at_attn else tgt2421        tgt2 = self.self_attn(q, k, v=tgt2)422        tgt = tgt + self.dropout1(tgt2)423        return tgt424 425    def _forward_ca(self, tgt, memory, query_pos, pos, num_k_exclude_rope=0):426        kwds = {}427        if num_k_exclude_rope > 0:428            assert isinstance(self.cross_attn_image, RoPEAttention)429            kwds = {"num_k_exclude_rope": num_k_exclude_rope}430 431        # Cross-Attention432        tgt2 = self.norm2(tgt)433        tgt2 = self.cross_attn_image(434            q=tgt2 + query_pos if self.pos_enc_at_cross_attn_queries else tgt2,435            k=memory + pos if self.pos_enc_at_cross_attn_keys else memory,436            v=memory,437            **kwds,438        )439        tgt = tgt + self.dropout2(tgt2)440        return tgt441 442    def forward(443        self,444        tgt,445        memory,446        pos: Optional[Tensor] = None,447        query_pos: Optional[Tensor] = None,448        num_k_exclude_rope: int = 0,449    ) -> torch.Tensor:450 451        # Self-Attn, Cross-Attn452        tgt = self._forward_sa(tgt, query_pos)453        tgt = self._forward_ca(tgt, memory, query_pos, pos, num_k_exclude_rope)454        # MLP455        tgt2 = self.norm3(tgt)456        tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))457        tgt = tgt + self.dropout3(tgt2)458        return tgt459 460 461class MemoryAttention(nn.Module):462    def __init__(463        self,464        d_model: int,465        pos_enc_at_input: bool,466        layer: nn.Module,467        num_layers: int,468        batch_first: bool = True,  # Do layers expect batch first input?469    ):470        super().__init__()471        self.d_model = d_model472        self.layers = get_clones(layer, num_layers)473        self.num_layers = num_layers474        self.norm = nn.LayerNorm(d_model)475        self.pos_enc_at_input = pos_enc_at_input476        self.batch_first = batch_first477 478    def forward(479        self,480        curr: torch.Tensor,  # self-attention inputs481        memory: torch.Tensor,  # cross-attention inputs482        curr_pos: Optional[Tensor] = None,  # pos_enc for self-attention inputs483        memory_pos: Optional[Tensor] = None,  # pos_enc for cross-attention inputs484        num_obj_ptr_tokens: int = 0,  # number of object pointer *tokens*485    ):486        if isinstance(curr, list):487            assert isinstance(curr_pos, list)488            assert len(curr) == len(curr_pos) == 1489            curr, curr_pos = (490                curr[0],491                curr_pos[0],492            )493 494        assert (495            curr.shape[1] == memory.shape[1]496        ), "Batch size must be the same for curr and memory"497 498        output = curr499        if self.pos_enc_at_input and curr_pos is not None:500            output = output + 0.1 * curr_pos501 502        if self.batch_first:503            # Convert to batch first504            output = output.transpose(0, 1)505            curr_pos = curr_pos.transpose(0, 1)506            memory = memory.transpose(0, 1)507            memory_pos = memory_pos.transpose(0, 1)508 509        for layer in self.layers:510            kwds = {}511            if isinstance(layer.cross_attn_image, RoPEAttention):512                kwds = {"num_k_exclude_rope": num_obj_ptr_tokens}513 514            output = layer(515                tgt=output,516                memory=memory,517                pos=memory_pos,518                query_pos=curr_pos,519                **kwds,520            )521        normed_output = self.norm(output)522 523        if self.batch_first:524            # Convert back to seq first525            normed_output = normed_output.transpose(0, 1)526            curr_pos = curr_pos.transpose(0, 1)527 528        return normed_output529 530class MaskDownSampler(nn.Module):531    """532    Progressively downsample a mask by total_stride, each time by stride.533    Note that LayerNorm is applied per *token*, like in ViT.534 535    With each downsample (by a factor stride**2), channel capacity increases by the same factor.536    In the end, we linearly project to embed_dim channels.537    """538 539    def __init__(540        self,541        embed_dim=256,542        kernel_size=4,543        stride=4,544        padding=0,545        total_stride=16,546        activation=nn.GELU,547    ):548        super().__init__()549        num_layers = int(math.log2(total_stride) // math.log2(stride))550        assert stride**num_layers == total_stride551        self.encoder = nn.Sequential()552        mask_in_chans, mask_out_chans = 1, 1553        for _ in range(num_layers):554            mask_out_chans = mask_in_chans * (stride**2)555            self.encoder.append(556                nn.Conv2d(557                    mask_in_chans,558                    mask_out_chans,559                    kernel_size=kernel_size,560                    stride=stride,561                    padding=padding,562                )563            )564            self.encoder.append(LayerNorm2d(mask_out_chans))565            self.encoder.append(activation())566            mask_in_chans = mask_out_chans567 568        self.encoder.append(nn.Conv2d(mask_out_chans, embed_dim, kernel_size=1))569 570    def forward(self, x):571        return self.encoder(x)572 573 574# Lightly adapted from ConvNext (https://github.com/facebookresearch/ConvNeXt)575class CXBlock(nn.Module):576    r"""ConvNeXt Block. There are two equivalent implementations:577    (1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)578    (2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back579    We use (2) as we find it slightly faster in PyTorch580 581    Args:582        dim (int): Number of input channels.583        drop_path (float): Stochastic depth rate. Default: 0.0584        layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.585    """586 587    def __init__(588        self,589        dim,590        kernel_size=7,591        padding=3,592        drop_path=0.0,593        layer_scale_init_value=1e-6,594        use_dwconv=True,595    ):596        super().__init__()597        self.dwconv = nn.Conv2d(598            dim,599            dim,600            kernel_size=kernel_size,601            padding=padding,602            groups=dim if use_dwconv else 1,603        )  # depthwise conv604        self.norm = LayerNorm2d(dim, eps=1e-6)605        self.pwconv1 = nn.Linear(606            dim, 4 * dim607        )  # pointwise/1x1 convs, implemented with linear layers608        self.act = nn.GELU()609        self.pwconv2 = nn.Linear(4 * dim, dim)610        # self.gamma = (611        self.g_weight = (612            nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)613            if layer_scale_init_value > 0614            else None615        )616        self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()617 618    def forward(self, x):619        input = x620        x = self.dwconv(x)621        x = self.norm(x)622        x = x.permute(0, 2, 3, 1)  # (N, C, H, W) -> (N, H, W, C)623        x = self.pwconv1(x)624        x = self.act(x)625        x = self.pwconv2(x)626        if self.g_weight is not None:627            x = self.g_weight * x628        x = x.permute(0, 3, 1, 2)  # (N, H, W, C) -> (N, C, H, W)629 630        x = input + self.drop_path(x)631        return x632 633 634class Fuser(nn.Module):635    def __init__(self, layer, num_layers, dim=None, input_projection=False):636        super().__init__()637        self.proj = nn.Identity()638        self.layers = get_clones(layer, num_layers)639 640        if input_projection:641            assert dim is not None642            self.proj = nn.Conv2d(dim, dim, kernel_size=1)643 644    def forward(self, x):645        # normally x: (N, C, H, W)646        x = self.proj(x)647        for layer in self.layers:648            x = layer(x)649        return x650 651 652class MemoryEncoder(nn.Module):653    def __init__(654        self,655        out_dim,656        mask_downsampler,657        fuser,658        position_encoding,659        in_dim=256,  # in_dim of pix_feats660    ):661        super().__init__()662 663        self.mask_downsampler = mask_downsampler664 665        self.pix_feat_proj = nn.Conv2d(in_dim, in_dim, kernel_size=1)666        self.fuser = fuser667        self.position_encoding = position_encoding668        self.out_proj = nn.Identity()669        if out_dim != in_dim:670            self.out_proj = nn.Conv2d(in_dim, out_dim, kernel_size=1)671 672    def forward(673        self,674        pix_feat: torch.Tensor,675        masks: torch.Tensor,676        skip_mask_sigmoid: bool = False,677    ) -> Tuple[torch.Tensor, torch.Tensor]:678        ## Process masks679        # sigmoid, so that less domain shift from gt masks which are bool680        if not skip_mask_sigmoid:681            masks = F.sigmoid(masks)682        masks = self.mask_downsampler(masks)683 684        ## Fuse pix_feats and downsampled masks685        # in case the visual features are on CPU, cast them to CUDA686        pix_feat = pix_feat.to(masks.device)687 688        x = self.pix_feat_proj(pix_feat)689        x = x + masks690        x = self.fuser(x)691        x = self.out_proj(x)692 693        pos = self.position_encoding(x).to(x.dtype)694 695        return {"vision_features": x, "vision_pos_enc": [pos]}696 697 698class ImageEncoder(nn.Module):699    def __init__(700        self,701        trunk: nn.Module,702        neck: nn.Module,703        scalp: int = 0,704    ):705        super().__init__()706        self.trunk = trunk707        self.neck = neck708        self.scalp = scalp709        assert (710            self.trunk.channel_list == self.neck.backbone_channel_list711        ), f"Channel dims of trunk and neck do not match. Trunk: {self.trunk.channel_list}, neck: {self.neck.backbone_channel_list}"712 713    def forward(self, sample: torch.Tensor):714        # Forward through backbone715        features, pos = self.neck(self.trunk(sample))716        if self.scalp > 0:717            # Discard the lowest resolution features718            features, pos = features[: -self.scalp], pos[: -self.scalp]719 720        src = features[-1]721        output = {722            "vision_features": src,723            "vision_pos_enc": pos,724            "backbone_fpn": features,725        }726        return output727 728 729class FpnNeck(nn.Module):730    """731    A modified variant of Feature Pyramid Network (FPN) neck732    (we remove output conv and also do bicubic interpolation similar to ViT733    pos embed interpolation)734    """735 736    def __init__(737        self,738        position_encoding: nn.Module,739        d_model: int,740        backbone_channel_list: List[int],741        kernel_size: int = 1,742        stride: int = 1,743        padding: int = 0,744        fpn_interp_model: str = "bilinear",745        fuse_type: str = "sum",746        fpn_top_down_levels: Optional[List[int]] = None,747    ):748        """Initialize the neck749        :param trunk: the backbone750        :param position_encoding: the positional encoding to use751        :param d_model: the dimension of the model752        :param neck_norm: the normalization to use753        """754        super().__init__()755        self.position_encoding = position_encoding756        self.convs = nn.ModuleList()757        self.backbone_channel_list = backbone_channel_list758        for dim in backbone_channel_list:759            current = nn.Sequential()760            current.add_module(761                "conv",762                nn.Conv2d(763                    in_channels=dim,764                    out_channels=d_model,765                    kernel_size=kernel_size,766                    stride=stride,767                    padding=padding,768                ),769            )770 771            self.convs.append(current)772        self.fpn_interp_model = fpn_interp_model773        assert fuse_type in ["sum", "avg"]774        self.fuse_type = fuse_type775 776        # levels to have top-down features in its outputs777        # e.g. if fpn_top_down_levels is [2, 3], then only outputs of level 2 and 3778        # have top-down propagation, while outputs of level 0 and level 1 have only779        # lateral features from the same backbone level.780        if fpn_top_down_levels is None:781            # default is to have top-down features on all levels782            fpn_top_down_levels = range(len(self.convs))783        self.fpn_top_down_levels = list(fpn_top_down_levels)784 785    def forward(self, xs: List[torch.Tensor]):786 787        out = [None] * len(self.convs)788        pos = [None] * len(self.convs)789        assert len(xs) == len(self.convs)790        # fpn forward pass791        # see https://github.com/facebookresearch/detectron2/blob/main/detectron2/modeling/backbone/fpn.py792        prev_features = None793        # forward in top-down order (from low to high resolution)794        n = len(self.convs) - 1795        for i in range(n, -1, -1):796            x = xs[i]797            lateral_features = self.convs[n - i](x)798            if i in self.fpn_top_down_levels and prev_features is not None:799                top_down_features = F.interpolate(800                    prev_features.to(dtype=torch.float32),801                    scale_factor=2.0,802                    mode=self.fpn_interp_model,803                    align_corners=(804                        None if self.fpn_interp_model == "nearest" else False805                    ),806                    antialias=False,807                )808                prev_features = lateral_features + top_down_features809                if self.fuse_type == "avg":810                    prev_features /= 2811            else:812                prev_features = lateral_features813            x_out = prev_features814            out[i] = x_out815            pos[i] = self.position_encoding(x_out).to(x_out.dtype)816 817        return out, pos818 819def window_partition(x, window_size):820    """821    Partition into non-overlapping windows with padding if needed.822    Args:823        x (tensor): input tokens with [B, H, W, C].824        window_size (int): window size.825    Returns:826        windows: windows after partition with [B * num_windows, window_size, window_size, C].827        (Hp, Wp): padded height and width before partition828    """829    B, H, W, C = x.shape830 831    pad_h = (window_size - H % window_size) % window_size832    pad_w = (window_size - W % window_size) % window_size833    if pad_h > 0 or pad_w > 0:834        x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h))835    Hp, Wp = H + pad_h, W + pad_w836 837    x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C)838    windows = (839        x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)840    )841    return windows, (Hp, Wp)842 843 844def window_unpartition(windows, window_size, pad_hw, hw):845    """846    Window unpartition into original sequences and removing padding.847    Args:848        x (tensor): input tokens with [B * num_windows, window_size, window_size, C].849        window_size (int): window size.850        pad_hw (Tuple): padded height and width (Hp, Wp).851        hw (Tuple): original height and width (H, W) before padding.852    Returns:853        x: unpartitioned sequences with [B, H, W, C].854    """855    Hp, Wp = pad_hw856    H, W = hw857    B = windows.shape[0] // (Hp * Wp // window_size // window_size)858    x = windows.view(859        B, Hp // window_size, Wp // window_size, window_size, window_size, -1860    )861    x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1)862 863    if Hp > H or Wp > W:864        x = x[:, :H, :W, :].contiguous()865    return x866 867 868class PatchEmbed(nn.Module):869    """870    Image to Patch Embedding.871    """872 873    def __init__(874        self,875        kernel_size: Tuple[int, ...] = (7, 7),876        stride: Tuple[int, ...] = (4, 4),877        padding: Tuple[int, ...] = (3, 3),878        in_chans: int = 3,879        embed_dim: int = 768,880    ):881        """882        Args:883            kernel_size (Tuple): kernel size of the projection layer.884            stride (Tuple): stride of the projection layer.885            padding (Tuple): padding size of the projection layer.886            in_chans (int): Number of input image channels.887            embed_dim (int):  embed_dim (int): Patch embedding dimension.888        """889        super().__init__()890        self.proj = nn.Conv2d(891            in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding892        )893 894    def forward(self, x: torch.Tensor) -> torch.Tensor:895        x = self.proj(x)896        # B C H W -> B H W C897        x = x.permute(0, 2, 3, 1)898        return x899 900def do_pool(x: torch.Tensor, pool: nn.Module, norm: nn.Module = None) -> torch.Tensor:901    if pool is None:902        return x903    # (B, H, W, C) -> (B, C, H, W)904    x = x.permute(0, 3, 1, 2)905    x = pool(x)906    # (B, C, H', W') -> (B, H', W', C)907    x = x.permute(0, 2, 3, 1)908    if norm:909        x = norm(x)910 911    return x912 913 914class MultiScaleAttention(nn.Module):915    def __init__(916        self,917        dim: int,918        dim_out: int,919        num_heads: int,920        q_pool: nn.Module = None,921    ):922        super().__init__()923 924        self.dim = dim925        self.dim_out = dim_out926 927        self.num_heads = num_heads928        head_dim = dim_out // num_heads929        self.scale = head_dim**-0.5930 931        self.q_pool = q_pool932        self.qkv = nn.Linear(dim, dim_out * 3)933        self.proj = nn.Linear(dim_out, dim_out)934 935    def forward(self, x: torch.Tensor) -> torch.Tensor:936        B, H, W, _ = x.shape937        # qkv with shape (B, H * W, 3, nHead, C)938        qkv = self.qkv(x).reshape(B, H * W, 3, self.num_heads, -1)939        # q, k, v with shape (B, H * W, nheads, C)940        q, k, v = torch.unbind(qkv, 2)941 942        # Q pooling (for downsample at stage changes)943        if self.q_pool:944            q = do_pool(q.reshape(B, H, W, -1), self.q_pool)945            H, W = q.shape[1:3]  # downsampled shape946            q = q.reshape(B, H * W, self.num_heads, -1)947 948        # Torch's SDPA expects [B, nheads, H*W, C] so we transpose949        x = F.scaled_dot_product_attention(950            q.transpose(1, 2),951            k.transpose(1, 2),952            v.transpose(1, 2),953        )954        # Transpose back955        x = x.transpose(1, 2)956        x = x.reshape(B, H, W, -1)957 958        x = self.proj(x)959 960        return x961 962 963class MultiScaleBlock(nn.Module):964    def __init__(965        self,966        dim: int,967        dim_out: int,968        num_heads: int,969        mlp_ratio: float = 4.0,970        drop_path: float = 0.0,971        norm_layer: Union[nn.Module, str] = "LayerNorm",972        q_stride: Tuple[int, int] = None,973        act_layer: nn.Module = nn.GELU,974        window_size: int = 0,975    ):976        super().__init__()977 978        if isinstance(norm_layer, str):979            norm_layer = partial(getattr(nn, norm_layer), eps=1e-6)980 981        self.dim = dim982        self.dim_out = dim_out983        self.norm1 = norm_layer(dim)984 985        self.window_size = window_size986 987        self.pool, self.q_stride = None, q_stride988        if self.q_stride:989            self.pool = nn.MaxPool2d(990                kernel_size=q_stride, stride=q_stride, ceil_mode=False991            )992 993        self.attn = MultiScaleAttention(994            dim,995            dim_out,996            num_heads=num_heads,997            q_pool=self.pool,998        )999        self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()1000 1001        self.norm2 = norm_layer(dim_out)1002        self.mlp = MLP(1003            dim_out,1004            int(dim_out * mlp_ratio),1005            dim_out,1006            num_layers=2,1007            activation=act_layer,1008        )1009 1010        if dim != dim_out:1011            self.proj = nn.Linear(dim, dim_out)1012 1013    def forward(self, x: torch.Tensor) -> torch.Tensor:1014        shortcut = x  # B, H, W, C1015        x = self.norm1(x)1016 1017        # Skip connection1018        if self.dim != self.dim_out:1019            shortcut = do_pool(self.proj(x), self.pool)1020 1021        # Window partition1022        window_size = self.window_size1023        if window_size > 0:1024            H, W = x.shape[1], x.shape[2]1025            x, pad_hw = window_partition(x, window_size)1026 1027        # Window Attention + Q Pooling (if stage change)1028        x = self.attn(x)1029        if self.q_stride:1030            # Shapes have changed due to Q pooling1031            window_size = self.window_size // self.q_stride[0]1032            H, W = shortcut.shape[1:3]1033 1034            pad_h = (window_size - H % window_size) % window_size1035            pad_w = (window_size - W % window_size) % window_size1036            pad_hw = (H + pad_h, W + pad_w)1037 1038        # Reverse window partition1039        if self.window_size > 0:1040            x = window_unpartition(x, window_size, pad_hw, (H, W))1041 1042        x = shortcut + self.drop_path(x)1043        # MLP1044        x = x + self.drop_path(self.mlp(self.norm2(x)))1045        return x1046 1047 1048class Hiera(nn.Module):1049    """1050    Reference: https://arxiv.org/abs/2306.009891051    """1052 1053    def __init__(1054        self,1055        embed_dim: int = 96,  # initial embed dim1056        num_heads: int = 1,  # initial number of heads1057        drop_path_rate: float = 0.0,  # stochastic depth1058        q_pool: int = 3,  # number of q_pool stages1059        q_stride: Tuple[int, int] = (2, 2),  # downsample stride bet. stages1060        stages: Tuple[int, ...] = (2, 3, 16, 3),  # blocks per stage1061        dim_mul: float = 2.0,  # dim_mul factor at stage shift1062        head_mul: float = 2.0,  # head_mul factor at stage shift1063        window_pos_embed_bkg_spatial_size: Tuple[int, int] = (14, 14),1064        # window size per stage, when not using global att.1065        window_spec: Tuple[int, ...] = (1066            8,1067            4,1068            14,1069            7,1070        ),1071        # global attn in these blocks1072        global_att_blocks: Tuple[int, ...] = (1073            12,1074            16,1075            20,1076        ),1077        return_interm_layers=True,  # return feats from every stage1078    ):1079        super().__init__()1080 1081        assert len(stages) == len(window_spec)1082        self.window_spec = window_spec1083 1084        depth = sum(stages)1085        self.q_stride = q_stride1086        self.stage_ends = [sum(stages[:i]) - 1 for i in range(1, len(stages) + 1)]1087        assert 0 <= q_pool <= len(self.stage_ends[:-1])1088        self.q_pool_blocks = [x + 1 for x in self.stage_ends[:-1]][:q_pool]1089        self.return_interm_layers = return_interm_layers1090 1091        self.patch_embed = PatchEmbed(1092            embed_dim=embed_dim,1093        )1094        # Which blocks have global att?1095        self.global_att_blocks = global_att_blocks1096 1097        # Windowed positional embedding (https://arxiv.org/abs/2311.05613)1098        self.window_pos_embed_bkg_spatial_size = window_pos_embed_bkg_spatial_size1099        self.pos_embed = nn.Parameter(1100            torch.zeros(1, embed_dim, *self.window_pos_embed_bkg_spatial_size)1101        )1102        self.pos_embed_window = nn.Parameter(1103            torch.zeros(1, embed_dim, self.window_spec[0], self.window_spec[0])1104        )1105 1106        dpr = [1107            x.item() for x in torch.linspace(0, drop_path_rate, depth)1108        ]  # stochastic depth decay rule1109 1110        cur_stage = 11111        self.blocks = nn.ModuleList()1112 1113        for i in range(depth):1114            dim_out = embed_dim1115            # lags by a block, so first block of1116            # next stage uses an initial window size1117            # of previous stage and final window size of current stage1118            window_size = self.window_spec[cur_stage - 1]1119 1120            if self.global_att_blocks is not None:1121                window_size = 0 if i in self.global_att_blocks else window_size1122 1123            if i - 1 in self.stage_ends:1124                dim_out = int(embed_dim * dim_mul)1125                num_heads = int(num_heads * head_mul)1126                cur_stage += 11127 1128            block = MultiScaleBlock(1129                dim=embed_dim,1130                dim_out=dim_out,1131                num_heads=num_heads,1132                drop_path=dpr[i],1133                q_stride=self.q_stride if i in self.q_pool_blocks else None,1134                window_size=window_size,1135            )1136 1137            embed_dim = dim_out1138            self.blocks.append(block)1139 1140        self.channel_list = (1141            [self.blocks[i].dim_out for i in self.stage_ends[::-1]]1142            if return_interm_layers1143            else [self.blocks[-1].dim_out]1144        )1145 1146    def _get_pos_embed(self, hw: Tuple[int, int]) -> torch.Tensor:1147        h, w = hw1148        window_embed = self.pos_embed_window1149        pos_embed = F.interpolate(self.pos_embed, size=(h, w), mode="bicubic")1150        pos_embed = pos_embed + window_embed.tile(1151            [x // y for x, y in zip(pos_embed.shape, window_embed.shape)]1152        )1153        pos_embed = pos_embed.permute(0, 2, 3, 1)1154        return pos_embed1155 1156    def forward(self, x: torch.Tensor) -> List[torch.Tensor]:1157        x = self.patch_embed(x)1158        # x: (B, H, W, C)1159 1160        # Add pos embed1161        x = x + self._get_pos_embed(x.shape[1:3])1162 1163        outputs = []1164        for i, blk in enumerate(self.blocks):1165            x = blk(x)1166            if (i == self.stage_ends[-1]) or (1167                i in self.stage_ends and self.return_interm_layers1168            ):1169                feats = x.permute(0, 3, 1, 2)1170                outputs.append(feats)1171 1172        return outputs1173 1174class TwoWayTransformer(nn.Module):1175    def __init__(1176        self,1177        depth: int,1178        embedding_dim: int,1179        num_heads: int,1180        mlp_dim: int,1181        activation: Type[nn.Module] = nn.ReLU,1182        attention_downsample_rate: int = 2,1183    ) -> None:1184        """1185        A transformer decoder that attends to an input image using1186        queries whose positional embedding is supplied.1187 1188        Args:1189          depth (int): number of layers in the transformer1190          embedding_dim (int): the channel dimension for the input embeddings1191          num_heads (int): the number of heads for multihead attention. Must1192            divide embedding_dim1193          mlp_dim (int): the channel dimension internal to the MLP block1194          activation (nn.Module): the activation to use in the MLP block1195        """1196        super().__init__()1197        self.depth = depth1198        self.embedding_dim = embedding_dim1199        self.num_heads = num_heads1200        self.mlp_dim = mlp_dim

Showing the first 1,200 of 4103 lines. Download the file for the rest.