{"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":"%%capture\n!pip install -qq torch==1.7.1+cu110 torchvision==0.8.2+cu110 torchaudio==0.7.2 -f https://download.pytorch.org/whl/torch_stable.html\n!pip install -qq git+https://github.com/qubvel/segmentation_models.pytorch\n!pip install -qq timm==0.4.12\n!pip install -qq einops\n\nfrom skimage import io, filters, transform \nimport tifffile as tiff\nimport albumentations as A\n \nimport os\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport segmentation_models_pytorch as smp","metadata":{"papermill":{"duration":13.182907,"end_time":"2022-07-12T04:01:53.673664","exception":false,"start_time":"2022-07-12T04:01:40.490757","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-16T07:24:53.481823Z","iopub.execute_input":"2022-09-16T07:24:53.482448Z","iopub.status.idle":"2022-09-16T07:26:00.776944Z","shell.execute_reply.started":"2022-09-16T07:24:53.482413Z","shell.execute_reply":"2022-09-16T07:26:00.774292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    IMAGE_SIZE = 1536\n    TRAIN_BATCH_SIZE = 1\n    VALID_BATCH_SIZE = 2*TRAIN_BATCH_SIZE\n    EPOCHS = 30\n    DEVICE = \"cuda\"\n    OUTPUT = \"/hubmap/\"\n    MODEL_PATH = \"nvidia/segformer-b5-finetuned-cityscapes-1024-1024\"\n    FOLDS = 5\n    LR = 5e-5\n    CSV_PATH = \"../input/hubmap-folds/train_unsplit_data.csv\"\n    MEAN =  [0.78036435, 0.75635034, 0.77327976]\n    STD = [0.24925208, 0.26279064, 0.258655 ] ","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:00.786166Z","iopub.execute_input":"2022-09-16T07:26:00.790806Z","iopub.status.idle":"2022-09-16T07:26:00.805805Z","shell.execute_reply.started":"2022-09-16T07:26:00.790719Z","shell.execute_reply":"2022-09-16T07:26:00.803579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:00.812291Z","iopub.execute_input":"2022-09-16T07:26:00.817105Z","iopub.status.idle":"2022-09-16T07:26:13.981167Z","shell.execute_reply.started":"2022-09-16T07:26:00.817010Z","shell.execute_reply":"2022-09-16T07:26:13.979176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport torch\nimport numpy as np\nimport tifffile as tiff\n\nfrom skimage import io, transform, filters \n\nclass HubDataset(torch.utils.data.Dataset):\n    def __init__(self, image_path, mask_path, pixel_size=None, augmentations=None):\n        self.image_path = image_path\n        self.mask_path = mask_path\n        self.augmentations = augmentations\n        self.pixel_size = pixel_size\n        \n    def __len__(self):\n        return len(self.image_path)\n    \n    def __getitem__(self,item):\n        image = tiff.imread(self.image_path[item])\n        mask = io.imread(self.mask_path[item])\n        mask = mask.reshape(mask.shape[0],mask.shape[1],1)\n        \n        image = image.astype(np.float32)/255\n        mask  = mask.astype(np.float32)/255\n\n        # s = self.pixel_size/0.4 * (config.IMAGE_SIZE/image.shape[0])\n        ## resize\n        #image = cv2.resize(image,dsize=None, fx=s,fy=s,interpolation=cv2.INTER_LINEAR)\n        #mask  = cv2.resize(mask, dsize=None, fx=s,fy=s,interpolation=cv2.INTER_LINEAR)\n        image = cv2.resize(image, dsize=(config.IMAGE_SIZE, config.IMAGE_SIZE),interpolation=cv2.INTER_LINEAR)\n        mask = cv2.resize(mask, dsize=(config.IMAGE_SIZE, config.IMAGE_SIZE),interpolation=cv2.INTER_LINEAR)\n\n        # image = image - np.min(image)\n        # image = image / np.max(image)\n        if self.augmentations is not None:\n            image, mask = self.augmentations(image, mask)\n            \n                \n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        \n        mask = mask.reshape(mask.shape[0],mask.shape[1],1)  ##  small problem soving: that adds channels to remove bug\n        mask = np.transpose(mask, (2, 0, 1)).astype(np.float32)\n        \n        return {\n            \"image\": torch.tensor(image, dtype=torch.float),\n            \"mask\" : torch.tensor(mask, dtype=torch.float),\n        }","metadata":{"papermill":{"duration":0.017041,"end_time":"2022-07-12T04:01:53.717711","exception":false,"start_time":"2022-07-12T04:01:53.70067","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-09-16T07:26:42.564810Z","iopub.execute_input":"2022-09-16T07:26:42.565245Z","iopub.status.idle":"2022-09-16T07:26:42.582518Z","shell.execute_reply.started":"2022-09-16T07:26:42.565212Z","shell.execute_reply":"2022-09-16T07:26:42.581206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_metric(_mask1,_mask2):\n    \n    batch_size = _mask1.shape[0]\n    dice_total = 0.0\n    \n    for idx in range(batch_size):\n        mask1 = _mask1[idx].reshape(_mask1[idx].shape[1],_mask1[idx].shape[2]).cpu().numpy()\n        mask2 = _mask2[idx].reshape(_mask2[idx].shape[1],_mask2[idx].shape[2]).cpu().numpy()\n        intersect = np.sum(mask1*mask2)\n        fsum = np.sum(mask1)\n        ssum = np.sum(mask2)\n        eps = 1e-7 ##  for empty masks the numerator should not be divivded by zero\n        dice = (2 * intersect + eps) / (fsum + ssum + eps)\n        dice = np.mean(dice)\n        \n        dice_total += dice\n        \n    final_dice = dice_total/batch_size\n    final_dice = round(final_dice, 4)  # for easy reading till 4 decimal places\n    return final_dice  ","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:42.988100Z","iopub.execute_input":"2022-09-16T07:26:42.988784Z","iopub.status.idle":"2022-09-16T07:26:42.998769Z","shell.execute_reply.started":"2022-09-16T07:26:42.988751Z","shell.execute_reply":"2022-09-16T07:26:42.997151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nJaccardLoss = smp.losses.JaccardLoss(mode='binary',smooth=1.0)\nDiceLoss = smp.losses.DiceLoss(mode='binary', smooth=1.0)\nFocalLoss = smp.losses.FocalLoss(mode=\"binary\")\nBCELoss     = smp.losses.SoftBCEWithLogitsLoss()\nLovaszLoss  = smp.losses.LovaszLoss(mode='binary', per_image=False)\nTverskyLoss = smp.losses.TverskyLoss(mode='binary', log_loss=False)\n\ndef criterion(y_pred, y_true):\n    return 0.5*BCELoss(y_pred, y_true) + 0.5*DiceLoss(y_pred, y_true)\n\nscaler = torch.cuda.amp.GradScaler()\n\ndef train(model,train_loader,device,optimizer):\n    model.train()\n    running_train_loss = 0.0\n    for data in train_loader:\n        inputs = data['image']\n        masks = data['mask']\n\n        ## CUDA\n        inputs = inputs.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        ## forward \n        with torch.cuda.amp.autocast():\n            outputs = model(inputs,)\n            loss = criterion(outputs, masks)\n\n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_train_loss +=loss.item()\n        \n    train_loss_value = running_train_loss/len(train_loader)\n    print(f'train custom Loss is {train_loss_value}')\n\n@torch.no_grad()    \ndef evaluation(model,valid_loader,device,):\n    model.eval()\n    running_dice_score = 0.0\n    running_val_loss = 0.0\n    for data in valid_loader:\n        inputs = data['image']\n        masks = data['mask']\n        \n        inputs = inputs.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        output = model(inputs,)\n        running_val_loss +=  criterion(output, masks)\n        \n        output = torch.sigmoid(output)\n        running_dice_score += dice_metric(masks, output)\n\n    val_loss = running_val_loss/len(valid_loader) \n    dice_score = running_dice_score/len(valid_loader)\n        \n    print(f'valid custom Loss is {val_loss}')\n        \n    return dice_score","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:43.372563Z","iopub.execute_input":"2022-09-16T07:26:43.375635Z","iopub.status.idle":"2022-09-16T07:26:43.488331Z","shell.execute_reply.started":"2022-09-16T07:26:43.375598Z","shell.execute_reply":"2022-09-16T07:26:43.486662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\n\n\ndef do_random_flip(image, mask):\n    if np.random.rand() > 0.5:\n        image = cv2.flip(image, 0)\n        mask = cv2.flip(mask, 0)\n    if np.random.rand() > 0.5:\n        image = cv2.flip(image, 1)\n        mask = cv2.flip(mask, 1)\n    if np.random.rand() > 0.5:\n        image = image.transpose(1, 0, 2)\n        mask = mask.transpose(1, 0)\n\n    image = np.ascontiguousarray(image)\n    mask = np.ascontiguousarray(mask)\n    return image, mask\n\n\ndef do_random_rot90(image, mask):\n    r = np.random.choice(\n        [\n            0,\n            cv2.ROTATE_90_CLOCKWISE,\n            cv2.ROTATE_90_COUNTERCLOCKWISE,\n            cv2.ROTATE_180,\n        ]\n    )\n    if r == 0:\n        return image, mask\n    else:\n        image = cv2.rotate(image, r)\n        mask = cv2.rotate(mask, r)\n        return image, mask\n\n\n# crop ##----\ndef do_crop(image, mask, size, xy=(0, 0)):\n    height, width = image.shape[:2]\n    x, y = xy\n    if x is None:\n        x = (width - size) // 2\n    if y is None:\n        y = (height - size) // 2\n\n    image = image[y : y + size, x : x + size]\n    mask = mask[y : y + size, x : x + size]\n    return image, mask\n\n\ndef do_random_crop(image, mask, size):\n    height, width = image.shape[:2]\n    x = np.random.choice(width - size) if width > size else 0\n    y = np.random.choice(height - size) if height > size else 0\n    image = image[y : y + size, x : x + size]\n    mask = mask[y : y + size, x : x + size]\n    return image, mask\n\n\n# transform ##----\ndef do_random_rotate_scale(image, mask, angle=30, scale=[0.8, 1.2]):\n    angle = np.random.uniform(-angle, angle)\n    scale = np.random.uniform(*scale) if scale is not None else 1\n\n    height, width = image.shape[:2]\n    center = (height // 2, width // 2)\n\n    transform = cv2.getRotationMatrix2D(center, angle, scale)\n    image = cv2.warpAffine(\n        image,\n        transform,\n        (width, height),\n        flags=cv2.INTER_LINEAR,\n        borderMode=cv2.BORDER_CONSTANT,\n        borderValue=(0, 0, 0),\n    )\n    mask = cv2.warpAffine(\n        mask,\n        transform,\n        (width, height),\n        flags=cv2.INTER_LINEAR,\n        borderMode=cv2.BORDER_CONSTANT,\n        borderValue=0,\n    )\n    return image, mask\n\n\n# noise\ndef do_random_noise(image, mask, mag=0.1):\n    height, width = image.shape[:2]\n    noise = np.random.uniform(-1, 1, (height, width, 1)) * mag\n    image = image + noise\n    image = np.clip(image, 0, 1)\n    return image, mask\n\n\n# https://openreview.net/pdf?id=rkBBChjiG\n# <todo> mixup/cutout\n\n# intensity\ndef do_random_contast(image, mask, mag=0.3):\n    alpha = 1 + np.random.uniform(-1, 1) * mag\n    image = image * alpha\n    image = np.clip(image, 0, 1)\n    return image, mask\n\n\ndef do_random_hsv(image, mask, mag=[0.15, 0.25, 0.25]):\n    image = (image * 255).astype(np.uint8)\n    hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)\n\n    h = hsv[:, :, 0].astype(np.float32)  # hue\n    s = hsv[:, :, 1].astype(np.float32)  # saturation\n    v = hsv[:, :, 2].astype(np.float32)  # value\n    h = (h * (1 + np.random.uniform(-1, 1) * mag[0])) % 180\n    s = s * (1 + np.random.uniform(-1, 1) * mag[1])\n    v = v * (1 + np.random.uniform(-1, 1) * mag[2])\n\n    hsv[:, :, 0] = np.clip(h, 0, 180).astype(np.uint8)\n    hsv[:, :, 1] = np.clip(s, 0, 255).astype(np.uint8)\n    hsv[:, :, 2] = np.clip(v, 0, 255).astype(np.uint8)\n    image = cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)\n    image = image.astype(np.float32) / 255\n    return image, mask\n\n\ndef do_gray(image, mask):\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)\n    image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)\n    return image, mask\n\n\ndef valid_augment5(image, mask):\n    # image, mask  = do_crop(image, mask, image_size, xy=(None,None))\n    return image, mask\n\n\ndef train_augment5a(image, mask):\n\n    image, mask = do_random_flip(image, mask)\n    image, mask = do_random_rot90(image, mask)\n\n    for fn in np.random.choice(\n        [\n            lambda image, mask: (image, mask),\n            lambda image, mask: do_random_noise(image, mask, mag=0.1),\n            lambda image, mask: do_random_contast(image, mask, mag=0.25),\n            lambda image, mask: do_random_hsv(image, mask, mag=[0.30, 0.30, 0]),\n        ],\n        2,\n    ):\n        image, mask = fn(image, mask)\n\n    for fn in np.random.choice(\n        [\n            lambda image, mask: (image, mask),\n            lambda image, mask: do_random_rotate_scale(\n                image, mask, angle=45, scale=[0.5, 2]\n            ),\n        ],\n        1,\n    ):\n        image, mask = fn(image, mask)\n\n    return image, mask\n\n\ndef train_augment5b(image, mask):\n    image, mask = do_random_flip(image, mask)\n    image, mask = do_random_rot90(image, mask)\n\n    for fn in np.random.choice(\n        [\n            lambda image, mask: (image, mask),\n            lambda image, mask: do_random_noise(image, mask, mag=0.1),\n            lambda image, mask: do_random_contast(image, mask, mag=0.40),\n            lambda image, mask: do_random_hsv(image, mask, mag=[0.40, 0.40, 0]),\n        ],\n        2,\n    ):\n        image, mask = fn(image, mask)\n\n    for fn in np.random.choice(\n        [\n            lambda image, mask: (image, mask),\n            lambda image, mask: do_random_rotate_scale(\n                image, mask, angle=45, scale=[0.50, 2.0]\n            ),\n        ],\n        1,\n    ):\n        image, mask = fn(image, mask)\n\n    return image, mask","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:43.767207Z","iopub.execute_input":"2022-09-16T07:26:43.768734Z","iopub.status.idle":"2022-09-16T07:26:43.807594Z","shell.execute_reply.started":"2022-09-16T07:26:43.768601Z","shell.execute_reply":"2022-09-16T07:26:43.805805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/hubmap-coat/')\n\nfrom coat import *\nfrom daformer import *\nfrom helper import *","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:44.166039Z","iopub.execute_input":"2022-09-16T07:26:44.166441Z","iopub.status.idle":"2022-09-16T07:26:44.198195Z","shell.execute_reply.started":"2022-09-16T07:26:44.166408Z","shell.execute_reply":"2022-09-16T07:26:44.196905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\nclass MixUpSample(nn.Module):\n    def __init__(self, scale_factor=4):\n        super().__init__()\n        self.mixing = nn.Parameter(torch.tensor(0.5))\n        self.scale_factor = scale_factor\n\n    def forward(self, x):\n        x = self.mixing * F.interpolate(\n            x, scale_factor=self.scale_factor, mode=\"bilinear\", align_corners=False\n        ) + (1 - self.mixing) * F.interpolate(\n            x, scale_factor=self.scale_factor, mode=\"nearest\"\n        )\n        return x\n    \nclass Net(nn.Module):\n    \n    def __init__(self,\n                 encoder=coat_lite_medium,\n                 decoder=daformer_conv3x3,\n                 encoder_cfg={},\n                 decoder_cfg={},\n                 ):\n        \n        super(Net, self).__init__()\n        decoder_dim = decoder_cfg.get('decoder_dim', 320)\n\n        self.encoder = encoder\n\n        self.rgb = RGB()\n\n        encoder_dim = self.encoder.embed_dims\n        # [64, 128, 320, 512]\n\n        self.decoder = decoder(\n            encoder_dim=encoder_dim,\n            decoder_dim=decoder_dim,\n        )\n#         self.logit = nn.Sequential(\n#             nn.Conv2d(decoder_dim, 1, kernel_size=1),\n#             nn.Upsample(scale_factor = 4, mode='bilinear', align_corners=False),\n#         )\n        self.logit = nn.Conv2d(decoder_dim, 1, kernel_size=1)\n        self.mixup = MixUpSample()\n    def forward(self, x):\n\n        x = self.rgb(x)\n\n        B, C, H, W = x.shape\n        encoder = self.encoder(x)\n\n        last, decoder = self.decoder(encoder)\n        logits = self.logit(last)\n        \n        upsampled_logits = self.mixup(logits)\n        \n        return upsampled_logits","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:44.565541Z","iopub.execute_input":"2022-09-16T07:26:44.566023Z","iopub.status.idle":"2022-09-16T07:26:44.581282Z","shell.execute_reply.started":"2022-09-16T07:26:44.565974Z","shell.execute_reply":"2022-09-16T07:26:44.579554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ### encoder try other notebook to change the Parallel blocks class\n# class coat_parallel_small_plus1 (CoaT):\n#     def __init__(self, **kwargs):\n#         super(coat_parallel_small_plus1, self).__init__(\n#             patch_size=4,\n#             embed_dims=[152, 320, 320, 320, 320],\n#             serial_depths=[2, 2, 2, 2, 2],\n#             parallel_depth=6,\n#             num_heads=8,\n#             mlp_ratios=[4, 4, 4, 4, 4],\n#             pretrain ='coat_small_7479cf9b.pth',\n#             **kwargs)","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:46.740073Z","iopub.execute_input":"2022-09-16T07:26:46.740536Z","iopub.status.idle":"2022-09-16T07:26:46.746867Z","shell.execute_reply.started":"2022-09-16T07:26:46.740462Z","shell.execute_reply":"2022-09-16T07:26:46.745229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def HubmapModel():\n    encoder = coat_lite_medium()\n    checkpoint = '../input/hubmap-coat-medium/coat_lite_medium_384x384_f9129688.pth'\n    checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)\n    state_dict = checkpoint['model']\n    encoder.load_state_dict(state_dict,strict=False)\n    \n    net = Net(encoder=encoder).cuda()\n    \n    return net","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:47.456588Z","iopub.execute_input":"2022-09-16T07:26:47.457058Z","iopub.status.idle":"2022-09-16T07:26:47.464686Z","shell.execute_reply.started":"2022-09-16T07:26:47.457008Z","shell.execute_reply":"2022-09-16T07:26:47.463058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## MAIN","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\n\n\n\n## PRINT CONFIG ##\nprint(\"configuration :\")\nprint(f\" -- FOLDS : {config.FOLDS}\")\nprint(f\" -- MODEL : {config.MODEL_PATH}\")\nprint(f\" -- LR : {config.LR}\")\nprint(f\" -- TRAIN_BATCH_SIZE  : {config.TRAIN_BATCH_SIZE}\")\nprint(f\" -- VALID_BATCH_SIZE  : {config.VALID_BATCH_SIZE}\")\nprint(f\" -- EPOCHS  : {config.EPOCHS}\")\nprint(f\" -- CSV_PATH  : {config.CSV_PATH}\")\n\n\ndf = pd.read_csv(config.CSV_PATH)\n\nprint(f\"read csv- - - {config.CSV_PATH}\")\n\nfor fold in {0}:\n\n    best_score = 0.0\n\n    model = HubmapModel()\n    model.to(\"cuda\")\n\n    df_train = df[df.fold != fold].reset_index(drop=True)\n    df_valid = df[df.fold == fold].reset_index(drop=True)\n\n    df_train = df_train.drop(columns=\"fold\")\n    df_valid = df_valid.drop(columns=\"fold\")\n\n    train_ids = df_train.id.values.tolist()\n    valid_ids = df_valid.id.values.tolist()\n\n    train_images = [\n        os.path.join(\n            \"../input/hubmap-organ-segmentation/train_images\", str(i) + \".tiff\"\n        )\n        for i in train_ids\n    ]\n    train_masks = [\n        os.path.join(\n            \"../input/hubmap-hpa-2022-maskdataset/hubmap_2022_MaskDataset\",\n            str(i) + \".png\",\n        )\n        for i in train_ids\n    ]\n\n    valid_images = [\n        os.path.join(\n            \"../input/hubmap-organ-segmentation/train_images\", str(i) + \".tiff\"\n        )\n        for i in valid_ids\n    ]\n    valid_masks = [\n        os.path.join(\n            \"../input/hubmap-hpa-2022-maskdataset/hubmap_2022_MaskDataset\",\n            str(i) + \".png\",\n        )\n        for i in valid_ids\n    ]\n\n    train_dataset = HubDataset(\n        image_path=train_images, mask_path=train_masks, augmentations=train_augment5a\n    )\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset, batch_size=config.TRAIN_BATCH_SIZE, shuffle=True, pin_memory=True\n    )\n    valid_dataset = HubDataset(\n        image_path=valid_images, mask_path=valid_masks, augmentations=valid_augment5\n    )\n    valid_loader = torch.utils.data.DataLoader(\n        valid_dataset,\n        batch_size=config.VALID_BATCH_SIZE,\n        shuffle=False,\n        pin_memory=True,\n    )\n\n    optimizer = torch.optim.Adam(\n        model.parameters(),\n        lr=config.LR,\n    )\n    scheduler = torch.optim.lr_scheduler.ExponentialLR(\n        optimizer, gamma=0.95, verbose=True\n    )\n\n    print(\n        f\"============================== FOLD -- {fold} ==============================\"\n    )\n\n    for epoch in range(config.EPOCHS):\n        print(f\"==================== Epoch -- {epoch} ====================\")\n        train(\n            model=model,\n            train_loader=train_loader,\n            device=config.DEVICE,\n            optimizer=optimizer,\n        )\n        scheduler.step()\n        dice_score = evaluation(\n            model=model,\n            valid_loader=valid_loader,\n            device=config.DEVICE,\n        )\n\n        print(f\"validation DICE Metric={dice_score}\")\n\n        if dice_score > best_score:\n            best_score = dice_score\n            torch.save(model.state_dict(), \"model-\" + str(fold) + \".pth\")","metadata":{"execution":{"iopub.status.busy":"2022-09-16T07:26:49.833641Z","iopub.execute_input":"2022-09-16T07:26:49.835828Z"},"trusted":true},"execution_count":null,"outputs":[]}]}