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

Method forward

models/aios/transformer.py:580–693  ·  view source on GitHub ↗

Input: - src: [bs, sum(hi*wi), 256] - pos: pos embed for src. [bs, sum(hi*wi), 256] - spatial_shapes: h,w of each level [num_level, 2] - level_start_index: [num_level] start point of level in sum(hi*wi). - valid_ratios: [bs, num_le

(self,
                src: Tensor,
                pos: Tensor,
                spatial_shapes: Tensor,
                level_start_index: Tensor,
                valid_ratios: Tensor,
                key_padding_mask: Tensor,
                ref_token_index: Optional[Tensor] = None,
                ref_token_coord: Optional[Tensor] = None)

Source from the content-addressed store, hash-verified

578 return reference_points
579
580 def forward(self,
581 src: Tensor,
582 pos: Tensor,
583 spatial_shapes: Tensor,
584 level_start_index: Tensor,
585 valid_ratios: Tensor,
586 key_padding_mask: Tensor,
587 ref_token_index: Optional[Tensor] = None,
588 ref_token_coord: Optional[Tensor] = None):
589 """
590 Input:
591 - src: [bs, sum(hi*wi), 256]
592 - pos: pos embed for src. [bs, sum(hi*wi), 256]
593 - spatial_shapes: h,w of each level [num_level, 2]
594 - level_start_index: [num_level] start point of level in sum(hi*wi).
595 - valid_ratios: [bs, num_level, 2]
596 - key_padding_mask: [bs, sum(hi*wi)]
597
598 - ref_token_index: bs, nq
599 - ref_token_coord: bs, nq, 4
600 Intermedia:
601 - reference_points: [bs, sum(hi*wi), num_level, 2]
602 Outpus:
603 - output: [bs, sum(hi*wi), 256]
604 """
605 # pdb.set_trace()
606 if self.two_stage_type in [
607 'no', 'standard', 'enceachlayer', 'enclayer1'
608 ]:
609 assert ref_token_index is None
610
611 output = src
612
613 # preparation and reshape
614 if self.num_layers > 0:
615 if self.deformable_encoder:
616 reference_points = self.get_reference_points(spatial_shapes,
617 valid_ratios,
618 device=src.device)
619 # import pdb; pdb.set_trace()
620
621 intermediate_output = []
622 intermediate_ref = []
623 if ref_token_index is not None:
624 out_i = torch.gather(
625 output, 1,
626 ref_token_index.unsqueeze(-1).repeat(1, 1, self.d_model))
627 intermediate_output.append(out_i)
628 intermediate_ref.append(ref_token_coord)
629
630 # intermediate_coord = []
631 # main process
632 for layer_id, layer in enumerate(self.layers):
633 # main process
634 dropflag = False
635 if self.enc_layer_dropout_prob is not None:
636 prob = random.random()
637 if prob < self.enc_layer_dropout_prob[layer_id]:

Callers

nothing calls this directly

Calls 4

get_reference_pointsMethod · 0.95
maxMethod · 0.80
randomMethod · 0.45

Tested by

no test coverage detected