| 381 | |
| 382 | class SimpleMFFPN(SimpleFPN): |
| 383 | def __init__(self, |
| 384 | vis_token_dim, |
| 385 | mask_dim, |
| 386 | backbone, # placeholder |
| 387 | bn_group, |
| 388 | pixel_decoder_cfg=None): |
| 389 | super(SimpleFPN, self).__init__() |
| 390 | self.embed_dim = backbone.embed_dim |
| 391 | self.mask_dim = mask_dim |
| 392 | self.vis_token_dim = vis_token_dim |
| 393 | self.pixel_decoder_cfg = pixel_decoder_cfg |
| 394 | |
| 395 | fpn1 = nn.Sequential( |
| 396 | nn.ConvTranspose2d(self.embed_dim, self.mask_dim, kernel_size=2, stride=2), |
| 397 | Norm2d(self.mask_dim), |
| 398 | nn.GELU(), |
| 399 | nn.ConvTranspose2d(self.mask_dim, self.mask_dim, kernel_size=2, stride=2), |
| 400 | ) |
| 401 | |
| 402 | fpn2 = nn.Sequential( |
| 403 | nn.ConvTranspose2d(self.embed_dim, self.vis_token_dim, kernel_size=2, stride=2), |
| 404 | # in compliance with decoder dim request |
| 405 | Norm2d(self.vis_token_dim), |
| 406 | ) |
| 407 | |
| 408 | fpn3 = nn.Sequential( |
| 409 | # in compliance with decoder dim request |
| 410 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 411 | Norm2d(self.vis_token_dim), |
| 412 | ) |
| 413 | |
| 414 | fpn4 = nn.Sequential( |
| 415 | nn.MaxPool2d(kernel_size=2, stride=2), |
| 416 | # in compliance with decoder dim request |
| 417 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 418 | Norm2d(self.vis_token_dim), |
| 419 | ) |
| 420 | |
| 421 | self.fpns = nn.ModuleList([fpn2, fpn3, fpn4]) |
| 422 | |
| 423 | if self.pixel_decoder_cfg is None: |
| 424 | self.mask_features = fpn1 |
| 425 | else: |
| 426 | raise |
| 427 | |
| 428 | self.maskformer_num_feature_levels = 3 # always use 3 scales |
| 429 | |
| 430 | def forward(self, features): |
| 431 | x = features['backbone_output'] |