Team Ai
Apppublic

michaelcreatesstuff/llm-grounded-diffusion

sourceHugging Faceupdated 3y agoView on Hugging Face
2likes
sam.py201 linesDownload Raw Back to models
1import gc2import matplotlib.pyplot as plt3import numpy as np4import torch5import torch.nn.functional as F6from models import torch_device7from transformers import SamModel, SamProcessor8import utils9import cv210from scipy import ndimage11 12def load_sam():13    sam_model = SamModel.from_pretrained("facebook/sam-vit-base").to(torch_device)14    sam_processor = SamProcessor.from_pretrained("facebook/sam-vit-base")15 16    sam_model_dict = dict(17        sam_model = sam_model, sam_processor = sam_processor18    )19 20    return sam_model_dict21 22# Not fully backward compatible with the previous implementation23# Reference: lmdv2/notebooks/gen_masked_latents_multi_object_ref_ca_loss_modular.ipynb24def sam(sam_model_dict, image, input_points=None, input_boxes=None, target_mask_shape=None, return_numpy=True):25    """target_mask_shape: (h, w)"""26    sam_model, sam_processor = sam_model_dict['sam_model'], sam_model_dict['sam_processor']27    28    if input_boxes and isinstance(input_boxes[0], tuple):29        # Convert tuple to list30        input_boxes = [list(input_box) for input_box in input_boxes]31        32    if input_boxes and input_boxes[0] and isinstance(input_boxes[0][0], tuple):33        # Convert tuple to list34        input_boxes = [[list(input_box) for input_box in input_boxes_item] for input_boxes_item in input_boxes]35    36    with torch.no_grad():37        with torch.autocast(torch_device):38            inputs = sam_processor(image, input_points=input_points, input_boxes=input_boxes, return_tensors="pt").to(torch_device)39            outputs = sam_model(**inputs)40        masks = sam_processor.image_processor.post_process_masks(41            outputs.pred_masks.cpu().float(), inputs["original_sizes"].cpu(), inputs["reshaped_input_sizes"].cpu()42        )43        conf_scores = outputs.iou_scores.cpu().numpy()[0,0]44        del inputs, outputs45    46    gc.collect()47    torch.cuda.empty_cache()48    49    if return_numpy:50        masks = [F.interpolate(masks_item.type(torch.float), target_mask_shape, mode='bilinear').type(torch.bool).numpy() for masks_item in masks]51    else:52        masks = [F.interpolate(masks_item.type(torch.float), target_mask_shape, mode='bilinear').type(torch.bool) for masks_item in masks]53 54    return masks, conf_scores55 56def sam_point_input(sam_model_dict, image, input_points, **kwargs):57    return sam(sam_model_dict, image, input_points=input_points, **kwargs)58    59def sam_box_input(sam_model_dict, image, input_boxes, **kwargs):60    return sam(sam_model_dict, image, input_boxes=input_boxes, **kwargs)61 62def get_iou_with_resize(mask, masks, masks_shape):63    masks = np.array([cv2.resize(mask.astype(np.uint8) * 255, masks_shape[::-1], cv2.INTER_LINEAR).astype(bool) for mask in masks])64    return utils.iou(mask, masks)65 66def select_mask(masks, conf_scores, coarse_ious=None, rule="largest_over_conf", discourage_mask_below_confidence=0.85, discourage_mask_below_coarse_iou=0.2, verbose=False):67    """masks: numpy bool array"""68    mask_sizes = masks.sum(axis=(1, 2))69    70    # Another possible rule: iou with the attention mask71    if rule == "largest_over_conf":72        # Use the largest segmentation73        # Discourage selecting masks with conf too low or coarse iou is too low74        max_mask_size = np.max(mask_sizes)75        if coarse_ious is not None:76            scores = mask_sizes - (conf_scores < discourage_mask_below_confidence) * max_mask_size - (coarse_ious < discourage_mask_below_coarse_iou) * max_mask_size77        else:78            scores = mask_sizes - (conf_scores < discourage_mask_below_confidence) * max_mask_size79        if verbose:80            print(f"mask_sizes: {mask_sizes}, scores: {scores}")81    else:82        raise ValueError(f"Unknown rule: {rule}")83 84    mask_id = np.argmax(scores)85    mask = masks[mask_id]86    87    selection_conf = conf_scores[mask_id]88    89    if coarse_ious is not None:90        selection_coarse_iou = coarse_ious[mask_id]91    else:92        selection_coarse_iou = None93 94    if verbose:95        # print(f"Confidences: {conf_scores}")96        print(f"Selected a mask with confidence: {selection_conf}, coarse_iou: {selection_coarse_iou}")97 98    if verbose:99        plt.figure(figsize=(10, 8))100        # plt.suptitle("After SAM")101        for ind in range(3):102            plt.subplot(1, 3, ind+1)103            # This is obtained before resize.104            plt.title(f"Mask {ind}, score {scores[ind]}, conf {conf_scores[ind]:.2f}, iou {coarse_ious[ind] if coarse_ious is not None else None:.2f}")105            plt.imshow(masks[ind])106        plt.tight_layout()107        plt.show()108 109    return mask, selection_conf110 111def preprocess_mask(token_attn_np_smooth, mask_th, n_erode_dilate_mask=0):112    token_attn_np_smooth_normalized = token_attn_np_smooth - token_attn_np_smooth.min()113    token_attn_np_smooth_normalized /= token_attn_np_smooth_normalized.max()114    mask_thresholded = token_attn_np_smooth_normalized > mask_th115    116    if n_erode_dilate_mask:117        mask_thresholded = ndimage.binary_erosion(mask_thresholded, iterations=n_erode_dilate_mask)118        mask_thresholded = ndimage.binary_dilation(mask_thresholded, iterations=n_erode_dilate_mask)119    120    return mask_thresholded121 122# The overall pipeline to refine the attention mask123def sam_refine_attn(sam_input_image, token_attn_np, model_dict, height, width, H, W, use_box_input, gaussian_sigma, mask_th_for_box, n_erode_dilate_mask_for_box, mask_th_for_point, discourage_mask_below_confidence, discourage_mask_below_coarse_iou, verbose):124    125    # token_attn_np is for visualizations126    token_attn_np_smooth = ndimage.gaussian_filter(token_attn_np, sigma=gaussian_sigma)127 128    # (w, h)129    mask_size_scale = height // token_attn_np_smooth.shape[1], width // token_attn_np_smooth.shape[0]130 131    if use_box_input:132        # box input133        mask_binary = preprocess_mask(token_attn_np_smooth, mask_th_for_box, n_erode_dilate_mask=n_erode_dilate_mask_for_box)134 135        input_boxes = utils.binary_mask_to_box(mask_binary, w_scale=mask_size_scale[0], h_scale=mask_size_scale[1])136        input_boxes = [input_boxes]137 138        masks, conf_scores = sam_box_input(model_dict, image=sam_input_image, input_boxes=input_boxes, target_mask_shape=(H, W))139    else:140        # point input141        mask_binary = preprocess_mask(token_attn_np_smooth, mask_th_for_point, n_erode_dilate_mask=0)142 143        # Uses the max coordinate only144        max_coord = np.unravel_index(token_attn_np_smooth.argmax(), token_attn_np_smooth.shape)145        # print("max_coord:", max_coord)146        input_points = [[[max_coord[1] * mask_size_scale[1], max_coord[0] * mask_size_scale[0]]]]147 148        masks, conf_scores = sam_point_input(model_dict, image=sam_input_image, input_points=input_points, target_mask_shape=(H, W))149        150    if verbose:151        plt.title("Coarse binary mask (for box for box input and for iou)")152        plt.imshow(mask_binary)153        plt.show()154    155    coarse_ious = get_iou_with_resize(mask_binary, masks, masks_shape=mask_binary.shape)156 157    mask_selected, conf_score_selected = select_mask(masks, conf_scores, coarse_ious=coarse_ious, 158                                                         rule="largest_over_conf", 159                                                         discourage_mask_below_confidence=discourage_mask_below_confidence, 160                                                         discourage_mask_below_coarse_iou=discourage_mask_below_coarse_iou,161                                                         verbose=True)162 163    return mask_selected, conf_score_selected164 165def sam_refine_box(sam_input_image, box, *args, **kwargs):166    sam_input_images, boxes = [sam_input_image], [box]167    return sam_refine_boxes(sam_input_images, boxes, *args, **kwargs)168 169def sam_refine_boxes(sam_input_images, boxes, model_dict, height, width, H, W, discourage_mask_below_confidence, discourage_mask_below_coarse_iou, verbose):170    # (w, h)171    input_boxes = [[utils.scale_proportion(box, H=height, W=width) for box in boxes_item] for boxes_item in boxes]172 173    masks, conf_scores = sam_box_input(model_dict, image=sam_input_images, input_boxes=input_boxes, target_mask_shape=(H, W))174    175    mask_selected_batched_list, conf_score_selected_batched_list = [], []176    177    for boxes_item, masks_item in zip(boxes, masks):178        mask_selected_list, conf_score_selected_list = [], []179        for box, three_masks in zip(boxes_item, masks_item):180            mask_binary = utils.proportion_to_mask(box, H, W, return_np=True)181            if verbose:182                # Also the box is the input for SAM183                plt.title("Binary mask from input box (for iou)")184                plt.imshow(mask_binary)185                plt.show()186                        187            coarse_ious = get_iou_with_resize(mask_binary, three_masks, masks_shape=mask_binary.shape)188 189            mask_selected, conf_score_selected = select_mask(three_masks, conf_scores, coarse_ious=coarse_ious, 190                                                                rule="largest_over_conf", 191                                                                discourage_mask_below_confidence=discourage_mask_below_confidence, 192                                                                discourage_mask_below_coarse_iou=discourage_mask_below_coarse_iou,193                                                                verbose=True)194 195            mask_selected_list.append(mask_selected)196            conf_score_selected_list.append(conf_score_selected)197        mask_selected_batched_list.append(mask_selected_list)198        conf_score_selected_batched_list.append(conf_score_selected_list)199    200    return mask_selected_batched_list, conf_score_selected_batched_list201