CoolFace
Apppublic

WompUniversity/Inpaint-Anything-no-errors

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
fill_anything.py128 linesDownload Raw Back to root
1import cv22import sys3import argparse4import numpy as np5import torch6from pathlib import Path7from matplotlib import pyplot as plt8from typing import Any, Dict, List9 10from sam_segment import predict_masks_with_sam11from stable_diffusion_inpaint import fill_img_with_sd12from utils import load_img_to_array, save_array_to_img, dilate_mask, \13    show_mask, show_points14 15 16def setup_args(parser):17    parser.add_argument(18        "--input_img", type=str, required=True,19        help="Path to a single input img",20    )21    parser.add_argument(22        "--point_coords", type=float, nargs='+', required=True,23        help="The coordinate of the point prompt, [coord_W coord_H].",24    )25    parser.add_argument(26        "--point_labels", type=int, nargs='+', required=True,27        help="The labels of the point prompt, 1 or 0.",28    )29    parser.add_argument(30        "--text_prompt", type=str, required=True,31        help="Text prompt",32    )33    parser.add_argument(34        "--dilate_kernel_size", type=int, default=None,35        help="Dilate kernel size. Default: None",36    )37    parser.add_argument(38        "--output_dir", type=str, required=True,39        help="Output path to the directory with results.",40    )41    parser.add_argument(42        "--sam_model_type", type=str,43        default="vit_h", choices=['vit_h', 'vit_l', 'vit_b'],44        help="The type of sam model to load. Default: 'vit_h"45    )46    parser.add_argument(47        "--sam_ckpt", type=str, required=True,48        help="The path to the SAM checkpoint to use for mask generation.",49    )50    parser.add_argument(51        "--seed", type=int,52        help="Specify seed for reproducibility.",53    )54    parser.add_argument(55        "--deterministic", action="store_true",56        help="Use deterministic algorithms for reproducibility.",57    )58 59 60 61if __name__ == "__main__":62    """Example usage:63    python fill_anything.py \64        --input_img FA_demo/FA1_dog.png \65        --point_coords 750 500 \66        --point_labels 1 \67        --text_prompt "a teddy bear on a bench" \68        --dilate_kernel_size 15 \69        --output_dir ./results \70        --sam_model_type "vit_h" \71        --sam_ckpt sam_vit_h_4b8939.pth 72    """73    parser = argparse.ArgumentParser()74    setup_args(parser)75    args = parser.parse_args(sys.argv[1:])76    device = "cuda" if torch.cuda.is_available() else "cpu"77 78    img = load_img_to_array(args.input_img)79 80    masks, _, _ = predict_masks_with_sam(81        img,82        [args.point_coords],83        args.point_labels,84        model_type=args.sam_model_type,85        ckpt_p=args.sam_ckpt,86        device=device,87    )88    masks = masks.astype(np.uint8) * 25589 90    # dilate mask to avoid unmasked edge effect91    if args.dilate_kernel_size is not None:92        masks = [dilate_mask(mask, args.dilate_kernel_size) for mask in masks]93 94    # visualize the segmentation results95    img_stem = Path(args.input_img).stem96    out_dir = Path(args.output_dir) / img_stem97    out_dir.mkdir(parents=True, exist_ok=True)98    for idx, mask in enumerate(masks):99        # path to the results100        mask_p = out_dir / f"mask_{idx}.png"101        img_points_p = out_dir / f"with_points.png"102        img_mask_p = out_dir / f"with_{Path(mask_p).name}"103 104        # save the mask105        save_array_to_img(mask, mask_p)106 107        # save the pointed and masked image108        dpi = plt.rcParams['figure.dpi']109        height, width = img.shape[:2]110        plt.figure(figsize=(width/dpi/0.77, height/dpi/0.77))111        plt.imshow(img)112        plt.axis('off')113        show_points(plt.gca(), [args.point_coords], args.point_labels,114                    size=(width*0.04)**2)115        plt.savefig(img_points_p, bbox_inches='tight', pad_inches=0)116        show_mask(plt.gca(), mask, random_color=False)117        plt.savefig(img_mask_p, bbox_inches='tight', pad_inches=0)118        plt.close()119 120    # fill the masked image121    for idx, mask in enumerate(masks):122        if args.seed is not None:123            torch.manual_seed(args.seed)124        mask_p = out_dir / f"mask_{idx}.png"125        img_filled_p = out_dir / f"filled_with_{Path(mask_p).name}"126        img_filled = fill_img_with_sd(127            img, mask, args.text_prompt, device=device)128        save_array_to_img(img_filled, img_filled_p)