| 349 | @TRANSFORMER_LAYER_SEQUENCE.register_module() |
| 350 | class LATRTransformerDecoder(TransformerLayerSequence): |
| 351 | def __init__(self, |
| 352 | *args, embed_dims=None, |
| 353 | post_norm_cfg=dict(type='LN'), |
| 354 | enlarge_length=10, |
| 355 | M_decay_ratio=10, |
| 356 | num_query=None, |
| 357 | num_anchor_per_query=None, |
| 358 | anchor_y_steps=None, |
| 359 | **kwargs): |
| 360 | super(LATRTransformerDecoder, self).__init__(*args, **kwargs) |
| 361 | if post_norm_cfg is not None: |
| 362 | self.post_norm = build_norm_layer(post_norm_cfg, |
| 363 | self.embed_dims)[1] |
| 364 | else: |
| 365 | self.post_norm = None |
| 366 | |
| 367 | self.num_query = num_query |
| 368 | self.num_anchor_per_query = num_anchor_per_query |
| 369 | self.anchor_y_steps = anchor_y_steps |
| 370 | self.num_points_per_anchor = len(anchor_y_steps) // num_anchor_per_query |
| 371 | |
| 372 | self.embed_dims = embed_dims |
| 373 | self.gflat_pred_layer = nn.Sequential( |
| 374 | nn.Conv2d(embed_dims + 4, embed_dims, 3, stride=2, padding=1, bias=False), |
| 375 | nn.BatchNorm2d(embed_dims), |
| 376 | nn.ReLU(True), |
| 377 | nn.Conv2d(embed_dims, embed_dims, 3, stride=2, padding=1, bias=False), |
| 378 | nn.BatchNorm2d(embed_dims), |
| 379 | nn.ReLU(True), |
| 380 | nn.AdaptiveAvgPool2d(1), |
| 381 | nn.Conv2d(embed_dims, embed_dims, 1), |
| 382 | nn.BatchNorm2d(embed_dims), |
| 383 | nn.ReLU(True), |
| 384 | nn.Conv2d(embed_dims, embed_dims // 4, 1), |
| 385 | nn.BatchNorm2d(embed_dims // 4), |
| 386 | nn.ReLU(True), |
| 387 | nn.Conv2d(embed_dims // 4, 2, 1)) |
| 388 | |
| 389 | self.position_encoder = nn.Sequential( |
| 390 | nn.Conv2d(3, self.embed_dims*4, kernel_size=1, stride=1, padding=0), |
| 391 | nn.ReLU(), |
| 392 | nn.Conv2d(self.embed_dims*4, self.embed_dims, kernel_size=1, stride=1, padding=0), |
| 393 | ) |
| 394 | self.M_decay_ratio = M_decay_ratio |
| 395 | self.enlarge_length = enlarge_length |
| 396 | |
| 397 | def init_weights(self): |
| 398 | super().init_weights() |