{"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":"Fine-tuning UNet with fastai v1","metadata":{}},{"cell_type":"code","source":"import sklearn\nprint(sklearn.__version__)\n!pip install torch==1.4.0\n!pip install torchvision==0.5.0\n!pip install fastai==1.0.60\n!pip install numpy==1.18.2\n!pip install pandas==0.25.3\n!pip install sklearn==0.22.2\n\n","metadata":{"execution":{"iopub.status.busy":"2022-05-25T02:33:21.934542Z","iopub.execute_input":"2022-05-25T02:33:21.935005Z","iopub.status.idle":"2022-05-25T02:33:51.762086Z","shell.execute_reply.started":"2022-05-25T02:33:21.934968Z","shell.execute_reply":"2022-05-25T02:33:51.761289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom tqdm import tqdm\nfrom skimage.morphology import label, binary_opening, disk\nfrom fastai import vision, basic_data, layers, metrics\nfrom fastai.callbacks import hooks\nimport torch \nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nimport numpy as np \nimport pandas as pd \n\ndata_root = '../input/airbus-ship-detection/'\npath_train = os.path.join(data_root,'train_v2')\npath_test = os.path.join(data_root,'test_v2')\n\nimg_shape = (768, 768)\n\n# Booleans\nUSE_SELF_ATTENTION = True\nUSE_UNET34_AIRBUS = False\nUSE_FULL_RES_PRED = False\nUSE_FT_DECODE = True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-05-07T07:10:39.577467Z","iopub.execute_input":"2022-05-07T07:10:39.577765Z","iopub.status.idle":"2022-05-07T07:10:39.5901Z","shell.execute_reply.started":"2022-05-07T07:10:39.577737Z","shell.execute_reply":"2022-05-07T07:10:39.589233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get groundtruth and remove empty images","metadata":{}},{"cell_type":"code","source":"# Get dataframe with label\nmasks_df = pd.read_csv(os.path.join(data_root, 'train_ship_segmentations_v2.csv'))\nmasks_df = masks_df[~masks_df['ImageId'].isin(['6384c3e78.jpg'])]  # remove corrupted image\nmasks_df = masks_df.dropna() # remove images withtout ships\nunique_img_ids_df = masks_df.groupby('ImageId').size().reset_index(name='counts')\nunique_img_ids_df = unique_img_ids_df.drop(columns=['counts'])","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:39.614697Z","iopub.execute_input":"2022-05-07T07:10:39.615322Z","iopub.status.idle":"2022-05-07T07:10:40.224316Z","shell.execute_reply.started":"2022-05-07T07:10:39.615268Z","shell.execute_reply":"2022-05-07T07:10:40.223553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Losses","metadata":{}},{"cell_type":"code","source":"def cuda(x): \n    return x.cuda() if torch.cuda.is_available() else x\n\ndef make_one_hot(labels, c=2):\n    one_hot = cuda(torch.FloatTensor(labels.size(0), c, labels.size(2), labels.size(3)).zero_())\n    target = one_hot.scatter_(1, labels.data, 1)\n    target = cuda(Variable(target))\n    return target\n\ndef dice_loss(input, target):\n    # Input is of shape N,C,H,W\n    smooth = 1\n    batch_size = input.size(0)\n    input = F.softmax(input, dim=1)\n    # Since we have only 2 classes transform it to N,H,W and treat as sigmoid\n    input = input.view(batch_size, 2, -1)[:, 1, :]\n    target = make_one_hot(target).view(batch_size, 2, -1)[:, 1, :]\n\n    inter = torch.sum(input * target) + smooth\n    union = torch.sum(input) + torch.sum(target) + smooth\n\n    return -torch.log(2.0 * inter / union)\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma):\n        super(FocalLoss, self).__init__()\n        self.gamma = gamma\n\n    def forward(self, input, target):\n        # One-hot encode target\n        target = target.squeeze(1)\n    \n        input = input.view(input.size(0),input.size(1),-1) # N,C,H,W => N,C,H*W\n        input = input.transpose(1,2)                       # N,C,H*W => N,H*W,C\n        input = input.contiguous().view(-1,input.size(2))  # N,H*W,C => N*H*W,C\n        target = target.view(-1,1)\n\n        logpt = F.log_softmax(input, dim=1)\n        logpt = logpt.gather(1,target)\n        logpt = logpt.view(-1)\n        pt = Variable(logpt.data.exp())\n\n        loss = -1 * (1-pt)**self.gamma * logpt\n        return loss.mean()\n    \nclass MixedLoss(nn.Module):\n    def __init__(self, alpha, gamma):\n        super().__init__()\n        self.alpha = alpha\n        self.focal = FocalLoss(gamma)\n        \n    def forward(self, input, target):\n        return self.alpha * self.focal(input, target) + dice_loss(input, target)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:40.226914Z","iopub.execute_input":"2022-05-07T07:10:40.22725Z","iopub.status.idle":"2022-05-07T07:10:40.25042Z","shell.execute_reply.started":"2022-05-07T07:10:40.227208Z","shell.execute_reply":"2022-05-07T07:10:40.249378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics","metadata":{}},{"cell_type":"code","source":"def dice(input, target):\n    \"\"\"Dice metric.\"\"\"\n    input = input.argmax(dim=1).float()\n    target = target.squeeze(1).float()\n    \n    smooth = 1.\n    numerator = 2. * (input * target).sum()\n    denumerator = (input + target).sum()\n    return (numerator + smooth) / (denumerator + smooth)\n\ndef IoU(input, target):\n    input = input.argmax(dim=1).float()\n    target = target.squeeze(1).float()\n    \n    smooth = 1.\n    intersection = (input * target).sum()\n    union = (input + target).sum() - intersection\n    return (intersection + smooth) / (union + smooth)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:40.252809Z","iopub.execute_input":"2022-05-07T07:10:40.253321Z","iopub.status.idle":"2022-05-07T07:10:40.265543Z","shell.execute_reply.started":"2022-05-07T07:10:40.253279Z","shell.execute_reply":"2022-05-07T07:10:40.264874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data loader","metadata":{}},{"cell_type":"code","source":"def open_mask(fn):\n    masks = masks_df[masks_df['ImageId'] == str(os.path.split(fn)[1])]['EncodedPixels'].tolist()\n    masks = \" \".join(str(x) for x in masks) # convert list to string\n    mask_img = vision.image.open_mask_rle(masks, shape=(768, 768))\n    return vision.ImageSegment(mask_img.data.T.permute(2,0,1).float())\n     \nclass SegmentationLabelList(vision.ImageList):\n    _processor=vision.data.SegmentationProcessor\n    def __init__(self, items:basic_data.Iterator, classes:basic_data.Collection=None, **kwargs):\n        super().__init__(items, **kwargs)\n        self.copy_new.append('classes')\n        self.classes, self.loss_func = classes, layers.CrossEntropyFlat(axis=1)\n\n    def open(self, fn): return open_mask(fn)\n    def analyze_pred(self, pred, thresh:float=0.5): return pred.argmax(dim=0)[None]\n    def reconstruct(self, t:basic_data.Tensor): return vision.ImageSegment(t)\n\nclass SegmentationItemList(vision.ImageList):\n    _label_cls, _square_show_res = SegmentationLabelList, False\n    \ndef get_data(df, bs=16, img_size=(256, 256)):\n    # Do not augment since we have a large dataset anyway, only resize\n    tfms = ([vision.transform.crop_pad()], [vision.transform.crop_pad()])\n    \n    # Build DataBunch\n    return (SegmentationItemList.from_df(unique_img_ids_df, path=path_train)\n            .split_by_rand_pct(0.2)\n            .label_from_func(lambda x: x, classes=[0, 1])\n            .transform(tfms, size=img_size, tfm_y=True)\n            .add_test(vision.Path(path_test).ls(), tfm_y=False)\n            .databunch(path=data_root, bs=bs)\n            .normalize(vision.imagenet_stats))","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:40.267265Z","iopub.execute_input":"2022-05-07T07:10:40.26765Z","iopub.status.idle":"2022-05-07T07:10:40.29017Z","shell.execute_reply.started":"2022-05-07T07:10:40.267604Z","shell.execute_reply":"2022-05-07T07:10:40.28941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Build databunch\ndata = get_data(masks_df)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:40.295507Z","iopub.execute_input":"2022-05-07T07:10:40.295906Z","iopub.status.idle":"2022-05-07T07:10:41.654055Z","shell.execute_reply.started":"2022-05-07T07:10:40.295875Z","shell.execute_reply":"2022-05-07T07:10:41.653264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.show_batch()","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:41.656552Z","iopub.execute_input":"2022-05-07T07:10:41.656824Z","iopub.status.idle":"2022-05-07T07:10:44.520821Z","shell.execute_reply.started":"2022-05-07T07:10:41.656789Z","shell.execute_reply":"2022-05-07T07:10:44.517987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = data.train_ds\nx, y = ds[0]","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:44.522361Z","iopub.execute_input":"2022-05-07T07:10:44.52294Z","iopub.status.idle":"2022-05-07T07:10:44.671881Z","shell.execute_reply.started":"2022-05-07T07:10:44.522898Z","shell.execute_reply":"2022-05-07T07:10:44.671006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:44.679683Z","iopub.execute_input":"2022-05-07T07:10:44.680041Z","iopub.status.idle":"2022-05-07T07:10:44.74441Z","shell.execute_reply.started":"2022-05-07T07:10:44.680005Z","shell.execute_reply":"2022-05-07T07:10:44.743491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:44.745854Z","iopub.execute_input":"2022-05-07T07:10:44.746382Z","iopub.status.idle":"2022-05-07T07:10:44.768368Z","shell.execute_reply.started":"2022-05-07T07:10:44.746326Z","shell.execute_reply":"2022-05-07T07:10:44.767678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Build model","metadata":{}},{"cell_type":"markdown","source":"预测","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nmodel = vision.models.resnet34\nlearner = vision.unet_learner(data, model, loss_func=MixedLoss(10., 2.), metrics=[dice, IoU], self_attention=USE_SELF_ATTENTION)\nprint(learner.path)\nlearner.path = Path('../input')\nprint(learner.path)\nlearner.load('unet_5ep')\nlearner.path = Path('../input/airbus-ship-detection')\ni = 10 #1、10、11、17\nimg = learner.data.test_ds[i][0]\npre = learner.predict(img)[0]","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:44.76965Z","iopub.execute_input":"2022-05-07T07:10:44.77011Z","iopub.status.idle":"2022-05-07T07:10:47.118751Z","shell.execute_reply.started":"2022-05-07T07:10:44.770075Z","shell.execute_reply":"2022-05-07T07:10:47.11794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:47.121253Z","iopub.execute_input":"2022-05-07T07:10:47.121821Z","iopub.status.idle":"2022-05-07T07:10:47.163367Z","shell.execute_reply.started":"2022-05-07T07:10:47.121783Z","shell.execute_reply":"2022-05-07T07:10:47.162473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pre","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:47.165033Z","iopub.execute_input":"2022-05-07T07:10:47.165326Z","iopub.status.idle":"2022-05-07T07:10:47.184296Z","shell.execute_reply.started":"2022-05-07T07:10:47.165293Z","shell.execute_reply":"2022-05-07T07:10:47.183487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.show_batch()","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:47.185903Z","iopub.execute_input":"2022-05-07T07:10:47.186195Z","iopub.status.idle":"2022-05-07T07:10:49.626058Z","shell.execute_reply.started":"2022-05-07T07:10:47.186152Z","shell.execute_reply":"2022-05-07T07:10:49.625327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = vision.models.resnet50\nlearner = vision.unet_learner(data, model, loss_func=MixedLoss(10., 2.), metrics=[dice, IoU], self_attention=USE_SELF_ATTENTION)\nif USE_UNET34_AIRBUS == True:\n    learner.path = Path('../input')\n    learner.load('unet_5ep')\n    learner.path = Path('../input/airbus-ship-detection')\nlearner.model_dir = '/kaggle/working'","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:49.627224Z","iopub.execute_input":"2022-05-07T07:10:49.627605Z","iopub.status.idle":"2022-05-07T07:10:52.487331Z","shell.execute_reply.started":"2022-05-07T07:10:49.627574Z","shell.execute_reply":"2022-05-07T07:10:52.486664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = learner.data.test_ds[17][0]\npre = learner.predict(img)[0]","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:52.489264Z","iopub.execute_input":"2022-05-07T07:10:52.489687Z","iopub.status.idle":"2022-05-07T07:10:52.596177Z","shell.execute_reply.started":"2022-05-07T07:10:52.48965Z","shell.execute_reply":"2022-05-07T07:10:52.595415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:52.601442Z","iopub.execute_input":"2022-05-07T07:10:52.601725Z","iopub.status.idle":"2022-05-07T07:10:52.65983Z","shell.execute_reply.started":"2022-05-07T07:10:52.601695Z","shell.execute_reply":"2022-05-07T07:10:52.659129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pre","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:52.66107Z","iopub.execute_input":"2022-05-07T07:10:52.661539Z","iopub.status.idle":"2022-05-07T07:10:52.682418Z","shell.execute_reply.started":"2022-05-07T07:10:52.661498Z","shell.execute_reply":"2022-05-07T07:10:52.681636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# Find optimal LR\nlearner.lr_find()\nlearner.recorder.plot(suggestion=True)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:52.684532Z","iopub.execute_input":"2022-05-07T07:10:52.684836Z","iopub.status.idle":"2022-05-07T07:10:56.246555Z","shell.execute_reply.started":"2022-05-07T07:10:52.684799Z","shell.execute_reply":"2022-05-07T07:10:56.245158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train on 3 epoch\nlearner.fit_one_cycle(3, max_lr=5e-4)\nlearner.recorder.plot_losses()\nlearner.recorder.plot_lr(show_moms=True)\nlearner.save('unet50_3ep')","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.247849Z","iopub.status.idle":"2022-05-07T07:10:56.248515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train on 5 more epoch with model unfreezed\nlearner.unfreeze()\nlearner.lr_find()\nlearner.recorder.plot(suggestion=True)","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.249749Z","iopub.status.idle":"2022-05-07T07:10:56.250411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learner.fit_one_cycle(2, max_lr=slice(1e-6, 1e-5))\nlearner.recorder.plot_losses()\nlearner.recorder.plot_metrics()\nlearner.recorder.plot_lr(show_moms=True)\nlearner.save('unet50_5ep') #保存模型","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learner.show_results()","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.253403Z","iopub.status.idle":"2022-05-07T07:10:56.254048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing","metadata":{}},{"cell_type":"code","source":"def get_test_masks():\n    \n    if USE_FULL_RES_PRED == True:\n        learner.data = get_data(masks_df, img_size=(768, 768))\n\n    pred_masks = []\n    for x, y in tqdm(learner.data.test_ds):\n        _, _, output = learner.predict(x) # network output 2x256x256 or 2x768x768\n        \n        if USE_FULL_RES_PRED == False:\n            upsampler = torch.nn.Upsample(scale_factor=3, mode='bilinear', align_corners=False) # 768/256 = factor of 3\n            output = upsampler(output.unsqueeze(0)).squeeze(0)  # 2x256x256\n        \n        probs = F.softmax(output, dim=0)  # 2x256x256 or 2x768x768\n        mask_tensor = probs.argmax(dim=0)   # 256x256 or 768x768 (hot tensor)\n        \n        labels = label(mask_tensor)\n        pred_masks.append([vision.image.rle_encode((labels.T)==k) for k in np.unique(labels[labels>0])])\n    return pred_masks\n\ndef get_test_masks_opening():\n    pred_masks = []\n    for x, y in tqdm(learner.data.test_ds):\n        _, _, output = learner.predict(x) # network output 2x256x256 or 2x768x768\n\n        if USE_FULL_RES_PRED == False:\n            upsampler = torch.nn.Upsample(scale_factor=3, mode='bilinear', align_corners=False) # 768/256 = factor of 3\n            output = upsampler(output.unsqueeze(0)).squeeze(0)  # 2x256x256\n\n        probs = F.softmax(output, dim=0)  # 2x256x256 or 2x768x768\n        mask_tensor = probs.argmax(dim=0)   # 256x256 or 768x768 (hot tensor)\n\n        mask_tensor = binary_opening(mask_tensor, disk(2))\n        labels = label(mask_tensor)\n        pred_masks.append([vision.image.rle_encode((labels.T)==k) for k in np.unique(labels[labels>0])])\n    return pred_masks\n    \ndef create_submission_df(test_masks):\n    \"\"\"Create submission dataframe.\"\"\"\n    test_ids = list(map(lambda x: x.name, learner.data.test_dl.dataset.items))\n    img_masks = list(zip(test_ids, test_masks))\n    flat_img_masks = [] \n    for img, masks in img_masks:\n        if len(masks) > 0:\n            for mask in masks:\n                flat_img_masks.append([img, mask])\n        else:\n            flat_img_masks.append([img, None])\n    df = pd.DataFrame(flat_img_masks, columns=['ImageId', 'EncodedPixels'])\n    return df","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.255233Z","iopub.status.idle":"2022-05-07T07:10:56.255891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_masks = get_test_masks_opening()\ndf_submission = create_submission_df(test_masks)\ndf_submission.to_csv('submission_wo_clf.csv', header=True, index=False)\n\nfrom IPython.display import FileLink\nFileLink('submission_wo_clf.csv')","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.257048Z","iopub.status.idle":"2022-05-07T07:10:56.257701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Remove false positives with a pre-trained classifier","metadata":{}},{"cell_type":"code","source":"# Get dataframe with label\nclf_df = pd.read_csv(os.path.join('../input', 'clfairbus/clf_256_test_preds.csv'))\n\nfor i, row in clf_df.iterrows():\n    if row['Label'] == 0:\n        df_submission.loc[df_submission['ImageId'] == row['ImageId'], 'EncodedPixels'] = None","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.258861Z","iopub.status.idle":"2022-05-07T07:10:56.259496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission.to_csv('submission_w_clf.csv', header=True, index=False)\nFileLink('submission_w_clf.csv')","metadata":{"execution":{"iopub.status.busy":"2022-05-07T07:10:56.260667Z","iopub.status.idle":"2022-05-07T07:10:56.261303Z"},"trusted":true},"execution_count":null,"outputs":[]}]}