MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / forward

Method forward

models/aios/transformer.py:2304–2536  ·  view source on GitHub ↗
(self,
                srcs,
                masks,
                refpoint_embed,
                pos_embeds,
                tgt,
                attn_mask=None,
                attn_mask2=None,
                attn_mask3=None)

Source from the content-addressed store, hash-verified

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))

Callers

nothing calls this directly

Calls 4

get_valid_ratioMethod · 0.95
maxMethod · 0.80
detachMethod · 0.45

Tested by

no test coverage detected