{"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":"# Definitions\n* **Rle -** Run-Length Encoding is a form of lossless data compression in which runs of data (sequences in which the same data value occurs in many consecutive data elements) are stored as a single data value and count, rather than as the original run. This is most efficient on data that contains many such runs, for example, simple graphic images such as icons, line drawings, Conway's Game of Life, and animations. For files that do not have many runs, RLE could increase the file size. For example: the sequence `wwwwaaad` will be encoded to `w4a3d1`. ***Source:*** https://www.kaggle.com/code/susnato/understanding-run-length-encoding-and-decoding\n* **F2 Score -** The F2 score is a variant of the F-score (also known as the F1 score) which places more importance on recall than precision. It's a measure of a test's accuracy that considers both the precision and the recall of the test to compute the score.","metadata":{"id":"dgIAhpEqTOel"}},{"cell_type":"markdown","source":"# Imports","metadata":{"id":"qceXT4tPTiSg"}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch --quiet\n!pip -q install cairosvg==2.5.2 --quiet \n!pip -q install reportlab==3.5.65 --quiet\n!pip -q install cssutils==2.2.0 --quiet","metadata":{"id":"OyE_nN59Zc1V","outputId":"fd031966-006a-40ff-89bd-3780ca3d8880","execution":{"iopub.status.busy":"2023-05-27T05:30:09.859727Z","iopub.execute_input":"2023-05-27T05:30:09.860079Z","iopub.status.idle":"2023-05-27T05:31:46.641813Z","shell.execute_reply.started":"2023-05-27T05:30:09.860051Z","shell.execute_reply":"2023-05-27T05:31:46.640554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport albumentations as A\nimport segmentation_models_pytorch as smp\nimport warnings\nimport glob\nimport segmentation_models_pytorch.utils.metrics\nimport segmentation_models_pytorch.utils.train\nimport builtins\nimport pickle\n\nfrom tqdm import tqdm\n\nfrom torch.autograd import Variable\nfrom torch.utils.data import Dataset, DataLoader, SubsetRandomSampler\nfrom torchvision import transforms, utils\nfrom torch.optim import Adam\nfrom torch.optim.lr_scheduler import OneCycleLR\n\nfrom albumentations.pytorch import ToTensorV2\nfrom segmentation_models_pytorch.encoders import get_preprocessing_fn\n\nfrom PIL import Image\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"id":"6RcXzjG3TMuO","execution":{"iopub.status.busy":"2023-05-27T05:38:36.042247Z","iopub.execute_input":"2023-05-27T05:38:36.043196Z","iopub.status.idle":"2023-05-27T05:38:36.051509Z","shell.execute_reply.started":"2023-05-27T05:38:36.043159Z","shell.execute_reply":"2023-05-27T05:38:36.050187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{"id":"UPBoTA_MT0l1"}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"id":"kvSB81UWTyzz","execution":{"iopub.status.busy":"2023-05-27T05:31:49.043084Z","iopub.execute_input":"2023-05-27T05:31:49.043497Z","iopub.status.idle":"2023-05-27T05:31:49.072878Z","shell.execute_reply.started":"2023-05-27T05:31:49.043460Z","shell.execute_reply":"2023-05-27T05:31:49.071384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    # ============== set paths =============\n    root_dir = '/kaggle/input/airbus-ship-detection'\n    csv_file = 'train_ship_segmentations_v2.csv'\n    saved_model_path = '/kaggle/input/airbus-ship-detection-unet-weights/Unet_Weights.pth'\n    model_save_path = '/kaggle/working/Unet_Weights.pth'\n\n    # ============== model cfg =============\n    model_name = 'Unet'\n    backbone = 'efficientnet-b0'\n    \n    # ============== training cfg =============\n    size = 224\n    batch_size = 32 # 64\n    epochs = 5 # 30\n    lr = 1e-3\n    weight_decay = 1e-5\n    seed = 42\n\n    # ============== loss cfg =============\n    mode = 'binary'\n    gamma = 5.0\n\n    # ============== fixed cfg =============\n    num_workers = 2\n\n    # ============== augmentation =============\n    torch_aug = [\n        A.Resize(width=224, height=224),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2()\n    ]\n    \n    train_aug = [\n        A.Resize(width=224, height=224),\n        A.HorizontalFlip(p=0.5),\n        A.Rotate(limit=45, p=0.5),\n        A.RandomBrightnessContrast(p=0.5),\n        A.RandomResizedCrop(height=256, width=256, scale=(0.8, 1.0), ratio=(0.9, 1.1), p=0.5),\n        A.Affine(scale=1.0, rotate=0, translate_percent=0, shear=0, p=0.5),\n        A.GaussNoise(var_limit=(10.0, 50.0), p=0.5),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2()\n    ]\n    \n    test_aug = [\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.Rotate(limit=45, p=0.5),\n        A.RandomBrightnessContrast(p=0.5),\n        A.RandomResizedCrop(height=256, width=256, scale=(0.8, 1.0), ratio=(0.9, 1.1), p=0.5),\n        A.Affine(scale=1.0, rotate=0, translate_percent=0, shear=0, p=0.5),\n        A.GaussNoise(var_limit=(10.0, 50.0), p=0.5),\n    ]\n\n    exclude_list = {\n        '6384c3e78.jpg','13703f040.jpg', '14715c06d.jpg',  '33e0ff2d5.jpg',\n        '4d4e09f2a.jpg', '877691df8.jpg', '8b909bb20.jpg', 'a8d99130e.jpg', \n        'ad55c3143.jpg', 'c8260c541.jpg', 'd6c7f17c7.jpg', 'dc3e7c901.jpg',\n        'e44dffe88.jpg', 'ef87bad36.jpg', 'f083256d8.jpg'\n    }","metadata":{"id":"m8Rfj9QNT16v","execution":{"iopub.status.busy":"2023-05-27T05:54:41.246343Z","iopub.execute_input":"2023-05-27T05:54:41.246903Z","iopub.status.idle":"2023-05-27T05:54:41.261986Z","shell.execute_reply.started":"2023-05-27T05:54:41.246869Z","shell.execute_reply":"2023-05-27T05:54:41.259825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(CFG.seed)","metadata":{"id":"hOFRzj74Ef66","execution":{"iopub.status.busy":"2023-05-27T05:31:49.092018Z","iopub.execute_input":"2023-05-27T05:31:49.092826Z","iopub.status.idle":"2023-05-27T05:31:49.107993Z","shell.execute_reply.started":"2023-05-27T05:31:49.092792Z","shell.execute_reply":"2023-05-27T05:31:49.107107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper Functions","metadata":{"id":"dIAjE97bTrfJ"}},{"cell_type":"code","source":"# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.T.flatten()\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)\n\ndef rle_decode(mask_rle, shape=(768, 768)):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T  # Needed to align to RLE direction","metadata":{"id":"s9b-ezVvTnVf","execution":{"iopub.status.busy":"2023-05-27T05:31:49.109673Z","iopub.execute_input":"2023-05-27T05:31:49.110648Z","iopub.status.idle":"2023-05-27T05:31:49.119765Z","shell.execute_reply.started":"2023-05-27T05:31:49.110614Z","shell.execute_reply":"2023-05-27T05:31:49.118885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def imshow(img):\n    npimg = img.numpy()\n    plt.imshow(np.transpose(npimg, (1, 2, 0)))\n\n\ndef plot_images_masks(images, masks):\n    num_images = images.shape[0]\n\n    fig, axs = plt.subplots(num_images, 2, figsize=(5, num_images*2))\n    \n    for i in range(num_images):\n        img = images[i].permute(1, 2, 0).numpy()  # move channels to last dimension\n        mask = masks[i].squeeze(0).numpy()  # remove channels dimension\n        img = np.clip(img, 0, 1)\n        mask = np.clip(mask, 0, 1)\n        \n        if num_images == 1:\n            axs[0].imshow(img, cmap='gray')\n            axs[0].axis('off')\n            axs[0].set_title('Image')\n\n            axs[1].imshow(mask, cmap='gray')\n            axs[1].axis('off')\n            axs[1].set_title('Mask')\n        else:\n            axs[i, 0].imshow(img, cmap='gray')\n            axs[i, 0].axis('off')\n            axs[i, 0].set_title('Image')\n\n            axs[i, 1].imshow(mask, cmap='gray')\n            axs[i, 1].axis('off')\n            axs[i, 1].set_title('Mask')\n    \n    plt.tight_layout()\n    plt.show()\n\n\ndef plot_images_masks_and_preds(images, masks, preds):\n    num_images = images.shape[0]\n\n    fig, axs = plt.subplots(num_images, 3, figsize=(5, num_images*2))\n    \n    for i in range(num_images):\n        img = images[i].permute(1, 2, 0).numpy()  # move channels to last dimension\n        mask = masks[i].squeeze(0).numpy()  # remove channels dimension\n        pred = preds[i].permute(1, 2, 0).numpy()  # move channels to last dimension\n        img = np.clip(img, 0, 1)\n        mask = np.clip(mask, 0, 1)\n        pred = np.clip(pred, 0, 1)\n        \n        if num_images == 1:\n            axs[0].imshow(img)\n            axs[0].axis('off')\n            axs[0].set_title('Image')\n\n            axs[1].imshow(mask)\n            axs[1].axis('off')\n            axs[1].set_title('Mask')\n\n            axs[2].imshow(preds)\n            axs[2].axis('off')\n            axs[2].set_title('pred mask')\n        else:\n            axs[i, 0].imshow(img)\n            axs[i, 0].axis('off')\n            axs[i, 0].set_title('Image')\n\n            axs[i, 1].imshow(mask)\n            axs[i, 1].axis('off')\n            axs[i, 1].set_title('Mask')\n\n            axs[i, 2].imshow(pred)\n            axs[i, 2].axis('off')\n            axs[i, 2].set_title('pred mask')\n    \n    plt.tight_layout()\n    plt.show()\n\n\ndef plot_learning_curve(train_loss_history, val_loss_history):\n    plt.figure()\n    plt.plot(train_loss_history, label='train')\n    plt.plot(val_loss_history, label='val')\n    plt.title('Training and validation loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.show()\n\n\ndef masks_as_image(in_mask_list):\n    # Take the individual ship masks and create a single mask array for all ships\n    all_masks = np.zeros((768, 768), dtype = np.int16)\n    #if isinstance(in_mask_list, list):\n    for mask in in_mask_list:\n        if isinstance(mask, str):\n            all_masks += rle_decode(mask)\n    return np.expand_dims(all_masks, -1)","metadata":{"id":"ELuQRMJ2Ttck","execution":{"iopub.status.busy":"2023-05-27T05:31:49.121235Z","iopub.execute_input":"2023-05-27T05:31:49.121862Z","iopub.status.idle":"2023-05-27T05:31:49.142787Z","shell.execute_reply.started":"2023-05-27T05:31:49.121827Z","shell.execute_reply":"2023-05-27T05:31:49.141830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{"id":"c_IhTISBT8iC"}},{"cell_type":"code","source":"class Airbus(Dataset):\n    \"\"\" Airbus Ship dataset \"\"\"\n    \n    def __init__(self, csv_file, root_dir, transform=None, train=True):\n        \"\"\"\n        Arguments:\n            csv_file (string): Path to the csv file with annotations.\n            root_dir (string): Directory with all the images.\n            transform (callable, optional): Optional transform to be applied\n                on a sample.\n            train (bool): Indicate whether train or test set.\n        \"\"\"\n        self.images_frame = pd.read_csv(os.path.join(root_dir, csv_file)).fillna(-1)\n        self.root_dir = root_dir\n        self.transform = transform\n        self.train = train\n        if train: \n            self.data_path = 'train_v2'\n        else:\n            self.data_path = 'test_v2'\n        self.filenames = self.get_file_names(self.images_frame)\n    \n    \n    def get_file_names(self, images_frame):\n        dataset_filenames = []\n        filenames = glob.glob(os.path.join(self.root_dir, self.data_path,'*.jpg'))\n        \n        for fn in filenames:\n            if fn in CFG.exclude_list:\n                continue\n            dataset_filenames.append(fn)\n        return dataset_filenames\n    \n    \n    def get_mask(self, image, image_id):\n        height, width = image.shape[:2]\n        img_rle_masks = self.images_frame.loc[self.images_frame['ImageId'] == image_id, 'EncodedPixels'].tolist()\n        all_masks = np.zeros((height, width))\n        \n        if img_rle_masks == [-1]:\n            return all_masks\n        for rle_mask in img_rle_masks:\n            all_masks += rle_decode(rle_mask)\n                             \n        return all_masks\n    \n    def __len__(self):\n        return len(self.filenames)\n    \n    \n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n        \n        image_name = self.filenames[idx]\n        image = np.array(Image.open(image_name))\n        image_id = self.filenames[idx].split('/')[-1]\n        \n        if self.train:\n            mask = self.get_mask(image, image_id)\n            sample = {'image': image, 'mask': mask}\n        else:\n            mask = np.zeros_like(image)\n            sample = {'image': image}\n            \n        if self.transform:\n            sample = self.transform(image=image, mask=mask)\n            \n        if self.train:\n            return sample['image'], sample['mask'][np.newaxis, :, :]\n        \n        return sample['image']","metadata":{"id":"bf05d3gVT4ot","execution":{"iopub.status.busy":"2023-05-27T05:31:49.145545Z","iopub.execute_input":"2023-05-27T05:31:49.145889Z","iopub.status.idle":"2023-05-27T05:31:49.160694Z","shell.execute_reply.started":"2023-05-27T05:31:49.145856Z","shell.execute_reply":"2023-05-27T05:31:49.159557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = A.Compose(CFG.train_aug)\nval_transform = A.Compose(CFG.torch_aug)","metadata":{"id":"Ni7rH8DiT9PK","execution":{"iopub.status.busy":"2023-05-27T05:31:49.162277Z","iopub.execute_input":"2023-05-27T05:31:49.162684Z","iopub.status.idle":"2023-05-27T05:31:49.175115Z","shell.execute_reply.started":"2023-05-27T05:31:49.162652Z","shell.execute_reply":"2023-05-27T05:31:49.174143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_set = Airbus(CFG.csv_file, CFG.root_dir, val_transform)\n\n# Split the dataset into training and validation sets\ntrain_ratio = 0.8  \nval_ratio = 1 - train_ratio  \n\ntrain_size = int(train_ratio * len(train_set))\nval_size = len(train_set) - train_size\n\ntrain_dataset, val_dataset = torch.utils.data.random_split(train_set, [train_size, val_size])","metadata":{"id":"kxDyz5g0UAcO","execution":{"iopub.status.busy":"2023-05-27T05:31:49.178999Z","iopub.execute_input":"2023-05-27T05:31:49.179297Z","iopub.status.idle":"2023-05-27T05:31:55.915627Z","shell.execute_reply.started":"2023-05-27T05:31:49.179267Z","shell.execute_reply":"2023-05-27T05:31:55.914616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True)\n\nval_loader = DataLoader(val_dataset, batch_size=CFG.batch_size, shuffle=False)\n\ndata_loaders = {\n    'train': train_loader,\n    'val': val_loader\n}","metadata":{"id":"JqW931ugbOVa","execution":{"iopub.status.busy":"2023-05-27T05:31:55.917278Z","iopub.execute_input":"2023-05-27T05:31:55.917645Z","iopub.status.idle":"2023-05-27T05:31:55.924393Z","shell.execute_reply.started":"2023-05-27T05:31:55.917609Z","shell.execute_reply":"2023-05-27T05:31:55.923395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_set)","metadata":{"id":"NoLa1fK0s4kG","outputId":"33ea6d90-95b5-41b6-d19b-6915e810c18e","execution":{"iopub.status.busy":"2023-05-27T05:31:55.925884Z","iopub.execute_input":"2023-05-27T05:31:55.926383Z","iopub.status.idle":"2023-05-27T05:31:55.936318Z","shell.execute_reply.started":"2023-05-27T05:31:55.926352Z","shell.execute_reply":"2023-05-27T05:31:55.935358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"WnBrTPZ_UF1L"}},{"cell_type":"code","source":"model = smp.Unet(\n    encoder_name = CFG.backbone, \n    encoder_weights = 'imagenet',\n    in_channels = 3,\n    classes = 1,\n    activation = 'sigmoid'\n)","metadata":{"id":"PdCIAHCEUFX7","execution":{"iopub.status.busy":"2023-05-27T05:31:55.937559Z","iopub.execute_input":"2023-05-27T05:31:55.937965Z","iopub.status.idle":"2023-05-27T05:31:57.502844Z","shell.execute_reply.started":"2023-05-27T05:31:55.937931Z","shell.execute_reply":"2023-05-27T05:31:57.501895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.path.exists(CFG.saved_model_path):\n    print(model.load_state_dict(torch.load(CFG.saved_model_path, map_location=device)))","metadata":{"id":"FRhmMwE1BuJs","outputId":"e99f07e7-235a-4c26-a511-0337d992056d","execution":{"iopub.status.busy":"2023-05-27T05:54:55.071102Z","iopub.execute_input":"2023-05-27T05:54:55.071980Z","iopub.status.idle":"2023-05-27T05:54:58.474341Z","shell.execute_reply.started":"2023-05-27T05:54:55.071943Z","shell.execute_reply":"2023-05-27T05:54:58.473374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{"id":"h_Bx6sjkWIDP"}},{"cell_type":"code","source":"criterion = smp.utils.losses.DiceLoss()","metadata":{"id":"VUTi6pl1eI6H","execution":{"iopub.status.busy":"2023-05-27T05:56:21.741284Z","iopub.execute_input":"2023-05-27T05:56:21.741987Z","iopub.status.idle":"2023-05-27T05:56:21.749162Z","shell.execute_reply.started":"2023-05-27T05:56:21.741949Z","shell.execute_reply":"2023-05-27T05:56:21.748229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer","metadata":{"id":"xCS0qbVObmEm"}},{"cell_type":"code","source":"optimizer = Adam(model.parameters(), lr=1e-4, weight_decay=CFG.weight_decay)","metadata":{"id":"l-z1am2fbnTW","execution":{"iopub.status.busy":"2023-05-27T05:56:22.031097Z","iopub.execute_input":"2023-05-27T05:56:22.031760Z","iopub.status.idle":"2023-05-27T05:56:22.038718Z","shell.execute_reply.started":"2023-05-27T05:56:22.031725Z","shell.execute_reply":"2023-05-27T05:56:22.037812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metric","metadata":{}},{"cell_type":"code","source":"metrics = [segmentation_models_pytorch.utils.metrics.IoU(threshold=0.5)]","metadata":{"execution":{"iopub.status.busy":"2023-05-27T05:56:22.733116Z","iopub.execute_input":"2023-05-27T05:56:22.733806Z","iopub.status.idle":"2023-05-27T05:56:22.738686Z","shell.execute_reply.started":"2023-05-27T05:56:22.733770Z","shell.execute_reply":"2023-05-27T05:56:22.737633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_IMAGES = 48000\nindices = range(NUM_IMAGES)\n\ntrain_subset_sampler = SubsetRandomSampler(indices)\ntest_subset_sampler = SubsetRandomSampler(indices)\n\nsubset_train_dataloader = DataLoader(train_dataset,\n                                     batch_size=64,\n                                     sampler=train_subset_sampler,\n                                     num_workers=4\n                                     )\nsubset_val_dataloader = DataLoader(val_dataset,\n                                   batch_size=64,\n                                   sampler=test_subset_sampler,\n                                   num_workers=4)\n\nsubset_loaders = {\n    'train': subset_train_dataloader,\n    'val': subset_val_dataloader\n}","metadata":{"id":"28GJEAf04SyG","execution":{"iopub.status.busy":"2023-05-27T05:56:23.988994Z","iopub.execute_input":"2023-05-27T05:56:23.989360Z","iopub.status.idle":"2023-05-27T05:56:23.995568Z","shell.execute_reply.started":"2023-05-27T05:56:23.989330Z","shell.execute_reply":"2023-05-27T05:56:23.994587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Scheduler","metadata":{}},{"cell_type":"code","source":"scheduler = OneCycleLR(optimizer, 1e-4, epochs=CFG.epochs, steps_per_epoch=16)","metadata":{"execution":{"iopub.status.busy":"2023-05-27T05:56:39.598042Z","iopub.execute_input":"2023-05-27T05:56:39.598401Z","iopub.status.idle":"2023-05-27T05:56:39.604010Z","shell.execute_reply.started":"2023-05-27T05:56:39.598371Z","shell.execute_reply":"2023-05-27T05:56:39.603122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"id":"hhR9_1FBbkAM"}},{"cell_type":"code","source":"def train_model(train_epoch, dataloader, scheduler, num_epochs=10):\n    train_logs = []\n    max_score = 0\n    \n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch + 1}/{num_epochs}')\n        \n        train_log = train_epoch.run(dataloader)\n        scheduler.step(epoch)\n        \n        if max_score < train_log['iou_score']:\n            torch.save(model.state_dict(), CFG.model_save_path)\n            max_score = train_log['iou_score']\n        \n        train_logs.append(train_log)\n        \n    return train_log","metadata":{"id":"m1DEaKshbKEE","execution":{"iopub.status.busy":"2023-05-27T05:56:42.859525Z","iopub.execute_input":"2023-05-27T05:56:42.860492Z","iopub.status.idle":"2023-05-27T05:56:42.867749Z","shell.execute_reply.started":"2023-05-27T05:56:42.860437Z","shell.execute_reply":"2023-05-27T05:56:42.866624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_epoch = segmentation_models_pytorch.utils.train.TrainEpoch(\n    model, \n    loss=criterion, \n    metrics=metrics, \n    optimizer=optimizer,\n    device=device,\n    verbose=True,\n)","metadata":{"id":"iH5jI3jlblLP","outputId":"3a119d3d-a422-4338-8c16-49a38400f734","execution":{"iopub.status.busy":"2023-05-27T05:56:46.741598Z","iopub.execute_input":"2023-05-27T05:56:46.741945Z","iopub.status.idle":"2023-05-27T05:56:46.771421Z","shell.execute_reply.started":"2023-05-27T05:56:46.741916Z","shell.execute_reply":"2023-05-27T05:56:46.770439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_model(train_epoch, subset_train_dataloader, scheduler, num_epochs=5)","metadata":{"execution":{"iopub.status.busy":"2023-05-26T20:57:05.305319Z","iopub.execute_input":"2023-05-26T20:57:05.305990Z","iopub.status.idle":"2023-05-26T23:12:35.930205Z","shell.execute_reply.started":"2023-05-26T20:57:05.305959Z","shell.execute_reply":"2023-05-26T23:12:35.928902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (images, masks) in enumerate(train_loader):\n    images = images.to(device)\n    masks = masks.to(device)\n    y_pred = model(images)\n\n    plot_images_masks_and_preds(images.detach().cpu(),\n                                masks.detach().cpu(),\n                                y_pred.detach().cpu())\n    \n    break","metadata":{"id":"qoAzsk4U-2cL","outputId":"98afe7cc-8321-4786-a62f-377bf7175704","execution":{"iopub.status.busy":"2023-05-27T05:57:59.418320Z","iopub.execute_input":"2023-05-27T05:57:59.419015Z","iopub.status.idle":"2023-05-27T05:58:10.634292Z","shell.execute_reply.started":"2023-05-27T05:57:59.418981Z","shell.execute_reply":"2023-05-27T05:58:10.633462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, (images, masks) in enumerate(val_loader):\n    images = images.to(device)\n    masks = masks.to(device)\n    y_pred = model(images)\n\n    plot_images_masks_and_preds(images.detach().cpu(),\n                                masks.detach().cpu(),\n                                y_pred.detach().cpu())\n    \n    break","metadata":{"id":"l3BLu6vQHkoc","outputId":"5008dfbc-97ad-4052-add8-a68055b31861","execution":{"iopub.status.busy":"2023-05-27T05:58:45.655145Z","iopub.execute_input":"2023-05-27T05:58:45.655645Z","iopub.status.idle":"2023-05-27T05:58:55.029716Z","shell.execute_reply.started":"2023-05-27T05:58:45.655611Z","shell.execute_reply":"2023-05-27T05:58:55.028622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test","metadata":{"id":"iuV-B5WNHlnn"}},{"cell_type":"code","source":"class TTA:\n    def __init__(self, model, transform, device='cpu'):\n        self.model = model\n        self.transform = transform\n        self.device = device\n\n    def __call__(self, images, num_augmentations=5):\n        # Expecting images to be of shape [B, C, H, W]\n\n        tta_preds = []\n        for _ in range(num_augmentations):\n            augmented_images = self.transform(images)\n\n            preds = self._predict(augmented_images)\n            preds_inv = self._predict(F.hflip(augmented_images))\n\n            tta_preds.extend([preds, preds_inv])\n\n        tta_preds = torch.stack(tta_preds)\n        averaged_preds = torch.mean(tta_preds, dim=0)\n        return averaged_preds\n\n    def _predict(self, images):\n        with torch.no_grad():\n            preds = self.model(images.to(self.device))\n        return preds.cpu()\n","metadata":{"id":"3vjwqOJoHnYc","execution":{"iopub.status.busy":"2023-05-27T05:59:14.367370Z","iopub.execute_input":"2023-05-27T05:59:14.368083Z","iopub.status.idle":"2023-05-27T05:59:14.375612Z","shell.execute_reply.started":"2023-05-27T05:59:14.368047Z","shell.execute_reply":"2023-05-27T05:59:14.374659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_transforms = A.Compose(CFG.test_aug)","metadata":{"id":"ZThtB4luHoed","execution":{"iopub.status.busy":"2023-05-27T05:59:17.326260Z","iopub.execute_input":"2023-05-27T05:59:17.326613Z","iopub.status.idle":"2023-05-27T05:59:17.333091Z","shell.execute_reply.started":"2023-05-27T05:59:17.326585Z","shell.execute_reply":"2023-05-27T05:59:17.330264Z"},"trusted":true},"execution_count":null,"outputs":[]}]}