{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"!pip install git+https://github.com/qubvel/segmentation_models.pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc\nimport os\nimport random\nimport time\nimport warnings\nwarnings.simplefilter(\"ignore\")\n\n\nfrom albumentations import *\nfrom albumentations.pytorch import ToTensor\nimport cv2\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import KFold\nimport tifffile as tiff\nimport torch\nimport torch.backends.cudnn as cudnn\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom tqdm import tqdm_notebook as tqdm","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def seed_everything(seed=2**3):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(121)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fold = 0\nnfolds = 5\nreduce = 4\nsz = 256\n\nBATCH_SIZE = 16\nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nNUM_WORKERS = 4\nNUM_EPOCHS = 45\nSEED = 2020\nTH = 0.39\n\nDATA = '../input/hubmap-kidney-segmentation/test/'\nLABELS = '../input/hubmap-kidney-segmentation/train.csv'\nMASKS = '../input/hubmap-256x256/masks'\nTRAIN = '../input/hubmap-256x256/train'\ndf_sample = pd.read_csv('../input/hubmap-kidney-segmentation/sample_submission.csv')\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def rle_encode_less_memory(img):\n    pixels=img.T.flatten()\n    pixels[0]=0\n    pixels[-1]=0\n    runs = np.where(pixels[1:] != pixels[:-1])[0]+2\n    runs[1::2]-=runs[::2]\n    return ' '.join(str(x) for x in runs)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"mean = np.array([0.65459856,0.48386562,0.69428385])\nstd = np.array([0.15167958,0.23584107,0.13146145])\n\ndef img2tensor(img, dtype:np.dtype=np.float32):\n    if img.ndim==2: \n        img=np.expand_dims(img, 2)\n    img=np.transpose(img, (2, 0, 1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\nclass HuBMAPDataset(Dataset):\n    def __init__(self, fold=fold, train=True, tfms=None):\n        ids = pd.read_csv(LABELS).id.values\n        kf = KFold(n_splits=nfolds, random_state=SEED, shuffle=True)\n        ids=set(ids[list(kf.split(ids))[fold][0 if train else 1]])\n        self.fnames=[fname for fname in os.listdir(TRAIN) if fname.split('_')[0] in ids]\n        self.train = train\n        self.tfms = tfms\n    \n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        imgs=cv2.cvtColor(cv2.imread(os.path.join(TRAIN, fname)), cv2.COLOR_BGR2RGB)\n        masks=cv2.imread(os.path.join(MASKS, fname), cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented=self.tfms(image=imgs, mask=masks)\n            imgs, masks=augmented['image'], augmented['mask']\n        return img2tensor((imgs/255.0-mean)/std), img2tensor(masks)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_augmentation(p=1.0):\n    return Compose([\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            IAAPiecewiseAffine(p=0.3),\n        ], p=0.3),\n        OneOf([\n            HueSaturationValue(10, 15, 10),\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),\n        ], p=0.3),\n    ], p=p)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# ds = HuBMAPDataset(tfms=get_augmentation())\n# dl = DataLoader(ds, batch_size=16, shuffle=True, num_workers=NUM_WORKERS)\n# imgs, masks = next(iter(dl))\n# print(imgs.shape)\n# print(masks.shape)\n\n# plt.figure(figsize=(16, 16))\n# for i, (img, mask) in enumerate(zip(imgs, masks)):\n#     img = ((img.permute(1, 2, 0)*std + mean) * 255.0).numpy().astype(np.uint8)\n#     plt.subplot(4, 4, i+1)\n#     plt.imshow(img, vmin=0, vmax=255)\n#     plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n#     plt.axis('off')\n#     plt.subplots_adjust(wspace=None, hspace=None)\n# plt.show()\n\n# del ds, dl, imgs, masks","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, self).__init__()\n    \n    def forward(self, inputs, targets, smooth=1):\n        #\n        inputs = F.sigmoid(inputs)\n        #flatten\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        #element_wise production to get intersection score\n        intersection = (inputs*targets).sum()\n        dice_score = (2*intersection + smooth) / (inputs.sum() + targets.sum() + smooth)\n        return 1 - dice_score","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def UnetPlusPlus():\n    return smp.Unet(\n        encoder_name='efficientnet-b7',\n        encoder_weights='imagenet',\n        in_channels=3,\n        classes=1\n    )","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def UnetResNext():\n    return smp.Unet(\n        encoder_name='se_resnext50_32x4d',\n        encoder_weights='imagenet',\n        in_channels=3,\n        classes=1\n    )","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import segmentation_models_pytorch as smp\ndef UnetDenseNet():\n    return smp.Unet(\n    encoder_name='densenet201',\n    encoder_weights='imagenet',\n    in_channels=3,\n    classes=1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_one_epoch(fold, model, dataloader_train, dataloader_valid, optimizer, loss_function):\n    #training phase\n    model.train()\n    train_loss = 0\n    for i, (imgs, masks) in enumerate(dataloader_train):\n        optimizer.zero_grad()\n        imgs = imgs.to(DEVICE)\n        masks = masks.to(DEVICE)\n        #forward pass\n        outputs = model(imgs)\n        #cal loss and backward\n        loss = loss_function(outputs, masks)\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item()\n    train_loss /= len(dataloader_train)\n    \n    #validating phase\n    model.eval()\n    valid_loss = 0\n    with torch.no_grad():\n        for i, (imgs, masks) in enumerate(dataloader_valid):\n            imgs = imgs.to(DEVICE)\n            masks = masks.to(DEVICE)\n            outputs = model(imgs)\n            loss = loss_function(outputs, masks)\n            valid_loss += loss.item()\n    valid_loss /=len(dataloader_valid)\n    print(f'FOLD: {fold + 1}, EPOCH: {epoch + 1} - train loss: {train_loss} -  valid_loss: {valid_loss}')\n    return train_loss, valid_loss\n\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"best_valid_loss = 0\nfor fold in range(nfolds):\n    ds_t = HuBMAPDataset(fold=fold, train=True, tfms=get_augmentation())\n    ds_v = HuBMAPDataset(fold=fold, train=False)\n    dataloader_t = torch.utils.data.DataLoader(ds_t, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\n    dataloader_v = torch.utils.data.DataLoader(ds_v, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\n    model = UnetDenseNet().to(DEVICE)\n    diceloss = DiceLoss()\n    optimizer = torch.optim.Adam([\n        {'params': model.decoder.parameters(), 'lr': 1e-3},\n        {'params': model.encoder.parameters(), 'lr': 1e-3},\n    ])\n\n#     scheduler = optim.lr_scheduler.OneCycleLR(optimizer=optimizer, \n#                                               pct_start=0.1, \n#                                               div_factor=1e-3, \n#                                               max_lr=1e-2, \n#                                               epochs=NUM_EPOCHS, \n#                                               steps_per_epoch=len(dataloader_t))\n    train_loss = 0\n    valid_loss = 0\n\n    for epoch in tqdm(range(NUM_EPOCHS)):\n        train_loss, valid_loss = train_one_epoch(fold, model, dataloader_t, dataloader_v, optimizer, diceloss)\n    \n    torch.save(model.state_dict(), f'model_fold_{fold}.pth')\n    if best_valid_loss == 0:\n        best_valid_loss = valid_loss\n    if best_valid_loss >= valid_loss:\n        best_valid_loss = valid_loss\n        torch.save(model, 'best_unet_model.pth')\n    \n    gc.collect()\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}