{"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":"# Imports","metadata":{"id":"XhFxAengE7NL"}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch --quiet","metadata":{"id":"Od4ncF77E7NN","outputId":"d5445211-7324-41f0-af86-be7b0649dbaa","execution":{"iopub.status.busy":"2023-06-04T11:06:29.473599Z","iopub.execute_input":"2023-06-04T11:06:29.474613Z","iopub.status.idle":"2023-06-04T11:06:47.929847Z","shell.execute_reply.started":"2023-06-04T11:06:29.474573Z","shell.execute_reply":"2023-06-04T11:06:47.928654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport os\nimport gc\nimport glob\nimport json\nimport multiprocessing as mp\nimport warnings\nimport time\n\nimport albumentations as A\n\nimport matplotlib.pyplot as plt\nimport PIL.Image as Image\nimport cv2\n\nimport numpy as np\nimport pandas as pd\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.utils.data as thd\n\nimport segmentation_models_pytorch as smp\nimport segmentation_models_pytorch.utils as smp_utils\nimport segmentation_models_pytorch.utils.losses as smp_losses\n\nfrom torchvision import transforms\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau, OneCycleLR\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom segmentation_models_pytorch.encoders import get_preprocessing_fn\n\nfrom pathlib import Path\n\nfrom sklearn.metrics import fbeta_score\nfrom sklearn.exceptions import UndefinedMetricWarning\n\nfrom scipy.ndimage.filters import gaussian_filter1d\n\nfrom albumentations.pytorch import ToTensorV2\n\nfrom tqdm import tqdm\n\nwarnings.simplefilter('ignore')","metadata":{"id":"fQyhetQME7NO","outputId":"2aa9ab30-33b8-4e3a-9f88-8ef99eded152","execution":{"iopub.status.busy":"2023-06-04T11:06:47.932237Z","iopub.execute_input":"2023-06-04T11:06:47.932903Z","iopub.status.idle":"2023-06-04T11:06:54.968228Z","shell.execute_reply.started":"2023-06-04T11:06:47.932863Z","shell.execute_reply":"2023-06-04T11:06:54.967284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{"id":"SYc4Kg5rE7NO"}},{"cell_type":"code","source":"class CFG:\n    # ============== set paths =============\n    base_path = \"/kaggle/input\"\n    base_output_path = '/kaggle/working/'\n    comp_folder_name = 'vesuvius-challenge-ink-detection'\n    \n    comp_dataset_path = os.path.join(base_path, comp_folder_name)\n    train_dataset_path = os.path.join(comp_dataset_path, 'train')\n    test_dataset_path = os.path.join(comp_dataset_path, 'test')\n    \n    save_base_path = '/kaggle/working/'\n    trained_models_path_effiecient_net = '/kaggle/input/vesuvius-models/Unet_EfficientNetB4_model.pth'\n    trained_models_path_effiecient_netb6 = '/kaggle/input/vesuvius-models-efficientnetb6/Unet_EfficientNetB6_model.pth'\n    trained_models_path_reg_net = '/kaggle/input/vesuvius-unet-regnety32/Unet_RegNet_model.pth'\n    \n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = 'Unet'\n    backbone = 'efficientnet-b6'\n#     backbone = 'efficientnet-b4'\n#     backbone = 'se_resnext50_32x4d'\n#     backbone = 'resnet3d'\n#     backbone = 'efficientnet-b0'\n    \n    z_chans = 65\n    in_chans = 6 # 65\n    \n    # ============== training cfg =============\n    size = 384\n    tile_size = size\n    stride = size // 4\n    \n    batch_size = 16 # 32\n    \n    epochs = 5 # 30    \n    \n    lr = 1e-3\n\n    # ============== fold =============\n    valid_id = 1\n\n    # ============== fixed =============\n    weight_decay = 1e-3 # 1e-4\n\n    num_workers = 4\n\n    seed = 42\n\n    # ============== augmentation =============\n    train_aug_list = [\n        A.Resize(size, size),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.75),\n        A.ShiftScaleRotate(p=0.75),\n        A.OneOf([\n                A.GaussNoise(var_limit=[10, 50]),\n                A.GaussianBlur(),\n                A.MotionBlur(),\n                ], p=0.4),\n        A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n#         A.CoarseDropout(max_holes=1, max_width=int(size * 0.3), max_height=int(size * 0.3), \n#                         mask_fill_value=0, p=0.5),\n        # A.Cutout(max_h_size=int(size * 0.6),\n        #          max_w_size=int(size * 0.6), num_holes=1, p=1.0),\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\n    \n    valid_aug_list = [\n        A.Resize(size, size),\n        A.Normalize(\n            mean= [0] * in_chans,\n            std= [1] * in_chans\n        ),\n        ToTensorV2(transpose_mask=True)\n    ]\n    \n    test_aug_list = [\n        A.HorizontalFlip(),\n        A.VerticalFlip(),\n        ToTensorV2(),\n    ]","metadata":{"id":"NcD3gE7fE7NQ","execution":{"iopub.status.busy":"2023-06-04T11:06:54.969884Z","iopub.execute_input":"2023-06-04T11:06:54.970225Z","iopub.status.idle":"2023-06-04T11:06:55.183693Z","shell.execute_reply.started":"2023-06-04T11:06:54.970192Z","shell.execute_reply":"2023-06-04T11:06:55.182653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=None, cudnn_deterministic=True):\n    if seed is None:\n        seed = 42\n\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = cudnn_deterministic\n    torch.backends.cudnn.benchmark = False\n    \nset_seed(CFG.seed)","metadata":{"id":"kw6M4CcYE7NQ","execution":{"iopub.status.busy":"2023-06-04T11:06:55.186734Z","iopub.execute_input":"2023-06-04T11:06:55.187390Z","iopub.status.idle":"2023-06-04T11:06:55.204582Z","shell.execute_reply.started":"2023-06-04T11:06:55.187354Z","shell.execute_reply":"2023-06-04T11:06:55.203513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"id":"GAotprlOE7NR","execution":{"iopub.status.busy":"2023-06-04T11:06:55.206458Z","iopub.execute_input":"2023-06-04T11:06:55.206833Z","iopub.status.idle":"2023-06-04T11:06:55.236336Z","shell.execute_reply.started":"2023-06-04T11:06:55.206801Z","shell.execute_reply":"2023-06-04T11:06:55.235319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper","metadata":{"id":"xiiA4cqTE7NS"}},{"cell_type":"code","source":"# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    # pixels = (pixels >= thr).astype(int)\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"id":"KbriZyQ1E7NT","execution":{"iopub.status.busy":"2023-06-04T11:06:55.238071Z","iopub.execute_input":"2023-06-04T11:06:55.238413Z","iopub.status.idle":"2023-06-04T11:06:55.247464Z","shell.execute_reply.started":"2023-06-04T11:06:55.238383Z","shell.execute_reply":"2023-06-04T11:06:55.246436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_learning_curves(logs_df, valid=False):\n    fig, axis = plt.subplots(1, 3, figsize=(24, 8))\n\n    fig.suptitle('Learning curves')\n    axis[0].plot(logs_df.index.tolist(), logs_df['dice_loss'], label='Train Loss')\n    axis[0].set_ylabel('Dice Loss')\n    axis[0].set_xlabel('Epoch')\n    \n    axis[1].plot(logs_df.index.tolist(), logs_df['iou_score'], label='Train IOU')\n    axis[1].set_ylabel('IOU')\n    axis[1].set_xlabel('Epoch')\n    \n    axis[2].plot(logs_df.index.tolist(), logs_df['fscore'], label='Train Fscore')\n    axis[2].set_ylabel('Fscore')\n    axis[2].set_xlabel('Epoch')\n    \n    if valid:\n        axis[0].plot(logs_df.index.tolist(), logs_df['val_dice_loss'], label='Validation Loss')\n        axis[1].plot(logs_df.index.tolist(), logs_df['val_iou_score'], label='Validation IOU')\n        axis[2].plot(logs_df.index.tolist(), logs_df['val_fscore'], label='Validation Fscore')\n        \n    axis[0].legend()\n    axis[1].legend()\n    axis[2].legend()\n    \n    \ndef plot_pred_and_gt_mask(mask_pred, valid_mask_gt, thresh_hold):\n    fig, axes = plt.subplots(1, 3, figsize=(15, 8))\n    axes[0].imshow(valid_mask_gt)\n    axes[1].imshow(mask_pred)\n    axes[2].imshow((mask_pred>=thresh_hold).astype(int))\n\n\ndef plot_learning_curves_with_trend_line(logs_df, valid=False):\n    fig, axis = plt.subplots(1, 3, figsize=(24, 8))\n    fig.suptitle('smoothened Learning curves')\n    \n    # Plotting training curves\n    train_loss_smooth = gaussian_filter1d(logs_df['dice_loss'], sigma=2)\n    train_iou_smooth = gaussian_filter1d(logs_df['iou_score'], sigma=2)\n    train_fscore_smooth = gaussian_filter1d(logs_df['fscore'], sigma=2)\n    \n    axis[0].plot(logs_df.index.tolist(), train_loss_smooth, label='Train Loss')\n    axis[1].plot(logs_df.index.tolist(), train_iou_smooth, label='Train IOU')\n    axis[2].plot(logs_df.index.tolist(), train_fscore_smooth, label='Train Fscore')\n    \n    # Plotting validation curves if valid is True\n    if valid:\n        val_loss_smooth = gaussian_filter1d(logs_df['val_dice_loss'], sigma=2)\n        val_iou_smooth = gaussian_filter1d(logs_df['val_iou_score'], sigma=2)\n        val_fscore_smooth = gaussian_filter1d(logs_df['val_fscore'], sigma=2)\n        \n        axis[0].plot(logs_df.index.tolist(), val_loss_smooth, label='Validation Loss')\n        axis[1].plot(logs_df.index.tolist(), val_iou_smooth, label='Validation IOU')\n        axis[2].plot(logs_df.index.tolist(), val_fscore_smooth, label='Validation Fscore')\n    \n    # Add trend lines\n    axis[0].plot(logs_df.index.tolist(), train_loss_smooth, label='Train Loss Trend', linestyle='--', color='gray')\n    axis[1].plot(logs_df.index.tolist(), train_iou_smooth, label='Train IOU Trend', linestyle='--', color='gray')\n    axis[2].plot(logs_df.index.tolist(), train_fscore_smooth, label='Train Fscore Trend', linestyle='--', color='gray')\n    \n    if valid:\n        axis[0].plot(logs_df.index.tolist(), val_loss_smooth, label='Validation Loss Trend', linestyle='--', color='gray')\n        axis[1].plot(logs_df.index.tolist(), val_iou_smooth, label='Validation IOU Trend', linestyle='--', color='gray')\n        axis[2].plot(logs_df.index.tolist(), val_fscore_smooth, label='Validation Fscore Trend', linestyle='--', color='gray')\n        \n    axis[0].set_ylabel('Dice Loss')\n    axis[0].set_xlabel('Epoch')\n    axis[1].set_ylabel('IOU')\n    axis[1].set_xlabel('Epoch')\n    axis[2].set_ylabel('Fscore')\n    axis[2].set_xlabel('Epoch')\n    \n    axis[0].legend()\n    axis[1].legend()\n    axis[2].legend()","metadata":{"id":"Uw6x2DJSE7NT","execution":{"iopub.status.busy":"2023-06-04T11:06:55.249333Z","iopub.execute_input":"2023-06-04T11:06:55.249705Z","iopub.status.idle":"2023-06-04T11:06:55.272148Z","shell.execute_reply.started":"2023-06-04T11:06:55.249674Z","shell.execute_reply":"2023-06-04T11:06:55.271132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_pred_mask(model, dataloader, valid_mask_gt):\n    mask_pred = np.zeros_like(valid_mask_gt)\n    mask_count = np.zeros_like(valid_mask_gt)\n    \n    model.eval()\n    \n    for step, (images, _) in tqdm(enumerate(dataloader), total=len(dataloader)):\n        images = images.to(device)\n        batch_size = images.size(0)\n\n        with torch.no_grad():\n#             y_preds = tta(images).cpu().numpy()\n            # y_preds = model(images).cpu().numpy()\n            y_preds = TTA(images, model).cpu().numpy()\n\n        start_idx = step * CFG.batch_size\n        end_idx = start_idx + batch_size\n        for i, (x1, y1, x2, y2) in enumerate(valid_xyxys[start_idx:end_idx]):\n            mask_pred[y1:y2, x1:x2] += y_preds[i].reshape(mask_pred[y1:y2, x1:x2].shape)\n            mask_count[y1:y2, x1:x2] += np.ones((CFG.tile_size, CFG.tile_size))\n            \n    mask_pred /= mask_count\n    return mask_pred\n\n\ndef find_best_threshold(model, pred_mask, valid_mask_gt):\n    thresholds = np.arange(0.5, 1.05, 0.05)\n    best_score = 0.0\n    best_threshold = 0.0\n\n    model.eval()\n\n    for threshold in thresholds:\n        with torch.no_grad():\n            temp_mask = (pred_mask >= threshold).astype(int)\n            temp_mask_tensor = torch.from_numpy(temp_mask)\n            valid_mask_gt_tensor = torch.from_numpy(valid_mask_gt)\n            \n            score = smp.utils.functional.f_score(temp_mask_tensor, valid_mask_gt_tensor, beta=0.5)\n            print(f\"Threashold: {threshold}, Score: {score}\")\n            \n            if score > best_score:\n                best_score = score\n                best_threshold = threshold\n    \n    return best_threshold","metadata":{"id":"nHFq3vMxE7Nd","execution":{"iopub.status.busy":"2023-06-04T11:16:11.714815Z","iopub.execute_input":"2023-06-04T11:16:11.715183Z","iopub.status.idle":"2023-06-04T11:16:11.729235Z","shell.execute_reply.started":"2023-06-04T11:16:11.715154Z","shell.execute_reply":"2023-06-04T11:16:11.727905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Set up data","metadata":{"id":"nX2V_tY1E7NU"}},{"cell_type":"markdown","source":"**Data setup functions**","metadata":{"id":"Tun2DXuwE7NU"}},{"cell_type":"code","source":"def read_image_mask(fragment_id):\n    z_dim = CFG.in_chans\n    z_mid = 65 // 2 \n    z_start, z_end = z_mid - z_dim // 2, z_mid + z_dim // 2\n    indx = range(z_start, z_end)\n\n    images = []\n    for i in tqdm(indx):\n        image = cv2.imread(f\"{CFG.train_dataset_path}/{fragment_id}/surface_volume/{i:02}.tif\", 0)\n\n        pad0 = (CFG.tile_size - image.shape[0] % CFG.tile_size)\n        pad1 = (CFG.tile_size - image.shape[1] % CFG.tile_size)\n\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0)\n        images.append(image)\n\n    images = np.stack(images, axis=2)\n    mask = cv2.imread(f\"{CFG.train_dataset_path}/{fragment_id}/inklabels.png\", 0)\n\n    mask = np.pad(mask, [(0, pad0), (0, pad1)], constant_values=0)\n    mask = mask.astype('float32')\n    mask /= 255.0\n\n    return images, mask\n\n\ndef slice_fragment_to_subvolumes(fragment_ids):\n    train_images = []\n    train_masks = []\n\n    valid_images = []\n    valid_masks = []\n    valid_xyxys = []\n\n    for fragment_id in range(1, 4):\n\n        image, mask = read_image_mask(fragment_id)\n            \n        x1_list = list(range(0, image.shape[1]-CFG.tile_size+1, CFG.stride))\n        y1_list = list(range(0, image.shape[0]-CFG.tile_size+1, CFG.stride))\n\n        for y1 in y1_list:\n            for x1 in x1_list:\n                y2 = y1 + CFG.tile_size\n                x2 = x1 + CFG.tile_size\n                \n                sliced_image = image[y1:y2, x1:x2]\n                sliced_mask = mask[y1:y2, x1:x2, None]\n                \n\n                if fragment_id == CFG.valid_id:\n                    valid_images.append(sliced_image)\n                    valid_masks.append(sliced_mask)\n                    valid_xyxys.append([x1, y1, x2, y2])\n                else:\n                    train_images.append(sliced_image)\n                    train_masks.append(sliced_mask)\n\n    return train_images, train_masks, valid_images, valid_masks, valid_xyxys","metadata":{"id":"emUTc6CxE7NU","execution":{"iopub.status.busy":"2023-06-04T11:06:55.287712Z","iopub.execute_input":"2023-06-04T11:06:55.288272Z","iopub.status.idle":"2023-06-04T11:06:55.302069Z","shell.execute_reply.started":"2023-06-04T11:06:55.288240Z","shell.execute_reply":"2023-06-04T11:06:55.301407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Dataset class**","metadata":{"id":"1cNIc_fVE7NU"}},{"cell_type":"code","source":"class SubvolumeDataset(thd.Dataset):\n    def __init__(\n        self,\n        images,\n        masks,\n        transform=None,\n    ):\n        self.images = images\n        self.masks = masks\n        self.transform = transform\n\n    \n    def __len__(self):\n        return len(self.images)\n\n    \n    def __getitem__(self, index):\n        images = self.images[index]\n        mask = self.masks[index]\n            \n        if self.transform:\n            data = self.transform(image=images, mask=mask)\n            images = data['image']\n            mask = data['mask']\n            \n        return images, mask","metadata":{"id":"wpl7KV1tE7NU","execution":{"iopub.status.busy":"2023-06-04T11:06:55.306440Z","iopub.execute_input":"2023-06-04T11:06:55.306969Z","iopub.status.idle":"2023-06-04T11:06:55.314551Z","shell.execute_reply.started":"2023-06-04T11:06:55.306935Z","shell.execute_reply":"2023-06-04T11:06:55.313977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = Path(CFG.train_dataset_path)\nall_fragments = sorted([f.name for f in train_path.iterdir()])\n\ntrain_trasnforms = A.Compose(CFG.train_aug_list)\nval_transforms = A.Compose(CFG.valid_aug_list)","metadata":{"id":"wYdniqomE7NV","execution":{"iopub.status.busy":"2023-06-04T11:06:55.315864Z","iopub.execute_input":"2023-06-04T11:06:55.316400Z","iopub.status.idle":"2023-06-04T11:06:55.330031Z","shell.execute_reply.started":"2023-06-04T11:06:55.316369Z","shell.execute_reply":"2023-06-04T11:06:55.329135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Slices 3d fragments to [6, 384, 384] images and [384, 384] mask pairs**","metadata":{"id":"s28Is2smE7NV"}},{"cell_type":"code","source":"train_images, train_masks, valid_images, valid_masks, valid_xyxys = slice_fragment_to_subvolumes(all_fragments)","metadata":{"id":"rC9110NEE7NV","outputId":"f723dcec-c2e7-4077-ff92-e67f4ec99af1","execution":{"iopub.status.busy":"2023-06-04T11:06:55.331336Z","iopub.execute_input":"2023-06-04T11:06:55.331924Z","iopub.status.idle":"2023-06-04T11:07:37.955110Z","shell.execute_reply.started":"2023-06-04T11:06:55.331893Z","shell.execute_reply":"2023-06-04T11:07:37.954169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Create train and validation datasets**","metadata":{"id":"qhCYOmQ9E7NW"}},{"cell_type":"code","source":"train_dataset = SubvolumeDataset(train_images, train_masks, train_trasnforms)\nval_dataset = SubvolumeDataset(valid_images, valid_masks, val_transforms)\n\n\ntrain_data_loader = thd.DataLoader(train_dataset,\n                               batch_size=CFG.batch_size,\n                               shuffle=True,\n                               num_workers=CFG.num_workers, pin_memory=True, drop_last=True,\n                              )\n\nval_data_loader = thd.DataLoader(val_dataset,\n                               batch_size=CFG.batch_size,\n                               shuffle=False,\n                               num_workers=CFG.num_workers, pin_memory=True, drop_last=True,\n                              )\n\ndataloaders = {\n    'train': train_data_loader,\n    'val': val_data_loader\n}","metadata":{"id":"yrs2tOgWE7NW","execution":{"iopub.status.busy":"2023-06-04T11:07:37.957454Z","iopub.execute_input":"2023-06-04T11:07:37.958278Z","iopub.status.idle":"2023-06-04T11:07:37.965915Z","shell.execute_reply.started":"2023-06-04T11:07:37.958232Z","shell.execute_reply":"2023-06-04T11:07:37.964762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Sanity check ","metadata":{"id":"xB7NTjIZE7NW"}},{"cell_type":"code","source":"plot_dataset = SubvolumeDataset(train_images, train_masks)\n\nindex = 500\nprint(f\"Sub Volume image shape = {plot_dataset[index][0].shape}\")\nprint(f\"Sub Volume mask shape = {plot_dataset[index][1].shape}\")\nprint(f\"dataset len = {len(plot_dataset)}\")\n\ntransform = CFG.train_aug_list\ntransform = A.Compose(\n    [t for t in transform if not isinstance(t, (A.Normalize, ToTensorV2))])\n\nplot_count = 0\nfor i in range(5000, 7000):\n\n    image, mask = plot_dataset[i]\n    data = transform(image=image, mask=mask)\n    aug_image = data['image']\n    aug_mask = data['mask']\n    \n    if mask.sum() == 0.0:\n        continue\n\n    fig, axes = plt.subplots(1, 4, figsize=(15, 8))\n    axes[0].imshow(image[..., 0], cmap=\"gray\")\n    axes[1].imshow(mask, cmap=\"gray\")\n    axes[2].imshow(aug_image[..., 0], cmap=\"gray\")\n    axes[3].imshow(aug_mask, cmap=\"gray\")\n\n    plot_count += 1\n    if plot_count == 3:\n        break\n        \ndel plot_dataset\ngc.collect()","metadata":{"id":"0zPi6Rt9E7NX","outputId":"70debc10-3176-44b7-86f4-3f56101e20e8","execution":{"iopub.status.busy":"2023-06-04T11:07:37.968393Z","iopub.execute_input":"2023-06-04T11:07:37.968733Z","iopub.status.idle":"2023-06-04T11:07:40.360702Z","shell.execute_reply.started":"2023-06-04T11:07:37.968707Z","shell.execute_reply":"2023-06-04T11:07:40.359758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"O_RwjMJVE7NX"}},{"cell_type":"code","source":"class InkDetector(torch.nn.Module):\n    def __init__(self, cfg, weight='imagenet'):\n        super().__init__()\n        self.cfg = cfg\n        self.model = smp.Unet(\n            encoder_name = cfg.backbone, \n            encoder_weights = weight,\n            in_channels = cfg.in_chans,\n            classes = cfg.target_size,\n            activation = 'sigmoid',\n        )\n\n    def forward(self, image):\n        output = self.model(image)\n        return output","metadata":{"id":"9LBmNNYYE7NX","execution":{"iopub.status.busy":"2023-06-04T11:07:40.362117Z","iopub.execute_input":"2023-06-04T11:07:40.362669Z","iopub.status.idle":"2023-06-04T11:07:40.370675Z","shell.execute_reply.started":"2023-06-04T11:07:40.362618Z","shell.execute_reply":"2023-06-04T11:07:40.369688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{"id":"TONi9ZGDE7NY"}},{"cell_type":"code","source":"loss = smp_losses.DiceLoss()","metadata":{"id":"fS_0leW8E7NY","execution":{"iopub.status.busy":"2023-06-04T11:07:40.372175Z","iopub.execute_input":"2023-06-04T11:07:40.372752Z","iopub.status.idle":"2023-06-04T11:07:40.381982Z","shell.execute_reply.started":"2023-06-04T11:07:40.372705Z","shell.execute_reply":"2023-06-04T11:07:40.380852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics","metadata":{"id":"m2gFzlQJE7NY"}},{"cell_type":"code","source":"metrics = [\n    smp_utils.metrics.IoU(threshold=0.5),\n    smp_utils.metrics.Fscore(beta=0.5)\n]","metadata":{"id":"jnG6s8q7E7NZ","execution":{"iopub.status.busy":"2023-06-04T11:07:40.383616Z","iopub.execute_input":"2023-06-04T11:07:40.384004Z","iopub.status.idle":"2023-06-04T11:07:40.393343Z","shell.execute_reply.started":"2023-06-04T11:07:40.383973Z","shell.execute_reply":"2023-06-04T11:07:40.392429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"id":"iaRDYh89E7Nb"}},{"cell_type":"code","source":"def train_model(model, model_epochs, dataloaders, scheduler, num_epochs=10, track=False, path=CFG.trained_models_path_effiecient_netb6):\n    logs = {\n        'dice_loss': [],\n        'iou_score': [],\n        'fscore': [],\n        'val_dice_loss': [],\n        'val_iou_score': [],\n        'val_fscore': []\n    }\n    \n    max_score = 0\n    \n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch + 1}/{num_epochs}')\n        \n        # Train\n        train_epoch = model_epochs['train']\n        train_log = train_epoch.run(dataloaders['train'])\n        scheduler.step(epoch)\n        \n        if max_score < train_log['fscore']:\n            torch.save(model.state_dict(), os.path.join(CFG.save_base_path, f'{CFG.backbone}.pth'))\n            torch.save(model, path)\n            max_score = train_log['fscore']\n        \n        for scr, res in train_log.items():\n            logs[scr].append(res)\n        \n        if track:\n            # Validate\n            val_epoch = model_epochs['val']\n            val_log = val_epoch.run(dataloaders['val'])\n\n            for scr, res in val_log.items():\n                logs[f'val_{scr}'].append(res)\n        \n    return logs","metadata":{"id":"DLyNk-siE7Nb","execution":{"iopub.status.busy":"2023-06-04T11:07:40.394723Z","iopub.execute_input":"2023-06-04T11:07:40.395137Z","iopub.status.idle":"2023-06-04T11:07:40.406795Z","shell.execute_reply.started":"2023-06-04T11:07:40.395104Z","shell.execute_reply":"2023-06-04T11:07:40.405746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Validation","metadata":{"id":"luFlZeWKE7Nc"}},{"cell_type":"code","source":"VALID = False\nfragment_id = CFG.valid_id\n\nvalid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"/train/{fragment_id}/inklabels.png\", 0)\nvalid_mask_gt = valid_mask_gt / 255\npad0 = (CFG.tile_size - valid_mask_gt.shape[0] % CFG.tile_size)\npad1 = (CFG.tile_size - valid_mask_gt.shape[1] % CFG.tile_size)\nvalid_mask_gt = np.pad(valid_mask_gt, [(0, pad0), (0, pad1)], constant_values=0)","metadata":{"id":"dtnP8ktNE7Nc","execution":{"iopub.status.busy":"2023-06-04T11:07:40.408325Z","iopub.execute_input":"2023-06-04T11:07:40.408712Z","iopub.status.idle":"2023-06-04T11:07:41.014466Z","shell.execute_reply.started":"2023-06-04T11:07:40.408681Z","shell.execute_reply":"2023-06-04T11:07:41.013468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TTA","metadata":{"id":"moFTonjlQ_pg"}},{"cell_type":"code","source":"def TTA(x:torch.Tensor,model:nn.Module):\n    #x.shape=(batch,c,h,w)\n    shape=x.shape\n    x=[x,*[torch.rot90(x,k=i,dims=(-2,-1)) for i in range(1,4)]]\n    x=torch.cat(x,dim=0)\n    x=model(x)\n    x=torch.sigmoid(x)\n    x=x.reshape(4,shape[0],*shape[2:])\n    x=[torch.rot90(x[i],k=-i,dims=(-2,-1)) for i in range(4)]\n    x=torch.stack(x,dim=0)\n    return x.mean(0)","metadata":{"id":"ifOqPaZtE7Nc","execution":{"iopub.status.busy":"2023-06-04T11:07:41.016075Z","iopub.execute_input":"2023-06-04T11:07:41.016439Z","iopub.status.idle":"2023-06-04T11:07:41.024493Z","shell.execute_reply.started":"2023-06-04T11:07:41.016406Z","shell.execute_reply":"2023-06-04T11:07:41.023597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EfficientNetB4","metadata":{"id":"tDnLKHl-eGWW"}},{"cell_type":"code","source":"if os.path.exists(CFG.trained_models_path_effiecient_net):\n   model = torch.load(CFG.trained_models_path_effiecient_net, map_location=device)\n   print(\"ok\")\nelse:\n    model = InkDetector(CFG)\n    \nmodel.to(device);","metadata":{"id":"98pdguTBE7NY","execution":{"iopub.status.busy":"2023-06-04T11:07:41.026079Z","iopub.execute_input":"2023-06-04T11:07:41.026716Z","iopub.status.idle":"2023-06-04T11:07:45.413401Z","shell.execute_reply.started":"2023-06-04T11:07:41.026678Z","shell.execute_reply":"2023-06-04T11:07:45.412388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer","metadata":{"id":"P0VhhUGFE7NZ"}},{"cell_type":"code","source":"optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)","metadata":{"id":"Fkv7UVBIE7Na","execution":{"iopub.status.busy":"2023-06-04T11:07:45.414992Z","iopub.execute_input":"2023-06-04T11:07:45.415333Z","iopub.status.idle":"2023-06-04T11:07:45.424450Z","shell.execute_reply.started":"2023-06-04T11:07:45.415302Z","shell.execute_reply":"2023-06-04T11:07:45.423305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scheduler = OneCycleLR(optimizer, max_lr=5e-3, epochs=CFG.epochs, steps_per_epoch=len(train_dataset)//CFG.batch_size)","metadata":{"id":"7xVc91t2E7Na","execution":{"iopub.status.busy":"2023-06-04T11:08:27.039253Z","iopub.execute_input":"2023-06-04T11:08:27.039616Z","iopub.status.idle":"2023-06-04T11:08:27.044249Z","shell.execute_reply.started":"2023-06-04T11:08:27.039587Z","shell.execute_reply":"2023-06-04T11:08:27.043363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{"id":"m_a9zOkNgdlr"}},{"cell_type":"code","source":"TRACK_EFFICIENT = False\nTRAIN_EFFICIENT = False\nVALID_EFFICIENT = False","metadata":{"id":"dkAosfjRTM54","execution":{"iopub.status.busy":"2023-06-04T11:08:31.142555Z","iopub.execute_input":"2023-06-04T11:08:31.143581Z","iopub.status.idle":"2023-06-04T11:08:31.147499Z","shell.execute_reply.started":"2023-06-04T11:08:31.143548Z","shell.execute_reply":"2023-06-04T11:08:31.146305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_epoch = smp_utils.train.TrainEpoch(\n    model, \n    loss=loss, \n    metrics=metrics, \n    optimizer=optimizer,\n    device=device,\n    verbose=True,\n)\n\nval_epoch = smp_utils.train.ValidEpoch (\n    model, \n    loss=loss, \n    metrics=metrics, \n    device=device,\n    verbose=True,\n)\n\nmodel_epochs = {\n    'train': train_epoch,\n    'val': val_epoch\n}","metadata":{"id":"9pWTpcFtym6b","execution":{"iopub.status.busy":"2023-06-04T11:08:31.326307Z","iopub.execute_input":"2023-06-04T11:08:31.326675Z","iopub.status.idle":"2023-06-04T11:08:31.348205Z","shell.execute_reply.started":"2023-06-04T11:08:31.326639Z","shell.execute_reply":"2023-06-04T11:08:31.347079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_EFFICIENT:\n    logs = train_model(model_epochs, sub_dataloaders, scheduler, CFG.epochs, TRACK_EFFICIENT)  ","metadata":{"id":"bFuQLAwME7Nb","execution":{"iopub.status.busy":"2023-06-04T11:08:31.472491Z","iopub.execute_input":"2023-06-04T11:08:31.472899Z","iopub.status.idle":"2023-06-04T11:08:31.479252Z","shell.execute_reply.started":"2023-06-04T11:08:31.472865Z","shell.execute_reply":"2023-06-04T11:08:31.477947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_EFFICIENT:\n    if not TRACK_EFFICIENT:\n        train_logs = {}\n        train_logs['dice_loss'] = logs['dice_loss']\n        train_logs['iou_score'] = logs['iou_score']\n        train_logs['fscore'] = logs['fscore']\n        df = pd.DataFrame(train_logs)\n    else:\n        df = pd.DataFrame(logs)\n    plot_learning_curves(df, TRACK)","metadata":{"id":"Hw-E3sfOE7Nb","execution":{"iopub.status.busy":"2023-06-04T11:08:31.626500Z","iopub.execute_input":"2023-06-04T11:08:31.627364Z","iopub.status.idle":"2023-06-04T11:08:31.635125Z","shell.execute_reply.started":"2023-06-04T11:08:31.627321Z","shell.execute_reply":"2023-06-04T11:08:31.634171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"id":"o3W3R1diE7Nc","execution":{"iopub.status.busy":"2023-06-04T11:08:31.768502Z","iopub.execute_input":"2023-06-04T11:08:31.769173Z","iopub.status.idle":"2023-06-04T11:08:32.022461Z","shell.execute_reply.started":"2023-06-04T11:08:31.769139Z","shell.execute_reply":"2023-06-04T11:08:32.021522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if VALID:\n    pred_mask = create_pred_mask(model, val_data_loader, valid_mask_gt)\n    plot_pred_and_gt_mask(pred_mask, valid_mask_gt, 0.55)\n    torch.save(pred_mask, '/content/drive/MyDrive/Kaggle/vesuvius/efficient_net_pred_mask')","metadata":{"id":"2vLt9Jg4E7Nd","execution":{"iopub.status.busy":"2023-06-04T11:08:32.024564Z","iopub.execute_input":"2023-06-04T11:08:32.025113Z","iopub.status.idle":"2023-06-04T11:08:32.031318Z","shell.execute_reply.started":"2023-06-04T11:08:32.025078Z","shell.execute_reply":"2023-06-04T11:08:32.030293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RegNet_X_3_2_GF","metadata":{"id":"niCcWwa1Vg5R"}},{"cell_type":"code","source":"if os.path.exists(CFG.trained_models_path_reg_net):\n    model_unet_regnet = torch.load(CFG.trained_models_path_reg_net, map_location=device)\n    print(\"ok\")\nelse:\n    model_unet_regnet = smp.Unet(\n                encoder_name = CFG.backbone, \n                encoder_weights = 'imagenet',\n                in_channels = CFG.in_chans,\n                classes = CFG.target_size,\n                activation = 'sigmoid',\n    )","metadata":{"id":"_3i6BE61X5NL","execution":{"iopub.status.busy":"2023-06-04T11:08:32.200133Z","iopub.execute_input":"2023-06-04T11:08:32.200742Z","iopub.status.idle":"2023-06-04T11:08:33.730452Z","shell.execute_reply.started":"2023-06-04T11:08:32.200706Z","shell.execute_reply":"2023-06-04T11:08:33.729386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer","metadata":{"id":"ZjfAn2u2ZNOQ"}},{"cell_type":"code","source":"# optimizer = Adam(model_unet_regnet.parameters(), lr=1e-1, weight_decay=1e-4)\noptimizer = SGD(model_unet_regnet.parameters(), momentum=0.9, lr=0.5, weight_decay=2e-3)","metadata":{"id":"8ICEKAx8ZNOR","execution":{"iopub.status.busy":"2023-06-04T11:08:33.732515Z","iopub.execute_input":"2023-06-04T11:08:33.733050Z","iopub.status.idle":"2023-06-04T11:08:33.741820Z","shell.execute_reply.started":"2023-06-04T11:08:33.733016Z","shell.execute_reply":"2023-06-04T11:08:33.740591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Scheduler","metadata":{"id":"uquFmKuGZNOS"}},{"cell_type":"code","source":"scheduler = OneCycleLR(optimizer, max_lr=5e-2, epochs=CFG.epochs, steps_per_epoch=len(train_dataset)//CFG.batch_size)\n# scheduler = CosineAnnealingLR(optimizer, T_max=len(train_dataset)//CFG.batch_size)","metadata":{"id":"pBI2rTCcebCT","execution":{"iopub.status.busy":"2023-06-04T11:08:33.744968Z","iopub.execute_input":"2023-06-04T11:08:33.745951Z","iopub.status.idle":"2023-06-04T11:08:33.754487Z","shell.execute_reply.started":"2023-06-04T11:08:33.745919Z","shell.execute_reply":"2023-06-04T11:08:33.753499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{"id":"JWpGDMoZednc"}},{"cell_type":"code","source":"TRAIN_REGNET = False\nREG_TRACK = False\nREG_VALID = False\nLOAD_MASK = False\nPLOT_CURVES = False","metadata":{"id":"pw4w0cimTWj9","execution":{"iopub.status.busy":"2023-06-04T11:09:29.813424Z","iopub.execute_input":"2023-06-04T11:09:29.813798Z","iopub.status.idle":"2023-06-04T11:09:29.818305Z","shell.execute_reply.started":"2023-06-04T11:09:29.813766Z","shell.execute_reply":"2023-06-04T11:09:29.817332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"regnet_train_epoch = smp_utils.train.TrainEpoch(\n    model_unet_regnet, \n    loss=loss, \n    metrics=metrics, \n    optimizer=optimizer,\n    device=device,\n    verbose=True,\n)\n\nregnet_val_epoch = smp_utils.train.ValidEpoch (\n    model_unet_regnet, \n    loss=loss, \n    metrics=metrics, \n    device=device,\n    verbose=True,\n)\n\nregnet_model_epochs = {\n    'train': regnet_train_epoch,\n    'val': regnet_val_epoch\n}","metadata":{"id":"oAxkc8GsY3EB","execution":{"iopub.status.busy":"2023-06-04T11:08:33.766720Z","iopub.execute_input":"2023-06-04T11:08:33.767057Z","iopub.status.idle":"2023-06-04T11:08:33.788734Z","shell.execute_reply.started":"2023-06-04T11:08:33.767026Z","shell.execute_reply":"2023-06-04T11:08:33.787937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.path.exists('/content/drive/MyDrive/Kaggle/vesuvius/regnet_logs'):\n    all_logs = torch.load('/content/drive/MyDrive/Kaggle/vesuvius/regnet_logs')\n    print(\"ok\")\nelse:\n    all_logs = {\n    'dice_loss': [],\n    'iou_score': [],\n    'fscore': [],\n    'val_dice_loss': [],\n    'val_iou_score': [],\n    'val_fscore': [],\n}\n\nlogs = {\n    'dice_loss': [],\n    'iou_score': [],\n    'fscore': [],\n    'val_dice_loss': [],\n    'val_iou_score': [],\n    'val_fscore': [],\n}","metadata":{"id":"GZsOKbY6S0k7","execution":{"iopub.status.busy":"2023-06-04T11:08:33.836438Z","iopub.execute_input":"2023-06-04T11:08:33.836789Z","iopub.status.idle":"2023-06-04T11:08:33.843277Z","shell.execute_reply.started":"2023-06-04T11:08:33.836761Z","shell.execute_reply":"2023-06-04T11:08:33.842044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_REGNET:\n    logs = train_model(model_unet_regnet, regnet_model_epochs, dataloaders, scheduler, CFG.epochs, REG_TRACK, CFG.trained_models_path_reg_net)","metadata":{"id":"gkAlQby3YTrb","execution":{"iopub.status.busy":"2023-06-04T11:08:33.962469Z","iopub.execute_input":"2023-06-04T11:08:33.963166Z","iopub.status.idle":"2023-06-04T11:08:33.968235Z","shell.execute_reply.started":"2023-06-04T11:08:33.963129Z","shell.execute_reply":"2023-06-04T11:08:33.966943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if PLOT_CURVES:\n    if REG_TRACK:\n        all_logs['val_dice_loss'].extend(logs['val_dice_loss'])\n        all_logs['val_iou_score'].extend(logs['val_iou_score'])\n        all_logs['val_fscore'].extend(logs['val_fscore'])\n\n    all_logs['dice_loss'].extend(logs['dice_loss'])\n    all_logs['iou_score'].extend(logs['iou_score'])\n    all_logs['fscore'].extend(logs['fscore'])\n    df = pd.DataFrame(all_logs)\n    torch.save(all_logs, '/kaggle/working/regnet_logs')\n    plot_learning_curves(df, REG_TRACK)\n    plot_learning_curves_with_trend_line(df, REG_TRACK)","metadata":{"id":"wBZVY_YpBC6G","execution":{"iopub.status.busy":"2023-06-04T11:09:33.351606Z","iopub.execute_input":"2023-06-04T11:09:33.351997Z","iopub.status.idle":"2023-06-04T11:09:33.360876Z","shell.execute_reply.started":"2023-06-04T11:09:33.351965Z","shell.execute_reply":"2023-06-04T11:09:33.359799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"id":"1hAH_SdMaIEK","execution":{"iopub.status.busy":"2023-06-04T11:09:35.473920Z","iopub.execute_input":"2023-06-04T11:09:35.474282Z","iopub.status.idle":"2023-06-04T11:09:35.731753Z","shell.execute_reply.started":"2023-06-04T11:09:35.474253Z","shell.execute_reply":"2023-06-04T11:09:35.730542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if REG_VALID:\n    if LOAD_MASK:\n        pred_mask = torch.load('/content/drive/MyDrive/Kaggle/vesuvius/regnet_pred_mask')\n        print(\"ok\")\n    else:\n        pred_mask = create_pred_mask(model_unet_regnet, val_data_loader, valid_mask_gt)\n        torch.save(pred_mask, '/content/drive/MyDrive/Kaggle/vesuvius/regnet_pred_mask')\n        \n    plot_pred_and_gt_mask(pred_mask, valid_mask_gt, 0.65)    ","metadata":{"id":"kDovxAAtaIEK","execution":{"iopub.status.busy":"2023-06-04T11:09:45.955024Z","iopub.execute_input":"2023-06-04T11:09:45.955480Z","iopub.status.idle":"2023-06-04T11:09:45.964990Z","shell.execute_reply.started":"2023-06-04T11:09:45.955442Z","shell.execute_reply":"2023-06-04T11:09:45.964042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EfficientNetB6","metadata":{"id":"Fmcw0AAy6p-6"}},{"cell_type":"code","source":"if os.path.exists(CFG.trained_models_path_effiecient_netb6):\n    model_unet_eff_net_b6 = torch.load(CFG.trained_models_path_effiecient_netb6, map_location=device)\n    print(\"ok\")\nelse:\n    model_unet_eff_net_b6 = smp.Unet(\n                encoder_name = 'efficientnet-b6', \n                encoder_weights = 'imagenet',\n                in_channels = CFG.in_chans,\n                classes = CFG.target_size,\n                activation = 'sigmoid',\n    )","metadata":{"outputId":"356f0904-a180-46b3-95ea-88d26e2e06de","id":"txkexoxw6p-7","execution":{"iopub.status.busy":"2023-06-04T11:09:50.020484Z","iopub.execute_input":"2023-06-04T11:09:50.021107Z","iopub.status.idle":"2023-06-04T11:09:52.347599Z","shell.execute_reply.started":"2023-06-04T11:09:50.021074Z","shell.execute_reply":"2023-06-04T11:09:52.346563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimizer","metadata":{"id":"usJ9nX4A6p-9"}},{"cell_type":"code","source":"# optimizer = Adam(model_unet_eff_net_b6.parameters(), lr=1e-3, weight_decay=1e-3)\noptimizer = SGD(model_unet_eff_net_b6.parameters(), momentum=0.9, lr=1e-3, weight_decay=2e-3)","metadata":{"id":"1CJWz7FC6p--","execution":{"iopub.status.busy":"2023-06-04T11:09:53.448311Z","iopub.execute_input":"2023-06-04T11:09:53.448731Z","iopub.status.idle":"2023-06-04T11:09:53.457812Z","shell.execute_reply.started":"2023-06-04T11:09:53.448690Z","shell.execute_reply":"2023-06-04T11:09:53.456901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Scheduler","metadata":{"id":"uDdnSYqs6p--"}},{"cell_type":"code","source":"scheduler = OneCycleLR(optimizer, max_lr=5e-2, epochs=CFG.epochs, steps_per_epoch=len(train_dataset)//CFG.batch_size)\n# scheduler = CosineAnnealingLR(optimizer, T_max=len(train_dataset)//CFG.batch_size)","metadata":{"id":"OdwIp-x46p-_","execution":{"iopub.status.busy":"2023-06-04T11:09:54.928060Z","iopub.execute_input":"2023-06-04T11:09:54.928417Z","iopub.status.idle":"2023-06-04T11:09:54.933177Z","shell.execute_reply.started":"2023-06-04T11:09:54.928388Z","shell.execute_reply":"2023-06-04T11:09:54.932272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{"id":"YlEGQtTL6p-_"}},{"cell_type":"code","source":"TRAIN_EFFNETB6 = False\nEFFNETB6_TRACK = True\nEFFNETB6_VALID = True\nEFFNETB6_LOAD_MASK = True\nPLOT_CURVES = True","metadata":{"id":"tyKxAd2a6p_A","execution":{"iopub.status.busy":"2023-06-04T11:09:56.514063Z","iopub.execute_input":"2023-06-04T11:09:56.514426Z","iopub.status.idle":"2023-06-04T11:09:56.519381Z","shell.execute_reply.started":"2023-06-04T11:09:56.514396Z","shell.execute_reply":"2023-06-04T11:09:56.518169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"effnet_b6_train_epoch = smp_utils.train.TrainEpoch(\n    model_unet_eff_net_b6, \n    loss=loss, \n    metrics=metrics, \n    optimizer=optimizer,\n    device=device,\n    verbose=True,\n)\n\neffnet_b6_val_epoch = smp_utils.train.ValidEpoch (\n    model_unet_eff_net_b6, \n    loss=loss, \n    metrics=metrics, \n    device=device,\n    verbose=True,\n)\n\neffnet_b6_model_epochs = {\n    'train': effnet_b6_train_epoch,\n    'val': effnet_b6_val_epoch\n}","metadata":{"id":"m1Fj5q2M6p_A","execution":{"iopub.status.busy":"2023-06-04T11:09:57.333996Z","iopub.execute_input":"2023-06-04T11:09:57.334571Z","iopub.status.idle":"2023-06-04T11:09:57.361897Z","shell.execute_reply.started":"2023-06-04T11:09:57.334534Z","shell.execute_reply":"2023-06-04T11:09:57.360984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.path.exists('/kaggle/input/efficient-net-b6-train-logs/effb6_logs'):\n    all_logs = torch.load('/kaggle/input/efficient-net-b6-train-logs/effb6_logs')\n    print(\"ok\")\nelse:\n    all_logs = {\n    'dice_loss': [],\n    'iou_score': [],\n    'fscore': [],\n    'val_dice_loss': [],\n    'val_iou_score': [],\n    'val_fscore': [],\n}\n\nlogs = {\n    'dice_loss': [],\n    'iou_score': [],\n    'fscore': [],\n    'val_dice_loss': [],\n    'val_iou_score': [],\n    'val_fscore': [],\n}","metadata":{"outputId":"77f4d362-587e-4820-89fe-0070847396a0","id":"CIeOefCv6p_B","execution":{"iopub.status.busy":"2023-06-04T11:14:22.245682Z","iopub.execute_input":"2023-06-04T11:14:22.246421Z","iopub.status.idle":"2023-06-04T11:14:22.258941Z","shell.execute_reply.started":"2023-06-04T11:14:22.246386Z","shell.execute_reply":"2023-06-04T11:14:22.257993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataloaders['train'])","metadata":{"id":"Lsx_3x0ZB2Qb","outputId":"cd89d461-55fb-450b-9186-3c2adaf2ed91","execution":{"iopub.status.busy":"2023-06-04T11:14:25.464390Z","iopub.execute_input":"2023-06-04T11:14:25.464771Z","iopub.status.idle":"2023-06-04T11:14:25.470662Z","shell.execute_reply.started":"2023-06-04T11:14:25.464739Z","shell.execute_reply":"2023-06-04T11:14:25.469580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_EFFNETB6:\n    logs = train_model(model_unet_eff_net_b6, effnet_b6_model_epochs, dataloaders, scheduler, CFG.epochs, EFFNETB6_TRACK, CFG.trained_models_path_effiecient_netb6)","metadata":{"outputId":"24696926-380d-4b8b-b620-7532dd5057ad","id":"JJRGGYZv6p_C","execution":{"iopub.status.busy":"2023-06-04T11:14:28.088058Z","iopub.execute_input":"2023-06-04T11:14:28.088415Z","iopub.status.idle":"2023-06-04T11:14:28.092998Z","shell.execute_reply.started":"2023-06-04T11:14:28.088387Z","shell.execute_reply":"2023-06-04T11:14:28.091921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if PLOT_CURVES:\n    if EFFNETB6_TRACK:\n        all_logs['val_dice_loss'].extend(logs['val_dice_loss'])\n        all_logs['val_iou_score'].extend(logs['val_iou_score'])\n        all_logs['val_fscore'].extend(logs['val_fscore'])\n\n    all_logs['dice_loss'].extend(logs['dice_loss'])\n    all_logs['iou_score'].extend(logs['iou_score'])\n    all_logs['fscore'].extend(logs['fscore'])\n    df = pd.DataFrame(all_logs)\n    torch.save(all_logs, '/kaggle/working/effb6_logs')\n    \n    plot_learning_curves(df, EFFNETB6_TRACK)\n    plot_learning_curves_with_trend_line(df, EFFNETB6_TRACK)","metadata":{"id":"ywAj4DVX6p_D","execution":{"iopub.status.busy":"2023-06-04T11:15:03.631329Z","iopub.execute_input":"2023-06-04T11:15:03.631718Z","iopub.status.idle":"2023-06-04T11:15:05.065398Z","shell.execute_reply.started":"2023-06-04T11:15:03.631686Z","shell.execute_reply":"2023-06-04T11:15:05.064541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"outputId":"54a84594-616a-4454-ff6f-cc2d34a1598d","id":"VWSQF4v-6p_D","execution":{"iopub.status.busy":"2023-06-04T11:15:15.333475Z","iopub.execute_input":"2023-06-04T11:15:15.333838Z","iopub.status.idle":"2023-06-04T11:15:15.703581Z","shell.execute_reply.started":"2023-06-04T11:15:15.333807Z","shell.execute_reply":"2023-06-04T11:15:15.702569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if EFFNETB6_VALID:\n    if EFFNETB6_LOAD_MASK:\n        pred_mask = torch.load('/kaggle/input/efficient-net-b6-train-logs/effnetb6_pred_mask')\n        print(\"ok\")\n    else:\n        pred_mask = create_pred_mask(model_unet_eff_net_b6, val_data_loader, valid_mask_gt)\n        torch.save(pred_mask, '/kaggle/input/efficient-net-b6-train-logs/effnetb6_pred_mask')\n        \n    best_th = find_best_threshold(model_unet_eff_net_b6, pred_mask, valid_mask_gt)\n    plot_pred_and_gt_mask(pred_mask, valid_mask_gt, best_th)","metadata":{"outputId":"7de64b00-5f82-41b2-ee11-455c399d12c1","id":"PZ7mhmGH6p_E","execution":{"iopub.status.busy":"2023-06-04T11:16:17.398392Z","iopub.execute_input":"2023-06-04T11:16:17.398799Z","iopub.status.idle":"2023-06-04T11:16:44.529654Z","shell.execute_reply.started":"2023-06-04T11:16:17.398767Z","shell.execute_reply":"2023-06-04T11:16:44.528689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Resources\n1. [Segmentation Models Pytorch](https://smp.readthedocs.io/en/latest/models.html#id9)\n2. [UNet](https://www.youtube.com/watch?v=oLvmLJkmXuc&list=PLhhyoLH6IjfwqKKZhVLp7diKFxTmj4Q6s&index=12)\n2. [ResNet vs EfficientNet](https://medium.com/@enrico.randellini/image-classification-resnet-vs-efficientnet-vs-efficientnet-v2-vs-compact-convolutional-c205838bbf49)\n3. [Papers with Code](https://paperswithcode.com/)\n4. [Visual Guide to learning rate schedulers](https://towardsdatascience.com/a-visual-guide-to-learning-rate-schedulers-in-pytorch-24bbb262c863)\n1. [2.5d segmentation baseline [training]](https://www.kaggle.com/code/tanakar/2-5d-segmentaion-baseline-training)\n3. [Vesuvius Challenge - 3D ResNet Training](https://www.kaggle.com/code/samfc10/vesuvius-challenge-3d-resnet-training)\n3. [2.5d segmentation baseline [inference]](https://www.kaggle.com/code/yoyobar/3d-resnet-baseline-inference)\n","metadata":{"id":"L_DKWX026p_F"}},{"cell_type":"code","source":"","metadata":{"id":"ZJUMFAuA2Cr9"},"execution_count":null,"outputs":[]}]}