(self, input_mask: torch.Tensor, has_mask_input: torch.Tensor)
| 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) |
| 71 | mask_embedding = mask_embedding + ( |
| 72 | 1 - has_mask_input |
| 73 | ) * self.model.prompt_encoder.no_mask_embed.weight.reshape(1, -1, 1, 1) |
| 74 | return mask_embedding |
| 75 | |
| 76 | def mask_postprocessing(self, masks: torch.Tensor, orig_im_size: torch.Tensor) -> torch.Tensor: |
| 77 | masks = F.interpolate( |