(self,
srcs,
masks,
refpoint_embed,
pos_embeds,
tgt,
attn_mask=None,
attn_mask2=None,
attn_mask3=None)
| 264 | |
| 265 | # srcs: features; refpoint_embed: |
| 266 | def forward(self, |
| 267 | srcs, |
| 268 | masks, |
| 269 | refpoint_embed, |
| 270 | pos_embeds, |
| 271 | tgt, |
| 272 | attn_mask=None, |
| 273 | attn_mask2=None, |
| 274 | attn_mask3=None): |
| 275 | # pdb.set_trace() |
| 276 | # prepare input for encoder |
| 277 | src_flatten = [] |
| 278 | mask_flatten = [] |
| 279 | lvl_pos_embed_flatten = [] |
| 280 | spatial_shapes = [] |
| 281 | for lvl, (src, mask, pos_embed) in enumerate( |
| 282 | zip(srcs, masks, pos_embeds)): # for feature level |
| 283 | bs, c, h, w = src.shape |
| 284 | spatial_shape = (h, w) |
| 285 | spatial_shapes.append(spatial_shape) |
| 286 | |
| 287 | src = src.flatten(2).transpose(1, 2) # bs, hw, c |
| 288 | mask = mask.flatten(1) # bs, hw |
| 289 | pos_embed = pos_embed.flatten(2).transpose(1, 2) # bs, hw, c |
| 290 | if self.num_feature_levels > 1 and self.level_embed is not None: |
| 291 | lvl_pos_embed = pos_embed + self.level_embed[lvl].view( |
| 292 | 1, 1, -1) # level_embed[lvl]: [256] |
| 293 | else: |
| 294 | lvl_pos_embed = pos_embed |
| 295 | lvl_pos_embed_flatten.append(lvl_pos_embed) |
| 296 | src_flatten.append(src) |
| 297 | mask_flatten.append(mask) |
| 298 | src_flatten = torch.cat(src_flatten, 1) # bs, \sum{hxw}, c |
| 299 | mask_flatten = torch.cat(mask_flatten, 1) # bs, \sum{hxw} |
| 300 | lvl_pos_embed_flatten = torch.cat(lvl_pos_embed_flatten, |
| 301 | 1) # bs, \sum{hxw}, c |
| 302 | spatial_shapes = torch.as_tensor(spatial_shapes, |
| 303 | dtype=torch.long, |
| 304 | device=src_flatten.device) |
| 305 | level_start_index = torch.cat((spatial_shapes.new_zeros( |
| 306 | (1, )), spatial_shapes.prod(1).cumsum(0)[:-1])) |
| 307 | valid_ratios = torch.stack([self.get_valid_ratio(m) for m in masks], 1) |
| 308 | # two stage |
| 309 | if self.two_stage_type in ['early', 'combine']: |
| 310 | output_memory, output_proposals = gen_encoder_output_proposals( |
| 311 | src_flatten, mask_flatten, spatial_shapes) |
| 312 | output_memory = self.enc_output_norm_backbone( |
| 313 | self.enc_output_backbone(output_memory)) |
| 314 | |
| 315 | # gather boxes |
| 316 | topk = self.num_queries |
| 317 | enc_outputs_class = self.encoder.class_embed[0](output_memory) |
| 318 | enc_topk_proposals = torch.topk(enc_outputs_class.max(-1)[0], |
| 319 | topk, |
| 320 | dim=1)[1] # bs, nq |
| 321 | enc_refpoint_embed = torch.gather( |
| 322 | output_proposals, 1, |
| 323 | enc_topk_proposals.unsqueeze(-1).repeat(1, 1, 4)) |
nothing calls this directly
no test coverage detected