| 764 | self.maskformer_num_feature_levels = num_feature_levels # always use 3 scales |
| 765 | |
| 766 | def forward(self, features): |
| 767 | x, m = features['backbone_output'].decompose() |
| 768 | backbone = self.backbone[0] |
| 769 | Hp, Wp = x.shape[-2:] |
| 770 | if self.pos_mode == 'simple_interpolate': |
| 771 | pos = backbone.pos_embed.reshape(1, backbone.patch_embed.patch_shape[0], |
| 772 | backbone.patch_embed.patch_shape[1], |
| 773 | backbone.pos_embed.size(2)).permute(0, 3, 1, 2) |
| 774 | else: |
| 775 | pos = backbone.pos_embed.reshape(1, backbone.patch_embed.patch_shape[0], |
| 776 | backbone.patch_embed.patch_shape[1], |
| 777 | backbone.pos_embed.size(2))[:, :Hp, :Wp, :].permute(0, 3, 1, 2) |
| 778 | |
| 779 | out = [op(x) for op in self.fpns][::-1] # [r4, r3, r2, r1] |
| 780 | num_cur_levels = 0 |
| 781 | multi_scale_features = [] |
| 782 | multi_scale_masks = [] |
| 783 | multi_scale_poss = [] |
| 784 | for o in out: |
| 785 | if num_cur_levels < self.maskformer_num_feature_levels: |
| 786 | mask_o = F.interpolate(m[None].float(), size=o.shape[-2:]).to(torch.bool)[0] |
| 787 | multi_scale_features.append(o) |
| 788 | multi_scale_masks.append(mask_o) |
| 789 | pos_l = F.interpolate(pos[None], size=o.shape[-3:], mode='trilinear', align_corners=False)[0] |
| 790 | multi_scale_poss.append(pos_l) |
| 791 | num_cur_levels += 1 |
| 792 | |
| 793 | features.update({'neck_output': {'mask_features': None, |
| 794 | 'multi_scale_features': multi_scale_features, |
| 795 | 'multi_scale_masks': multi_scale_masks, |
| 796 | 'multi_scale_pos': multi_scale_poss}}) |
| 797 | return features |