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

Method forward

models/aios/transformer.py:266–498  ·  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

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

Callers

nothing calls this directly

Calls 4

get_valid_ratioMethod · 0.95
maxMethod · 0.80
detachMethod · 0.45

Tested by

no test coverage detected