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)
| 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]: |
nothing calls this directly
no test coverage detected