Encodes map, may output a global feature, feature map, or both.
| 295 | |
| 296 | |
| 297 | class MapEncoder(nn.Module): |
| 298 | """Encodes map, may output a global feature, feature map, or both.""" |
| 299 | def __init__( |
| 300 | self, |
| 301 | model_arch: str, |
| 302 | input_image_shape: tuple = (3, 224, 224), |
| 303 | global_feature_dim=None, |
| 304 | grid_feature_dim=None, |
| 305 | ) -> None: |
| 306 | super(MapEncoder, self).__init__() |
| 307 | self.return_global_feat = global_feature_dim is not None |
| 308 | self.return_grid_feat = grid_feature_dim is not None |
| 309 | encoder = base_models.RasterizedMapEncoder( |
| 310 | model_arch=model_arch, |
| 311 | input_image_shape=input_image_shape, |
| 312 | feature_dim=global_feature_dim |
| 313 | ) |
| 314 | self.input_image_shape = input_image_shape |
| 315 | # build graph for extracting intermediate features |
| 316 | feat_nodes = { |
| 317 | 'map_model.layer1': 'layer1', |
| 318 | 'map_model.layer2': 'layer2', |
| 319 | 'map_model.layer3': 'layer3', |
| 320 | 'map_model.layer4': 'layer4', |
| 321 | 'map_model.fc' : 'fc', |
| 322 | } |
| 323 | self.encoder_heads = create_feature_extractor(encoder, feat_nodes) |
| 324 | if self.return_grid_feat: |
| 325 | encoder_channels = list(encoder.feature_channels().values()) |
| 326 | input_shape_scale = encoder.feature_scales()["layer4"] |
| 327 | self.decoder = MapGridDecoder( |
| 328 | input_shape=(encoder_channels[-1], input_image_shape[1]*input_shape_scale, input_image_shape[2]*input_shape_scale), |
| 329 | encoder_channels=encoder_channels[:-1], |
| 330 | output_channel=grid_feature_dim, |
| 331 | batchnorm=True, |
| 332 | ) |
| 333 | self.encoder_feat_scales = list(encoder.feature_scales().values()) |
| 334 | |
| 335 | def feat_map_out_dim(self, H, W): |
| 336 | dim_scale = self.encoder_feat_scales[-4] # decoder has 3 upsampling |
| 337 | return (H * dim_scale, W * dim_scale ) |
| 338 | |
| 339 | def forward(self, map_inputs, encoder_feats=None): |
| 340 | if encoder_feats is None: |
| 341 | encoder_feats = self.encoder_heads(map_inputs) |
| 342 | fc_out = encoder_feats['fc'] if self.return_global_feat else None |
| 343 | encoder_feats = [encoder_feats[k] for k in ["layer1", "layer2", "layer3", "layer4"]] |
| 344 | feat_map_out = None |
| 345 | if self.return_grid_feat: |
| 346 | feat_map_out = self.decoder.forward(feat_to_decode=encoder_feats[-1], |
| 347 | encoder_feats=encoder_feats[:-1]) |
| 348 | return fc_out, feat_map_out |
| 349 | |
| 350 | from tbsim.models.base_models import Up, ConvBlock, IdentityBlock |
| 351 |