(
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,
)
| 1771 | self.hw_face_bbox = nn.Embedding(1, 2) |
| 1772 | |
| 1773 | def forward( |
| 1774 | self, |
| 1775 | tgt, |
| 1776 | memory, |
| 1777 | tgt_mask: Optional[Tensor] = None, |
| 1778 | tgt_mask2: Optional[Tensor] = None, |
| 1779 | tgt_mask3: Optional[Tensor] = None, |
| 1780 | memory_mask: Optional[Tensor] = None, |
| 1781 | tgt_key_padding_mask: Optional[Tensor] = None, |
| 1782 | memory_key_padding_mask: Optional[Tensor] = None, |
| 1783 | pos: Optional[Tensor] = None, |
| 1784 | refpoints_unsigmoid: Optional[Tensor] = None, # num_queries, bs, 2 |
| 1785 | # for memory |
| 1786 | level_start_index: Optional[Tensor] = None, # num_levels |
| 1787 | spatial_shapes: Optional[Tensor] = None, # bs, num_levels, 2 |
| 1788 | valid_ratios: Optional[Tensor] = None, |
| 1789 | ): |
| 1790 | output = tgt |
| 1791 | |
| 1792 | intermediate = [] |
| 1793 | reference_points = refpoints_unsigmoid.sigmoid() |
| 1794 | ref_points = [reference_points] |
| 1795 | |
| 1796 | effect_num_dn = self.num_dn if self.training else 0 |
| 1797 | inter_select_number = self.num_group |
| 1798 | for layer_id, layer in enumerate(self.layers): |
| 1799 | if self.deformable_decoder: |
| 1800 | if reference_points.shape[-1] == 4: |
| 1801 | reference_points_input = reference_points[:, :, None] \ |
| 1802 | * torch.cat([valid_ratios, valid_ratios], -1)[None, :] # nq, bs, nlevel, 4 |
| 1803 | else: |
| 1804 | assert reference_points.shape[-1] == 2 |
| 1805 | reference_points_input = reference_points[:, :, |
| 1806 | None] * valid_ratios[ |
| 1807 | None, :] |
| 1808 | query_sine_embed = gen_sineembed_for_position( |
| 1809 | reference_points_input[:, :, 0, :] |
| 1810 | ) # convert the position query from bbox to sine/cosin embend |
| 1811 | else: |
| 1812 | query_sine_embed = gen_sineembed_for_position( |
| 1813 | reference_points) # nq, bs, 256*2 |
| 1814 | reference_points_input = None |
| 1815 | |
| 1816 | raw_query_pos = self.ref_point_head( |
| 1817 | query_sine_embed) # nq, bs, 256 |
| 1818 | pos_scale = self.query_scale( |
| 1819 | output) if self.query_scale is not None else 1 # ? |
| 1820 | query_pos = pos_scale * raw_query_pos |
| 1821 | if not self.deformable_decoder: |
| 1822 | query_sine_embed = query_sine_embed[ |
| 1823 | ..., :self.d_model] * self.query_pos_sine_scale(output) |
| 1824 | |
| 1825 | # modulated HW attentions |
| 1826 | if not self.deformable_decoder and self.modulate_hw_attn: |
| 1827 | refHW_cond = self.ref_anchor_head( |
| 1828 | output).sigmoid() # nq, bs, 2 |
| 1829 | query_sine_embed[..., self.d_model // 2:] *= ( |
| 1830 | refHW_cond[..., 0] / |
nothing calls this directly
no test coverage detected