(self,
vis_token_dim,
mask_dim,
backbone, # placeholder
bn_group,
pixel_decoder_cfg=None)
| 83 | |
| 84 | class SimpleFPN(nn.Module): |
| 85 | def __init__(self, |
| 86 | vis_token_dim, |
| 87 | mask_dim, |
| 88 | backbone, # placeholder |
| 89 | bn_group, |
| 90 | pixel_decoder_cfg=None): |
| 91 | super(SimpleFPN, self).__init__() |
| 92 | self.embed_dim = backbone.embed_dim |
| 93 | self.mask_dim = mask_dim |
| 94 | self.vis_token_dim = vis_token_dim |
| 95 | self.pixel_decoder_cfg = pixel_decoder_cfg |
| 96 | |
| 97 | fpn1 = nn.Sequential( |
| 98 | nn.ConvTranspose2d(self.embed_dim, self.embed_dim, kernel_size=2, stride=2), |
| 99 | Norm2d(self.embed_dim), |
| 100 | nn.GELU(), |
| 101 | nn.ConvTranspose2d(self.embed_dim, self.embed_dim, kernel_size=2, stride=2), |
| 102 | # in compliance with decoder dim request |
| 103 | Norm2d(self.embed_dim), |
| 104 | nn.GELU(), |
| 105 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 106 | Norm2d(self.vis_token_dim), |
| 107 | ) |
| 108 | |
| 109 | fpn2 = nn.Sequential( |
| 110 | nn.ConvTranspose2d(self.embed_dim, self.embed_dim, kernel_size=2, stride=2), |
| 111 | # in compliance with decoder dim request |
| 112 | Norm2d(self.embed_dim), |
| 113 | nn.GELU(), |
| 114 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 115 | Norm2d(self.vis_token_dim), |
| 116 | ) |
| 117 | |
| 118 | fpn3 = nn.Sequential( |
| 119 | # in compliance with decoder dim request |
| 120 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 121 | Norm2d(self.vis_token_dim), |
| 122 | ) |
| 123 | |
| 124 | fpn4 = nn.Sequential( |
| 125 | nn.MaxPool2d(kernel_size=2, stride=2), |
| 126 | # in compliance with decoder dim request |
| 127 | nn.Conv2d(self.embed_dim, self.vis_token_dim, kernel_size=1, stride=1, padding=0), |
| 128 | Norm2d(self.vis_token_dim), |
| 129 | ) |
| 130 | |
| 131 | self.fpns = nn.ModuleList([fpn1, fpn2, fpn3, fpn4]) |
| 132 | |
| 133 | if self.pixel_decoder_cfg is None: |
| 134 | self.mask_features = nn.Conv2d(self.vis_token_dim, self.mask_dim, kernel_size=1, stride=1, padding=0) |
| 135 | c2_xavier_fill(self.mask_features) |
| 136 | else: |
| 137 | input_shape = {name: ShapeSpec(channels=self.vis_token_dim, stride=[4, 8, 16, 32][i]) |
| 138 | for i, name in enumerate(["fpn1", "fpn2", "fpn3", "fpn4"])} |
| 139 | self.pixel_decoder = MSDeformAttnPixelDecoder(conv_dim=self.vis_token_dim, |
| 140 | input_shape=input_shape, |
| 141 | mask_dim=self.mask_dim, |
| 142 | **pixel_decoder_cfg) |
nothing calls this directly
no test coverage detected