Predicts masks. See 'forward' for more details.
(
self,
image_embeddings: torch.Tensor,
image_pe: torch.Tensor,
sparse_prompt_embeddings: torch.Tensor,
dense_prompt_embeddings: torch.Tensor,
repeat_image: bool,
high_res_features: Optional[List[torch.Tensor]] = None,
)
| 166 | return masks, iou_pred, sam_tokens_out, object_score_logits |
| 167 | |
| 168 | def predict_masks( |
| 169 | self, |
| 170 | image_embeddings: torch.Tensor, |
| 171 | image_pe: torch.Tensor, |
| 172 | sparse_prompt_embeddings: torch.Tensor, |
| 173 | dense_prompt_embeddings: torch.Tensor, |
| 174 | repeat_image: bool, |
| 175 | high_res_features: Optional[List[torch.Tensor]] = None, |
| 176 | ) -> Tuple[torch.Tensor, torch.Tensor]: |
| 177 | """Predicts masks. See 'forward' for more details.""" |
| 178 | # Concatenate output tokens |
| 179 | s = 0 |
| 180 | if self.pred_obj_scores: |
| 181 | output_tokens = torch.cat( |
| 182 | [ |
| 183 | self.obj_score_token.weight, |
| 184 | self.iou_token.weight, |
| 185 | self.mask_tokens.weight, |
| 186 | ], |
| 187 | dim=0, |
| 188 | ) |
| 189 | s = 1 |
| 190 | else: |
| 191 | output_tokens = torch.cat( |
| 192 | [self.iou_token.weight, self.mask_tokens.weight], dim=0 |
| 193 | ) |
| 194 | output_tokens = output_tokens.unsqueeze(0).expand( |
| 195 | sparse_prompt_embeddings.size(0), -1, -1 |
| 196 | ) |
| 197 | tokens = torch.cat((output_tokens, sparse_prompt_embeddings), dim=1) |
| 198 | |
| 199 | # Expand per-image data in batch direction to be per-mask |
| 200 | if repeat_image: |
| 201 | src = torch.repeat_interleave(image_embeddings, tokens.shape[0], dim=0) |
| 202 | else: |
| 203 | assert image_embeddings.shape[0] == tokens.shape[0] |
| 204 | src = image_embeddings |
| 205 | src = src + dense_prompt_embeddings |
| 206 | assert ( |
| 207 | image_pe.size(0) == 1 |
| 208 | ), "image_pe should have size 1 in batch dim (from `get_dense_pe()`)" |
| 209 | pos_src = torch.repeat_interleave(image_pe, tokens.shape[0], dim=0) |
| 210 | b, c, h, w = src.shape |
| 211 | |
| 212 | # Run the transformer |
| 213 | hs, src = self.transformer(src, pos_src, tokens) |
| 214 | iou_token_out = hs[:, s, :] |
| 215 | mask_tokens_out = hs[:, s + 1 : (s + 1 + self.num_mask_tokens), :] |
| 216 | |
| 217 | # Upscale mask embeddings and predict masks using the mask tokens |
| 218 | src = src.transpose(1, 2).view(b, c, h, w) |
| 219 | if not self.use_high_res_features: |
| 220 | upscaled_embedding = self.output_upscaling(src) |
| 221 | else: |
| 222 | dc1, ln1, act1, dc2, act2 = self.output_upscaling |
| 223 | feat_s0, feat_s1 = high_res_features |
| 224 | upscaled_embedding = act1(ln1(dc1(src) + feat_s1)) |
| 225 | upscaled_embedding = act2(dc2(upscaled_embedding) + feat_s0) |