(self,
vis_token_dim,
mask_dim,
backbone, # placeholder
bn_group,
activation='gelu',
task_sp_list=(),
maskformer_num_feature_levels=3,
pixel_decoder_cfg=None)
| 167 | |
| 168 | class MoreSimpleFPN(SimpleFPN): |
| 169 | def __init__(self, |
| 170 | vis_token_dim, |
| 171 | mask_dim, |
| 172 | backbone, # placeholder |
| 173 | bn_group, |
| 174 | activation='gelu', |
| 175 | task_sp_list=(), |
| 176 | maskformer_num_feature_levels=3, |
| 177 | pixel_decoder_cfg=None): |
| 178 | super(SimpleFPN, self).__init__() |
| 179 | self.task_sp_list = task_sp_list |
| 180 | |
| 181 | self.embed_dim = backbone.embed_dim |
| 182 | self.mask_dim = mask_dim |
| 183 | self.vis_token_dim = vis_token_dim |
| 184 | self.pixel_decoder_cfg = pixel_decoder_cfg |
| 185 | |
| 186 | fpn1 = nn.Sequential( |
| 187 | nn.ConvTranspose2d(self.embed_dim, self.embed_dim, kernel_size=2, stride=2), |
| 188 | Norm2d(self.embed_dim), |
| 189 | _get_activation(activation), |
| 190 | nn.ConvTranspose2d(self.embed_dim, self.vis_token_dim, kernel_size=2, stride=2), |
| 191 | # in compliance with decoder dim request |
| 192 | Norm2d(self.vis_token_dim), |
| 193 | ) |
| 194 | |
| 195 | fpn2 = nn.Sequential( |
| 196 | nn.ConvTranspose2d(self.embed_dim, self.vis_token_dim, kernel_size=2, stride=2), |
| 197 | # in compliance with decoder dim request |
| 198 | Norm2d(self.vis_token_dim), |
| 199 | ) |
| 200 | |
| 201 | fpn3 = nn.Sequential( |
| 202 | # in compliance with decoder dim request |
| 203 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 204 | Norm2d(self.vis_token_dim), |
| 205 | ) |
| 206 | |
| 207 | fpn4 = nn.Sequential( |
| 208 | nn.MaxPool2d(kernel_size=2, stride=2), |
| 209 | # in compliance with decoder dim request |
| 210 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 211 | Norm2d(self.vis_token_dim), |
| 212 | ) |
| 213 | |
| 214 | self.fpns = nn.ModuleList([fpn1, fpn2, fpn3, fpn4]) |
| 215 | |
| 216 | if self.pixel_decoder_cfg is None: |
| 217 | self.mask_features = nn.Conv2d(self.vis_token_dim, self.mask_dim, kernel_size=1, stride=1, padding=0) |
| 218 | c2_xavier_fill(self.mask_features) |
| 219 | else: |
| 220 | input_shape = {name: ShapeSpec(channels=self.vis_token_dim, stride=[4, 8, 16, 32][i]) |
| 221 | for i, name in enumerate(["fpn1", "fpn2", "fpn3", "fpn4"])} |
| 222 | self.pixel_decoder = MSDeformAttnPixelDecoder(conv_dim=self.vis_token_dim, |
| 223 | input_shape=input_shape, |
| 224 | mask_dim=self.mask_dim, |
| 225 | **pixel_decoder_cfg) |
| 226 |
nothing calls this directly
no test coverage detected