{"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":"code","source":"# Libraries\nimport os\nfrom os.path import join\nfrom tqdm import tqdm\nimport random\n\nimport numpy as np\nimport pandas as pd\npd.set_option('display.max_rows', 500)\npd.set_option('display.max_columns', 500)\npd.set_option('display.width', 1000)\nimport matplotlib.pyplot as plt\nplt.rcParams.update({'font.size': 18})\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset, sampler\n\nfrom albumentations import (HorizontalFlip, VerticalFlip, ShiftScaleRotate, Normalize, Resize, Compose, GaussNoise)\nfrom albumentations.pytorch import ToTensorV2\n\ndef initialize_seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    \ninitialize_seeds(2021)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-11-17T08:55:53.197715Z","iopub.execute_input":"2021-11-17T08:55:53.197987Z","iopub.status.idle":"2021-11-17T08:55:53.20851Z","shell.execute_reply.started":"2021-11-17T08:55:53.197957Z","shell.execute_reply":"2021-11-17T08:55:53.207701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load Dataset\n\n### Understand the Structure of the Dataset\n\n\n*   train - train images in PNG format\n*   The training annotations -> run length encoded masks\n*   Images -> PNG format (The number of images is small, but the number of annotated objects is quite high.)\n*   Test set -> 240 images\n\n### Files\n\n**train.csv** - IDs and masks for all training objects. None of this metadata is provided for the test set.\n* id - unique identifier for object\n* annotation - run length encoded pixels for the identified neuronal cell\n* width - source image width\n* height - source image height\n* cell_type - the cell line\n* plate_time - time plate was created\n* sample_date - date sample was created\n* sample_id - sample identifier\n* elapsed_timedelta - time since first image taken of sample\n\n**sample_submission.csv** - a sample submission file in the correct format\n\n**train** - train images in PNG format\n\n**test** - test images in PNG format. Only a few test set images are available for download; the remainder can only be accessed by your notebooks when you submit.\n\n**train_semi_supervised** - unlabeled images offered in case you want to use additional data for a semi-supervised approach.\n\n**LIVECell_dataset_2021** - A mirror of the data from the LIVECell dataset. LIVECell is the predecessor dataset to this competition. You will find extra data for the SH-SHY5Y cell line, plus several other cell lines not covered in the competition dataset that may be of interest for transfer learning.","metadata":{}},{"cell_type":"code","source":"DATA_PATH             = '../input/sartorius-cell-instance-segmentation'\nSAMPLE_SUBMISSION     = join(DATA_PATH,'train')\nTRAIN_CSV             = join(DATA_PATH,'train.csv')\nTRAIN_PATH            = join(DATA_PATH,'train')\nTEST_PATH             = join(DATA_PATH,'test')\n\ndf_train = pd.read_csv(TRAIN_CSV)\nprint(f'Training Set Shape: {df_train.shape} - {df_train[\"id\"].nunique()} \\\nImages - Memory Usage: {df_train.memory_usage().sum() / 1024 ** 2:.2f} MB')","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:53.210387Z","iopub.execute_input":"2021-11-17T08:55:53.211194Z","iopub.status.idle":"2021-11-17T08:55:53.626633Z","shell.execute_reply.started":"2021-11-17T08:55:53.211159Z","shell.execute_reply":"2021-11-17T08:55:53.625652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(mask_rle, shape, color=1):\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    '''\n    s = mask_rle.split()\n    # print(s)\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.float32)\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n    return img.reshape(shape)\n\ndef build_masks(df_train, image_id, input_shape):\n    height, width = input_shape\n    labels = df_train[df_train[\"id\"] == image_id][\"annotation\"].tolist()\n    mask = np.zeros((height, width))\n    for label in labels:\n        mask += rle_decode(label, shape=(height, width))\n    mask = mask.clip(0, 1)\n    return np.array(mask)","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:53.629383Z","iopub.execute_input":"2021-11-17T08:55:53.629688Z","iopub.status.idle":"2021-11-17T08:55:53.643537Z","shell.execute_reply.started":"2021-11-17T08:55:53.629631Z","shell.execute_reply":"2021-11-17T08:55:53.642443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training dataset & data loader","metadata":{}},{"cell_type":"code","source":"class CellDataset(Dataset):\n    def __init__(self, df: pd.core.frame.DataFrame, train:bool):\n        self.IMAGE_RESIZE = (224, 224)\n        self.RESNET_MEAN = (0.485, 0.456, 0.406)\n        self.RESNET_STD = (0.229, 0.224, 0.225)\n        self.df = df\n        self.base_path = TRAIN_PATH\n        self.gb = self.df.groupby('id')\n        self.transforms = Compose([Resize(self.IMAGE_RESIZE[0],  self.IMAGE_RESIZE[1]), \n                                   Normalize(mean=self.RESNET_MEAN, std= self.RESNET_STD, p=1), \n                                   HorizontalFlip(p=0.5),\n                                   VerticalFlip(p=0.5)])\n        \n        # Split train and val set\n        all_image_ids = np.array(df_train.id.unique())\n        print(len(all_image_ids))\n        np.random.seed(42)\n        iperm = np.random.permutation(len(all_image_ids))\n        num_train_samples = int(len(all_image_ids) * 0.9)\n        print(num_train_samples)\n\n        if train:\n            self.image_ids = all_image_ids[iperm[:num_train_samples]]\n        else:\n             self.image_ids = all_image_ids[iperm[num_train_samples:]] \n\n    def __getitem__(self, idx: int) -> dict:\n\n        image_id = self.image_ids[idx]\n        df = self.gb.get_group(image_id)\n\n        # Read image\n        image_path = os.path.join(self.base_path, image_id + \".png\")\n        image = cv2.imread(image_path)\n\n        # Create the mask\n        mask = build_masks(df_train, image_id, input_shape=(520, 704))\n        mask = (mask >= 1).astype('float32')\n        augmented = self.transforms(image=image, mask=mask)\n        image = augmented['image']\n        mask = augmented['mask']\n        # print(np.moveaxis(image,0,2).shape)\n        return np.moveaxis(np.array(image),2,0), mask.reshape((1, self.IMAGE_RESIZE[0], self.IMAGE_RESIZE[1]))\n\n\n    def __len__(self):\n        return len(self.image_ids)","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:53.645751Z","iopub.execute_input":"2021-11-17T08:55:53.646784Z","iopub.status.idle":"2021-11-17T08:55:53.666021Z","shell.execute_reply.started":"2021-11-17T08:55:53.646743Z","shell.execute_reply":"2021-11-17T08:55:53.665104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = CellDataset(df_train, train=True)\ndl_train = DataLoader(ds_train, batch_size=16, num_workers=2, pin_memory=True, shuffle=False)\n\nds_val = CellDataset(df_train, train=False)\ndl_val = DataLoader(ds_val, batch_size=4, num_workers=2, pin_memory=True, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:53.667212Z","iopub.execute_input":"2021-11-17T08:55:53.667689Z","iopub.status.idle":"2021-11-17T08:55:53.698928Z","shell.execute_reply.started":"2021-11-17T08:55:53.667654Z","shell.execute_reply":"2021-11-17T08:55:53.698116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(dl_train))","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:53.702564Z","iopub.execute_input":"2021-11-17T08:55:53.7028Z","iopub.status.idle":"2021-11-17T08:55:53.711766Z","shell.execute_reply.started":"2021-11-17T08:55:53.70277Z","shell.execute_reply":"2021-11-17T08:55:53.710719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot simages and mask from dataloader\nbatch = next(iter(dl_train))\nimages, masks = batch\nprint(f\"image shape: {images.shape},\\nmask shape:{masks.shape},\\nbatch len: {len(batch)}\")\n\nplt.figure(figsize=(10, 5))\n\nplt.subplot(1, 3, 1)\nplt.imshow(images[1][1])\nplt.title('Original image')\n\nplt.subplot( 1, 3, 2)\nplt.imshow(masks[1][0])\nplt.title('Mask')\n\nplt.subplot( 1, 3, 3)\nplt.imshow(images[1][1])\nplt.imshow(masks[1][0],alpha=0.2)\nplt.title('Both')\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:53.715864Z","iopub.execute_input":"2021-11-17T08:55:53.716138Z","iopub.status.idle":"2021-11-17T08:55:57.242919Z","shell.execute_reply.started":"2021-11-17T08:55:53.716103Z","shell.execute_reply":"2021-11-17T08:55:57.242223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Unet Model","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, inChannel, outChannel):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential(nn.Conv2d(inChannel, outChannel, 3, padding=1),\n                                 nn.BatchNorm2d(outChannel),\n                                 nn.ReLU(inplace=True),\n                                 nn.Conv2d(outChannel, outChannel, 3, padding=1),\n                                 nn.BatchNorm2d(outChannel),\n                                 nn.ReLU(inplace=True))\n\n    def forward(self, x):\n        x = self.conv(x)\n        return x\n    \nclass Down(nn.Module):\n    \n    def __init__(self, inChannel, outChannel):\n        super(Down, self).__init__()\n        self.conv = nn.Sequential(nn.MaxPool2d(kernel_size=2),\n                                 DoubleConv(inChannel, outChannel))\n\n    def forward(self, x):\n        x = self.conv(x)\n        return x\n    \nclass Up(nn.Module):\n    \n    def __init__(self, inChannel, outChannel):\n        super(Up, self).__init__()\n        self.upsample = nn.Sequential(nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),\n                                      nn.Conv2d(inChannel, inChannel//2, kernel_size=3, padding=1),\n                                      nn.ReLU(inplace=True))\n        self.conv = DoubleConv(inChannel, outChannel)\n        \n    def forward(self, x, skipX):\n        # 转置卷积\n        x = self.upsample(x)\n        x = torch.cat((x, skipX), dim=1)\n        x = self.conv(x)\n        return x\n    \n    \nclass Unet(nn.Module):\n    \n    def __init__(self):\n        super(Unet, self).__init__()\n        self.donw1 = DoubleConv(3, 64)\n        self.donw2 = Down(64, 128)\n        self.donw3 = Down(128, 256)\n        self.donw4 = Down(256, 512)\n        self.donw5 = Down(512, 1024)\n        self.up4 = Up(1024, 512)\n        self.up3 = Up(512, 256)\n        self.up2 = Up(256, 128)\n        self.up1 = Up(128, 64)\n        self.oneMult = nn.Conv2d(64, 1, kernel_size=1)\n    \n    def forward(self, x):\n        \n        # x: (batchSize, channel, h, w)\n        # 5次下采样，4次上采样，oneMult1*1卷积，不改变size改变通道数\n        x1 = self.donw1(x)\n        x2 = self.donw2(x1)\n        x3 = self.donw3(x2)\n        x4 = self.donw4(x3)\n        x = self.donw5(x4)\n        x = self.up4(x, x4)\n        x = self.up3(x, x3)\n        x = self.up2(x, x2)\n        x = self.up1(x, x1)\n        x = self.oneMult(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:57.244514Z","iopub.execute_input":"2021-11-17T08:55:57.244968Z","iopub.status.idle":"2021-11-17T08:55:57.262313Z","shell.execute_reply.started":"2021-11-17T08:55:57.244929Z","shell.execute_reply":"2021-11-17T08:55:57.261672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define the network","metadata":{}},{"cell_type":"code","source":"# put on GPU here if you have it\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n#net.to(device);  # remove semi-colon to see net structure\nprint(f\"The device is {device}!!\")","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:57.266254Z","iopub.execute_input":"2021-11-17T08:55:57.266535Z","iopub.status.idle":"2021-11-17T08:55:57.277857Z","shell.execute_reply.started":"2021-11-17T08:55:57.266501Z","shell.execute_reply":"2021-11-17T08:55:57.277192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Unet()\nmodel.to(device)\n\n# 保存整个模型\n# torch.save(model, path_model)\n\n# 保存模型参数\n# net_state_dict = model.state_dict()\n# torch.save(net_state_dict, path_state_dict)\n\n# 损失函数&优化器\ncriterion = nn.BCEWithLogitsLoss().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4)","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:57.27892Z","iopub.execute_input":"2021-11-17T08:55:57.279669Z","iopub.status.idle":"2021-11-17T08:55:57.586853Z","shell.execute_reply.started":"2021-11-17T08:55:57.279629Z","shell.execute_reply":"2021-11-17T08:55:57.586093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train the network","metadata":{}},{"cell_type":"code","source":"# 早停函数\nclass EarlyStopping:\n    \"\"\"Early stops the training if validation loss doesn't improve after a given patience.\"\"\"\n    def __init__(self, patience=7, verbose=False, delta=0): # 当验证集损失在连续7次训练周期中都没有得到降低时，停止模型训练，以防止模型过拟合\n        \"\"\"\n        Args:\n            patience (int): How long to wait after last time validation loss improved.\n                            Default: 7\n            verbose (bool): If True, prints a message for each validation loss improvement. \n                            Default: False\n            delta (float): Minimum change in the monitored quantity to qualify as an improvement.\n                            Default: 0\n        \"\"\"\n        self.patience = patience\n        self.verbose = verbose\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.val_loss_min = np.Inf\n        self.delta = delta\n\n    def __call__(self, val_loss, model):\n\n        score = -val_loss\n\n        if self.best_score is None:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n        elif score < self.best_score + self.delta:\n            self.counter += 1\n            print(f'EarlyStopping counter: {self.counter} out of {self.patience}')\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n            self.counter = 0\n\n    def save_checkpoint(self, val_loss, model):\n        '''Saves model when validation loss decrease.'''\n        if self.verbose:\n            print(f'Validation loss decreased ({self.val_loss_min:.6f} --> {val_loss:.6f}).  Saving model ...')\n        torch.save(model.state_dict(), 'checkpoint.pt')\t# 这里会存储迄今最优模型的参数\n        self.val_loss_min = val_loss","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:57.588071Z","iopub.execute_input":"2021-11-17T08:55:57.588464Z","iopub.status.idle":"2021-11-17T08:55:57.59867Z","shell.execute_reply.started":"2021-11-17T08:55:57.588428Z","shell.execute_reply":"2021-11-17T08:55:57.597938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_curve = list()\nvalid_curve = list()\n\ndef train(loader, model, criterion, optimizer):\n    model.train()\n    runningLoss = 0\n    with tqdm(total=len(dl_train)) as tq:\n        for i, (img, target) in enumerate(loader):\n            img = img.to(device)\n            target = target.to(device)\n            output = model(img)\n            loss = criterion(output, target)\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n            runningLoss += loss.item() * img.size(0)\n            tq.update(1)\n        train_curve.append(loss.item()) \n            \n    return runningLoss / len(loader.dataset)\n\ndef valid(loader, model, criterion):\n    model.eval()\n    runningLoss = 0\n    totalIou = 0\n    with torch.no_grad():\n        with tqdm(total=len(loader)) as tq:\n            for i, (img, target) in enumerate(loader):\n                img = img.to(device)\n                target = target.to(device)\n                output = model(img)\n                loss = criterion(output, target)\n                runningLoss += loss.item() * img.size(0)\n                \n                iou = computeIOU(output, target)\n                totalIou += iou * img.size(0)\n                tq.update(1)\n            valid_curve.append(loss.item()) \n                \n    return runningLoss / len(loader.dataset), totalIou / len(loader.dataset)\n    \ndef computeIOU(output, mask):\n    pred = torch.zeros(output.size()).cuda()\n    pred[output > 0] = 1\n    pred = pred.to(torch.uint8)\n    mask = mask.to(torch.uint8)\n    intersection = (pred & mask)\n    union = (pred | mask)\n    return torch.sum(intersection).item() / torch.sum(union).item()","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:57.599891Z","iopub.execute_input":"2021-11-17T08:55:57.600409Z","iopub.status.idle":"2021-11-17T08:55:57.616722Z","shell.execute_reply.started":"2021-11-17T08:55:57.600371Z","shell.execute_reply":"2021-11-17T08:55:57.615963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lowestLoss = None\n# checkpoint_interval = 5\npatience = 7\nearly_stopping = EarlyStopping(patience, verbose=True)\n\nfor epoch in range(50):# 以设置大一些，希望通过 early stopping 来结束模型训练\n    trainLoss = train(dl_train, model, criterion, optimizer)\n    validLoss, iou = valid(dl_val, model, criterion)\n    print('Epoch ：{}'.format(epoch))\n    print('Train Loss : {}    Valid Loss : {}'.format(trainLoss, validLoss))\n    print('IOU : {}'.format(iou))\n    \n    if  lowestLoss is None or validLoss < lowestLoss:\n        lowestLoss = validLoss\n        torch.save({'model_state_dict': model.state_dict(),\n                   'trainLoss': trainLoss,\n                   'validLoss': validLoss,\n                   'optimizer': optimizer,\n                   'iou': iou},\n                  f'./checkpoint.pth')\n    \n    early_stopping(validLoss, model)\n    # 若满足早停要求\n    if early_stopping.early_stop:\n        print(\"Early stopping\")\n        # 结束模型训练\n        break\n    \ntrain_x = range(len(train_curve))\ntrain_y = train_curve\n\ntrain_iters = len(dl_train)\n# valid_x = np.arange(1, len(valid_curve)+1) * train_iters # 由于valid中记录的是epochloss，需要对记录点进行转换到iterations\nvalid_y = valid_curve\n\nplt.plot(train_x, train_y, label='Train')\nplt.plot(train_x, valid_y, label='Valid')\n\nplt.legend(loc='upper right')\nplt.ylabel('loss value')\nplt.xlabel('Iteration')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-11-17T08:55:57.618827Z","iopub.execute_input":"2021-11-17T08:55:57.619324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# checkpoint为字典型\n# model.load_state_dict(torch.load('./checkpoint.pth'))\ntorch.load('./checkpoint.pt')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test dataset & dataloader","metadata":{}},{"cell_type":"code","source":"class TestCellDataset(Dataset):\n    def __init__(self):\n        self.test_path = TEST_PATH\n        self.IMAGE_RESIZE = (224, 224)\n        self.RESNET_MEAN = (0.485, 0.456, 0.406)\n        self.RESNET_STD = (0.229, 0.224, 0.225)\n        \n        # I am not sure if they adapt the sample submission csv or only the test folder\n        # I am using the test folders as the ground truth for the images to predict, which should be always right\n        # The sample csv is ignored\n        self.image_ids = [f[:-4]for f in os.listdir(self.test_path)]\n        self.num_samples = len(self.image_ids)\n        self.transform = Compose([Resize(self.IMAGE_RESIZE[0], self.IMAGE_RESIZE[1]), Normalize(mean=self.RESNET_MEAN, std=self.RESNET_STD, p=1), ToTensorV2()])\n\n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        path = os.path.join(self.test_path, image_id + \".png\")\n        image = cv2.imread(path)\n        image = self.transform(image=image)['image']\n        return {'image': image, 'id': image_id}\n\n    def __len__(self):\n        return self.num_samples","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del dl_train, ds_train, optimizer","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_test = TestCellDataset()\ndl_test = DataLoader(ds_test, batch_size=16, shuffle=False, num_workers=2, pin_memory=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Post processsing","metadata":{}},{"cell_type":"code","source":"# separate different components of the prediction mask\ndef post_process(probability, threshold=0.5, min_size=300):\n    mask = cv2.threshold(probability, threshold, 1, cv2.THRESH_BINARY)[1]\n    num_component, component = cv2.connectedComponents(mask.astype(np.uint8))\n    predictions = []\n    for c in range(1, num_component):\n        p = (component == c)\n        if p.sum() > min_size:\n            a_prediction = np.zeros((520, 704), np.float32)\n            a_prediction[p] = 1\n            predictions.append(a_prediction)\n    return predictions\n\ndef rle_encoding(x):\n    dots = np.where(x.flatten() == 1)[0]\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if (b>prev+1): run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return ' '.join(map(str, run_lengths))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_is_run_length(mask_rle):\n    if not mask_rle:\n        return True\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    start_prev = starts[0]\n    ok = True\n    for start in starts[1:]:\n        ok = ok and start > start_prev\n        start_prev = start\n        if not ok:\n            return False\n    return True\n\ndef create_empty_submission():\n    fs = os.listdir(\"../input/sartorius-cell-instance-segmentation/test\")\n    df = pd.DataFrame([(f[:-4], \"\") for f in fs], columns=['id', 'predicted'])\n    df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict & submit","metadata":{}},{"cell_type":"code","source":"model.eval()\n\nsubmission = []\nfor i, batch in enumerate(tqdm(dl_test)):\n    preds = torch.sigmoid(model(batch['image'].cuda()))\n    preds = preds.detach().cpu().numpy()[:, 0, :, :] # (batch_size, 1, size, size) -> (batch_size, size, size)\n    for image_id, probability_mask in zip(batch['id'], preds):\n        try:\n            #if probability_mask.shape != IMAGE_RESIZE:\n            #    probability_mask = cv2.resize(probability_mask, dsize=IMAGE_RESIZE, interpolation=cv2.INTER_LINEAR)\n            probability_mask = cv2.resize(probability_mask, dsize=(704, 520), interpolation=cv2.INTER_LINEAR)\n            predictions = post_process(probability_mask)\n            for prediction in predictions:\n                #plt.imshow(prediction)\n                #plt.show()\n                try:\n                    submission.append((image_id, rle_encoding(prediction)))\n                except:\n                    print(\"Error in RL encoding\")\n        except Exception as e:\n            print(f\"Exception for img: {image_id}: {e}\")\n        \n        # Fill images with no predictions\n        image_ids = [image_id for image_id, preds in submission]\n        if image_id not in image_ids:\n            submission.append((image_id, \"\"))\n            \ndf_submission = pd.DataFrame(submission, columns=['id', 'predicted'])\ndf_submission.to_csv('submission.csv', index=False)\n\nif df_submission['predicted'].apply(check_is_run_length).mean() != 1:\n    print(\"Check run lenght failed\")\n    create_empty_submission()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}