(self,
srcs,
masks,
refpoint_embed,
pos_embeds,
tgt,
attn_mask=None,
attn_mask2=None,
attn_mask3=None)
| 2302 | |
| 2303 | # srcs: features; refpoint_embed: |
| 2304 | def forward(self, |
| 2305 | srcs, |
| 2306 | masks, |
| 2307 | refpoint_embed, |
| 2308 | pos_embeds, |
| 2309 | tgt, |
| 2310 | attn_mask=None, |
| 2311 | attn_mask2=None, |
| 2312 | attn_mask3=None): |
| 2313 | # pdb.set_trace() |
| 2314 | # prepare input for encoder |
| 2315 | src_flatten = [] |
| 2316 | mask_flatten = [] |
| 2317 | lvl_pos_embed_flatten = [] |
| 2318 | spatial_shapes = [] |
| 2319 | for lvl, (src, mask, pos_embed) in enumerate( |
| 2320 | zip(srcs, masks, pos_embeds)): # for feature level |
| 2321 | bs, c, h, w = src.shape |
| 2322 | spatial_shape = (h, w) |
| 2323 | spatial_shapes.append(spatial_shape) |
| 2324 | |
| 2325 | src = src.flatten(2).transpose(1, 2) # bs, hw, c |
| 2326 | mask = mask.flatten(1) # bs, hw |
| 2327 | pos_embed = pos_embed.flatten(2).transpose(1, 2) # bs, hw, c |
| 2328 | if self.num_feature_levels > 1 and self.level_embed is not None: |
| 2329 | lvl_pos_embed = pos_embed + self.level_embed[lvl].view( |
| 2330 | 1, 1, -1) # level_embed[lvl]: [256] |
| 2331 | else: |
| 2332 | lvl_pos_embed = pos_embed |
| 2333 | lvl_pos_embed_flatten.append(lvl_pos_embed) |
| 2334 | src_flatten.append(src) |
| 2335 | mask_flatten.append(mask) |
| 2336 | src_flatten = torch.cat(src_flatten, 1) # bs, \sum{hxw}, c |
| 2337 | mask_flatten = torch.cat(mask_flatten, 1) # bs, \sum{hxw} |
| 2338 | lvl_pos_embed_flatten = torch.cat(lvl_pos_embed_flatten, |
| 2339 | 1) # bs, \sum{hxw}, c |
| 2340 | spatial_shapes = torch.as_tensor(spatial_shapes, |
| 2341 | dtype=torch.long, |
| 2342 | device=src_flatten.device) |
| 2343 | level_start_index = torch.cat((spatial_shapes.new_zeros( |
| 2344 | (1, )), spatial_shapes.prod(1).cumsum(0)[:-1])) |
| 2345 | valid_ratios = torch.stack([self.get_valid_ratio(m) for m in masks], 1) |
| 2346 | # two stage |
| 2347 | if self.two_stage_type in ['early', 'combine']: |
| 2348 | output_memory, output_proposals = gen_encoder_output_proposals( |
| 2349 | src_flatten, mask_flatten, spatial_shapes) |
| 2350 | output_memory = self.enc_output_norm_backbone( |
| 2351 | self.enc_output_backbone(output_memory)) |
| 2352 | |
| 2353 | # gather boxes |
| 2354 | topk = self.num_queries |
| 2355 | enc_outputs_class = self.encoder.class_embed[0](output_memory) |
| 2356 | enc_topk_proposals = torch.topk(enc_outputs_class.max(-1)[0], |
| 2357 | topk, |
| 2358 | dim=1)[1] # bs, nq |
| 2359 | enc_refpoint_embed = torch.gather( |
| 2360 | output_proposals, 1, |
| 2361 | enc_topk_proposals.unsqueeze(-1).repeat(1, 1, 4)) |
nothing calls this directly
no test coverage detected