MCPcopy Create free account
hub / github.com/WeChatCV/WeVisionOne / MaskDINOEncoder

Class MaskDINOEncoder

WeVisionOne/pixel_decoder/maskdino_encoder.py:135–353  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

133
134@SEM_SEG_HEADS_REGISTRY.register()
135class 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),

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected