(
self,
image_embeddings: torch.Tensor,
point_coords: torch.Tensor,
point_labels: torch.Tensor,
mask_input: torch.Tensor,
has_mask_input: torch.Tensor,
orig_im_size: torch.Tensor,
)
| 106 | |
| 107 | @torch.no_grad() |
| 108 | def forward( |
| 109 | self, |
| 110 | image_embeddings: torch.Tensor, |
| 111 | point_coords: torch.Tensor, |
| 112 | point_labels: torch.Tensor, |
| 113 | mask_input: torch.Tensor, |
| 114 | has_mask_input: torch.Tensor, |
| 115 | orig_im_size: torch.Tensor, |
| 116 | ): |
| 117 | sparse_embedding = self._embed_points(point_coords, point_labels) |
| 118 | dense_embedding = self._embed_masks(mask_input, has_mask_input) |
| 119 | |
| 120 | masks, scores = self.model.mask_decoder.predict_masks( |
| 121 | image_embeddings=image_embeddings, |
| 122 | image_pe=self.model.prompt_encoder.get_dense_pe(), |
| 123 | sparse_prompt_embeddings=sparse_embedding, |
| 124 | dense_prompt_embeddings=dense_embedding, |
| 125 | ) |
| 126 | |
| 127 | if self.use_stability_score: |
| 128 | scores = calculate_stability_score( |
| 129 | masks, self.model.mask_threshold, self.stability_score_offset |
| 130 | ) |
| 131 | |
| 132 | if self.return_single_mask: |
| 133 | masks, scores = self.select_masks(masks, scores, point_coords.shape[1]) |
| 134 | |
| 135 | upscaled_masks = self.mask_postprocessing(masks, orig_im_size) |
| 136 | |
| 137 | if self.return_extra_metrics: |
| 138 | stability_scores = calculate_stability_score( |
| 139 | upscaled_masks, self.model.mask_threshold, self.stability_score_offset |
| 140 | ) |
| 141 | areas = (upscaled_masks > self.model.mask_threshold).sum(-1).sum(-1) |
| 142 | return upscaled_masks, scores, stability_scores, areas, masks |
| 143 | |
| 144 | return upscaled_masks, scores, masks |
nothing calls this directly
no test coverage detected