| 133 | |
| 134 | @SEM_SEG_HEADS_REGISTRY.register() |
| 135 | class MaskDINOEncoder(nn.Module): |
| 136 | @configurable |
| 137 | def __init__( |
| 138 | self, |
| 139 | input_shape: Dict[str, ShapeSpec], |
| 140 | *, |
| 141 | transformer_dropout: float, |
| 142 | transformer_nheads: int, |
| 143 | transformer_dim_feedforward: int, |
| 144 | transformer_enc_layers: int, |
| 145 | conv_dim: int, |
| 146 | mask_dim: int, |
| 147 | norm: Optional[Union[str, Callable]] = None, |
| 148 | # deformable transformer encoder args |
| 149 | transformer_in_features: List[str], |
| 150 | common_stride: int, |
| 151 | num_feature_levels: int, |
| 152 | total_num_feature_levels: int, |
| 153 | feature_order: str, |
| 154 | ): |
| 155 | super().__init__() |
| 156 | transformer_input_shape = { |
| 157 | k: v for k, v in input_shape.items() if k in transformer_in_features |
| 158 | } |
| 159 | # this is the input shape of pixel decoder |
| 160 | input_shape = sorted(input_shape.items(), key=lambda x: x[1].stride) |
| 161 | self.in_features = [k for k, v in input_shape] # starting from "res2" to "res5" |
| 162 | self.feature_strides = [v.stride for k, v in input_shape] |
| 163 | self.feature_channels = [v.channels for k, v in input_shape] |
| 164 | self.feature_order = feature_order |
| 165 | |
| 166 | if feature_order == "low2high": |
| 167 | transformer_input_shape = sorted(transformer_input_shape.items(), key=lambda x: -x[1].stride) |
| 168 | else: |
| 169 | transformer_input_shape = sorted(transformer_input_shape.items(), key=lambda x: x[1].stride) |
| 170 | self.transformer_in_features = [k for k, v in transformer_input_shape] # starting from "res2" to "res5" |
| 171 | transformer_in_channels = [v.channels for k, v in transformer_input_shape] |
| 172 | self.transformer_feature_strides = [v.stride for k, v in transformer_input_shape] # to decide extra FPN layers |
| 173 | |
| 174 | self.maskdino_num_feature_levels = num_feature_levels # always use 3 scales |
| 175 | self.total_num_feature_levels = total_num_feature_levels |
| 176 | self.common_stride = common_stride |
| 177 | |
| 178 | self.transformer_num_feature_levels = len(self.transformer_in_features) |
| 179 | self.low_resolution_index = transformer_in_channels.index(max(transformer_in_channels)) |
| 180 | self.high_resolution_index = 0 if self.feature_order == 'low2high' else -1 |
| 181 | if self.transformer_num_feature_levels > 1: |
| 182 | input_proj_list = [] |
| 183 | for in_channels in transformer_in_channels[::-1]: |
| 184 | input_proj_list.append(nn.Sequential( |
| 185 | nn.Conv2d(in_channels, conv_dim, kernel_size=1), |
| 186 | nn.GroupNorm(32, conv_dim), |
| 187 | )) |
| 188 | # input projectino for downsample |
| 189 | in_channels = max(transformer_in_channels) |
| 190 | for _ in range(self.total_num_feature_levels - self.transformer_num_feature_levels): # exclude the res2 |
| 191 | input_proj_list.append(nn.Sequential( |
| 192 | nn.Conv2d(in_channels, conv_dim, kernel_size=3, stride=2, padding=1), |
nothing calls this directly
no outgoing calls
no test coverage detected