{"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":"## summary\n\n* 2.5d segmentation\n    *  segmentation_models_pytorch \n    *  Unet\n* use only 6 slices in the middle\n* slide inference","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nimport pickle\nfrom torch.utils.data import DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport warnings\nimport sys\nimport pandas as pd\nimport os\nimport gc\nimport sys\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nimport cv2\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport argparse\nimport importlib\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam, SGD, AdamW\n\nimport datetime","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:41:39.766964Z","iopub.execute_input":"2023-05-17T13:41:39.767949Z","iopub.status.idle":"2023-05-17T13:41:42.678506Z","shell.execute_reply.started":"2023-05-17T13:41:39.767912Z","shell.execute_reply":"2023-05-17T13:41:42.677337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sys.path.append('/kaggle/input/pretrainedmodels/pretrainedmodels-0.7.4')\n# sys.path.append('/kaggle/input/efficientnet-pytorch/EfficientNet-PyTorch-master')\n# sys.path.append('/kaggle/input/timm-pytorch-image-models/pytorch-image-models-master')\n# sys.path.append('/kaggle/input/segmentation-models-pytorch/segmentation_models.pytorch-master')","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:41:42.680705Z","iopub.execute_input":"2023-05-17T13:41:42.681684Z","iopub.status.idle":"2023-05-17T13:41:42.688239Z","shell.execute_reply.started":"2023-05-17T13:41:42.681643Z","shell.execute_reply":"2023-05-17T13:41:42.686919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation_models_pytorch warmup_scheduler segment_anything","metadata":{"_kg_hide-input":true,"scrolled":true,"execution":{"iopub.status.busy":"2023-05-17T13:41:42.691548Z","iopub.execute_input":"2023-05-17T13:41:42.693150Z","iopub.status.idle":"2023-05-17T13:42:04.744191Z","shell.execute_reply.started":"2023-05-17T13:41:42.693110Z","shell.execute_reply":"2023-05-17T13:42:04.742729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom torch.utils.data import DataLoader, Dataset\nimport cv2\nimport torch\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:04.747980Z","iopub.execute_input":"2023-05-17T13:42:04.749067Z","iopub.status.idle":"2023-05-17T13:42:08.244150Z","shell.execute_reply.started":"2023-05-17T13:42:04.749017Z","shell.execute_reply":"2023-05-17T13:42:08.242926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## config","metadata":{}},{"cell_type":"code","source":"import os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nclass CFG:\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n\n    # comp_dir_path = './'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'vesuvius-challenge-ink-detection'\n    # comp_dataset_path = f'{comp_dir_path}datasets/{comp_folder_name}/'\n    comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'\n    \n    exp_name = 'vesuvius_2d_slide_exp002'\n\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = 'Unet'\n    backbone = 'efficientnet-b0'\n    # backbone = 'se_resnext50_32x4d'\n\n    in_chans = 3 # 65\n    # ============== training cfg =============\n    size = 1024\n    tile_size = 224\n    stride = tile_size // 2\n\n    train_batch_size = 4 # 32\n    valid_batch_size = train_batch_size * 2\n    use_amp = True\n\n    scheduler = 'GradualWarmupSchedulerV2'\n    # scheduler = 'CosineAnnealingLR'\n    epochs = 15 # 30\n\n    # adamW warmupあり\n    warmup_factor = 10\n    # lr = 1e-4 / warmup_factor\n    lr = 1e-4 / warmup_factor\n\n    # ============== fold =============\n    valid_id = 1\n\n    # objective_cv = 'binary'  # 'binary', 'multiclass', 'regression'\n    metric_direction = 'maximize'  # maximize, 'minimize'\n    # metrics = 'dice_coef'\n\n    # ============== fixed =============\n    pretrained = True\n    inf_weight = 'best'  # 'best'\n\n    min_lr = 1e-6\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n\n    print_freq = 50\n    num_workers = 4\n\n    seed = 42\n\n    # ============== set dataset path =============\n    print('set dataset path')\n\n    outputs_path = f'/kaggle/working/outputs/{comp_name}/{exp_name}/'\n\n    submission_dir = outputs_path + 'submissions/'\n    submission_path = submission_dir + f'submission_{exp_name}.csv'\n\n    model_dir = outputs_path + \\\n        f'{comp_name}-models/'\n\n    figures_dir = outputs_path + 'figures/'\n\n    log_dir = outputs_path + 'logs/'\n    log_path = log_dir + f'{exp_name}.txt'\n\n    # ============== augmentation =============\n    train_aug_list = [\n        # A.RandomResizedCrop(\n        #     size, size, scale=(0.85, 1.0)),\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","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.246130Z","iopub.execute_input":"2023-05-17T13:42:08.246549Z","iopub.status.idle":"2023-05-17T13:42:08.265924Z","shell.execute_reply.started":"2023-05-17T13:42:08.246503Z","shell.execute_reply":"2023-05-17T13:42:08.264836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## helper","metadata":{}},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.267504Z","iopub.execute_input":"2023-05-17T13:42:08.269410Z","iopub.status.idle":"2023-05-17T13:42:08.286496Z","shell.execute_reply.started":"2023-05-17T13:42:08.269368Z","shell.execute_reply":"2023-05-17T13:42:08.285326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_logger(log_file):\n    from logging import getLogger, INFO, FileHandler, Formatter, StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\ndef 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","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.287808Z","iopub.execute_input":"2023-05-17T13:42:08.291387Z","iopub.status.idle":"2023-05-17T13:42:08.300871Z","shell.execute_reply.started":"2023-05-17T13:42:08.291349Z","shell.execute_reply":"2023-05-17T13:42:08.299760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_dirs(cfg):\n    for dir in [cfg.model_dir, cfg.figures_dir, cfg.submission_dir, cfg.log_dir]:\n        os.makedirs(dir, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.302616Z","iopub.execute_input":"2023-05-17T13:42:08.302983Z","iopub.status.idle":"2023-05-17T13:42:08.315485Z","shell.execute_reply.started":"2023-05-17T13:42:08.302947Z","shell.execute_reply":"2023-05-17T13:42:08.314441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cfg_init(cfg, mode='train'):\n    set_seed(cfg.seed)\n    # set_env_name()\n    # set_dataset_path(cfg)\n\n    if mode == 'train':\n        make_dirs(cfg)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.316987Z","iopub.execute_input":"2023-05-17T13:42:08.317502Z","iopub.status.idle":"2023-05-17T13:42:08.326746Z","shell.execute_reply.started":"2023-05-17T13:42:08.317457Z","shell.execute_reply":"2023-05-17T13:42:08.325601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg_init(CFG)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nLogger = init_logger(log_file=CFG.log_path)\n\nLogger.info('\\n\\n-------- exp_info -----------------')\n# Logger.info(datetime.datetime.now().strftime('%Y年%m月%d日 %H:%M:%S'))","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.332464Z","iopub.execute_input":"2023-05-17T13:42:08.332758Z","iopub.status.idle":"2023-05-17T13:42:08.406973Z","shell.execute_reply.started":"2023-05-17T13:42:08.332719Z","shell.execute_reply":"2023-05-17T13:42:08.405888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## image, mask","metadata":{}},{"cell_type":"code","source":"def read_image_mask(fragment_id):\n\n    images = []\n\n    # idxs = range(65)\n    mid = 65 // 2\n    start = mid - CFG.in_chans // 2\n    end = mid + CFG.in_chans // 2 + 1\n    idxs = range(start, end)\n\n    for i in tqdm(idxs):\n        \n        image = cv2.imread(CFG.comp_dataset_path + f\"train/{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\n        images.append(image)\n    images = np.stack(images, axis=2)\n\n    mask = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/inklabels.png\", 0)\n    mask = np.pad(mask, [(0, pad0), (0, pad1)], constant_values=0)\n\n    mask = mask.astype('float32')\n    mask /= 255.0\n    \n    return images, mask","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.408589Z","iopub.execute_input":"2023-05-17T13:42:08.410997Z","iopub.status.idle":"2023-05-17T13:42:08.422039Z","shell.execute_reply.started":"2023-05-17T13:42:08.410956Z","shell.execute_reply":"2023-05-17T13:42:08.420979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, mask = read_image_mask(1)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:08.424058Z","iopub.execute_input":"2023-05-17T13:42:08.425073Z","iopub.status.idle":"2023-05-17T13:42:13.993939Z","shell.execute_reply.started":"2023-05-17T13:42:08.425034Z","shell.execute_reply":"2023-05-17T13:42:13.992790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images.shape, mask.shape","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:13.995609Z","iopub.execute_input":"2023-05-17T13:42:13.995999Z","iopub.status.idle":"2023-05-17T13:42:14.005971Z","shell.execute_reply.started":"2023-05-17T13:42:13.995958Z","shell.execute_reply":"2023-05-17T13:42:14.004835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_valid_dataset():\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                # xyxys.append((x1, y1, x2, y2))\n        \n                if fragment_id == CFG.valid_id:\n                    valid_images.append(image[y1:y2, x1:x2])\n                    valid_masks.append(mask[y1:y2, x1:x2, None])\n\n                    valid_xyxys.append([x1, y1, x2, y2])\n                else:\n                    train_images.append(image[y1:y2, x1:x2])\n                    train_masks.append(mask[y1:y2, x1:x2, None])\n\n    return train_images, train_masks, valid_images, valid_masks, valid_xyxys","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:14.008034Z","iopub.execute_input":"2023-05-17T13:42:14.008786Z","iopub.status.idle":"2023-05-17T13:42:14.021347Z","shell.execute_reply.started":"2023-05-17T13:42:14.008747Z","shell.execute_reply":"2023-05-17T13:42:14.020210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images, train_masks, valid_images, valid_masks, valid_xyxys = get_train_valid_dataset()","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:14.022814Z","iopub.execute_input":"2023-05-17T13:42:14.023443Z","iopub.status.idle":"2023-05-17T13:42:34.599751Z","shell.execute_reply.started":"2023-05-17T13:42:14.023331Z","shell.execute_reply":"2023-05-17T13:42:34.598687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_images)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.601354Z","iopub.execute_input":"2023-05-17T13:42:34.601768Z","iopub.status.idle":"2023-05-17T13:42:34.608453Z","shell.execute_reply.started":"2023-05-17T13:42:34.601727Z","shell.execute_reply":"2023-05-17T13:42:34.607237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_xyxys = np.stack(valid_xyxys)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.610127Z","iopub.execute_input":"2023-05-17T13:42:34.610836Z","iopub.status.idle":"2023-05-17T13:42:34.634117Z","shell.execute_reply.started":"2023-05-17T13:42:34.610793Z","shell.execute_reply":"2023-05-17T13:42:34.633051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_xyxys.shape","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.636060Z","iopub.execute_input":"2023-05-17T13:42:34.636457Z","iopub.status.idle":"2023-05-17T13:42:34.645365Z","shell.execute_reply.started":"2023-05-17T13:42:34.636416Z","shell.execute_reply":"2023-05-17T13:42:34.644149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## dataset","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom torch.utils.data import DataLoader, Dataset\nimport cv2\nimport torch\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.647120Z","iopub.execute_input":"2023-05-17T13:42:34.647618Z","iopub.status.idle":"2023-05-17T13:42:34.654687Z","shell.execute_reply.started":"2023-05-17T13:42:34.647580Z","shell.execute_reply":"2023-05-17T13:42:34.653448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(cfg.train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(cfg.valid_aug_list)\n\n    # print(aug)\n    return aug\n\nclass CustomDataset(Dataset):\n    def __init__(self, images, cfg, labels=None, transform=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        # return len(self.df)\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        label = self.labels[idx]\n\n        if self.transform:\n            data = self.transform(image=image, mask=label)\n            image = data['image']\n            label = data['mask']\n\n        return image, torch.tensor(label, dtype=torch.long)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.656402Z","iopub.execute_input":"2023-05-17T13:42:34.657012Z","iopub.status.idle":"2023-05-17T13:42:34.669256Z","shell.execute_reply.started":"2023-05-17T13:42:34.656970Z","shell.execute_reply":"2023-05-17T13:42:34.668168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CustomDataset(\n    train_images, CFG, labels=train_masks, transform=get_transforms(data='train', cfg=CFG))\nvalid_dataset = CustomDataset(\n    valid_images, CFG, labels=valid_masks, transform=get_transforms(data='valid', cfg=CFG))\n\ntrain_loader = DataLoader(train_dataset,\n                          batch_size=CFG.train_batch_size,\n                          shuffle=True,\n                          num_workers=CFG.num_workers, pin_memory=True, drop_last=True,\n                          )\nvalid_loader = DataLoader(valid_dataset,\n                          batch_size=CFG.valid_batch_size,\n                          shuffle=False,\n                          num_workers=CFG.num_workers, pin_memory=True, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.670689Z","iopub.execute_input":"2023-05-17T13:42:34.671132Z","iopub.status.idle":"2023-05-17T13:42:34.685440Z","shell.execute_reply.started":"2023-05-17T13:42:34.671092Z","shell.execute_reply":"2023-05-17T13:42:34.684033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset[0][0].shape","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.687040Z","iopub.execute_input":"2023-05-17T13:42:34.687915Z","iopub.status.idle":"2023-05-17T13:42:34.890226Z","shell.execute_reply.started":"2023-05-17T13:42:34.687877Z","shell.execute_reply":"2023-05-17T13:42:34.889018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, label = train_dataset[100]","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.892243Z","iopub.execute_input":"2023-05-17T13:42:34.892918Z","iopub.status.idle":"2023-05-17T13:42:34.951502Z","shell.execute_reply.started":"2023-05-17T13:42:34.892877Z","shell.execute_reply":"2023-05-17T13:42:34.949177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.953214Z","iopub.execute_input":"2023-05-17T13:42:34.953593Z","iopub.status.idle":"2023-05-17T13:42:34.962194Z","shell.execute_reply.started":"2023-05-17T13:42:34.953553Z","shell.execute_reply":"2023-05-17T13:42:34.961102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_dataset = CustomDataset(\n#     train_images, CFG, labels=train_masks)\n\n# transform = CFG.train_aug_list\n# transform = A.Compose(\n#     [t for t in transform if not isinstance(t, (A.Normalize, ToTensorV2))])\n\n\n# plot_count = 0\n# for i in range(1000):\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:\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#     plt.savefig(CFG.figures_dir + f'aug_fold_{CFG.valid_id}_{plot_count}.png')\n\n#     plot_count += 1\n#     if plot_count == 5:\n#         break","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.963986Z","iopub.execute_input":"2023-05-17T13:42:34.964702Z","iopub.status.idle":"2023-05-17T13:42:34.974141Z","shell.execute_reply.started":"2023-05-17T13:42:34.964661Z","shell.execute_reply":"2023-05-17T13:42:34.973121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del plot_dataset\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.977300Z","iopub.execute_input":"2023-05-17T13:42:34.977639Z","iopub.status.idle":"2023-05-17T13:42:34.983896Z","shell.execute_reply.started":"2023-05-17T13:42:34.977610Z","shell.execute_reply":"2023-05-17T13:42:34.982808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"import torch\n# torch.multiprocessing.set_sharing_strategy('file_system')\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nimport torch.distributed as dist\nfrom torchvision import transforms\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import LearningRateMonitor, ModelCheckpoint\nfrom transformers.models.maskformer.modeling_maskformer import dice_loss, sigmoid_focal_loss\nfrom segment_anything import sam_model_registry","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:34.985762Z","iopub.execute_input":"2023-05-17T13:42:34.986185Z","iopub.status.idle":"2023-05-17T13:42:43.990870Z","shell.execute_reply.started":"2023-05-17T13:42:34.986134Z","shell.execute_reply":"2023-05-17T13:42:43.989431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SAMFinetuner(pl.LightningModule):\n\n    def __init__(\n            self,\n            model_type,\n            checkpoint_path,\n            freeze_image_encoder=False,\n            freeze_prompt_encoder=False,\n            freeze_mask_decoder=False,\n            batch_size=16,\n            learning_rate=1e-4,\n            weight_decay=1e-4,\n            train_dataset=None,\n            val_dataset=None,\n            metrics_interval=10,\n        ):\n        super(SAMFinetuner, self).__init__()\n\n        self.model_type = model_type\n        self.model = sam_model_registry[self.model_type](checkpoint=checkpoint_path)\n        self.model.to(device=self.device)\n        self.freeze_image_encoder = freeze_image_encoder\n        if freeze_image_encoder:\n            for param in self.model.image_encoder.parameters():\n                param.requires_grad = False\n        if freeze_prompt_encoder:\n            for param in self.model.prompt_encoder.parameters():\n                param.requires_grad = False\n        if freeze_mask_decoder:\n            for param in self.model.mask_decoder.parameters():\n                param.requires_grad = False\n        \n        self.batch_size = batch_size\n        self.learning_rate = learning_rate\n        self.weight_decay = weight_decay\n\n        self.train_dataset = train_dataset\n        self.val_dataset = val_dataset\n\n        self.train_metric = defaultdict(lambda: deque(maxlen=metrics_interval))\n\n        self.metrics_interval = metrics_interval\n\n    def forward(self, imgs, labels):\n        _, B, H, W = imgs.shape\n        features = self.model.image_encoder(imgs)\n        num_masks = B\n#         num_masks = sum([len(b) for b in bboxes])\n\n        loss_focal = loss_dice = loss_iou = 0.\n        predictions = []\n        tp, fp, fn, tn = [], [], [], []\n        for feature, label in zip(features, labels):\n            # Embed prompts\n            sparse_embeddings, dense_embeddings = self.model.prompt_encoder(\n                points=None,\n                boxes=None,\n                masks=None,\n            )\n            # Predict masks\n            low_res_masks, iou_predictions = self.model.mask_decoder(\n                image_embeddings=feature.unsqueeze(0),\n                image_pe=self.model.prompt_encoder.get_dense_pe(),\n                sparse_prompt_embeddings=sparse_embeddings,\n                dense_prompt_embeddings=dense_embeddings,\n                multimask_output=False,\n            )\n            # Upscale the masks to the original image resolution\n            masks = F.interpolate(\n                low_res_masks,\n                (H, W),\n                mode=\"bilinear\",\n                align_corners=False,\n            )\n            predictions.append(masks)\n            # Compute the iou between the predicted masks and the ground truth masks\n            batch_tp, batch_fp, batch_fn, batch_tn = smp.metrics.get_stats(\n                masks,\n                label.unsqueeze(1),\n                mode='binary',\n                threshold=0.5,\n            )\n            batch_iou = smp.metrics.iou_score(batch_tp, batch_fp, batch_fn, batch_tn)\n            # Compute the loss\n            masks = masks.squeeze(1).flatten(1)\n            label = label.flatten(1)\n            loss_focal += sigmoid_focal_loss(masks, label.float(), num_masks)\n            loss_dice += dice_loss(masks, label.float(), num_masks)\n            loss_iou += F.mse_loss(iou_predictions, batch_iou, reduction='sum') / num_masks\n            tp.append(batch_tp)\n            fp.append(batch_fp)\n            fn.append(batch_fn)\n            tn.append(batch_tn)\n        return {\n            'loss': 20. * loss_focal + loss_dice + loss_iou,  # SAM default loss\n            'loss_focal': loss_focal,\n            'loss_dice': loss_dice,\n            'loss_iou': loss_iou,\n            'predictions': predictions,\n            'tp': torch.cat(tp),\n            'fp': torch.cat(fp),\n            'fn': torch.cat(fn),\n            'tn': torch.cat(tn),\n        }\n    \n    def training_step(self, batch, batch_nb):\n        imgs, labels = batch\n        outputs = self(imgs, labels)\n\n        for metric in ['tp', 'fp', 'fn', 'tn']:\n            self.train_metric[metric].append(outputs[metric])\n\n        # aggregate step metics\n        step_metrics = [torch.cat(list(self.train_metric[metric])) for metric in ['tp', 'fp', 'fn', 'tn']]\n        per_mask_iou = smp.metrics.iou_score(*step_metrics, reduction=\"micro-imagewise\")\n        metrics = {\n            \"loss\": outputs[\"loss\"],\n            \"loss_focal\": outputs[\"loss_focal\"],\n            \"loss_dice\": outputs[\"loss_dice\"],\n            \"loss_iou\": outputs[\"loss_iou\"],\n            \"train_per_mask_iou\": per_mask_iou,\n        }\n        self.log_dict(metrics, prog_bar=True, rank_zero_only=True)\n        return metrics\n    \n    def validation_step(self, batch, batch_nb):\n        imgs, labels = batch\n        outputs = self(imgs, labels)\n        outputs.pop(\"predictions\")\n        return outputs\n    \n    def validation_epoch_end(self, outputs):\n        if NUM_GPUS > 1:\n            outputs = all_gather(outputs)\n            # the outputs are a list of lists, so flatten it\n            outputs = [item for sublist in outputs for item in sublist]\n        # aggregate step metics\n        step_metrics = [\n            torch.cat(list([x[metric].to(self.device) for x in outputs]))\n            for metric in ['tp', 'fp', 'fn', 'tn']]\n        # per mask IoU means that we first calculate IoU score for each mask\n        # and then compute mean over these scores\n        per_mask_iou = smp.metrics.iou_score(*step_metrics, reduction=\"micro-imagewise\")\n\n        metrics = {\"val_per_mask_iou\": per_mask_iou}\n        self.log_dict(metrics)\n        return metrics\n    \n    def configure_optimizers(self):\n        opt = torch.optim.AdamW(self.parameters(), lr=self.learning_rate, weight_decay=self.weight_decay)\n        def warmup_step_lr_builder(warmup_steps, milestones, gamma):\n            def warmup_step_lr(steps):\n                if steps < warmup_steps:\n                    lr_scale = (steps + 1.) / float(warmup_steps)\n                else:\n                    lr_scale = 1.\n                    for milestone in sorted(milestones):\n                        if steps >= milestone * self.trainer.estimated_stepping_batches:\n                            lr_scale *= gamma\n                return lr_scale\n            return warmup_step_lr\n        scheduler = torch.optim.lr_scheduler.LambdaLR(\n            opt,\n            warmup_step_lr_builder(250, [0.66667, 0.86666], 0.1)\n        )\n        return {\n            'optimizer': opt,\n            'lr_scheduler': {\n                'scheduler': scheduler,\n                'interval': \"step\",\n                'frequency': 1,\n            }\n        }\n    \n    def train_dataloader(self):\n        train_loader = DataLoader(\n            self.train_dataset,\n            batch_size=CFG.train_batch_size,\n            shuffle=True,\n            num_workers=CFG.num_workers, \n            pin_memory=True, \n            drop_last=True,\n        )\n        return train_loader\n    \n    def val_dataloader(self):\n        val_loader = DataLoader(\n            self.val_dataset,\n            batch_size=CFG.valid_batch_size,\n            shuffle=False,\n            num_workers=CFG.num_workers, \n            pin_memory=True, \n            drop_last=False\n        )\n        return val_loader","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:43.999002Z","iopub.execute_input":"2023-05-17T13:42:43.999727Z","iopub.status.idle":"2023-05-17T13:42:44.038902Z","shell.execute_reply.started":"2023-05-17T13:42:43.999689Z","shell.execute_reply":"2023-05-17T13:42:44.037741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ../input/segment-anything-models/sam_vit_b_01ec64.pth","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:44.040790Z","iopub.execute_input":"2023-05-17T13:42:44.041319Z","iopub.status.idle":"2023-05-17T13:42:45.007735Z","shell.execute_reply.started":"2023-05-17T13:42:44.041278Z","shell.execute_reply":"2023-05-17T13:42:45.006419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = SAMFinetuner(\n    \"vit_b\",\n    \"../input/segment-anything-models/sam_vit_b_01ec64.pth\",\n    freeze_image_encoder=True,\n    freeze_prompt_encoder=True,\n    freeze_mask_decoder=False,\n    train_dataset=train_dataset,\n    val_dataset=valid_dataset,\n    batch_size=4,\n    learning_rate=1e-4,\n    weight_decay=1e-2,\n    metrics_interval=50,\n)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:45.009593Z","iopub.execute_input":"2023-05-17T13:42:45.010258Z","iopub.status.idle":"2023-05-17T13:42:50.811872Z","shell.execute_reply.started":"2023-05-17T13:42:45.010208Z","shell.execute_reply":"2023-05-17T13:42:50.810765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"callbacks = [\n    LearningRateMonitor(logging_interval='step'),\n    ModelCheckpoint(\n        dirpath=\"./output\",\n        filename='{step}-{val_per_mask_iou:.2f}',\n        save_last=True,\n        save_top_k=1,\n        monitor=\"val_per_mask_iou\",\n        mode=\"max\",\n        save_weights_only=True,\n        every_n_train_steps=50,\n    ),\n]","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:50.813531Z","iopub.execute_input":"2023-05-17T13:42:50.813914Z","iopub.status.idle":"2023-05-17T13:42:50.823509Z","shell.execute_reply.started":"2023-05-17T13:42:50.813873Z","shell.execute_reply":"2023-05-17T13:42:50.822478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_GPUS = 1","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:50.826989Z","iopub.execute_input":"2023-05-17T13:42:50.827292Z","iopub.status.idle":"2023-05-17T13:42:50.837353Z","shell.execute_reply.started":"2023-05-17T13:42:50.827264Z","shell.execute_reply":"2023-05-17T13:42:50.836279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = pl.Trainer(\n#     strategy='ddp' if NUM_GPUS > 1 else None,\n#     strategy=\"dp\",\n    accelerator=\"cuda\",\n    devices=NUM_GPUS,\n    precision=16,\n    callbacks=callbacks,\n    max_epochs=-1,\n    max_steps=600,\n    val_check_interval=200,\n    check_val_every_n_epoch=None,\n    num_sanity_val_steps=0,\n)","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:50.838729Z","iopub.execute_input":"2023-05-17T13:42:50.839724Z","iopub.status.idle":"2023-05-17T13:42:50.903791Z","shell.execute_reply.started":"2023-05-17T13:42:50.839679Z","shell.execute_reply":"2023-05-17T13:42:50.902814Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# labels.dtype","metadata":{"execution":{"iopub.status.busy":"2023-05-17T13:42:50.905398Z","iopub.execute_input":"2023-05-17T13:42:50.905764Z","iopub.status.idle":"2023-05-17T13:42:50.910814Z","shell.execute_reply.started":"2023-05-17T13:42:50.905722Z","shell.execute_reply":"2023-05-17T13:42:50.909581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")\nfrom collections import deque\n\n\ntrainer.fit(model)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-05-17T13:42:59.394331Z","iopub.execute_input":"2023-05-17T13:42:59.394893Z"},"trusted":true},"execution_count":null,"outputs":[]}]}