(
self,
tgt,
memory,
tgt_mask: Optional[Tensor] = None,
tgt_mask2: Optional[Tensor] = None,
tgt_mask3: Optional[Tensor] = None,
memory_mask: Optional[Tensor] = None,
tgt_key_padding_mask: Optional[Tensor] = None,
memory_key_padding_mask: Optional[Tensor] = None,
pos: Optional[Tensor] = None,
refpoints_unsigmoid: Optional[Tensor] = None, # num_queries, bs, 2
# for memory
level_start_index: Optional[Tensor] = None, # num_levels
spatial_shapes: Optional[Tensor] = None, # bs, num_levels, 2
valid_ratios: Optional[Tensor] = None,
)
| 862 | self.face_keypoint_embed = nn.Embedding(self.num_face_points, d_model) |
| 863 | |
| 864 | def forward( |
| 865 | self, |
| 866 | tgt, |
| 867 | memory, |
| 868 | tgt_mask: Optional[Tensor] = None, |
| 869 | tgt_mask2: Optional[Tensor] = None, |
| 870 | tgt_mask3: Optional[Tensor] = None, |
| 871 | memory_mask: Optional[Tensor] = None, |
| 872 | tgt_key_padding_mask: Optional[Tensor] = None, |
| 873 | memory_key_padding_mask: Optional[Tensor] = None, |
| 874 | pos: Optional[Tensor] = None, |
| 875 | refpoints_unsigmoid: Optional[Tensor] = None, # num_queries, bs, 2 |
| 876 | # for memory |
| 877 | level_start_index: Optional[Tensor] = None, # num_levels |
| 878 | spatial_shapes: Optional[Tensor] = None, # bs, num_levels, 2 |
| 879 | valid_ratios: Optional[Tensor] = None, |
| 880 | ): |
| 881 | output = tgt |
| 882 | |
| 883 | intermediate = [] |
| 884 | reference_points = refpoints_unsigmoid.sigmoid() |
| 885 | ref_points = [reference_points] |
| 886 | |
| 887 | effect_num_dn = self.num_dn if self.training else 0 |
| 888 | inter_select_number = self.num_group |
| 889 | for layer_id, layer in enumerate(self.layers): |
| 890 | if self.deformable_decoder: |
| 891 | if reference_points.shape[-1] == 4: |
| 892 | reference_points_input = reference_points[:, :, None] \ |
| 893 | * torch.cat([valid_ratios, valid_ratios], -1)[None, :] # nq, bs, nlevel, 4 |
| 894 | else: |
| 895 | assert reference_points.shape[-1] == 2 |
| 896 | reference_points_input = reference_points[:, :, |
| 897 | None] * valid_ratios[ |
| 898 | None, :] |
| 899 | query_sine_embed = gen_sineembed_for_position( |
| 900 | reference_points_input[:, :, 0, :] |
| 901 | ) # convert the position query from bbox to sine/cosin embend |
| 902 | else: |
| 903 | query_sine_embed = gen_sineembed_for_position( |
| 904 | reference_points) # nq, bs, 256*2 |
| 905 | reference_points_input = None |
| 906 | |
| 907 | raw_query_pos = self.ref_point_head( |
| 908 | query_sine_embed) # nq, bs, 256 |
| 909 | pos_scale = self.query_scale( |
| 910 | output) if self.query_scale is not None else 1 # ? |
| 911 | query_pos = pos_scale * raw_query_pos |
| 912 | if not self.deformable_decoder: |
| 913 | query_sine_embed = query_sine_embed[ |
| 914 | ..., :self.d_model] * self.query_pos_sine_scale(output) |
| 915 | |
| 916 | # modulated HW attentions |
| 917 | if not self.deformable_decoder and self.modulate_hw_attn: |
| 918 | refHW_cond = self.ref_anchor_head( |
| 919 | output).sigmoid() # nq, bs, 2 |
| 920 | query_sine_embed[..., self.d_model // 2:] *= ( |
| 921 | refHW_cond[..., 0] / |
nothing calls this directly
no test coverage detected