| 49 | return transformed_size |
| 50 | |
| 51 | def _embed_points(self, point_coords: torch.Tensor, point_labels: torch.Tensor) -> torch.Tensor: |
| 52 | point_coords = point_coords + 0.5 |
| 53 | point_coords = point_coords / self.img_size |
| 54 | point_embedding = self.model.prompt_encoder.pe_layer._pe_encoding(point_coords) |
| 55 | point_labels = point_labels.unsqueeze(-1).expand_as(point_embedding) |
| 56 | |
| 57 | point_embedding = point_embedding * (point_labels != -1) |
| 58 | point_embedding = point_embedding + self.model.prompt_encoder.not_a_point_embed.weight * ( |
| 59 | point_labels == -1 |
| 60 | ) |
| 61 | |
| 62 | for i in range(self.model.prompt_encoder.num_point_embeddings): |
| 63 | point_embedding = point_embedding + self.model.prompt_encoder.point_embeddings[ |
| 64 | i |
| 65 | ].weight * (point_labels == i) |
| 66 | |
| 67 | return point_embedding |
| 68 | |
| 69 | def _embed_masks(self, input_mask: torch.Tensor, has_mask_input: torch.Tensor) -> torch.Tensor: |
| 70 | mask_embedding = has_mask_input * self.model.prompt_encoder.mask_downscaling(input_mask) |