MCPcopy Create free account
hub / github.com/NVlabs/CTG / MapEncoder

Class MapEncoder

tbsim/models/diffuser_helpers.py:297–348  ·  view source on GitHub ↗

Encodes map, may output a global feature, feature map, or both.

Source from the content-addressed store, hash-verified

295
296
297class 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
350from tbsim.models.base_models import Up, ConvBlock, IdentityBlock
351

Callers 3

__init__Method · 0.90
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected