ByteDance/Sa2VA-InternVL3-8B
6106
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