Predict masks for the given input prompts, using the currently set image. Arguments: point_coords (np.ndarray or None): A Nx2 array of point prompts to the model. Each point is in (X,Y) in pixels. point_labels (np.ndarray or None): A length N array o
(
self,
point_coords: Optional[np.ndarray] = None,
point_labels: Optional[np.ndarray] = None,
box: Optional[np.ndarray] = None,
mask_input: Optional[np.ndarray] = None,
multimask_output: bool = True,
return_logits: bool = False,
normalize_coords=True,
)
| 223 | return all_masks, all_ious, all_low_res_masks |
| 224 | |
| 225 | def predict( |
| 226 | self, |
| 227 | point_coords: Optional[np.ndarray] = None, |
| 228 | point_labels: Optional[np.ndarray] = None, |
| 229 | box: Optional[np.ndarray] = None, |
| 230 | mask_input: Optional[np.ndarray] = None, |
| 231 | multimask_output: bool = True, |
| 232 | return_logits: bool = False, |
| 233 | normalize_coords=True, |
| 234 | ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: |
| 235 | """ |
| 236 | Predict masks for the given input prompts, using the currently set image. |
| 237 | |
| 238 | Arguments: |
| 239 | point_coords (np.ndarray or None): A Nx2 array of point prompts to the |
| 240 | model. Each point is in (X,Y) in pixels. |
| 241 | point_labels (np.ndarray or None): A length N array of labels for the |
| 242 | point prompts. 1 indicates a foreground point and 0 indicates a |
| 243 | background point. |
| 244 | box (np.ndarray or None): A length 4 array given a box prompt to the |
| 245 | model, in XYXY format. |
| 246 | mask_input (np.ndarray): A low resolution mask input to the model, typically |
| 247 | coming from a previous prediction iteration. Has form 1xHxW, where |
| 248 | for SAM, H=W=256. |
| 249 | multimask_output (bool): If true, the model will return three masks. |
| 250 | For ambiguous input prompts (such as a single click), this will often |
| 251 | produce better masks than a single prediction. If only a single |
| 252 | mask is needed, the model's predicted quality score can be used |
| 253 | to select the best mask. For non-ambiguous prompts, such as multiple |
| 254 | input prompts, multimask_output=False can give better results. |
| 255 | return_logits (bool): If true, returns un-thresholded masks logits |
| 256 | instead of a binary mask. |
| 257 | normalize_coords (bool): If true, the point coordinates will be normalized to the range [0,1] and point_coords is expected to be wrt. image dimensions. |
| 258 | |
| 259 | Returns: |
| 260 | (np.ndarray): The output masks in CxHxW format, where C is the |
| 261 | number of masks, and (H, W) is the original image size. |
| 262 | (np.ndarray): An array of length C containing the model's |
| 263 | predictions for the quality of each mask. |
| 264 | (np.ndarray): An array of shape CxHxW, where C is the number |
| 265 | of masks and H=W=256. These low resolution logits can be passed to |
| 266 | a subsequent iteration as mask input. |
| 267 | """ |
| 268 | if not self._is_image_set: |
| 269 | raise RuntimeError("An image must be set with .set_image(...) before mask prediction.") |
| 270 | |
| 271 | # Transform input prompts |
| 272 | |
| 273 | mask_input, unnorm_coords, labels, unnorm_box = self._prep_prompts(point_coords, point_labels, box, mask_input, |
| 274 | normalize_coords) |
| 275 | |
| 276 | masks, iou_predictions, low_res_masks = self._predict( |
| 277 | unnorm_coords, |
| 278 | labels, |
| 279 | unnorm_box, |
| 280 | mask_input, |
| 281 | multimask_output, |
| 282 | return_logits=return_logits, |
nothing calls this directly
no test coverage detected