{"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":"## Imports","metadata":{}},{"cell_type":"code","source":"!pip install -qq git+https://github.com/qubvel/segmentation_models.pytorch\n!pip install -qq timm==0.4.12\n!pip install -qq einops","metadata":{"execution":{"iopub.status.busy":"2022-09-12T17:38:49.414713Z","iopub.execute_input":"2022-09-12T17:38:49.415672Z","iopub.status.idle":"2022-09-12T17:39:30.762467Z","shell.execute_reply.started":"2022-09-12T17:38:49.415622Z","shell.execute_reply":"2022-09-12T17:39:30.761090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport sys\nimport cv2\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nfrom tqdm.notebook import tqdm\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\nimport gc\nfrom sklearn.model_selection import StratifiedKFold\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport segmentation_models_pytorch as smp","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-12T17:39:30.764902Z","iopub.execute_input":"2022-09-12T17:39:30.765301Z","iopub.status.idle":"2022-09-12T17:39:30.773005Z","shell.execute_reply.started":"2022-09-12T17:39:30.765255Z","shell.execute_reply":"2022-09-12T17:39:30.771698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    epochs = 50\n    batch_size = 1\n    n_folds = 5\n    lr = 1e-6\n    img_size = 768\ndevice = ('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-09-12T17:39:30.774369Z","iopub.execute_input":"2022-09-12T17:39:30.774830Z","iopub.status.idle":"2022-09-12T17:39:30.784806Z","shell.execute_reply.started":"2022-09-12T17:39:30.774792Z","shell.execute_reply":"2022-09-12T17:39:30.783641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper Functions","metadata":{}},{"cell_type":"code","source":"def set_seed(seed = 42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print('> SEEDING DONE')\n    \nset_seed(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-09-12T17:39:30.788950Z","iopub.execute_input":"2022-09-12T17:39:30.789648Z","iopub.status.idle":"2022-09-12T17:39:30.799324Z","shell.execute_reply.started":"2022-09-12T17:39:30.789621Z","shell.execute_reply":"2022-09-12T17:39:30.798246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle2mask(mask_rle: str, shape=None, label: int = 0):\n    \"\"\"\n    mask_rle: run-length as string formatted (start length)\n    shape: (height,width) of array to return\n    Returns numpy array, 1 - mask, 0 - background\n\n    \"\"\"\n    rle = np.array(list(map(int, mask_rle.split())))\n    labels = np.zeros(shape).flatten()\n    \n    for start, end in zip(rle[::2], rle[1::2]):\n        labels[start:start+end] = label\n#         labels[start:start+end] = 1\n\n    return labels.reshape(shape).T  # Needed to align to RLE direction\n\n\ndef mask_to_rle(mask):\n    #Rescale image to original size\n    size = int(len(mask.flatten())**.5)\n    n = Image.fromarray(mask.reshape((size, size))*255.0)\n    n = np.array(n).astype(np.float32)\n    #Get pixels to flatten\n    pixels = n.T.flatten()\n    #Round the pixels using the half of the range of pixel value\n    pixels = (pixels-min(pixels) > ((max(pixels)-min(pixels))/2)).astype(int)\n    pixels = np.nan_to_num(pixels) #incase of zero-div-error\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0]\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:01.063969Z","iopub.execute_input":"2022-09-12T18:35:01.064337Z","iopub.status.idle":"2022-09-12T18:35:01.075292Z","shell.execute_reply.started":"2022-09-12T18:35:01.064305Z","shell.execute_reply":"2022-09-12T18:35:01.073891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_IMG_PATH = \"../input/hubmap-organ-segmentation/train_images\"\nTEST_IMG_PATH = \"../input/hubmap-organ-segmentation/test_images\"\n\nTRAIN_IMG_ANNOTATIONS = \"../input/hubmap-organ-segmentation/train_annotations\"\nTRAIN_IMG_INFO = \"../input/hubmap-organ-segmentation/train.csv\"","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:01.320808Z","iopub.execute_input":"2022-09-12T18:35:01.321158Z","iopub.status.idle":"2022-09-12T18:35:01.326181Z","shell.execute_reply.started":"2022-09-12T18:35:01.321128Z","shell.execute_reply":"2022-09-12T18:35:01.324945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(TRAIN_IMG_INFO)\ntrain.head(2)","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:01.551445Z","iopub.execute_input":"2022-09-12T18:35:01.552417Z","iopub.status.idle":"2022-09-12T18:35:01.939470Z","shell.execute_reply.started":"2022-09-12T18:35:01.552371Z","shell.execute_reply":"2022-09-12T18:35:01.938450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"organs = ['prostate', 'spleen', 'lung', 'kidney', 'largeintestine']\norgan_annotations = {}\nfor i, organ in enumerate(organs):\n    organ_annotations[organ] = i + 1","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:01.941457Z","iopub.execute_input":"2022-09-12T18:35:01.941861Z","iopub.status.idle":"2022-09-12T18:35:01.946903Z","shell.execute_reply.started":"2022-09-12T18:35:01.941824Z","shell.execute_reply":"2022-09-12T18:35:01.945964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\nimgs = []\nfnames = {}\n\nfor i, img in enumerate(tqdm(os.listdir(TRAIN_IMG_PATH))):\n    path = os.path.join(TRAIN_IMG_PATH, img)\n    img_number = img.split(\".\")[0]\n    imgs.append(path)\n    fnames[img_number] = path","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:02.001327Z","iopub.execute_input":"2022-09-12T18:35:02.001635Z","iopub.status.idle":"2022-09-12T18:35:02.142458Z","shell.execute_reply.started":"2022-09-12T18:35:02.001607Z","shell.execute_reply":"2022-09-12T18:35:02.141525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transformer(stage):\n    if stage == \"train\":\n        return A.Compose([\n            A.augmentations.crops.RandomResizedCrop(height=CFG.img_size, width=CFG.img_size),\n            A.HorizontalFlip(),\n            A.VerticalFlip(),\n            A.RandomRotate90(),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, \n                             border_mode=cv2.BORDER_REFLECT),\n            A.OneOf([\n                A.OpticalDistortion(p=0.3),\n                A.GridDistortion(p=.1),\n                A.PiecewiseAffine(p=0.3),\n            ], p=0.5),\n            A.OneOf([\n                A.HueSaturationValue(10,15,10),\n                A.CLAHE(clip_limit=2),\n                A.RandomBrightnessContrast(), \n            ], p=0.5),\n            A.Normalize()\n        ])\n    else:\n        return A.Compose([\n                A.Resize(CFG.img_size, CFG.img_size),\n                A.Normalize() \n            ])","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:02.528431Z","iopub.execute_input":"2022-09-12T18:35:02.529069Z","iopub.status.idle":"2022-09-12T18:35:02.539625Z","shell.execute_reply.started":"2022-09-12T18:35:02.529034Z","shell.execute_reply":"2022-09-12T18:35:02.538546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_df(df):\n    fig,ax = plt.subplots(1,2,figsize=(15,5))\n    ax[0].plot(df['train_loss'])\n    ax[0].plot(df['valid_loss'])\n    ax[0].legend()\n    ax[0].set_title('Loss')\n    ax[1].plot(df['train_dice'])\n    ax[1].plot(df['valid_dice'])\n    ax[1].legend()\n    ax[1].set_title('Dice')","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:02.577333Z","iopub.execute_input":"2022-09-12T18:35:02.577676Z","iopub.status.idle":"2022-09-12T18:35:02.583694Z","shell.execute_reply.started":"2022-09-12T18:35:02.577647Z","shell.execute_reply":"2022-09-12T18:35:02.582637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CV","metadata":{}},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=CFG.n_folds)\ntrain[\"fold\"] = -1\nfor fold, (train_idx, val_idx) in enumerate(skf.split(train, train[\"organ\"])):\n    train.loc[val_idx, \"fold\"] = fold\ntrain.groupby(\"fold\")[\"organ\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:02.674439Z","iopub.execute_input":"2022-09-12T18:35:02.674793Z","iopub.status.idle":"2022-09-12T18:35:02.710127Z","shell.execute_reply.started":"2022-09-12T18:35:02.674763Z","shell.execute_reply":"2022-09-12T18:35:02.709003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class HubmapDataset(Dataset):\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.ids = df.id.values\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, index):\n        img_id = self.ids[index]\n        img = fnames[str(img_id)]\n        \n        rle = self.df[self.df['id'] == img_id]['rle'].values[0]\n        height = self.df[self.df['id'] == img_id]['img_height'].values[0]\n        width = self.df[self.df['id'] == img_id]['img_width'].values[0]\n        organ = self.df[self.df['id'] == img_id]['organ'].values[0]\n        \n        img = np.asarray(Image.open(img))\n        mask = rle2mask(rle, shape=(height, width), label=organ_annotations[organ])\n        \n        if self.transforms is not None:\n            transformed = self.transforms(image=img, mask=np.expand_dims(mask, axis=2))\n            img, mask = transformed[\"image\"], transformed[\"mask\"]\n\n        if len(mask.shape) > 2:\n            mask = np.squeeze(mask, 2)\n        \n        return np.transpose(img, (2, 0, 1)), mask","metadata":{"execution":{"iopub.status.busy":"2022-09-12T18:35:02.726771Z","iopub.execute_input":"2022-09-12T18:35:02.727117Z","iopub.status.idle":"2022-09-12T18:35:02.737717Z","shell.execute_reply.started":"2022-09-12T18:35:02.727085Z","shell.execute_reply":"2022-09-12T18:35:02.736547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomLoss(nn.Module):\n    def __init__(self):\n        super(CustomLoss,self).__init__()\n        self.diceloss = smp.losses.DiceLoss(mode='binary')\n        self.binloss = smp.losses.SoftBCEWithLogitsLoss(reduction = 'mean' , smooth_factor = 0.1)\n        self.lovloss = smp.losses.LovaszLoss(mode='binary')\n        \n    def forward(self, output, mask, aux_losses):\n        dice = self.diceloss(output, mask)\n        bce = self.binloss(output, mask)\n        lov = self.lovloss(output, mask)\n        loss = dice\n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-09-12T20:06:34.464104Z","iopub.execute_input":"2022-09-12T20:06:34.464478Z","iopub.status.idle":"2022-09-12T20:06:34.471189Z","shell.execute_reply.started":"2022-09-12T20:06:34.464443Z","shell.execute_reply":"2022-09-12T20:06:34.470157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceCoef(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super().__init__()\n\n    def forward(self, y_pred, y_true, smooth=1.):\n        y_true = y_true.view(-1)\n        y_pred = y_pred.view(-1)\n        \n        #Round off y_pred\n        y_pred = torch.round((y_pred - y_pred.min()) / (y_pred.max() - y_pred.min()))\n        \n        intersection = (y_true * y_pred).sum()\n        dice = (2.0*intersection + smooth)/(y_true.sum() + y_pred.sum() + smooth)\n        \n        return dice","metadata":{"execution":{"iopub.status.busy":"2022-09-12T20:06:36.506198Z","iopub.execute_input":"2022-09-12T20:06:36.506588Z","iopub.status.idle":"2022-09-12T20:06:36.514224Z","shell.execute_reply.started":"2022-09-12T20:06:36.506554Z","shell.execute_reply":"2022-09-12T20:06:36.513193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def criterion_aux_loss(logit, mask):\n    mask = mask.reshape(mask.shape[0], 1, mask.shape[1], mask.shape[2])\n    mask = F.interpolate(mask, size=logit.shape[-2:], mode='nearest')\n    loss = F.binary_cross_entropy_with_logits(logit, mask)\n    return loss","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:33.848641Z","iopub.execute_input":"2022-09-12T19:43:33.849317Z","iopub.status.idle":"2022-09-12T19:43:33.855207Z","shell.execute_reply.started":"2022-09-12T19:43:33.849282Z","shell.execute_reply":"2022-09-12T19:43:33.854077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"sys.path.append('../input/hubmap-coat/')\nsys.path.append(\"/kaggle/working/\")\nsys.path.append(\"../input/hubmap-submit-06/\")\nfrom coat import *\nfrom daformer import *\nfrom helper import *\nfrom pvt_v2 import *\nfrom segmentation_models_pytorch.decoders.unet.decoder import UnetDecoder","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:34.149765Z","iopub.execute_input":"2022-09-12T19:43:34.150451Z","iopub.status.idle":"2022-09-12T19:43:34.158188Z","shell.execute_reply.started":"2022-09-12T19:43:34.150415Z","shell.execute_reply":"2022-09-12T19:43:34.155038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def CustomCoat(**kwargs):\n    model = CoaT(\n            patch_size=4,\n            embed_dims=[152, 320, 320, 320],\n            serial_depths=[2, 2, 2, 2],\n            parallel_depth=6,\n            num_heads=8,\n            mlp_ratios=[4, 4, 4, 4],\n            **kwargs)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:34.316366Z","iopub.execute_input":"2022-09-12T19:43:34.317061Z","iopub.status.idle":"2022-09-12T19:43:34.322830Z","shell.execute_reply.started":"2022-09-12T19:43:34.317024Z","shell.execute_reply":"2022-09-12T19:43:34.321581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RGB(nn.Module):\n    IMAGE_RGB_MEAN = [0.485, 0.456, 0.406] #[0.5, 0.5, 0.5]\n    IMAGE_RGB_STD  = [0.229, 0.224, 0.225] #[0.5, 0.5, 0.5]\n \n    def __init__(self,):\n        super(RGB, self).__init__()\n        self.register_buffer('mean', torch.zeros(1,3,1,1))\n        self.register_buffer('std', torch.ones(1,3,1,1))\n        self.mean.data = torch.FloatTensor(self.IMAGE_RGB_MEAN).view(self.mean.shape)\n        self.std.data = torch.FloatTensor(self.IMAGE_RGB_STD).view(self.std.shape)\n\n    def forward(self, x):\n        x = (x-self.mean)/self.std\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:34.476157Z","iopub.execute_input":"2022-09-12T19:43:34.477133Z","iopub.status.idle":"2022-09-12T19:43:34.485294Z","shell.execute_reply.started":"2022-09-12T19:43:34.477084Z","shell.execute_reply":"2022-09-12T19:43:34.484307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubmapModel(nn.Module):\n    def __init__(self, \n                 encoder=CustomCoat(),\n                 decoder=daformer_conv3x3,\n                 encoder_cfg={},\n                 decoder_cfg={}):\n        \n        super(HubmapModel, self).__init__()\n        \n        self.rgb = RGB()\n        \n        self.encoder = encoder\n        encoder_dim = self.encoder.embed_dims\n        decoder_dim = [256, 128, 64, 32, 16]\n        \n        self.output_type = ['loss', 'inference']\n    \n        conv_dim = 32\n        self.conv = nn.Sequential(\n            nn.Conv2d(3, 32, kernel_size=3, stride=2, padding=1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, 32, kernel_size=3, stride=1, padding=1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, conv_dim, kernel_size=3, stride=1, padding=1, bias=False)\n        )\n        \n        self.decoder = UnetDecoder(\n            encoder_channels=[0, conv_dim] + encoder_dim,\n            decoder_channels=decoder_dim,\n            n_blocks=5,\n            use_batchnorm=True,\n            center=False,\n            attention_type=None,\n        )\n        \n        self.logit = nn.Sequential(\n            nn.Conv2d(decoder_dim[-1], 1, kernel_size=1),\n        )\n        \n        self.aux = nn.ModuleList([\n            nn.Conv2d(encoder_dim[i], 1, kernel_size=1, padding=0) \n            for i in range(len(encoder_dim))\n        ])\n        \n        self.upsample = nn.ModuleList([\n            nn.Upsample(scale_factor=4*pow(2, i), mode='bilinear', align_corners=False)\n            for i in range(len(encoder_dim))\n        ])\n        \n    def forward(self, img, mask):\n        x = self.rgb(img)\n        \n        encoder = self.encoder(x)\n        a, out1, out2, out3 = encoder\n        \n        conv = self.conv(x)\n\n        if 1:\n            feature = encoder[::-1]  # reverse channels to start from head of encoder\n            head = feature[0]\n            skip = feature[1:] + [conv, None]\n            d = self.decoder.center(head)\n\n            decoder = []\n            for i, decoder_block in enumerate(self.decoder.blocks):\n                s = skip[i]\n                d = decoder_block(d, s)\n                decoder.append(d)\n            last = d\n#         print('decoder',[f.shape for f in decoder])\n        \n        logit = self.logit(last)\n        output = {}\n        \n        if 'loss' in self.output_type:\n            logit = logit.reshape(img.shape[0], CFG.img_size, CFG.img_size)\n            output['label_loss'] = F.binary_cross_entropy_with_logits(logit, mask)\n            for i in range(len(self.aux)):   \n                out = self.aux[i](encoder[i])\n                output[f'aux{i}_loss'] = criterion_aux_loss(out, mask)\n\n        if 'inference' in self.output_type:\n            probability_from_logit = torch.sigmoid(logit)\n            output['probability'] = probability_from_logit\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:34.636585Z","iopub.execute_input":"2022-09-12T19:43:34.637806Z","iopub.status.idle":"2022-09-12T19:43:35.041031Z","shell.execute_reply.started":"2022-09-12T19:43:34.637759Z","shell.execute_reply":"2022-09-12T19:43:35.040019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:35.043101Z","iopub.execute_input":"2022-09-12T19:43:35.043462Z","iopub.status.idle":"2022-09-12T19:43:35.214501Z","shell.execute_reply.started":"2022-09-12T19:43:35.043426Z","shell.execute_reply":"2022-09-12T19:43:35.213378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_optimizer_and_scheduler(dataloader):\n    optimizer = torch.optim.Adam([\n            {'params': model.decoder.parameters(), 'lr': 5e-5}, \n            {'params': model.encoder.parameters(), 'lr': 8e-5},  \n        ])\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=1e-4, epochs=CFG.epochs, steps_per_epoch=len(dataloader))\n    return optimizer, scheduler","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:35.216802Z","iopub.execute_input":"2022-09-12T19:43:35.217207Z","iopub.status.idle":"2022-09-12T19:43:35.225105Z","shell.execute_reply.started":"2022-09-12T19:43:35.217168Z","shell.execute_reply":"2022-09-12T19:43:35.224164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(dataloader, model, scheduler, optimizer):\n    model.train()\n    total_loss = 0\n    total_score = 0\n    pbar = tqdm(dataloader, total=len(dataloader))\n    \n    for img, mask in pbar:\n        optimizer.zero_grad()\n        img = img.to(device)\n        mask = mask.to(device)\n\n        outputs = model(img, mask)\n        aux_losses = [outputs[x] for x in outputs.keys() if 'aux' in x]\n        outputs = outputs['probability']\n        outputs = outputs.reshape(img.shape[0], CFG.img_size, CFG.img_size)\n        \n        loss = loss_func(outputs, mask, aux_losses)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n\n        epoch_loss = loss.item()\n        epoch_score = dice_coe(outputs, mask).item()\n        \n        pbar.set_postfix({\"loss\": epoch_loss, \"dice score\": epoch_score})\n        total_loss += epoch_loss\n        total_score += epoch_score\n    \n    total_loss /= len(dataloader)\n    total_score /= len(dataloader)\n    gc.collect()\n    torch.cuda.empty_cache()\n    return total_loss, total_score","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:35.318079Z","iopub.execute_input":"2022-09-12T19:43:35.318417Z","iopub.status.idle":"2022-09-12T19:43:35.328265Z","shell.execute_reply.started":"2022-09-12T19:43:35.318387Z","shell.execute_reply":"2022-09-12T19:43:35.327212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_one_epoch(dataloader, model):\n    model.eval()\n    with torch.no_grad():\n        total_loss = 0\n        total_score = 0\n        pbar = tqdm(dataloader, total=len(dataloader))\n\n        for img, mask in pbar:\n            img = img.to(device)\n            mask = mask.to(device)\n\n            outputs = model(img, mask)\n            aux_losses = [outputs[x] for x in outputs.keys() if 'aux' in x]\n            outputs = outputs['probability']\n            outputs = outputs.reshape(img.shape[0], CFG.img_size, CFG.img_size)\n\n            loss = loss_func(outputs, mask, aux_losses)\n            epoch_loss = loss.item()\n            epoch_score = dice_coe(outputs, mask).item()\n\n            pbar.set_postfix({\"loss\": epoch_loss, \"dice score\": epoch_score})\n            total_loss += epoch_loss\n            total_score += epoch_score\n\n        total_loss /= len(dataloader)\n        total_score /= len(dataloader)\n        gc.collect()\n        torch.cuda.empty_cache()\n        return total_loss, total_score","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:35.693101Z","iopub.execute_input":"2022-09-12T19:43:35.694022Z","iopub.status.idle":"2022-09-12T19:43:35.703997Z","shell.execute_reply.started":"2022-09-12T19:43:35.693972Z","shell.execute_reply":"2022-09-12T19:43:35.702708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fnc(train_dataloader, valid_dataloader, model, fold, optimizer, scheduler):\n    train_losses = []\n    train_scores = []\n    valid_losses = []\n    valid_scores = []\n    \n    best_loss = 999\n    best_score = -1\n    for epoch in range(CFG.epochs):\n        train_loss, train_score = train_one_epoch(train_dataloader, model, scheduler, optimizer)\n        valid_loss, valid_score = valid_one_epoch(valid_dataloader, model)\n        \n        train_losses.append(train_loss)\n        train_scores.append(train_score)\n        valid_losses.append(valid_loss)\n        valid_scores.append(valid_score)\n        \n        print(f\"-------- Epoch {epoch + 1} --------\")\n        print(\"Train Loss: \", train_loss)\n        print(\"Train Score: \", train_score)\n        print(\"Valid Loss: \", valid_loss)\n        print(\"Valid Score: \", valid_score)\n        \n        if valid_score > best_score:\n            best_score = valid_score\n            torch.save(model.state_dict(), f\"fold{fold}_best_score.pth\")\n            print(\"New Best Score\")\n        \n        if valid_loss < best_loss:\n            best_loss = valid_loss\n            torch.save(model.state_dict(), f\"fold{fold}_best_loss.pth\")\n            print(\"New Best Loss\")\n        print()\n        \n    column_names = ['train_loss','valid_loss','train_dice','valid_dice']\n    df = pd.DataFrame(np.stack([train_losses, valid_losses, train_scores, valid_scores],\n                               axis=1),columns=column_names)\n    display(df)\n    plot_df(df)","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:36.029810Z","iopub.execute_input":"2022-09-12T19:43:36.030178Z","iopub.status.idle":"2022-09-12T19:43:36.041327Z","shell.execute_reply.started":"2022-09-12T19:43:36.030148Z","shell.execute_reply":"2022-09-12T19:43:36.040185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = HubmapModel().to(device)\nloss_func = CustomLoss()\ndice_coe = DiceCoef()\n\nfor fold in range(4, 5):\n    print(\"*\"*10, f\"Fold: {fold}\", \"*\"*10)\n    train_df = train[train[\"fold\"] != fold]\n    valid_df = train[train[\"fold\"] == fold]\n    \n    train_dataset = HubmapDataset(train_df, transformer(\"train\"))\n    valid_dataset = HubmapDataset(valid_df, transformer(\"valid\"))\n    \n    train_dataloader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True)\n    valid_dataloader = DataLoader(valid_dataset, batch_size=CFG.batch_size, shuffle=False)\n\n    optimizer, scheduler = get_optimizer_and_scheduler(train_dataloader)\n    \n    train_fnc(train_dataloader, valid_dataloader, model, fold, optimizer, scheduler)","metadata":{"execution":{"iopub.status.busy":"2022-09-12T19:43:36.304493Z","iopub.execute_input":"2022-09-12T19:43:36.305097Z","iopub.status.idle":"2022-09-12T20:02:12.180746Z","shell.execute_reply.started":"2022-09-12T19:43:36.305040Z","shell.execute_reply":"2022-09-12T20:02:12.179240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}