(self,
mask_dim,
backbone, # placeholder
bn_group,
activation='gelu',
task_sp_list=(),
pixel_decoder_cfg=None,
mask_forward=True
)
| 229 | |
| 230 | class SimpleNeck(nn.Module): |
| 231 | def __init__(self, |
| 232 | mask_dim, |
| 233 | backbone, # placeholder |
| 234 | bn_group, |
| 235 | activation='gelu', |
| 236 | task_sp_list=(), |
| 237 | pixel_decoder_cfg=None, |
| 238 | mask_forward=True |
| 239 | ): |
| 240 | super(SimpleNeck, self).__init__() |
| 241 | self.task_sp_list = task_sp_list |
| 242 | |
| 243 | self.vis_token_dim = self.embed_dim = backbone.embed_dim |
| 244 | self.mask_dim = mask_dim |
| 245 | self.pixel_decoder_cfg = pixel_decoder_cfg |
| 246 | |
| 247 | self.mask_map = nn.Sequential( |
| 248 | nn.ConvTranspose2d(self.embed_dim, self.embed_dim, kernel_size=2, stride=2), |
| 249 | Norm2d(self.embed_dim), |
| 250 | _get_activation(activation), |
| 251 | nn.ConvTranspose2d(self.embed_dim, self.mask_dim, kernel_size=2, stride=2), |
| 252 | ) if mask_dim else False |
| 253 | |
| 254 | self.maskformer_num_feature_levels = 1 # always use 3 scales |
| 255 | |
| 256 | self.mask_forward = mask_forward |
| 257 | |
| 258 | def forward(self, features): |
| 259 | if self.mask_map and self.mask_forward: |
nothing calls this directly
no test coverage detected