{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":1176357,"sourceType":"datasetVersion","datasetId":667852},{"sourceId":5334514,"sourceType":"datasetVersion","datasetId":3098086}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# **Cat Background Removal Segment Anything**\n**with cat rectangle using yolo**","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y ray\n!pip uninstall -y numpy\n!pip install ultralytics\n!pip install -q --no-warn-conflicts \\\n    ultralytics==8.3.30 \\\n    albumentations==1.4.8 \\\n    lxml==5.3.0 \\\n    numpy==1.26.4 \\\n    tqdm==4.66.5\nbreak","metadata":{"trusted":true,"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**restart, run after**","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/facebookresearch/segment-anything.git","metadata":{"papermill":{"duration":18.845824,"end_time":"2023-05-04T07:54:48.820851","exception":false,"start_time":"2023-05-04T07:54:29.975027","status":"completed"},"tags":[],"trusted":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2026-02-10T09:57:57.352099Z","iopub.execute_input":"2026-02-10T09:57:57.352881Z","iopub.status.idle":"2026-02-10T09:58:03.374648Z","shell.execute_reply.started":"2026-02-10T09:57:57.352850Z","shell.execute_reply":"2026-02-10T09:58:03.373922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport random\nfrom PIL import Image\nimport torch\nimport numpy as np\n\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\n%matplotlib inline\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nfrom segment_anything import sam_model_registry \nfrom segment_anything import SamAutomaticMaskGenerator\nfrom segment_anything import SamPredictor\n\n!mkdir removal\n!mkdir masked\n!mkdir original\n!mkdir npy","metadata":{"papermill":{"duration":6.20077,"end_time":"2023-05-04T07:54:55.025887","exception":false,"start_time":"2023-05-04T07:54:48.825117","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:03.376303Z","iopub.execute_input":"2026-02-10T09:58:03.376640Z","iopub.status.idle":"2026-02-10T09:58:06.607884Z","shell.execute_reply.started":"2026-02-10T09:58:03.376598Z","shell.execute_reply":"2026-02-10T09:58:06.607143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_anns(anns, axes=None):\n    if len(anns) == 0:\n        return\n    if axes:\n        ax = axes\n    else:\n        ax = plt.gca()\n        ax.set_autoscale_on(False)\n    sorted_anns = sorted(anns, key=(lambda x: x['area']), reverse=True)\n    polygons = []\n    color = []\n    for ann in sorted_anns:\n        m = ann['segmentation']\n        img = np.ones((m.shape[0], m.shape[1], 3))\n        color_mask = np.random.random((1, 3)).tolist()[0]\n        for i in range(3):\n            img[:,:,i] = color_mask[i]\n            ax.imshow(np.dstack((img, m*1.0)))#m*0.5\n            \n        img2=np.dstack((img, m*1.0)) ### mask image\n        return img2 ### get mask image","metadata":{"papermill":{"duration":0.027679,"end_time":"2023-05-04T07:54:55.06011","exception":false,"start_time":"2023-05-04T07:54:55.032431","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:06.609113Z","iopub.execute_input":"2026-02-10T09:58:06.609790Z","iopub.status.idle":"2026-02-10T09:58:06.615807Z","shell.execute_reply.started":"2026-02-10T09:58:06.609759Z","shell.execute_reply":"2026-02-10T09:58:06.615246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sam_checkpoint = \"/kaggle/input/segment-anything-models/sam_vit_h_4b8939.pth\"\nmodel_type = \"vit_h\" #\ndevice = \"cpu\" #cpu,cuda\nsam = sam_model_registry[model_type](checkpoint=sam_checkpoint)\nsam.to(device=device)\nmask_generator1 = SamAutomaticMaskGenerator(sam)","metadata":{"_kg_hide-output":true,"papermill":{"duration":34.366681,"end_time":"2023-05-04T07:55:34.039013","exception":false,"start_time":"2023-05-04T07:54:59.672332","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:06.617817Z","iopub.execute_input":"2026-02-10T09:58:06.618234Z","iopub.status.idle":"2026-02-10T09:58:13.368593Z","shell.execute_reply.started":"2026-02-10T09:58:06.618205Z","shell.execute_reply":"2026-02-10T09:58:13.367924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"paths=[]\nfor dirname, _, filenames in os.walk('/kaggle/input/datasets/andrewmvd/animal-faces/afhq/train/cat'):\n    for filename in filenames:\n        paths+=[(os.path.join(dirname, filename))]\n        \nrandom.shuffle(paths)\npaths=paths[0:3]\nprint(paths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:13.369371Z","iopub.execute_input":"2026-02-10T09:58:13.369617Z","iopub.status.idle":"2026-02-10T09:58:15.168664Z","shell.execute_reply.started":"2026-02-10T09:58:13.369593Z","shell.execute_reply":"2026-02-10T09:58:15.167990Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"get rectangles of cat","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nfrom ultralytics import YOLO\nOUTPUT_DIR = \"removal\"\nCONF_THRES = 0.3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:15.169663Z","iopub.execute_input":"2026-02-10T09:58:15.169925Z","iopub.status.idle":"2026-02-10T09:58:15.342843Z","shell.execute_reply.started":"2026-02-10T09:58:15.169902Z","shell.execute_reply":"2026-02-10T09:58:15.342138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = YOLO(\"yolo11n.pt\")  \ncount = 0\n\nfor path in paths:\n    file=path.split('/')[-1]\n    img_pil = Image.open(path)\n    img = np.array(img_pil) \n    display(img_pil)\n    results = model(path, conf=CONF_THRES)\n    h, w, _ = img.shape\n    \n    for r in results:\n        for box in r.boxes:\n            cls_id = int(box.cls[0])\n            print(cls_id)\n            label = model.names[cls_id]\n            # Filter for \"cat\" only\n            if label != \"cat\":\n                continue\n            x1, y1, x2, y2 = map(int, box.xyxy[0])\n            print(x2-x1,y2-y1)\n            if x2-x1 > 30 and y2-y1 > 30:\n                # Crop the detected area\n                crop = img[y1:y2, x1:x2]\n                if crop.size == 0:\n                    continue\n                out_path = os.path.join(OUTPUT_DIR, file)\n                cv2.imwrite(out_path, crop)\n                count += 1\n            \n    print(path)\n    print(f\"Saved {count} cat images to '{OUTPUT_DIR}/'\")\n    print('================================'*2)\n    print()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:15.343785Z","iopub.execute_input":"2026-02-10T09:58:15.344039Z","iopub.status.idle":"2026-02-10T09:58:16.837223Z","shell.execute_reply.started":"2026-02-10T09:58:15.344017Z","shell.execute_reply":"2026-02-10T09:58:16.836640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:16.838057Z","iopub.execute_input":"2026-02-10T09:58:16.838312Z","iopub.status.idle":"2026-02-10T09:58:17.049260Z","shell.execute_reply.started":"2026-02-10T09:58:16.838290Z","shell.execute_reply":"2026-02-10T09:58:17.048493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"paths2=[]\nfor dirname, _, filenames in os.walk('removal'):\n    for filename in filenames:\n        paths2+=[(os.path.join(dirname, filename))]\n\nprint(paths2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:17.051547Z","iopub.execute_input":"2026-02-10T09:58:17.052143Z","iopub.status.idle":"2026-02-10T09:58:17.057814Z","shell.execute_reply.started":"2026-02-10T09:58:17.052102Z","shell.execute_reply":"2026-02-10T09:58:17.057083Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_mask(i):\n    path=paths2[i]\n    image = cv2.imread(path)\n    #image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    masks1 = mask_generator1.generate(image)\n    bgw = np.ones(image.shape)*255\n    \n    fig, ax = plt.subplots(1,3,figsize=(6,3))\n    maski=masks1[0]\n    maskiseg=maski['segmentation']    \n    box0=maski['bbox']\n    xc,yc,w,h = box0\n    x0=int(xc)\n    y0=int(yc)\n    x1=int(xc+w)\n    y1=int(yc+h)\n    \n    rect = patches.Rectangle( (xc,yc),w,h, linewidth=2, edgecolor='yellow', fill=False)\n    boximage=image[y0:y1,x0:x1,:]\n    boxmask=maskiseg[y0:y1,x0:x1]\n    stri=str(i).zfill(3)\n    cv2.imwrite(f'original/{stri}.png', boximage)\n    np.save(f'npy/{stri}.npy',boxmask)\n    #show_anns([masksi],ax[0])\n    ax[0].imshow(image)\n    ax[0].add_patch(rect)\n    ax[1].imshow(boximage)\n    ax[2].imshow(boxmask)\n    ax[0].set_title(f'mask{i}') \n    plt.show()  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:17.060274Z","iopub.execute_input":"2026-02-10T09:58:17.060549Z","iopub.status.idle":"2026-02-10T09:58:17.071076Z","shell.execute_reply.started":"2026-02-10T09:58:17.060528Z","shell.execute_reply":"2026-02-10T09:58:17.070348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(len(paths2)):\n    make_mask(i)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T09:58:17.072096Z","iopub.execute_input":"2026-02-10T09:58:17.072405Z","iopub.status.idle":"2026-02-10T10:16:40.124459Z","shell.execute_reply.started":"2026-02-10T09:58:17.072374Z","shell.execute_reply":"2026-02-10T10:16:40.123459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(len(paths2)):\n    \n    stri=str(i).zfill(3)\n    boximage=cv2.imread(f'original/{stri}.png')\n    maski=np.load(f'npy/{stri}.npy')\n\n    negative_img0 = np.tile(maski[:,:,np.newaxis],(1,1,3)).astype(int)\n    negative_img = negative_img0*255\n    positive_img0 = np.logical_not(negative_img)\n    positive_img = positive_img0.astype(np.uint8)*255\n    \n    masked_image3 = cv2.multiply(negative_img0, boximage, dtype=cv2.CV_8U)\n    masked_image3 = masked_image3.astype(np.uint8)\n    cv2.imwrite(f'masked/{stri}_masked.png', masked_image3)\n\n    plt.figure(figsize=(4,4))\n    plt.imshow(masked_image3)\n    plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-10T10:16:40.125833Z","iopub.execute_input":"2026-02-10T10:16:40.126108Z","iopub.status.idle":"2026-02-10T10:16:41.035261Z","shell.execute_reply.started":"2026-02-10T10:16:40.126083Z","shell.execute_reply":"2026-02-10T10:16:41.034050Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}