MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / forward

Method forward

utils/sam_utils/onnx.py:108–144  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 6

_embed_pointsMethod · 0.95
_embed_masksMethod · 0.95
select_masksMethod · 0.95
mask_postprocessingMethod · 0.95
predict_masksMethod · 0.80

Tested by

no test coverage detected