MCPcopy Create free account
hub / github.com/ImprintLab/Medical-SAM2 / predict_masks

Method predict_masks

sam2_train/modeling/sam/mask_decoder.py:168–245  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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)

Callers 1

forwardMethod · 0.95

Calls 1

catMethod · 0.80

Tested by

no test coverage detected