(self,
vis_token_dim,
mask_dim,
backbone, # placeholder
bn_group,
s_kernel,
pixel_decoder_cfg=None)
| 267 | |
| 268 | class ShuffleFPN(SimpleFPN): |
| 269 | def __init__(self, |
| 270 | vis_token_dim, |
| 271 | mask_dim, |
| 272 | backbone, # placeholder |
| 273 | bn_group, |
| 274 | s_kernel, |
| 275 | pixel_decoder_cfg=None): |
| 276 | super(SimpleFPN, self).__init__() |
| 277 | self.embed_dim = backbone.embed_dim |
| 278 | self.mask_dim = mask_dim |
| 279 | self.vis_token_dim = vis_token_dim |
| 280 | self.pixel_decoder_cfg = pixel_decoder_cfg |
| 281 | |
| 282 | fpn1 = nn.Sequential( |
| 283 | nn.Conv2d(self.embed_dim, self.vis_token_dim * 8, kernel_size=s_kernel, stride=1, padding=s_kernel // 2), |
| 284 | nn.PixelShuffle(2), |
| 285 | Norm2d(self.vis_token_dim*2), |
| 286 | nn.GELU(), |
| 287 | nn.Conv2d(self.vis_token_dim*2, self.vis_token_dim * 4, kernel_size=s_kernel, stride=1, padding=s_kernel // 2), |
| 288 | nn.PixelShuffle(2), |
| 289 | Norm2d(self.vis_token_dim), |
| 290 | ) |
| 291 | |
| 292 | fpn2 = nn.Sequential( |
| 293 | nn.Conv2d(self.embed_dim, self.vis_token_dim * 4, kernel_size=s_kernel, stride=1, padding=s_kernel // 2), |
| 294 | nn.PixelShuffle(2), |
| 295 | Norm2d(self.vis_token_dim), |
| 296 | ) |
| 297 | |
| 298 | fpn3 = nn.Sequential( |
| 299 | # in compliance with decoder dim request |
| 300 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 301 | Norm2d(self.vis_token_dim), |
| 302 | ) |
| 303 | |
| 304 | fpn4 = nn.Sequential( |
| 305 | nn.MaxPool2d(kernel_size=2, stride=2), |
| 306 | # in compliance with decoder dim request |
| 307 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 308 | Norm2d(self.vis_token_dim), |
| 309 | ) |
| 310 | |
| 311 | self.fpns = nn.ModuleList([fpn1, fpn2, fpn3, fpn4]) |
| 312 | |
| 313 | if self.pixel_decoder_cfg is None: |
| 314 | self.mask_features = nn.Conv2d(self.vis_token_dim, self.mask_dim, kernel_size=1, stride=1, padding=0) |
| 315 | c2_xavier_fill(self.mask_features) |
| 316 | else: |
| 317 | input_shape = {name: ShapeSpec(channels=self.vis_token_dim, stride=[4, 8, 16, 32][i]) |
| 318 | for i, name in enumerate(["fpn1", "fpn2", "fpn3", "fpn4"])} |
| 319 | self.pixel_decoder = MSDeformAttnPixelDecoder(conv_dim=self.vis_token_dim, |
| 320 | input_shape=input_shape, |
| 321 | mask_dim=self.mask_dim, |
| 322 | **pixel_decoder_cfg) |
| 323 | |
| 324 | self.maskformer_num_feature_levels = 3 # always use 3 scales |
| 325 | |
| 326 |
nothing calls this directly
no test coverage detected