{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":9988,"databundleVersionId":868324,"sourceType":"competition"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Test Task (Borys Melnychuk)**","metadata":{}},{"cell_type":"markdown","source":"Downloading the library that provide us pretrained UNet","metadata":{}},{"cell_type":"code","source":"!pip install git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-07-12T13:11:40.716347Z","iopub.execute_input":"2024-07-12T13:11:40.716694Z","iopub.status.idle":"2024-07-12T13:12:14.925804Z","shell.execute_reply.started":"2024-07-12T13:11:40.716665Z","shell.execute_reply":"2024-07-12T13:12:14.924784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms.v2 as T\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nimport segmentation_models_pytorch as smp\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom skimage.morphology import binary_opening, disk\nfrom skimage.measure import label, regionprops\nimport seaborn as sns\nimport cv2\nfrom tqdm import tqdm\n\nimport os\nfrom pathlib import Path\nfrom datetime import datetime\nimport json\nimport random","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:14.928325Z","iopub.execute_input":"2024-07-12T13:12:14.928790Z","iopub.status.idle":"2024-07-12T13:12:24.995325Z","shell.execute_reply.started":"2024-07-12T13:12:14.928737Z","shell.execute_reply":"2024-07-12T13:12:24.994527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Assigning some constants\nUNET_IMG_SIZE = 256\nIMG_SIZE = 768\nIMG_CHANNELS = 3\n\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD = [0.229, 0.224, 0.225]\n\nROOT_DIR = '/kaggle/input/airbus-ship-detection'\nSEGMENTATION_FILENAME = os.path.join(ROOT_DIR, 'train_ship_segmentations_v2.csv')\nTRAIN_DIR = os.path.join(ROOT_DIR, 'train_v2')\nTEST_DIR = os.path.join(ROOT_DIR, 'test_v2')\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:24.996415Z","iopub.execute_input":"2024-07-12T13:12:24.996958Z","iopub.status.idle":"2024-07-12T13:12:25.026964Z","shell.execute_reply.started":"2024-07-12T13:12:24.996931Z","shell.execute_reply":"2024-07-12T13:12:25.025821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seg_df = pd.read_csv(SEGMENTATION_FILENAME)\nseg_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:25.028311Z","iopub.execute_input":"2024-07-12T13:12:25.028904Z","iopub.status.idle":"2024-07-12T13:12:26.075472Z","shell.execute_reply.started":"2024-07-12T13:12:25.028865Z","shell.execute_reply":"2024-07-12T13:12:26.074567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Creating dataset with and without the ships. This will help us in EDA\nnoships_df = seg_df[seg_df['EncodedPixels'].isna()]\n\nships_df = seg_df[~seg_df['EncodedPixels'].isna()].reset_index(drop=True)\nships_df = ships_df.groupby(['ImageId']).agg({'EncodedPixels': ' '.join}).reset_index()\n\nunique_ships = ships_df['ImageId'].unique().tolist()","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:26.077874Z","iopub.execute_input":"2024-07-12T13:12:26.078251Z","iopub.status.idle":"2024-07-12T13:12:27.232563Z","shell.execute_reply.started":"2024-07-12T13:12:26.078224Z","shell.execute_reply":"2024-07-12T13:12:27.231706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Our masks are encoded in run-length format so we need the function to decode and encode this","metadata":{}},{"cell_type":"code","source":"def rle_decode(mask_rle, shape=(IMG_SIZE, IMG_SIZE)):   # (height,width) \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    \n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:27.233588Z","iopub.execute_input":"2024-07-12T13:12:27.233878Z","iopub.status.idle":"2024-07-12T13:12:27.240588Z","shell.execute_reply.started":"2024-07-12T13:12:27.233854Z","shell.execute_reply":"2024-07-12T13:12:27.239782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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 multi_rle_encode(img):\n    labels = label(img)\n    return [rle_encode(labels==k) for k in np.unique(labels[labels>0])]","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:27.241676Z","iopub.execute_input":"2024-07-12T13:12:27.242021Z","iopub.status.idle":"2024-07-12T13:12:27.253728Z","shell.execute_reply.started":"2024-07-12T13:12:27.241986Z","shell.execute_reply":"2024-07-12T13:12:27.252830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **EDA**","metadata":{}},{"cell_type":"code","source":"pie_df = pd.DataFrame({\n    'type': ['train', 'test'],\n    'count': [len(os.listdir(TRAIN_DIR)), len(os.listdir(TEST_DIR))]\n})\n\nfig, axes = plt.subplots(1, 2, figsize=(16, 9))\n\nnoships_num = seg_df[seg_df['EncodedPixels'].isna()]['ImageId'].nunique()\n\naxes[0].pie(pie_df['count'], labels=pie_df['type'], autopct='%1.1f%%')\naxes[0].set_title('Pie plot of proportion of train and test data')\n\naxes[1].bar(x=pie_df['type'], height=pie_df['count'], label='total number')\naxes[1].bar(x=['train'], height=[noships_num], color='r', label='no ships number')\naxes[1].legend()\naxes[1].set_title('Bar plot of of train and test data')","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:27.254805Z","iopub.execute_input":"2024-07-12T13:12:27.255123Z","iopub.status.idle":"2024-07-12T13:12:29.013636Z","shell.execute_reply.started":"2024-07-12T13:12:27.255094Z","shell.execute_reply":"2024-07-12T13:12:29.012800Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We are dealing with highly imbalanced dataset. 77% of our images don't contain any ships. So we will be training only using images that contain ships","metadata":{}},{"cell_type":"code","source":"np.random.seed(42)\n\nunique_img_ids = seg_df[~seg_df['EncodedPixels'].isna()].groupby(['ImageId']).size().reset_index(name='counts')\n\nnum_ship_imgs = len(ships_df)\nval_ids = np.random.randint(0, num_ship_imgs, (int(num_ship_imgs * 0.1), ))\ntrain_ids = np.setdiff1d(np.arange(num_ship_imgs), val_ids)\n\ntrain_df = ships_df.iloc[train_ids]\nval_df = ships_df.iloc[val_ids]\n\ntrain_df = pd.merge(train_df, unique_img_ids)\nval_df = pd.merge(val_df, unique_img_ids)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:29.097280Z","iopub.execute_input":"2024-07-12T13:12:29.097603Z","iopub.status.idle":"2024-07-12T13:12:29.162978Z","shell.execute_reply.started":"2024-07-12T13:12:29.097573Z","shell.execute_reply":"2024-07-12T13:12:29.162209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n\ntrain_df['counts'].hist(bins=train_df['counts'].max(), ax=axes[0])\nval_df['counts'].hist(bins=val_df['counts'].max(), ax=axes[1])\n\naxes[0].set_title('Histogram of ship number in train dataset')\naxes[0].set_xlabel('Ship number')\naxes[1].set_title('Histogram of ship number in test dataset')\naxes[1].set_xlabel('Ship number')\n\nfig.suptitle('Histograms of ship number in train and validation dataset')","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:29.164082Z","iopub.execute_input":"2024-07-12T13:12:29.164336Z","iopub.status.idle":"2024-07-12T13:12:29.741019Z","shell.execute_reply.started":"2024-07-12T13:12:29.164314Z","shell.execute_reply":"2024-07-12T13:12:29.740202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Ships number is distributed evenly evenly train and validation datasets","metadata":{}},{"cell_type":"code","source":"# Plotting the images without the ships\n\nrows, cols = 5, 5\nfig, axes = plt.subplots(rows, cols, figsize=(16, 16))\n\nk = 0\n\nfor i in range(rows):\n    for j in range(cols):\n        img_path = os.path.join(TRAIN_DIR, noships_df['ImageId'].unique()[k])\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        axes[i, j].imshow(img)\n        axes[i, j].axis('off')\n        axes[i, j].set_title(img_path.split('/')[-1])\n        k += 1\n\nfig.tight_layout()\nfig.suptitle('The images that don\\'t contain the ships')","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:29.742376Z","iopub.execute_input":"2024-07-12T13:12:29.743253Z","iopub.status.idle":"2024-07-12T13:12:36.903067Z","shell.execute_reply.started":"2024-07-12T13:12:29.743216Z","shell.execute_reply":"2024-07-12T13:12:36.901711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plotting the images with the ships\n\nrows, cols = 5, 5\nfig, axes = plt.subplots(rows, cols, figsize=(16, 16), constrained_layout=True)\n\nk = 0\n\nfor i in range(rows):\n    for j in range(cols):\n        img_path = os.path.join(TRAIN_DIR, ships_df['ImageId'][k])\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        axes[i, j].imshow(img)\n        \n        masks = rle_decode(ships_df['EncodedPixels'][k])\n        \n        axes[i, j].imshow(masks, alpha=.25)\n        axes[i, j].axis('off')\n        axes[i, j].set_title(img_path.split('/')[-1])\n        \n        k += 1\n        \nfig.suptitle('The images that contain the ships and their masks')","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:36.904492Z","iopub.execute_input":"2024-07-12T13:12:36.904871Z","iopub.status.idle":"2024-07-12T13:12:45.936715Z","shell.execute_reply.started":"2024-07-12T13:12:36.904812Z","shell.execute_reply":"2024-07-12T13:12:45.935294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Model development**","metadata":{}},{"cell_type":"markdown","source":"We need to write own data augmentation because we need to apply them to both image and mask","metadata":{}},{"cell_type":"code","source":"def clip(img, dtype, maxval):\n    return np.clip(img, 0, maxval).astype(dtype)\n\nclass DualCompose:\n    def __init__(self, transforms):\n        self.transforms = transforms\n\n    def __call__(self, x, mask=None):\n        for t in self.transforms:\n            x, mask = t(x, mask)\n        return x, mask\n\n\nclass VerticalFlip:\n    def __init__(self, prob=0.5):\n        self.prob = prob\n\n    def __call__(self, img, mask=None):\n        if random.random() < self.prob:\n            img = cv2.flip(img, 0)\n            if mask is not None:\n                mask = cv2.flip(mask, 0)\n        return img, mask\n\n\nclass HorizontalFlip:\n    def __init__(self, prob=0.5):\n        self.prob = prob\n\n    def __call__(self, img, mask=None):\n        if random.random() < self.prob:\n            img = cv2.flip(img, 1)\n            if mask is not None:\n                mask = cv2.flip(mask, 1)\n        return img, mask\n\nclass RandomCrop:\n    def __init__(self, size):\n        self.h = size[0]\n        self.w = size[1]\n\n    def __call__(self, img, mask=None):\n        height, width, _ = img.shape\n\n        h_start = np.random.randint(0, height - self.h)\n        w_start = np.random.randint(0, width - self.w)\n\n        img = img[h_start: h_start + self.h, w_start: w_start + self.w,:]\n\n        assert img.shape[0] == self.h\n        assert img.shape[1] == self.w\n\n        if mask is not None:\n            if mask.ndim == 2:\n                mask = np.expand_dims(mask, axis=2)\n            mask = mask[h_start: h_start + self.h, w_start: w_start + self.w,:]\n\n        return img, mask\n\nclass CenterCrop:\n    def __init__(self, size):\n        self.height = size[0]\n        self.width = size[1]\n\n    def __call__(self, img, mask=None):\n        h, w, c = img.shape\n        dy = (h - self.height) // 2\n        dx = (w - self.width) // 2\n        y1 = dy\n        y2 = y1 + self.height\n        x1 = dx\n        x2 = x1 + self.width\n        img = img[y1:y2, x1:x2,:]\n        if mask is not None:\n            if mask.ndim == 2:\n                mask = np.expand_dims(mask, axis=2)\n            mask = mask[y1:y2, x1:x2,:]\n\n        return img, mask","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:45.941740Z","iopub.execute_input":"2024-07-12T13:12:45.942074Z","iopub.status.idle":"2024-07-12T13:12:45.958610Z","shell.execute_reply.started":"2024-07-12T13:12:45.942048Z","shell.execute_reply":"2024-07-12T13:12:45.957793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = {\n    'train': DualCompose([\n        HorizontalFlip(.25),\n        VerticalFlip(.25), \n        RandomCrop((256,256,3))\n    ]),\n    'val': DualCompose([\n        CenterCrop((512,512,3))\n    ])\n}","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:45.959803Z","iopub.execute_input":"2024-07-12T13:12:45.960121Z","iopub.status.idle":"2024-07-12T13:12:45.972456Z","shell.execute_reply.started":"2024-07-12T13:12:45.960093Z","shell.execute_reply":"2024-07-12T13:12:45.971486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ShipDataset(Dataset):\n    def __init__(self, root_dir: str, df: pd.DataFrame, mode: str = 'train', transforms=None):\n        self.root_dir = root_dir\n        self.mode = mode\n        self.df = df\n        \n        self.img_list = self.df['ImageId'].unique()\n        \n        self.transforms = transforms\n        self.posttransform = T.Compose([\n            T.ToImage(),\n            T.ToDtype(torch.float32, scale=True),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)  # use mean and std from ImageNet \n        ]) # here we just converting to tensor and normalizing the image\n        \n    def __getitem__(self, idx):\n        img_name = self.img_list[idx]\n        img_path = os.path.join(self.root_dir, img_name)\n        \n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        if self.mode == 'test':\n            return self.posttransform(img), str(img_path)\n        \n        encoded_mask = self.df[self.df['ImageId'] == self.img_list[idx]]['EncodedPixels'].values[0]\n        mask = rle_decode(encoded_mask)\n        \n        if self.transforms is not None:\n            img, mask = self.transforms(img, mask)\n        \n        return self.posttransform(img), torch.from_numpy(np.moveaxis(mask, -1, 0)).float()\n        \n        \n    def __len__(self) -> int:\n        return len(self.img_list)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:45.973530Z","iopub.execute_input":"2024-07-12T13:12:45.973877Z","iopub.status.idle":"2024-07-12T13:12:45.985317Z","shell.execute_reply.started":"2024-07-12T13:12:45.973851Z","shell.execute_reply":"2024-07-12T13:12:45.984383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# computing dice and jaccard\ndef compute_metrics(pred, true, batch_size=16, threshold=0.5):\n    pred = pred.view(batch_size, -1)\n    true = true.view(batch_size, -1)\n    \n    pred = (pred > threshold).float()\n    true = (true > threshold).float()\n    \n    pred_sum = pred.sum(-1)\n    true_sum = true.sum(-1)\n    \n    neg_index = torch.nonzero(true_sum == 0)\n    pos_index = torch.nonzero(true_sum >= 1)\n    \n    dice_neg = (pred_sum == 0).float()\n    dice_pos = 2 * ((pred * true).sum(-1)) / ((pred + true).sum(-1))\n    \n    dice_neg = dice_neg[neg_index]\n    dice_pos = dice_pos[pos_index]\n    \n    dice = torch.cat([dice_pos, dice_neg])\n    jaccard = dice / (2 - dice)\n    \n    return dice, jaccard\n    \nclass metrics:\n    def __init__(self, batch_size=16, threshold=0.5):\n        self.threshold = threshold\n        self.batchsize = batch_size\n        self.dice = []\n        self.jaccard = []\n    def collect(self, pred, true):\n        pred = torch.sigmoid(pred)\n        dice, jaccard = compute_metrics(pred, true, batch_size=self.batchsize, threshold=self.threshold)\n        self.dice.extend(dice)\n        self.jaccard.extend(jaccard)\n    def get(self):\n        dice = np.nanmean(self.dice)\n        jaccard = np.nanmean(self.jaccard)\n        return dice, jaccard\n    \nclass BCEDiceWithLogitsLoss(nn.Module):\n    def __init__(self, dice_weight=1, smooth=1):\n        super().__init__()\n        self.bce = nn.BCEWithLogitsLoss()\n        self.dice_weight = dice_weight\n        self.smooth = smooth\n        \n    def __call__(self, outputs, targets):\n        if outputs.size() != targets.size():\n            raise ValueError(\"size mismatch, {} != {}\".format(outputs.size(), targets.size()))\n            \n        loss = self.bce(outputs, targets)\n\n        targets = (targets == 1.0).float()\n        targets = targets.view(-1)\n        outputs = F.sigmoid(outputs)\n        outputs = outputs.view(-1)\n\n        intersection = (outputs * targets).sum()\n        dice = 2.0 * (intersection + self.smooth)  / (targets.sum() + outputs.sum() + self.smooth)\n        \n        loss -= self.dice_weight * torch.log(dice) # try with 1- dice\n\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:45.986392Z","iopub.execute_input":"2024-07-12T13:12:45.986648Z","iopub.status.idle":"2024-07-12T13:12:46.001704Z","shell.execute_reply.started":"2024-07-12T13:12:45.986623Z","shell.execute_reply":"2024-07-12T13:12:46.000885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<img src=\"https://lmb.informatik.uni-freiburg.de/people/ronneber/u-net/u-net-architecture.png\" alt=\"unet architecture\" />","metadata":{}},{"cell_type":"markdown","source":"<center>Architecture of UNet</center>","metadata":{}},{"cell_type":"code","source":"# Loading pretrained Unet\nmodel = smp.Unet(\"resnet34\", encoder_weights=\"imagenet\", activation=None).to(DEVICE)\n\n# defining hyperparameters\nTRAIN_BATCH_SIZE = 16\nVAL_BATCH_SIZE = 4\nTEST_BATCH_SIZE = 2\n\nLR = 1e-4\nNUM_EPOCHS = 3\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = optim.Adam(model.parameters(), lr=LR)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:46.002596Z","iopub.execute_input":"2024-07-12T13:12:46.002855Z","iopub.status.idle":"2024-07-12T13:12:47.354985Z","shell.execute_reply.started":"2024-07-12T13:12:46.002834Z","shell.execute_reply":"2024-07-12T13:12:47.354091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# creating dataset and dataloaders for train, validation and test images\ntrain_dataset = ShipDataset(TRAIN_DIR, train_df, 'train', transforms['train'])\nval_dataset = ShipDataset(TRAIN_DIR, val_df, 'val', transforms['val'])\n\ntrain_loader = DataLoader(dataset=train_dataset, shuffle=True, batch_size=TRAIN_BATCH_SIZE, num_workers=0)\nval_loader = DataLoader(dataset=val_dataset, shuffle=True, batch_size=VAL_BATCH_SIZE, num_workers=0)\n\ntest_list_imgs = os.listdir(TEST_DIR)\n\nprint(f'{len(test_list_imgs)} test images')\n\ntest_df = pd.DataFrame({\n    'ImageId': test_list_imgs,\n    'EncodedPixels': None\n})\n\n\ntest_dataset = ShipDataset(TEST_DIR, test_df, 'test', None)\ntest_loader = DataLoader(test_dataset, shuffle=False, batch_size=TEST_BATCH_SIZE, num_workers=0)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:47.356034Z","iopub.execute_input":"2024-07-12T13:12:47.356289Z","iopub.status.idle":"2024-07-12T13:12:47.388236Z","shell.execute_reply.started":"2024-07-12T13:12:47.356268Z","shell.execute_reply":"2024-07-12T13:12:47.387321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training**","metadata":{}},{"cell_type":"code","source":"def train(\n    model: nn.Module, \n    criterion: nn.Module, \n    train_loader: DataLoader, \n    valid_loader: DataLoader, \n    optimizer: nn.Module,\n    train_batch_size: int = 16, \n    val_batch_size: int = 4, \n    n_epochs: int = 1, \n    fold: int = 1, \n    device: torch.device = 'cpu'\n):\n    model_path = Path(f'model_{fold}.pt')\n    \n    if model_path.exists():\n        state = torch.load(str(model_path))\n        epoch = state['epoch']\n        step = state['step']\n        model.load_state_dict(state['model'])\n        print(f'Restored model, epoch {epoch}, step {step}')\n    else:\n        epoch = 1\n        step = 0\n\n    save = lambda ep: torch.save({\n        'model': model.state_dict(),\n        'epoch': ep,\n        'step': step,\n    }, str(model_path))\n\n    report_each = 50\n    log = open('train_{fold}.log'.format(fold=fold),'at', encoding='utf8')\n    \n    model = model.to(device)\n\n    for epoch in range(epoch, n_epochs + 1):\n        model.train()\n        random.seed()\n        train_loop = tqdm(total=len(train_loader) *  train_batch_size)\n        train_loop.set_description(f'Epoch {epoch}')\n        losses = []\n        valid_metrics = metrics(batch_size=val_batch_size)  # for validation\n        \n        try:\n            mean_loss = 0\n            for i, (inputs, targets) in enumerate(train_loader):\n                inputs, targets = inputs.to(device), targets.to(device)\n                optimizer.zero_grad()\n                outputs = model.forward(inputs)\n                loss = criterion(outputs, targets)\n                batch_size = inputs.size(0)\n                loss.backward()\n                optimizer.step()\n                step += 1\n                train_loop.update(batch_size)\n                losses.append(loss.item())\n                mean_loss = np.mean(losses[-report_each:])\n                train_loop.set_postfix(loss='{:.5f}'.format(mean_loss))\n                if i and i % report_each == 0:\n                    write_event(log, step, loss=mean_loss)\n            write_event(log, step, loss=mean_loss)\n            train_loop.close()\n            save(epoch + 1)\n            \n            # Validation\n            comb_loss_metrics = validation(model, criterion, valid_loader, valid_metrics, device)\n            write_event(log, step, **comb_loss_metrics)\n\n        except KeyboardInterrupt:\n            train_loop.close()\n            print('Ctrl+C, saving snapshot')\n            save(epoch)\n            print('done.')\n            return\n        \ndef validation(\n    model: nn.Module, \n    criterion: nn.Module, \n    valid_loader: DataLoader, \n    metrics, \n    device: torch.device\n):\n    print(\"Validation\")\n    \n    losses = []\n    model.eval()\n    \n    for inputs, targets in valid_loader:\n        inputs, targets = inputs.to(device), targets.to(device)\n        outputs = model.forward(inputs)\n        loss = criterion(outputs, targets)\n        losses.append(loss.item())\n        metrics.collect(outputs.detach().cpu(), targets.detach().cpu()) # get metrics \n    \n    valid_loss = np.mean(losses)  # float\n    valid_dice, valid_jaccard = metrics.get() # float\n\n    print('Valid loss: {:.5f}, Jaccard: {:.5f}, Dice: {:.5f}'.format(valid_loss, valid_jaccard, valid_dice))\n    comb_loss_metrics = {'valid_loss': valid_loss, 'jaccard': valid_jaccard.item(), 'dice': valid_dice.item()}\n    \n    return comb_loss_metrics\n\ndef write_event(log, step: int, **data):\n    data['step'] = step\n    data['dt'] = datetime.now().isoformat()\n    log.write(json.dumps(data, sort_keys=True))\n    log.write('\\n')\n    log.flush()","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:47.389525Z","iopub.execute_input":"2024-07-12T13:12:47.389822Z","iopub.status.idle":"2024-07-12T13:12:47.410172Z","shell.execute_reply.started":"2024-07-12T13:12:47.389798Z","shell.execute_reply":"2024-07-12T13:12:47.409158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def imshow_gt_out(img, mask_gt, mask_out):\n    \"\"\"\n    Plots input image, ground truth and output mask\n    \"\"\"\n    img = img.numpy().transpose((1, 2, 0))\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    img = std * img + mean\n    img = np.clip(img, 0, 1)\n\n    mask_gt = mask_gt.numpy().transpose((1, 2, 0))\n    mask_gt = np.clip(mask_gt, 0, 1)\n\n    mask_out = mask_out.numpy().transpose((1, 2, 0))\n    mask_out = np.clip(mask_out, 0, 1)\n\n    fig, axs = plt.subplots(1,3, figsize=(10,30))\n    axs[0].imshow(img)\n    axs[0].axis('off')\n    axs[0].set_title(\"Input image\")\n    axs[1].imshow(mask_gt)\n    axs[1].axis('off')\n    axs[1].set_title(\"Ground truth\")\n    axs[2].imshow(mask_out)\n    axs[2].axis('off')\n    axs[2].set_title(\"Model output\")\n    plt.subplots_adjust(wspace=0, hspace=0)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:47.411400Z","iopub.execute_input":"2024-07-12T13:12:47.412001Z","iopub.status.idle":"2024-07-12T13:12:47.423523Z","shell.execute_reply.started":"2024-07-12T13:12:47.411970Z","shell.execute_reply":"2024-07-12T13:12:47.422681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run_id = 1\n\ntrain(\n    model=model,\n    criterion=criterion,\n    optimizer=optimizer,\n    train_loader=train_loader,\n    valid_loader=val_loader,\n    train_batch_size= TRAIN_BATCH_SIZE,\n    val_batch_size=VAL_BATCH_SIZE,\n    fold=run_id,\n    n_epochs = NUM_EPOCHS,\n    device=DEVICE\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T13:12:47.424878Z","iopub.execute_input":"2024-07-12T13:12:47.425276Z","iopub.status.idle":"2024-07-12T14:05:43.005124Z","shell.execute_reply.started":"2024-07-12T13:12:47.425218Z","shell.execute_reply":"2024-07-12T14:05:43.004138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Displaying how our model perform on validation dataset\nimages, ground_truth = next(iter(val_loader))\nground_truth = ground_truth.data.cpu()\nimages = images.to(DEVICE)\noutput = model(images)\noutput = ((output > 0).float()) * 255\n\nimages = images.data.cpu()\noutput = output.data.cpu()\nimshow_gt_out(torchvision.utils.make_grid(images, nrow=1),torchvision.utils.make_grid(ground_truth, nrow=1), torchvision.utils.make_grid(output, nrow=1))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-12T14:05:43.006288Z","iopub.execute_input":"2024-07-12T14:05:43.006550Z","iopub.status.idle":"2024-07-12T14:05:44.017909Z","shell.execute_reply.started":"2024-07-12T14:05:43.006528Z","shell.execute_reply":"2024-07-12T14:05:44.017036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot losses\nlog_file = f'train_{run_id}.log'\nlogs = pd.read_json(log_file, lines=True)\n\nplt.figure(figsize=(26,6))\nplt.subplot(1, 2, 1)\nplt.plot(logs.step[logs.loss.notnull()],\n            logs.loss[logs.loss.notnull()],\n            label=\"on training set\")\n\nplt.plot(logs.step[logs.valid_loss.notnull()],\n            logs.valid_loss[logs.valid_loss.notnull()],\n            label = \"on validation set\")\n         \nplt.title('Losses during the training')\nplt.xlabel('Step')\nplt.ylabel('BCEWithLogitsLoss')\nplt.legend()\nplt.tight_layout()\nplt.show();","metadata":{"execution":{"iopub.status.busy":"2024-07-12T14:05:44.019164Z","iopub.execute_input":"2024-07-12T14:05:44.019530Z","iopub.status.idle":"2024-07-12T14:05:44.449433Z","shell.execute_reply.started":"2024-07-12T14:05:44.019496Z","shell.execute_reply":"2024-07-12T14:05:44.448484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Evaluation**","metadata":{}},{"cell_type":"code","source":"# We need to collect all test images and pass them through the model to get the masks\nmodel.eval()\n\nrows = []\n\ntest_loop = enumerate(tqdm(test_loader, desc='Test'))\nfor batch_idx, (imgs, names) in test_loop:\n    imgs = imgs.to(DEVICE)\n    \n    with torch.no_grad():\n        outputs = model(imgs)\n        \n    for i, img_name in enumerate(names):\n        mask = F.sigmoid(outputs[i,0]).data.detach().cpu().numpy()\n        \n        seg_mask = binary_opening(mask > 0.5, disk(2))\n        encoded_masks = multi_rle_encode(seg_mask)\n        \n        if len(encoded_masks) > 0:\n            for encoded_mask in encoded_masks:\n                rows += [{'ImageId': img_name.split('/')[-1], 'EncodedPixels': encoded_mask}]\n        else:\n            rows += [{'ImageId': img_name.split('/')[-1], 'EncodedPixels': None}]\n            \nmodel.train()\n            \nsubmission_df = pd.DataFrame(rows)[['ImageId', 'EncodedPixels']]\nsubmission_df.to_csv('submission.csv', index=False)\nsubmission_df.sample(15)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T14:05:44.450745Z","iopub.execute_input":"2024-07-12T14:05:44.451425Z","iopub.status.idle":"2024-07-12T14:21:19.148794Z","shell.execute_reply.started":"2024-07-12T14:05:44.451391Z","shell.execute_reply":"2024-07-12T14:21:19.147897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plotting the the images and output masks\n\nrows, cols = 8, 8\nfig, axes = plt.subplots(rows, cols, figsize=(16, 16), constrained_layout=True)\n\nk = 0\n\ntemp_df =  submission_df.fillna('') \\\n    .groupby(['ImageId']) \\\n    .agg({'EncodedPixels': ' '.join}) \\\n    .reset_index()\n\nfor i in range(rows):\n    for j in range(cols):\n        img_path = os.path.join(TEST_DIR, temp_df['ImageId'][k])\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        axes[i, j].imshow(img)\n        \n        masks = rle_decode(temp_df['EncodedPixels'][k])\n        \n        axes[i, j].imshow(masks, alpha=.4)\n        axes[i, j].axis('off')\n        \n        \n        if temp_df['EncodedPixels'][k] == '':\n            axes[i, j].set_title(f\"predicted no ships\")\n        else:\n            axes[i, j].set_title(f\"predicted ships\")\n            \n        k += 1\n\nfig.suptitle('Output masks on the test images')","metadata":{"execution":{"iopub.status.busy":"2024-07-12T14:27:11.855286Z","iopub.execute_input":"2024-07-12T14:27:11.855647Z","iopub.status.idle":"2024-07-12T14:27:30.177375Z","shell.execute_reply.started":"2024-07-12T14:27:11.855617Z","shell.execute_reply":"2024-07-12T14:27:30.175390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can also try to use Unet with deeper encoder (like Resnet 50, 101, 152), use gaussian, laplace filter (to smooth the images) and in such way we may get better results. But for now we get Dice = 0.67961, Jaccard = 0.62513, Kaggle submission public score = 0.81563, public score 0.67999","metadata":{}}]}