{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# [HubMAP+HPA] Compare Masks and Predictions","metadata":{}},{"cell_type":"markdown","source":"## Setups","metadata":{}},{"cell_type":"code","source":"!pip install -U segmentation-models-pytorch -q","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport math\nimport random\nimport gc\nfrom pathlib import Path\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold\nfrom tqdm.notebook import tqdm\nimport segmentation_models_pytorch as smp\n\nimport matplotlib.pyplot as plt\n\nprint(f\"segmentation_models_pytorch {smp.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:24:14.569635Z","iopub.execute_input":"2022-09-07T12:24:14.57037Z","iopub.status.idle":"2022-09-07T12:24:15.544266Z","shell.execute_reply.started":"2022-09-07T12:24:14.570334Z","shell.execute_reply":"2022-09-07T12:24:15.542719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PLOT = True\nSEED = 43\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nN_FOLDS = 4\nWEIGHTS = [\n    \"../input/hubhpa-train-unet-effnet-b4-0/model_0.pth\",\n    # add your models here\n]\nIMG_SIZE = 640","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:22:44.565052Z","iopub.execute_input":"2022-09-07T12:22:44.565362Z","iopub.status.idle":"2022-09-07T12:22:44.635211Z","shell.execute_reply.started":"2022-09-07T12:22:44.565326Z","shell.execute_reply":"2022-09-07T12:22:44.63437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_score(y_true: np.ndarray, y_pred: np.ndarray, thres: float=0.5) -> float:\n    y_true = (y_true > thres).astype(int)\n    y_pred = (y_pred > thres).astype(int)\n    intersection = (y_true * y_pred).sum()\n    denominator = (y_true + y_pred).sum()\n    return 2 * intersection / (denominator + 1e-6)\n\ndef rle2mask(rle, size):\n    rle = np.array(list(map(int, rle.split())))\n    label = np.zeros((size*size), dtype=np.uint8)\n    for start, end in zip(rle[::2], rle[1::2]):\n        label[start:start+end] = 1\n    return label.reshape(size, size).T","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:35:22.535876Z","iopub.execute_input":"2022-09-07T12:35:22.536233Z","iopub.status.idle":"2022-09-07T12:35:22.547729Z","shell.execute_reply.started":"2022-09-07T12:35:22.536196Z","shell.execute_reply":"2022-09-07T12:35:22.546794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubHpaDataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame, \n        img_dir: Path,\n    ):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = A.Compose([\n            A.Resize(IMG_SIZE, IMG_SIZE),\n            A.CenterCrop(IMG_SIZE, IMG_SIZE),\n            A.Normalize(),\n            ToTensorV2(transpose_mask=True)\n        ])\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        path = self.img_dir / f\"{self.df.loc[idx, 'id']}.tiff\"\n        img = cv2.imread(str(path))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        mask = rle2mask(rle=self.df['rle'].iloc[idx], size=self.df['img_width'].iloc[idx])\n        mask = np.expand_dims(mask, axis=2)\n        res = self.transform(image=img, mask=mask)\n        return res['image'], res['mask']","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:35:14.897484Z","iopub.execute_input":"2022-09-07T12:35:14.898103Z","iopub.status.idle":"2022-09-07T12:35:14.906Z","shell.execute_reply.started":"2022-09-07T12:35:14.898065Z","shell.execute_reply":"2022-09-07T12:35:14.905227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cross Validation and Plot Results","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(\"../input/hubmap-organ-segmentation/train.csv\")\ndataset = HubHpaDataset(train, Path(\"../input/hubmap-organ-segmentation/train_images\"))\nskf = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:35:16.096266Z","iopub.execute_input":"2022-09-07T12:35:16.097136Z","iopub.status.idle":"2022-09-07T12:35:16.230174Z","shell.execute_reply.started":"2022-09-07T12:35:16.097078Z","shell.execute_reply":"2022-09-07T12:35:16.229304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot(idx, Y_pred, score):\n    iid = train.iloc[idx]['id']\n    h = train.iloc[idx]['img_height']\n    w = train.iloc[idx]['img_width']\n    organ = train.iloc[idx]['organ']\n    print(f\"ID: {iid} | Organ: {organ} | Dice: {score:.3f}\")\n    fig, axes = plt.subplots(1, 4, figsize=(14,4))\n    img = cv2.imread(f\"../input/hubmap-organ-segmentation/train_images/{iid}.tiff\")\n    axes[0].imshow(cv2.cvtColor(img, cv2.COLOR_BGR2RGB), vmin=0, vmax=255)\n    axes[0].set_title(\"Input\")\n    mask = rle2mask(train.iloc[idx]['rle'], size=h)\n    axes[1].imshow(mask, vmin=0, vmax=1)\n    axes[1].set_title(\"Mask\")\n    axes[2].imshow(Y_pred, vmin=0, vmax=1)\n    axes[2].set_title(\"Prediction\")\n    img = np.array([\n        mask,\n        cv2.resize(Y_pred, (h,w), interpolation=cv2.INTER_CUBIC).astype(float),\n        np.zeros((h,w))])\n    axes[3].imshow(img.transpose(1,2,0), vmin=0, vmax=1)\n    axes[3].set_title(\"Overlay\")\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if PLOT:\n    print(\"Yellow: True Positive, Red: False Negative, Green: False Positive\")\n    \nscores = []\nfor fold, (_, val_idxs) in enumerate(skf.split(X=train, y=train[\"organ\"])):\n    print(f\"===FOLD {fold}===\")\n    if fold >= len(WEIGHTS):\n        break\n    \n    model = torch.load(WEIGHTS[fold], map_location=\"cpu\")\n    model.to(DEVICE)\n    model.eval()\n    \n    _scores = []\n    for idx in val_idxs:\n        X, Y = dataset[idx]\n        X = X.to(DEVICE).unsqueeze(0)\n        with torch.no_grad():\n            Y_pred = model(X).squeeze(0).sigmoid().cpu().numpy()\n        _score = dice_score(Y.numpy(), Y_pred, thres=0.5)\n        _scores.append(_score)\n        if PLOT:\n            plot(idx, Y_pred[0], _score)\n    scores.append(np.mean(_scores))","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:36:27.429035Z","iopub.execute_input":"2022-09-07T12:36:27.429759Z","iopub.status.idle":"2022-09-07T12:36:44.358243Z","shell.execute_reply.started":"2022-09-07T12:36:27.429724Z","shell.execute_reply":"2022-09-07T12:36:44.357311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CV Score","metadata":{}},{"cell_type":"code","source":"print(f\"CV={np.mean(scores):.3f} ± {np.std(scores):.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-07T12:34:50.662369Z","iopub.status.idle":"2022-09-07T12:34:50.663218Z","shell.execute_reply.started":"2022-09-07T12:34:50.662976Z","shell.execute_reply":"2022-09-07T12:34:50.663002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}