Embeds different types of prompts, returning both sparse and dense embeddings. Arguments: points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates and labels to embed. boxes (torch.Tensor or none): boxes to embed masks
(
self,
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
boxes: Optional[torch.Tensor],
masks: Optional[torch.Tensor],
batch_size = -1,
)
| 138 | return self.point_embeddings[0].weight.device |
| 139 | |
| 140 | def forward( |
| 141 | self, |
| 142 | points: Optional[Tuple[torch.Tensor, torch.Tensor]], |
| 143 | boxes: Optional[torch.Tensor], |
| 144 | masks: Optional[torch.Tensor], |
| 145 | batch_size = -1, |
| 146 | ) -> Tuple[torch.Tensor, torch.Tensor]: |
| 147 | """ |
| 148 | Embeds different types of prompts, returning both sparse and dense |
| 149 | embeddings. |
| 150 | |
| 151 | Arguments: |
| 152 | points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates |
| 153 | and labels to embed. |
| 154 | boxes (torch.Tensor or none): boxes to embed |
| 155 | masks (torch.Tensor or none): masks to embed |
| 156 | |
| 157 | Returns: |
| 158 | torch.Tensor: sparse embeddings for the points and boxes, with shape |
| 159 | BxNx(embed_dim), where N is determined by the number of input points |
| 160 | and boxes. |
| 161 | torch.Tensor: dense embeddings for the masks, in the shape |
| 162 | Bx(embed_dim)x(embed_H)x(embed_W) |
| 163 | """ |
| 164 | if batch_size == -1: |
| 165 | bs = self._get_batch_size(points, boxes, masks) |
| 166 | else: |
| 167 | bs = batch_size |
| 168 | sparse_embeddings = torch.empty( |
| 169 | (bs, 0, self.embed_dim), device=self._get_device() |
| 170 | ) |
| 171 | if points is not None: |
| 172 | coords, labels = points |
| 173 | point_embeddings = self._embed_points(coords, labels, pad=(boxes is None)) |
| 174 | sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1) |
| 175 | if boxes is not None: |
| 176 | box_embeddings = self._embed_boxes(boxes) |
| 177 | sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1) |
| 178 | |
| 179 | if masks is not None: |
| 180 | dense_embeddings = self._embed_masks(masks) |
| 181 | else: |
| 182 | dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand( |
| 183 | bs, -1, self.image_embedding_size[0], self.image_embedding_size[1] |
| 184 | ) |
| 185 | |
| 186 | return sparse_embeddings, dense_embeddings |
nothing calls this directly
no test coverage detected