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

Method forward

sam2_train/modeling/sam/prompt_encoder.py:140–186  ·  view source on GitHub ↗

Embeds different types of prompts, returning both sparse and dense embeddings. Arguments: points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates and labels to embed. boxes (torch.Tensor or none): boxes to embed masks

(
        self,
        points: Optional[Tuple[torch.Tensor, torch.Tensor]],
        boxes: Optional[torch.Tensor],
        masks: Optional[torch.Tensor],
        batch_size = -1,
    )

Source from the content-addressed store, hash-verified

138 return self.point_embeddings[0].weight.device
139
140 def forward(
141 self,
142 points: Optional[Tuple[torch.Tensor, torch.Tensor]],
143 boxes: Optional[torch.Tensor],
144 masks: Optional[torch.Tensor],
145 batch_size = -1,
146 ) -> Tuple[torch.Tensor, torch.Tensor]:
147 """
148 Embeds different types of prompts, returning both sparse and dense
149 embeddings.
150
151 Arguments:
152 points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
153 and labels to embed.
154 boxes (torch.Tensor or none): boxes to embed
155 masks (torch.Tensor or none): masks to embed
156
157 Returns:
158 torch.Tensor: sparse embeddings for the points and boxes, with shape
159 BxNx(embed_dim), where N is determined by the number of input points
160 and boxes.
161 torch.Tensor: dense embeddings for the masks, in the shape
162 Bx(embed_dim)x(embed_H)x(embed_W)
163 """
164 if batch_size == -1:
165 bs = self._get_batch_size(points, boxes, masks)
166 else:
167 bs = batch_size
168 sparse_embeddings = torch.empty(
169 (bs, 0, self.embed_dim), device=self._get_device()
170 )
171 if points is not None:
172 coords, labels = points
173 point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
174 sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
175 if boxes is not None:
176 box_embeddings = self._embed_boxes(boxes)
177 sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
178
179 if masks is not None:
180 dense_embeddings = self._embed_masks(masks)
181 else:
182 dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
183 bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
184 )
185
186 return sparse_embeddings, dense_embeddings

Callers

nothing calls this directly

Calls 6

_get_batch_sizeMethod · 0.95
_get_deviceMethod · 0.95
_embed_pointsMethod · 0.95
_embed_boxesMethod · 0.95
_embed_masksMethod · 0.95
catMethod · 0.80

Tested by

no test coverage detected