{"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":"# DeepLabV3 + EfficientNet B4","metadata":{}},{"cell_type":"markdown","source":"# Setups & Configs","metadata":{}},{"cell_type":"code","source":"!pip install -U segmentation-models-pytorch -q","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:40.507920Z","iopub.execute_input":"2022-09-20T00:36:40.508213Z","iopub.status.idle":"2022-09-20T00:36:49.896229Z","shell.execute_reply.started":"2022-09-20T00:36:40.508153Z","shell.execute_reply":"2022-09-20T00:36:49.894904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom pathlib import Path\nfrom collections import Counter\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport segmentation_models_pytorch as smp\nfrom torchmetrics import MeanMetric, Dice\nfrom torch.optim import AdamW\nimport torch.cuda.amp as amp\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm.scheduler import CosineLRScheduler\n\nfrom sklearn.model_selection import StratifiedKFold\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\ndata_path = Path('/kaggle/input/hubmap-organ-segmentation')\nfiles = [f.name for f in list(data_path.glob('*'))]\nprint(files)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:49.901612Z","iopub.execute_input":"2022-09-20T00:36:49.904280Z","iopub.status.idle":"2022-09-20T00:36:53.300399Z","shell.execute_reply.started":"2022-09-20T00:36:49.904220Z","shell.execute_reply":"2022-09-20T00:36:53.299171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = True\nFOLD = 0\nN_FOLDS = 4\nIMG_SIZE = 704\nSEED = 43\nif DEBUG:\n    EPOCHS = 5\nelse:\n    EPOCHS = 100","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.301942Z","iopub.execute_input":"2022-09-20T00:36:53.303190Z","iopub.status.idle":"2022-09-20T00:36:53.308981Z","shell.execute_reply.started":"2022-09-20T00:36:53.303148Z","shell.execute_reply":"2022-09-20T00:36:53.307965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"def 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\n\ndef mask2rle(mask):\n    pixels = mask.T.flatten()   \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0]\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.312385Z","iopub.execute_input":"2022-09-20T00:36:53.312706Z","iopub.status.idle":"2022-09-20T00:36:53.321653Z","shell.execute_reply.started":"2022-09-20T00:36:53.312680Z","shell.execute_reply":"2022-09-20T00:36:53.320179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class HubmapHpaDataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame, \n        img_path: Path,\n        transform: callable = None, \n    ):\n        self.df = df\n        self.img_path = img_path\n        self.transform = transform\n        self.resize = int(IMG_SIZE*2)\n        self.weights = {\n            \"kidney\": 0.634,\n            \"prostate\": 0.808,\n            \"largeintestine\": 0.586,\n            \"spleen\": 1.69,\n            \"lung\": 1.87\n        }\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        \n        im = self.img_path / f\"{self.df['id'].iloc[idx]}.tiff\"\n        img = cv2.resize(cv2.imread(str(im)), (self.resize, self.resize))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        mask = cv2.resize(rle2mask(rle=self.df['rle'].iloc[idx], size=self.df['img_width'].iloc[idx]), \n                          (self.resize, self.resize))\n        mask = np.expand_dims(mask, axis=2)\n        if self.transform:\n            result = self.transform(image=img, mask=mask)\n            img, mask = result['image'], result['mask']\n            \n        return img, mask, self.weights[self.df.loc[idx, 'organ']]\n","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.323887Z","iopub.execute_input":"2022-09-20T00:36:53.324737Z","iopub.status.idle":"2022-09-20T00:36:53.334901Z","shell.execute_reply.started":"2022-09-20T00:36:53.324700Z","shell.execute_reply":"2022-09-20T00:36:53.333908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transform(augment: bool=False):\n    ops = []\n    if augment:\n        # geometric\n        ops += [\n            A.HorizontalFlip(p=0.5),\n            A.RandomRotate90(p=1),\n            A.ShiftScaleRotate(\n                shift_limit=0.0625, \n                scale_limit=[-0.2, 1], \n                rotate_limit=45,\n                p=0.9\n            )\n        ]\n        # color\n        ops += [\n            A.OneOf([\n                A.HueSaturationValue(hue_shift_limit=10,\n                                    sat_shift_limit=15,\n                                    val_shift_limit=10,\n                                    p=0.2),\n                A.CLAHE(clip_limit=2, p=0.2),\n                A.RandomBrightnessContrast(p=0.2),\n                A.ColorJitter(\n                    brightness=0.1, \n                    contrast=0.2, \n                    saturation=0.2, \n                    hue=0.5,\n                    p=0.4)\n            ]),\n        ]\n        # other\n        ops += [\n            A.ImageCompression(quality_lower=50, quality_upper=100, p=0.5),\n        ]\n        \n    ops += [\n            A.Resize(IMG_SIZE, IMG_SIZE),\n            A.CenterCrop(IMG_SIZE, IMG_SIZE),\n            A.Normalize(\n                mean=[0.82784054, 0.80224235, 0.8201268],\n                std=[0.16691974, 0.19275635, 0.17306973],\n            ),\n            ToTensorV2(transpose_mask=True)\n        ]\n    return A.Compose(ops)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.337355Z","iopub.execute_input":"2022-09-20T00:36:53.338100Z","iopub.status.idle":"2022-09-20T00:36:53.350200Z","shell.execute_reply.started":"2022-09-20T00:36:53.338063Z","shell.execute_reply":"2022-09-20T00:36:53.349204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_df = pd.read_csv(data_path / 'train.csv')\nall_df.head()\nFOLDS = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nfor fold, (train_idxs, val_idxs) in enumerate(FOLDS.split(X=all_df, y=all_df[\"organ\"])):\n    if fold != FOLD:\n        continue\n    train_df, valid_df = all_df.iloc[train_idxs].reset_index(drop=True), all_df.iloc[val_idxs].reset_index(drop=True)\n\nprint(len(train_df), len(valid_df))","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.353363Z","iopub.execute_input":"2022-09-20T00:36:53.353702Z","iopub.status.idle":"2022-09-20T00:36:53.497376Z","shell.execute_reply.started":"2022-09-20T00:36:53.353674Z","shell.execute_reply":"2022-09-20T00:36:53.496428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = data_path / 'train_images'\n\ntrain_dataset = HubmapHpaDataset(\n    df = train_df,\n    img_path = img_path,\n    transform = get_transform(augment=True)\n)\nvalid_dataset = HubmapHpaDataset(\n    df = valid_df,\n    img_path = img_path,\n    transform = get_transform(augment=False)\n)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.499255Z","iopub.execute_input":"2022-09-20T00:36:53.500021Z","iopub.status.idle":"2022-09-20T00:36:53.506454Z","shell.execute_reply.started":"2022-09-20T00:36:53.499984Z","shell.execute_reply":"2022-09-20T00:36:53.505412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Trainer","metadata":{}},{"cell_type":"code","source":"class Trainer:\n    def __init__(\n        self, \n        model,\n        train_loader, \n        valid_loader,\n        loss, \n        metric,\n        optimizer, \n        scheduler\n    ):\n        self.model = model\n        self.train_loader = train_loader\n        self.valid_loader = valid_loader\n        self.loss = loss\n        self.metric = metric\n        self.optimizer = optimizer\n        self.scaler = amp.GradScaler(enabled=True)\n        self.scheduler = scheduler\n        self.best_loss = 9999\n        self.best_metric = 0\n        self.history = {\n            'epochs': [],\n            'lr': [],\n            'train_loss': [],\n            'train_metric': [],\n            'valid_loss': [],\n            'valid_metric': [],\n        }\n        \n    def fit(self, epochs, device):\n        self.model.to(device)\n        for epoch in range(epochs):\n            self.history['epochs'].append(epoch)\n            \n            # Train epoch\n            self.model.train()\n            mean_loss = MeanMetric().to(device)\n            mean_metric = MeanMetric().to(device)\n            for x, y, w in tqdm(self.train_loader, desc=f\"Train epoch {epoch}\"):\n                x, y, w= x.to(device), y.to(device), w.to(device)\n                w /= w.sum()\n                \n                self.optimizer.zero_grad()\n                with amp.autocast(enabled=True):\n                    pred = self.model(x)\n                    loss = self.loss(pred, y.float(), weight=w.view(-1, 1, 1, 1))\n                self.scaler.scale(loss).backward()\n                self.scaler.unscale_(self.optimizer)\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n                \n                mean_metric.update(self.metric(pred, y.long()))\n                mean_loss.update(loss)\n                \n            self.history['train_loss'].append(mean_loss.compute().item())\n            self.history['train_metric'].append(mean_metric.compute().item())\n            self.history['lr'].append(self.scheduler.get_epoch_values(epoch))\n            self.scheduler.step(epoch+1)\n            \n            # Valid epoch\n            self.model.eval()\n            mean_loss = MeanMetric().to(device)\n            mean_metric = MeanMetric().to(device)\n            with torch.no_grad():\n                for x, y, w in tqdm(self.valid_loader, desc=f\"Valid epoch {epoch}\"):\n                    x, y, w= x.to(device), y.to(device), w.to(device)\n                    w /= w.sum()\n                    pred = self.model(x)\n                    loss = self.loss(pred, y.float(), weight=w.view(-1, 1, 1, 1))\n\n                    mean_metric.update(self.metric(pred, y.long()))\n                    mean_loss.update(loss)\n            \n            self.history['valid_loss'].append(mean_loss.compute().item())\n            self.history['valid_metric'].append(mean_metric.compute().item())\n            \n            self.save_if_best()\n            self.plot_test()\n    \n    def save_if_best(self):\n        if self.best_loss > self.history['valid_loss'][-1]:\n            self.best_loss = self.history['valid_loss'][-1]\n            print(f\"save best loss model: {self.best_loss:.4f}\")\n            torch.save(self.model, f'model_best_loss_{FOLD}.pth')\n        else:\n            print(f\"loss: {self.history['valid_loss'][-1]:.4f}\")\n\n        if self.best_metric < self.history['valid_metric'][-1]:\n            self.best_metric = self.history['valid_metric'][-1]\n            print(f\"save best metric model: {self.best_metric:.4f}\")\n            torch.save(self.model, f'model_best_metric_{FOLD}.pth')\n        else:\n            print(f\"metric: {self.history['valid_metric'][-1]:.4f}\")\n            \n    def plot_test(self):\n        transform = get_transform(augment=False)\n        img = cv2.imread(\"../input/hubmap-organ-segmentation/test_images/10078.tiff\")\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        X = transform(image=img)['image']\n        X = X.to(device)\n        with torch.no_grad():\n            Y = self.model(X.unsqueeze(0)).squeeze(0).sigmoid()\n        img = X.cpu().numpy().transpose(1,2,0)\n        pred = Y.cpu().numpy()[0]\n        pred_255 = (pred*255).astype(np.uint8)\n        thres = cv2.threshold(pred_255, 0, 255, cv2.THRESH_OTSU)[0] / 255\n        mask_pred = (pred > thres).astype(np.int8)\n        fig, axes = plt.subplots(1, 3, figsize=(9,3))\n        axes[0].imshow(img)\n        axes[0].set_title(\"Input\")\n        axes[1].imshow(pred)\n        axes[1].set_title(\"Pred\")\n        axes[2].imshow(mask_pred)\n        axes[2].set_title(\"Mask\")\n        plt.show()\n        \n    def plot_history(self):\n        fig, axes = plt.subplots(1, 3, figsize=(20, 6))\n        \n        axes[0].set_title('Loss')\n        axes[0].plot(self.history['epochs'], self.history['train_loss'], 'r', label='train')\n        axes[0].plot(self.history['epochs'], self.history['valid_loss'], 'g', label='valid')\n        axes[0].legend()\n        axes[0].grid()\n        \n        axes[1].set_title('Metric')\n        axes[1].plot(self.history['epochs'], self.history['train_metric'], 'r', label='train')\n        axes[1].plot(self.history['epochs'], self.history['valid_metric'], 'g', label='valid')\n        axes[1].legend()\n        axes[1].grid()\n\n        axes[2].set_title('Learning Rate')\n        axes[2].plot(self.history['epochs'], self.history['lr'], 'b')\n        axes[2].grid()\n\n        plt.plot()\n        plt.show()\n            ","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.508236Z","iopub.execute_input":"2022-09-20T00:36:53.508616Z","iopub.status.idle":"2022-09-20T00:36:53.538322Z","shell.execute_reply.started":"2022-09-20T00:36:53.508580Z","shell.execute_reply":"2022-09-20T00:36:53.537179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\nmodel = smp.DeepLabV3Plus(\n    encoder_name = 'timm-efficientnet-b4',\n    encoder_weights = 'noisy-student',\n    classes = 1,\n    activation = None,\n)\n\ntrain_loader = DataLoader(train_dataset, batch_size=7, shuffle=True, num_workers=2)\nvalid_loader = DataLoader(valid_dataset, batch_size=8, shuffle=False, num_workers=2)\n\nloss = F.binary_cross_entropy_with_logits\nmetric = Dice(\n    average='samples',\n    ignore_index=0,\n).to(device)\n\n\n# Copy setting of: https://www.kaggle.com/code/shionhonda/hubhpa-train-effnet-b7\noptimizer = AdamW(model.parameters(), lr=1e-3)\nscheduler = CosineLRScheduler(optimizer, t_initial=EPOCHS, lr_min=1e-5, \n                              warmup_t=EPOCHS//5, warmup_lr_init=1e-4, warmup_prefix=True)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:53.540114Z","iopub.execute_input":"2022-09-20T00:36:53.540513Z","iopub.status.idle":"2022-09-20T00:36:55.581258Z","shell.execute_reply.started":"2022-09-20T00:36:53.540478Z","shell.execute_reply":"2022-09-20T00:36:55.580228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    model=model,\n    train_loader=train_loader, \n    valid_loader=valid_loader,\n    loss=loss, \n    metric=metric,\n    optimizer=optimizer, \n    scheduler=scheduler,\n)\ndel model, optimizer","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:55.582488Z","iopub.execute_input":"2022-09-20T00:36:55.582850Z","iopub.status.idle":"2022-09-20T00:36:55.589198Z","shell.execute_reply.started":"2022-09-20T00:36:55.582812Z","shell.execute_reply":"2022-09-20T00:36:55.588327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(EPOCHS, device)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:36:55.590841Z","iopub.execute_input":"2022-09-20T00:36:55.591505Z","iopub.status.idle":"2022-09-20T00:53:25.913447Z","shell.execute_reply.started":"2022-09-20T00:36:55.591468Z","shell.execute_reply":"2022-09-20T00:53:25.912287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.plot_history()","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:53:25.918803Z","iopub.execute_input":"2022-09-20T00:53:25.919861Z","iopub.status.idle":"2022-09-20T00:53:26.377874Z","shell.execute_reply.started":"2022-09-20T00:53:25.919816Z","shell.execute_reply":"2022-09-20T00:53:26.376956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Best loss: {trainer.best_loss:.4f}\")\nprint(f\"Best dice: {trainer.best_metric:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:53:26.379452Z","iopub.execute_input":"2022-09-20T00:53:26.379801Z","iopub.status.idle":"2022-09-20T00:53:26.387087Z","shell.execute_reply.started":"2022-09-20T00:53:26.379765Z","shell.execute_reply":"2022-09-20T00:53:26.385952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict Test","metadata":{}},{"cell_type":"code","source":"del trainer\ntransform = get_transform(augment=False)\nimg = cv2.imread(\"../input/hubmap-organ-segmentation/test_images/10078.tiff\")\nimg = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nX = transform(image=img)['image']\nX = X.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:53:26.388579Z","iopub.execute_input":"2022-09-20T00:53:26.389620Z","iopub.status.idle":"2022-09-20T00:53:26.437086Z","shell.execute_reply.started":"2022-09-20T00:53:26.389584Z","shell.execute_reply":"2022-09-20T00:53:26.436105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.load(f'model_best_loss_{FOLD}.pth', map_location=\"cpu\")\nmodel.to(device)\nmodel.eval()\nwith torch.no_grad():\n    Y = model(X.unsqueeze(0)).squeeze(0).sigmoid()\nimg = X.cpu().numpy().transpose(1,2,0)\npred_best_loss = Y.cpu().numpy()[0]\npred_255 = (pred_best_loss*255).astype(np.uint8)\nthres = cv2.threshold(pred_255, 0, 255, cv2.THRESH_OTSU)[0] / 255\nmask_pred_best_loss = (pred_best_loss > thres).astype(np.int8)\ndel model\n\nmodel = torch.load(f'model_best_metric_{FOLD}.pth', map_location=\"cpu\")\nmodel.to(device)\nmodel.eval()\nwith torch.no_grad():\n    Y = model(X.unsqueeze(0)).squeeze(0).sigmoid()\npred_best_dice = Y.cpu().numpy()[0]\npred_255 = (pred_best_dice*255).astype(np.uint8)\nthres = cv2.threshold(pred_255, 0, 255, cv2.THRESH_OTSU)[0] / 255\nmask_pred_best_dice = (pred_best_dice > thres).astype(np.int8)\ndel model","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:53:26.439147Z","iopub.execute_input":"2022-09-20T00:53:26.439789Z","iopub.status.idle":"2022-09-20T00:53:27.054091Z","shell.execute_reply.started":"2022-09-20T00:53:26.439745Z","shell.execute_reply":"2022-09-20T00:53:27.052963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(2, 3, figsize=(12,8))\naxes[0, 0].imshow(img)\naxes[0, 0].set_title(\"Input\")\naxes[0, 1].imshow(pred_best_loss)\naxes[0, 1].set_title(\"Pred by Best Loss\")\naxes[0, 2].imshow(mask_pred_best_loss)\naxes[0, 2].set_title(\"Mask by Best Loss\")\naxes[1, 0].imshow(img)\naxes[1, 0].set_title(\"Input\")\naxes[1, 1].imshow(pred_best_dice)\naxes[1, 1].set_title(\"Pred by Best Dice\")\naxes[1, 2].imshow(mask_pred_best_dice)\naxes[1, 2].set_title(\"Mask by Best Dice\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-20T00:53:27.055592Z","iopub.execute_input":"2022-09-20T00:53:27.056294Z","iopub.status.idle":"2022-09-20T00:53:28.061666Z","shell.execute_reply.started":"2022-09-20T00:53:27.056238Z","shell.execute_reply":"2022-09-20T00:53:28.060820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}