{"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":"# About this notebook\n- **Version 1 :**\n    + PyTorch CoAT starter code\n    + 5 folds\n    + OneCycleLR scheduler\n    \n\n- **Version 3 :**\n    + PyTorch CoAT Model starter code :\n    + fixing bugs\n    + 5 folds\n    + OneCycleLR scheduler\n\nIf this notebook is helpful, feel free to upvote :)","metadata":{}},{"cell_type":"code","source":"#!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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-17T00:57:45.317453Z","iopub.execute_input":"2022-10-17T00:57:45.317941Z","iopub.status.idle":"2022-10-17T00:57:45.340498Z","shell.execute_reply.started":"2022-10-17T00:57:45.317827Z","shell.execute_reply":"2022-10-17T00:57:45.339530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install segmentation-models-pytorch","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-17T00:57:45.342487Z","iopub.execute_input":"2022-10-17T00:57:45.343129Z","iopub.status.idle":"2022-10-17T00:57:45.347912Z","shell.execute_reply.started":"2022-10-17T00:57:45.343093Z","shell.execute_reply":"2022-10-17T00:57:45.346672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install staintools\n#!pip install spams","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:57:45.464639Z","iopub.execute_input":"2022-10-17T00:57:45.465057Z","iopub.status.idle":"2022-10-17T00:57:45.469931Z","shell.execute_reply.started":"2022-10-17T00:57:45.465017Z","shell.execute_reply":"2022-10-17T00:57:45.468983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/staintools-offline/spams-2.6.5.4-cp37-cp37m-linux_x86_64.whl\n!pip install /kaggle/input/staintools-offline/staintools-2.1.2-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:57:45.954786Z","iopub.execute_input":"2022-10-17T00:57:45.955250Z","iopub.status.idle":"2022-10-17T00:58:06.831082Z","shell.execute_reply.started":"2022-10-17T00:57:45.955186Z","shell.execute_reply":"2022-10-17T00:58:06.829955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"../input/segmentation-models-pytorch\")\nimport segmentation_models_pytorch as smp\n\nprint(f\"Segmentation Models version: {smp.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:06.833475Z","iopub.execute_input":"2022-10-17T00:58:06.833865Z","iopub.status.idle":"2022-10-17T00:58:12.321915Z","shell.execute_reply.started":"2022-10-17T00:58:06.833827Z","shell.execute_reply":"2022-10-17T00:58:12.320927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"from fastai.vision.all import *\n\nimport os\n\nimport pandas as pd\n\nfrom matplotlib import pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:12.325016Z","iopub.execute_input":"2022-10-17T00:58:12.326226Z","iopub.status.idle":"2022-10-17T00:58:12.913105Z","shell.execute_reply.started":"2022-10-17T00:58:12.326163Z","shell.execute_reply":"2022-10-17T00:58:12.912143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(os.listdir('../input/hubmap-hpa-2022-maskdataset/hubmap_2022_MaskDataset'))","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-17T00:58:12.915877Z","iopub.execute_input":"2022-10-17T00:58:12.917582Z","iopub.status.idle":"2022-10-17T00:58:13.006360Z","shell.execute_reply.started":"2022-10-17T00:58:12.917552Z","shell.execute_reply":"2022-10-17T00:58:13.005402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/hubmap-organ-segmentation/train.csv')\ntest = pd.read_csv('../input/hubmap-organ-segmentation/test.csv')\nTRAIN = '../input/hubmap-2022-256x256-stain-normalization/train'\nTEST_PATH = '../input/hubmap-organ-segmentation/test_images/'\nMASKS = '../input/hubmap-2022-256x256/masks/'\nLABELS = '../input/hubmap-organ-segmentation/train.csv'\ndisplay(train.head())\ndisplay(test.head())","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:13.007800Z","iopub.execute_input":"2022-10-17T00:58:13.008152Z","iopub.status.idle":"2022-10-17T00:58:13.368343Z","shell.execute_reply.started":"2022-10-17T00:58:13.008118Z","shell.execute_reply":"2022-10-17T00:58:13.367402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(os.listdir(TRAIN))","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:13.369777Z","iopub.execute_input":"2022-10-17T00:58:13.370590Z","iopub.status.idle":"2022-10-17T00:58:13.596146Z","shell.execute_reply.started":"2022-10-17T00:58:13.370552Z","shell.execute_reply":"2022-10-17T00:58:13.595113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Vahadane Stain Normalization","metadata":{}},{"cell_type":"code","source":"import spams\nimport cv2 as cv\nimport spams\nimport matplotlib.pyplot as plt\nimport tifffile as tiff\nimport numpy as np\nimport staintools","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:13.597480Z","iopub.execute_input":"2022-10-17T00:58:13.597932Z","iopub.status.idle":"2022-10-17T00:58:14.150756Z","shell.execute_reply.started":"2022-10-17T00:58:13.597893Z","shell.execute_reply":"2022-10-17T00:58:14.149794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_image_cv(path):\n    \n    img = cv2.imread(path)\n    im = cv.cvtColor(im, cv.COLOR_BGR2RGB)\n\n    return im\n\ndef read_image(path):\n    \"\"\"\n    Read an image to RGB uint8\n    :param path:\n    :return im:\n    \"\"\"\n    im = tiff.imread(path)\n    # im = cv.imread(path)\n    # im = cv.cvtColor(im, cv.COLOR_BGR2RGB)\n    \n    return im \n\ndef show_colors(C):\n    \"\"\"\n    Shows rows of C as colors (RGB)\n    :param C:\n    :return:\n    \"\"\"\n    n = C.shape[0]\n    for i in range(n):\n        if C[i].max() > 1.0:\n            plt.plot([0, 1], [n - 1 - i, n - 1 - i], c=C[i] / 255, linewidth=20)\n        else:\n            plt.plot([0, 1], [n - 1 - i, n - 1 - i], c=C[i], linewidth=20)\n        plt.axis('off')\n        plt.axis([0, 1, -1, n])\n\ndef show(image, now=True, fig_size=(10, 10)):\n    \"\"\"\n    Show an image (np.array).\n    Caution! Rescales image to be in range [0,1].\n    :param image:\n    :param now:\n    :param fig_size:\n    :return:\n    \"\"\"\n    image = image.astype(np.float32)\n    m, M = image.min(), image.max()\n    if fig_size != None:\n        plt.rcParams['figure.figsize'] = (fig_size[0], fig_size[1])\n    plt.imshow((image - m) / (M - m), cmap='gray')\n    plt.axis('off')\n    if now == True:\n        plt.show()\n\ndef build_stack(tup):\n    \"\"\"\n    Build a stack of images from a tuple of images\n    :param tup:\n    :return:\n    \"\"\"\n    N = len(tup)\n    if len(tup[0].shape) == 3:\n        h, w, c = tup[0].shape\n        stack = np.zeros((N, h, w, c))\n    if len(tup[0].shape) == 2:\n        h, w = tup[0].shape\n        stack = np.zeros((N, h, w))\n    for i in range(N):\n        stack[i] = tup[i]\n    return stack\n\ndef patch_grid(ims, width=5, sub_sample=None, rand=False, save_name=None):\n    \"\"\"\n    Display a grid of patches\n    :param ims:\n    :param width:\n    :param sub_sample:\n    :param rand:\n    :return:\n    \"\"\"\n    N0 = np.shape(ims)[0]\n    if sub_sample == None:\n        N = N0\n        stack = ims\n    elif sub_sample != None and rand == False:\n        N = sub_sample\n        stack = ims[:N]\n    elif sub_sample != None and rand == True:\n        N = sub_sample\n        idx = np.random.choice(range(N), sub_sample, replace=False)\n        stack = ims[idx]\n    height = np.ceil(float(N) / width).astype(np.uint16)\n    plt.rcParams['figure.figsize'] = (18, (18 / width) * height)\n    plt.figure()\n    for i in range(N):\n        plt.subplot(height, width, i + 1)\n        im = stack[i]\n        show(im, now=False, fig_size=None)\n    if save_name != None:\n        plt.savefig(save_name)\n    plt.show()\n\ndef standardize_brightness(I):\n    \"\"\"\n    :param I:\n    :return:\n    \"\"\"\n    p = np.percentile(I, 90)\n    return np.clip(I * 255.0 / p, 0, 255).astype(np.uint8)\n\n\ndef remove_zeros(I):\n    \"\"\"\n    Remove zeros, replace with 1's.\n    :param I: uint8 array\n    :return:\n    \"\"\"\n    mask = (I == 0)\n    I[mask] = 1\n    return I\n\n\ndef RGB_to_OD(I):\n    \"\"\"\n    Convert from RGB to optical density\n    :param I:\n    :return:\n    \"\"\"\n    I = remove_zeros(I)\n    return -1 * np.log(I / 255)\n\n\ndef OD_to_RGB(OD):\n    \"\"\"\n    Convert from optical density to RGB\n    :param OD:\n    :return:\n    \"\"\"\n    return (255 * np.exp(-1 * OD)).astype(np.uint8)\n\n\ndef normalize_rows(A):\n    \"\"\"\n    Normalize rows of an array\n    :param A:\n    :return:\n    \"\"\"\n    return A / np.linalg.norm(A, axis=1)[:, None]\n\ndef notwhite_mask(I, thresh=0.8):\n    \"\"\"\n    Get a binary mask where true denotes 'not white'\n    :param I:\n    :param thresh:\n    :return:\n    \"\"\"\n    I_LAB = cv.cvtColor(I, cv.COLOR_RGB2LAB)\n    L = I_LAB[:, :, 0] / 255.0\n    return (L < thresh)\n\n\ndef sign(x):\n    \"\"\"\n    Returns the sign of x\n    :param x:\n    :return:\n    \"\"\"\n    if x > 0:\n        return +1\n    elif x < 0:\n        return -1\n    elif x == 0:\n        return 0\n\ndef get_concentrations(I, stain_matrix, lamda=0.01):\n    \"\"\"\n    Get concentrations, a npix x 2 matrix\n    :param I:\n    :param stain_matrix: a 2x3 stain matrix\n    :return:\n    \"\"\"\n    OD = RGB_to_OD(I).reshape((-1, 3))\n    return spams.lasso(OD.T, D=stain_matrix.T, mode=2, lambda1=lamda, pos=True).toarray().T\n\ndef get_stain_matrix(I, threshold=0.8, lamda=0.1):\n    \"\"\"\n    Get 2x3 stain matrix. First row H and second row E\n    :param I:\n    :param threshold:\n    :param lamda:\n    :return:\n    \"\"\"\n    mask = notwhite_mask(I, thresh=threshold).reshape((-1,))\n    OD = RGB_to_OD(I).reshape((-1, 3))\n    OD = OD[mask]\n    dictionary = spams.trainDL(OD.T, K=2, lambda1=lamda, mode=2, modeD=0, posAlpha=True, posD=True, verbose=False).T\n    if dictionary[0, 0] < dictionary[1, 0]:\n        dictionary = dictionary[[1, 0], :]\n    dictionary = normalize_rows(dictionary)\n    return dictionary\n\n##########################\n\nclass normalizer(object):\n    \"\"\"\n    A stain normalization object\n    \"\"\"\n\n    def __init__(self):\n        self.stain_matrix_target = None\n\n    def fit(self, target):\n        target = standardize_brightness(target)\n        self.stain_matrix_target = get_stain_matrix(target)\n\n    def target_stains(self):\n        return OD_to_RGB(self.stain_matrix_target)\n\n    def transform(self, I):\n        I = standardize_brightness(I)\n        stain_matrix_source = get_stain_matrix(I)\n        source_concentrations = get_concentrations(I, stain_matrix_source)\n        return (255 * np.exp(-1 * np.dot(source_concentrations, self.stain_matrix_target).reshape(I.shape))).astype(\n            np.uint8)\n\n    def hematoxylin(self, I):\n        I = standardize_brightness(I)\n        h, w, c = I.shape\n        stain_matrix_source = get_stain_matrix(I)\n        source_concentrations = get_concentrations(I, stain_matrix_source)\n        H = source_concentrations[:, 0].reshape(h, w)\n        H = np.exp(-1 * H)\n        return H\n\n\ndef assure_path_exists(path):\n    dir = os.path.dirname(path)\n    if not os.path.exists(dir):\n        try:\n            os.makedirs(dir)\n        except OSError as e:\n            if e.errno != errno.EEXIS:\n                raise\n\ndef stain_normalize(image, target):\n    '''\n    image:  Image to transform the stain of.\n    target: Target stain to be applied on the image. (RGB colored image) \n    '''\n    stain_norm = normalizer()\n    stain_norm.fit(target)\n    transformed_image = stain_norm.transform(image)\n    \n    return transformed_image","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:14.152451Z","iopub.execute_input":"2022-10-17T00:58:14.153038Z","iopub.status.idle":"2022-10-17T00:58:14.183641Z","shell.execute_reply.started":"2022-10-17T00:58:14.152994Z","shell.execute_reply":"2022-10-17T00:58:14.182648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntarget = tiff.imread(\"/kaggle/input/hubmap-organ-segmentation/test_images/10078.tiff\")\nplt.imshow(target)\nplt.show()\n\nto_transform = tiff.imread(\"/kaggle/input/hubmap-organ-segmentation/train_images/10392.tiff\")\nplt.imshow(to_transform)\nplt.show()\n\ntransformed2 = stain_normalize(to_transform, target)\nplt.imshow(transformed2)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:14.185235Z","iopub.execute_input":"2022-10-17T00:58:14.185629Z","iopub.status.idle":"2022-10-17T00:58:28.189355Z","shell.execute_reply.started":"2022-10-17T00:58:14.185591Z","shell.execute_reply":"2022-10-17T00:58:28.188399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport glob\nimport random\nimport time\nimport math\nimport pathlib\n\nimport torch\nimport torch.nn as nn\nimport albumentations as A\nimport skimage\nfrom contextlib import contextmanager\n\nimport cv2\nimport pandas as pd\nfrom torch import Tensor\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, DataLoader\nfrom albumentations import *\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.optim.optimizer import Optimizer\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nimport tqdm\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import KFold\nfrom sklearn.model_selection import StratifiedKFold\nimport torchvision\nfrom segmentation_models_pytorch.encoders import encoders\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport tifffile as tiff\n\n\nfrom torch.nn.utils import weight_norm, spectral_norm\n\n#sys.path.append('../input/hubmap-coat/')\n\n#from coat import *\n#from daformer import *\n#from helper import *\n\ntorch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:28.193546Z","iopub.execute_input":"2022-10-17T00:58:28.193829Z","iopub.status.idle":"2022-10-17T00:58:28.911370Z","shell.execute_reply.started":"2022-10-17T00:58:28.193802Z","shell.execute_reply":"2022-10-17T00:58:28.910131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:28.912904Z","iopub.execute_input":"2022-10-17T00:58:28.913318Z","iopub.status.idle":"2022-10-17T00:58:28.920569Z","shell.execute_reply.started":"2022-10-17T00:58:28.913282Z","shell.execute_reply":"2022-10-17T00:58:28.919489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    train=True\n    debug=False\n    ENCODER = 'resnet50'\n    DECODER = 'UnetPlusPlus'\n    EPOCHS = 60\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts','OneCycleLR']\n    # CosineAnnealingLR params\n    cosanneal_params={\n        'T_max':10,\n        'eta_min':1e-4*0.5,\n        'last_epoch':-1\n    }\n    #ReduceLROnPlateau params\n    reduce_params={\n        'mode':'min',\n        'factor':0.2,\n        'patience':5,\n        'eps':1e-6,\n        'verbose':True\n    }\n    # CosineAnnealingWarmRestarts params\n    cosanneal_res_params={\n        'T_0':10,\n        'eta_min':1e-6,\n        'T_mult':1,\n        'last_epoch':-1\n    }\n    # OneCycleLR params\n    onecycle_params={\n        'pct_start':0.1,\n        'div_factor':1e1,\n        'max_lr':1e-3,\n        'steps_per_epoch':3, \n        'epochs':3\n    }\n    preds_col = 'Prediction'\n    lr=1e-3\n    weight_decay=1e-4\n    fold = 0\n    nfolds = 5\n    imsize = 256\n    BATCH_SIZE = 32\n    print_freq=100\n    DEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\n    NUM_WORKERS = 4\n    SEED = 24\n    trn_folds=[0, 1, 2, 3, 4]\n    model_name = 'resnext101'\n    resolution = (256,256)\n    deepsupervision = True\n    clfhead = True\n    clf_threshold = None","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:28.923282Z","iopub.execute_input":"2022-10-17T00:58:28.923621Z","iopub.status.idle":"2022-10-17T00:58:28.993746Z","shell.execute_reply.started":"2022-10-17T00:58:28.923589Z","shell.execute_reply":"2022-10-17T00:58:28.992586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# functions to convert encoding to mask and mask to encoding\ndef enc2mask(encs, shape):\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for m,enc in enumerate(encs):\n        if isinstance(enc,np.float) and np.isnan(enc): continue\n        s = enc.split()\n        for i in range(len(s)//2):\n            start = int(s[2*i]) - 1\n            length = int(s[2*i+1])\n            img[start:start+length] = 1 + m\n    return img.reshape(shape).T\n\ndef mask2enc(mask, n=1):\n    pixels = mask.T.flatten()\n    encs = []\n    for i in range(1,n+1):\n        p = (pixels == i).astype(np.int8)\n        if p.sum() == 0: encs.append(np.nan)\n        else:\n            p = np.concatenate([[0], p, [0]])\n            runs = np.where(p[1:] != p[:-1])[0] + 1\n            runs[1::2] -= runs[::2]\n            encs.append(' '.join(str(x) for x in runs))\n    return encs","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:28.995829Z","iopub.execute_input":"2022-10-17T00:58:28.996576Z","iopub.status.idle":"2022-10-17T00:58:29.010719Z","shell.execute_reply.started":"2022-10-17T00:58:28.996531Z","shell.execute_reply":"2022-10-17T00:58:29.009303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def PREDS_rle(liste):\n    out=[]    \n    for i in liste:\n        i = i.cpu().detach()\n        i = np.array(i.permute(1,2,0))\n        out.append(mask2enc(i))\n    return out","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.012538Z","iopub.execute_input":"2022-10-17T00:58:29.013456Z","iopub.status.idle":"2022-10-17T00:58:29.025562Z","shell.execute_reply.started":"2022-10-17T00:58:29.013406Z","shell.execute_reply":"2022-10-17T00:58:29.024450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\nclass 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\n\nclass Dice_soft(Metric):\n    def __init__(self, axis=1):\n        self.axis = axis\n    def reset(self): self.inter,self.union = 0,0\n    def accumulate(self, pred, targ):\n        pred = torch.sigmoid(pred)\n        self.inter += (pred*targ).float().sum().item()\n        self.union += (pred+targ).float().sum().item()\n    @property\n    def value(self): return 2.0 * self.inter/self.union if self.union > 0 else None\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f'[{name}] start')\n    yield\n    LOGGER.info(f'[{name}] done in {time.time() - t0:.0f} s.')\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=CFG.SEED):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.SEED)","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.029249Z","iopub.execute_input":"2022-10-17T00:58:29.029524Z","iopub.status.idle":"2022-10-17T00:58:29.047072Z","shell.execute_reply.started":"2022-10-17T00:58:29.029500Z","shell.execute_reply":"2022-10-17T00:58:29.046127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{}},{"cell_type":"code","source":"folds = train.copy()\nkf = StratifiedKFold(n_splits=CFG.nfolds,random_state=CFG.SEED,shuffle=True)\nfor fold, (_, val_idx) in enumerate(kf.split(folds, y=folds[\"organ\"])):\n    folds.loc[val_idx, \"fold\"] = fold","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.049401Z","iopub.execute_input":"2022-10-17T00:58:29.050072Z","iopub.status.idle":"2022-10-17T00:58:29.069491Z","shell.execute_reply.started":"2022-10-17T00:58:29.050035Z","shell.execute_reply":"2022-10-17T00:58:29.068476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"imaging_measurements = {\n  'HPA': {\n    'pixel_size': {\n      'kidney': 0.4,\n      'prostate': 0.4,\n      'largeintestine': 0.4,\n      'spleen': 0.4,\n      'lung': 0.4\n    },\n    'tissue_thickness': {\n      'kidney': 4,\n      'prostate': 4,\n      'largeintestine': 4,\n      'spleen': 4,\n      'lung': 4\n    }\n  },\n  'Hubmap': {\n    'pixel_size': {\n      'kidney': 0.5,\n      'prostate': 6.263,\n      'largeintestine': 0.229,\n      'spleen': 0.4945,\n      'lung': 0.7562\n    },\n    'tissue_thickness': {\n      'kidney': 10,\n      'prostate': 5,\n      'largeintestine': 8,\n      'spleen': 4,\n      'lung': 5\n    }\n  }\n}","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.071226Z","iopub.execute_input":"2022-10-17T00:58:29.071614Z","iopub.status.idle":"2022-10-17T00:58:29.080712Z","shell.execute_reply.started":"2022-10-17T00:58:29.071580Z","shell.execute_reply":"2022-10-17T00:58:29.079803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"# ====================================================\n# Dataset\n# ====================================================\nclass HuBMAPDataset(torch.utils.data.Dataset):\n    def __init__(self, df, tfms=None):\n        self.df = df\n        ids = self.df.id.values\n        self.fnames = [fname for fname in os.listdir(TRAIN_PATH) if int(fname.split('_')[0]) in ids]\n        self.image_size = CFG.imsize\n        self.tfms = tfms\n        \n    def img2tensor(self, img,dtype:np.dtype=np.float32):\n        if img.ndim==2 : img = np.expand_dims(img,2)\n        img = np.transpose(img,(2,0,1)) # C , H , W\n        return torch.from_numpy(img.astype(dtype, copy=False))\n    \n    def __len__(self):\n        return len(self.fnames)\n    \n    def resize(self, img, interp):\n        return  cv2.resize(\n            img, (self.image_size, self.image_size), interpolation=interp)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(TRAIN_PATH + fname), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread((MASK_PATH + fname),cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        \n        return self.img2tensor(self.resize(img , cv2.INTER_NEAREST)) , self.img2tensor(self.resize(mask , cv2.INTER_NEAREST))\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.083904Z","iopub.execute_input":"2022-10-17T00:58:29.084306Z","iopub.status.idle":"2022-10-17T00:58:29.093477Z","shell.execute_reply.started":"2022-10-17T00:58:29.084278Z","shell.execute_reply":"2022-10-17T00:58:29.092326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\n\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : 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, df, alpha=0.15, tfms=None):\n        self.df = df\n        ids = self.df.id.values\n        self.fnames = [fname for fname in os.listdir(TRAIN) if int(fname.split('_')[0]) in ids]\n        self.tfms = tfms\n        self.data_source = self.df.data_source\n        self.alpha = alpha\n        \n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        \n        fname = self.fnames[idx]\n        # retrieve image index from df\n        index = self.df[self.df['id'] == int(fname.split('_')[0])].index\n        organ = self.df.iloc[index[0]]['organ']\n        #print(f'organ : {organ}')\n        \n        img = cv2.cvtColor(cv2.imread(os.path.join(TRAIN,fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(MASKS,fname),cv2.IMREAD_GRAYSCALE)\n        \n        # Normalize tissue thickness:\n        domain_tissue_thickness = imaging_measurements['HPA']['tissue_thickness'][organ]\n        target_tissue_thickness = imaging_measurements['Hubmap']['tissue_thickness'][organ]\n        tissue_thickness_scale_factor = target_tissue_thickness - domain_tissue_thickness\n        image_hsv = cv2.cvtColor(img, cv2.COLOR_RGB2HSV).astype(np.float32)\n        image_hsv[:, :, 1] *= (1 + (self.alpha * tissue_thickness_scale_factor))\n        image_hsv[:, :, 2] *= (1 - (self.alpha * tissue_thickness_scale_factor))\n        image_hsv = image_hsv.astype(np.uint8)\n        image_scaled = cv2.cvtColor(image_hsv, cv2.COLOR_HSV2RGB)\n        \n        #if organ != 'prostate':\n        # Normalize Pixel Size:\n        domain_pixel_size = imaging_measurements['HPA']['pixel_size'][organ]\n        target_pixel_size = imaging_measurements['Hubmap']['pixel_size'][organ]\n        pixel_size_scale_factor = domain_pixel_size / target_pixel_size\n        \n        image_resized = cv2.resize(\n            image_scaled,\n            dsize=None,\n            fx=pixel_size_scale_factor,\n            fy=pixel_size_scale_factor,\n            interpolation=cv2.INTER_CUBIC\n            )\n        \n        mask_resized = cv2.resize(\n            mask,\n            dsize=None,\n            fx=pixel_size_scale_factor,\n            fy=pixel_size_scale_factor,\n            interpolation=cv2.INTER_CUBIC\n            )\n        # Resize \n        image_resized = cv2.resize(image_resized,dsize=(img.shape[1],img.shape[0]),interpolation=cv2.INTER_CUBIC)\n        mask_resized = cv2.resize(mask_resized,dsize=(img.shape[1],img.shape[0]),interpolation=cv2.INTER_CUBIC)\n        #else:\n            #image_resized = image_scaled\n            #mask_resized = mask\n\n        \n        if self.tfms is not None:\n            augmented = self.tfms(image=image_resized ,mask=mask_resized)\n            img,mask = augmented['image'],augmented['mask']\n        return img2tensor(img/255.0),img2tensor(mask)","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.095116Z","iopub.execute_input":"2022-10-17T00:58:29.095880Z","iopub.status.idle":"2022-10-17T00:58:29.114262Z","shell.execute_reply.started":"2022-10-17T00:58:29.095763Z","shell.execute_reply":"2022-10-17T00:58:29.113328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"df = folds\nids = df.id.values\nfnames = [fname for fname in os.listdir(TRAIN) if int(fname.split('_')[0]) in ids]\nfname = fnames[1008]\nimg = cv2.imread(os.path.join(TRAIN,fname))\nmask = cv2.imread(os.path.join(MASKS,fname),cv2.IMREAD_GRAYSCALE)\n\nindex = df[df['id'] == int(fname.split('_')[0])].index\nprint(df.iloc[index[0]]['organ'])\n_, image_augmented = augment_image(\n    image=img,\n    domain_pixel_size=imaging_measurements['hpa']['pixel_size'][df.iloc[index[0]]['organ']],\n    target_pixel_size=imaging_measurements['hubmap']['pixel_size'][df.iloc[index[0]]['organ']],\n    domain_tissue_thickness=imaging_measurements['hpa']['tissue_thickness'][df.iloc[index[0]]['organ']],\n    target_tissue_thickness=imaging_measurements['hubmap']['tissue_thickness'][df.iloc[index[0]]['organ']],\n    alpha=0.15\n    )\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.115784Z","iopub.execute_input":"2022-10-17T00:58:29.116403Z","iopub.status.idle":"2022-10-17T00:58:29.131677Z","shell.execute_reply.started":"2022-10-17T00:58:29.116366Z","shell.execute_reply":"2022-10-17T00:58:29.130517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"def get_aug(p=1.0):\n    return A.Compose([\n    A.RandomRotate90(p=0.5),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.Transpose(p=0.5),\n\n    # brightness_aug\n    A.OneOf([\n        A.RandomBrightnessContrast(p=1.0, contrast_limit=(-0.2, 0.2), brightness_limit=(-0.1, 0.1)),\n        A.HueSaturationValue(p=1.0,always_apply=False,hue_shift_limit=(-30, 30),sat_shift_limit=(-90, 20),val_shift_limit=(-20, 20),),\n        A.RandomGamma(p=0.5, gamma_limit=(50, 200)),\n    ],\n        p=1.0,),\n     A.OneOf([\n        A.ChannelShuffle(p=0.1),\n        A.ColorJitter(p=0.1),\n        A.FancyPCA(p=0.1),\n        A.augmentations.transforms.CLAHE(p=0.1),\n    ],\n         p=1.0),\n        \n        \n    # distortion_aug\n    A.OneOf([\n        A.OpticalDistortion(p=0.3), \n        A.GridDistortion(p=0.3), \n        A.ElasticTransform(p=0.1)\n    ], \n        p=1.0),\n\n    # noise_aug\n    A.OneOf([\n        A.Blur(always_apply=False, p=1.0, blur_limit=(3, 7)),\n        A.GaussNoise(always_apply=False, p=1.0, var_limit=(10.0, 50.0)),\n        A.MultiplicativeNoise(always_apply=False, p=1.0, multiplier=(0.9, 1.1), per_channel=True, elementwise=True),\n    ],\n        p=1.0),\n        \n    #A.OneOf([\n    #    A.HueSaturationValue(10,15,10),\n    #    A.CLAHE(clip_limit=2),\n    #    A.RandomBrightnessContrast(),            \n    #], p=0.3),\n])","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.133013Z","iopub.execute_input":"2022-10-17T00:58:29.133680Z","iopub.status.idle":"2022-10-17T00:58:29.145237Z","shell.execute_reply.started":"2022-10-17T00:58:29.133618Z","shell.execute_reply":"2022-10-17T00:58:29.144234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"# ====================================================\n# Transforms\n# ====================================================\ndef transformer(p=1.0):\n    return A.Compose([\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.3),\n        A.OneOf([\n            A.HueSaturationValue(10,15,10),\n            A.CLAHE(clip_limit=2),\n            A.RandomBrightnessContrast(),            \n        ], p=0.3),\n    ], p=p)\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.146583Z","iopub.execute_input":"2022-10-17T00:58:29.147320Z","iopub.status.idle":"2022-10-17T00:58:29.161523Z","shell.execute_reply.started":"2022-10-17T00:58:29.147276Z","shell.execute_reply":"2022-10-17T00:58:29.160420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Viewing Data","metadata":{}},{"cell_type":"code","source":"train_folds = folds[folds['fold'] != 2]\ntrain_folds = train_folds.reset_index(drop=True)\nvalid_folds = folds[folds['fold'] == 2]\nvalid_folds = valid_folds.reset_index(drop=True)\n\nds = HuBMAPDataset(df=train_folds, tfms=get_aug())\nvld = HuBMAPDataset(df=valid_folds, tfms=get_aug())","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.163252Z","iopub.execute_input":"2022-10-17T00:58:29.163594Z","iopub.status.idle":"2022-10-17T00:58:29.199914Z","shell.execute_reply.started":"2022-10-17T00:58:29.163560Z","shell.execute_reply":"2022-10-17T00:58:29.199054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(folds['organ'].unique())\nkidney_data = folds[folds['organ'] == \"prostate\"]\nkidney_data = kidney_data.reset_index(drop=True)\nkidney_dataset = HuBMAPDataset(df=kidney_data, tfms=get_aug())","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.201774Z","iopub.execute_input":"2022-10-17T00:58:29.202033Z","iopub.status.idle":"2022-10-17T00:58:29.223153Z","shell.execute_reply.started":"2022-10-17T00:58:29.202010Z","shell.execute_reply":"2022-10-17T00:58:29.222230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example = kidney_dataset[4]\nprint(example[0].shape)\nprint(example[1].shape)\nplt.imshow(example[0].permute(1,2,0))\nplt.show()\nplt.imshow(example[1].permute(1,2,0))\nplt.show()\n\nplt.imshow(example[0].permute(1,2,0),vmin=0,vmax=255)\nplt.imshow(example[1].permute(1,2,0), alpha=0.1)","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.224598Z","iopub.execute_input":"2022-10-17T00:58:29.225087Z","iopub.status.idle":"2022-10-17T00:58:29.934015Z","shell.execute_reply.started":"2022-10-17T00:58:29.225052Z","shell.execute_reply":"2022-10-17T00:58:29.933148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nds = HuBMAPDataset(df=folds, tfms=get_aug())\ndl = torch.utils.data.DataLoader(ds,batch_size=64,shuffle=False)\nit = iter(dl)\nimgs,masks = next(it)","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:29.935379Z","iopub.execute_input":"2022-10-17T00:58:29.936355Z","iopub.status.idle":"2022-10-17T00:58:32.952026Z","shell.execute_reply.started":"2022-10-17T00:58:29.936320Z","shell.execute_reply":"2022-10-17T00:58:32.951158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    plt.subplot(8,8,i+1)\n    plt.imshow(img.permute(2,1,0),vmin=0,vmax=255)\n    plt.imshow(mask.permute(2,1,0), alpha=0.1)\n    plt.axis('off')\n    \ndel ds,dl,imgs,masks","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:32.956050Z","iopub.execute_input":"2022-10-17T00:58:32.958380Z","iopub.status.idle":"2022-10-17T00:58:37.371268Z","shell.execute_reply.started":"2022-10-17T00:58:32.958343Z","shell.execute_reply":"2022-10-17T00:58:37.369844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{}},{"cell_type":"code","source":"\"\"\"class 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        self.decoder_dim = decoder_dim\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\n    def forward(self, batch):\n        x = self.rgb(batch)\n        B, C, H, W = x.shape\n        encoder = self.encoder(x) # [1, 512, 12, 12]\n        last, decoder = self.decoder(encoder) # [1, 320, 96, 96]\n        logit = self.logit(last)\n\n        output = {}\n        probability_from_logit = torch.sigmoid(logit)\n\n        output['probability'] = probability_from_logit\n        return output\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.377272Z","iopub.execute_input":"2022-10-17T00:58:37.377701Z","iopub.status.idle":"2022-10-17T00:58:37.387610Z","shell.execute_reply.started":"2022-10-17T00:58:37.377666Z","shell.execute_reply":"2022-10-17T00:58:37.386516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"def init_model():\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-10-17T00:58:37.388934Z","iopub.execute_input":"2022-10-17T00:58:37.389618Z","iopub.status.idle":"2022-10-17T00:58:37.400954Z","shell.execute_reply.started":"2022-10-17T00:58:37.389584Z","shell.execute_reply":"2022-10-17T00:58:37.399979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"sample = ds[0]\nimg = sample[0]\nmask = sample[1]\nimg = img.unsqueeze(0)\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.402377Z","iopub.execute_input":"2022-10-17T00:58:37.403012Z","iopub.status.idle":"2022-10-17T00:58:37.411721Z","shell.execute_reply.started":"2022-10-17T00:58:37.402979Z","shell.execute_reply":"2022-10-17T00:58:37.410617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DECODERS = [\n    \"Unet\",\n    \"Linknet\",\n    \"FPN\",\n    \"PSPNet\",\n    \"DeepLabV3\",\n    \"DeepLabV3Plus\",\n    \"PAN\",\n    \"UnetPlusPlus\",\n]\nENCODERS = list(encoders.keys())\n\n\ndef define_model(\n    decoder_name,\n    encoder_name,\n    num_classes=1,\n    activation=None,\n    encoder_weights=\"imagenet\",\n):\n    \"\"\"\n    Loads a segmentation architecture.\n    Args:\n        decoder_name (str): Decoder name.\n        encoder_name (str): Encoder name.\n        num_classes (int, optional): Number of classes. Defaults to 1.\n        pretrained : pretrained original weights\n        activation (str or None, optional): Activation of the last layer. Defaults to None.\n        encoder_weights (str, optional): Pretrained weights. Defaults to \"imagenet\".\n    Returns:\n        torch model: Segmentation model.\n    \"\"\"\n    assert decoder_name in DECODERS, \"Decoder name not supported\"\n    assert encoder_name in ENCODERS, \"Encoder name not supported\"\n\n    decoder = getattr(smp, decoder_name)\n\n    model = decoder(\n        encoder_name,\n        encoder_weights=encoder_weights,\n        classes=num_classes,\n        activation=activation,\n    )\n    model.num_classes = num_classes\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.413535Z","iopub.execute_input":"2022-10-17T00:58:37.414298Z","iopub.status.idle":"2022-10-17T00:58:37.423261Z","shell.execute_reply.started":"2022-10-17T00:58:37.414262Z","shell.execute_reply":"2022-10-17T00:58:37.422220Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FPN(nn.Module):\n    def __init__(self, input_channels:list, output_channels:list):\n        super().__init__()\n        self.convs = nn.ModuleList(\n            [nn.Sequential(nn.Conv2d(in_ch, out_ch*2, kernel_size=3, padding=1),\n             nn.ReLU(inplace=True), nn.BatchNorm2d(out_ch*2),\n             nn.Conv2d(out_ch*2, out_ch, kernel_size=3, padding=1))\n            for in_ch, out_ch in zip(input_channels, output_channels)])\n\n    def forward(self, xs:list, last_layer):\n        hcs = [F.interpolate(c(x),scale_factor=2**(len(self.convs)-i),mode='bilinear')\n               for i,(c,x) in enumerate(zip(self.convs, xs))]\n        hcs.append(last_layer)\n        return torch.cat(hcs, dim=1)\n\nclass UnetBlock(nn.Module):\n    def __init__(self, up_in_c:int, x_in_c:int, nf:int=None, blur:bool=False,\n                 self_attention:bool=False, **kwargs):\n        super().__init__()\n        self.shuf = PixelShuffle_ICNR(up_in_c, up_in_c//2, blur=blur, **kwargs)\n        self.bn = nn.BatchNorm2d(x_in_c)\n        ni = up_in_c//2 + x_in_c\n        nf = nf if nf is not None else max(up_in_c//2,32)\n        self.conv1 = ConvLayer(ni, nf, norm_type=None, **kwargs)\n        self.conv2 = ConvLayer(nf, nf, norm_type=None,\n            xtra=SelfAttention(nf) if self_attention else None, **kwargs)\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, up_in:Tensor, left_in:Tensor) -> Tensor:\n        s = left_in\n        up_out = self.shuf(up_in)\n        cat_x = self.relu(torch.cat([up_out, self.bn(s)], dim=1))\n        return self.conv2(self.conv1(cat_x))\n\nclass _ASPPModule(nn.Module):\n    def __init__(self, inplanes, planes, kernel_size, padding, dilation, groups=1):\n        super().__init__()\n        self.atrous_conv = nn.Conv2d(inplanes, planes, kernel_size=kernel_size,\n                stride=1, padding=padding, dilation=dilation, bias=False, groups=groups)\n        self.bn = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU()\n\n        self._init_weight()\n\n    def forward(self, x):\n        x = self.atrous_conv(x)\n        x = self.bn(x)\n\n        return self.relu(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\nclass ASPP(nn.Module):\n    def __init__(self, inplanes=512, mid_c=256, dilations=[6, 12, 18, 24], out_c=None):\n        super().__init__()\n        self.aspps = [_ASPPModule(inplanes, mid_c, 1, padding=0, dilation=1)] + \\\n            [_ASPPModule(inplanes, mid_c, 3, padding=d, dilation=d,groups=4) for d in dilations]\n        self.aspps = nn.ModuleList(self.aspps)\n        self.global_pool = nn.Sequential(nn.AdaptiveMaxPool2d((1, 1)),\n                        nn.Conv2d(inplanes, mid_c, 1, stride=1, bias=False),\n                        nn.BatchNorm2d(mid_c), nn.ReLU())\n        out_c = out_c if out_c is not None else mid_c\n        self.out_conv = nn.Sequential(nn.Conv2d(mid_c*(2+len(dilations)), out_c, 1, bias=False),\n                                    nn.BatchNorm2d(out_c), nn.ReLU(inplace=True))\n        self.conv1 = nn.Conv2d(mid_c*(2+len(dilations)), out_c, 1, bias=False)\n        self._init_weight()\n\n    def forward(self, x):\n        x0 = self.global_pool(x)\n        xs = [aspp(x) for aspp in self.aspps]\n        x0 = F.interpolate(x0, size=xs[0].size()[2:], mode='bilinear', align_corners=True)\n        x = torch.cat([x0] + xs, dim=1)\n        return self.out_conv(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n\nimport torch.hub\n# hub_model = torch.hub.load(\n#     'moskomule/senet.pytorch',\n#     'se_resnet50',\n#     pretrained=True,)\n\nclass UneXt50(nn.Module):\n    def __init__(self, stride=1, **kwargs):\n        super().__init__()\n        #encoder\n        m = torch.hub.load('facebookresearch/semi-supervised-ImageNet1K-models',\n                            'resnext101_32x4d_swsl')\n        # m = torch.hub.load('facebookresearch/semi-supervised-ImageNet1K-models',\n        #                   'resnext50_32x4d_swsl')\n        # m = torch.hub.load(\n        #     'moskomule/senet.pytorch',\n        #     'se_resnet101',\n        #     pretrained=True,)\n\n        #m=torch.hub.load('zhanghang1989/ResNeSt', 'resnest50', pretrained=True)\n        self.enc0 = nn.Sequential(m.conv1, m.bn1, nn.ReLU(inplace=True))\n        self.enc1 = nn.Sequential(nn.MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1),\n                            m.layer1) #256\n        self.enc2 = m.layer2 #512\n        self.enc3 = m.layer3 #1024\n        self.enc4 = m.layer4 #2048\n        #aspp with customized dilatations\n        self.aspp = ASPP(2048,256,out_c=512,dilations=[stride*1,stride*2,stride*3,stride*4])\n        self.drop_aspp = nn.Dropout2d(0.5)\n        #decoder\n        self.dec4 = UnetBlock(512,1024,256)\n        self.dec3 = UnetBlock(256,512,128)\n        self.dec2 = UnetBlock(128,256,64)\n        self.dec1 = UnetBlock(64,64,32)\n        self.fpn = FPN([512,256,128,64],[16]*4)\n        self.drop = nn.Dropout2d(0.1)\n        self.final_conv = ConvLayer(32+16*4, 1, ks=1, norm_type=None, act_cls=None)\n\n    def forward(self, x):\n        enc0 = self.enc0(x)\n        enc1 = self.enc1(enc0)\n        enc2 = self.enc2(enc1)\n        enc3 = self.enc3(enc2)\n        enc4 = self.enc4(enc3)\n        enc5 = self.aspp(enc4)\n        dec3 = self.dec4(self.drop_aspp(enc5),enc3)\n        dec2 = self.dec3(dec3,enc2)\n        dec1 = self.dec2(dec2,enc1)\n        dec0 = self.dec1(dec1,enc0)\n        x = self.fpn([enc5, dec3, dec2, dec1], dec0)\n        x = self.final_conv(self.drop(x))\n        x = F.interpolate(x,scale_factor=2,mode='bilinear')\n        return x\n\n#split the model to encoder and decoder for fast.ai\nsplit_layers = lambda m: [list(m.enc0.parameters())+list(m.enc1.parameters())+\n                list(m.enc2.parameters())+list(m.enc3.parameters())+\n                list(m.enc4.parameters()),\n                list(m.aspp.parameters())+list(m.dec4.parameters())+\n                list(m.dec3.parameters())+list(m.dec2.parameters())+\n                list(m.dec1.parameters())+list(m.fpn.parameters())+\n                list(m.final_conv.parameters())]\n\n#model = UneXt50()","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.424834Z","iopub.execute_input":"2022-10-17T00:58:37.425487Z","iopub.status.idle":"2022-10-17T00:58:37.698482Z","shell.execute_reply.started":"2022-10-17T00:58:37.425442Z","shell.execute_reply":"2022-10-17T00:58:37.697355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"class SoftDiceLoss(nn.Module):\n    def __init__(self, smooth=1., dims=(-2,-1)):\n\n        super(SoftDiceLoss, self).__init__()\n        self.smooth = smooth\n        self.dims = dims\n    \n    def forward(self, x, y):\n\n        tp = (x * y).sum(self.dims)\n        fp = (x * (1 - y)).sum(self.dims)\n        fn = ((1 - x) * y).sum(self.dims)\n        \n        dc = (2 * tp + self.smooth) / (2 * tp + fp + fn + self.smooth)\n        dc = dc.mean()\n\n        return 1 - dc","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.700863Z","iopub.execute_input":"2022-10-17T00:58:37.701184Z","iopub.status.idle":"2022-10-17T00:58:37.713013Z","shell.execute_reply.started":"2022-10-17T00:58:37.701157Z","shell.execute_reply":"2022-10-17T00:58:37.711971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Loss function\n# ====================================================\nclass 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        \n    def forward(self, outputs, mask):\n        dice = self.diceloss(outputs,mask)\n        bce = self.binloss(outputs , mask)\n        loss = dice * 0.7 + bce * 0.3\n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.714729Z","iopub.execute_input":"2022-10-17T00:58:37.715143Z","iopub.status.idle":"2022-10-17T00:58:37.723182Z","shell.execute_reply.started":"2022-10-17T00:58:37.715109Z","shell.execute_reply":"2022-10-17T00:58:37.721949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\ndef train_fn(train_loader, model, criterion, metric , optimizer, scheduler,epoch, DEVICE):\n    \n    metric.reset()\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    train_loss=0.0\n    score=0.0\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    for step, data in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        img, mask = data\n        img = img.to(DEVICE)\n        mask = mask.to(DEVICE)\n        batch_size = img.size(0)\n    \n        outputs = model(img)\n        loss = criterion(outputs, mask)\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        #loss\n        #loss = loss.detach().item()\n        train_loss += loss.item()\n        #score += metric(outputs,mask).item()\n        metric.accumulate(outputs.detach(), mask)\n        # record loss\n        losses.update(loss.item(), batch_size)\n        #scores.update(metric(outputs,mask).item(), batch_size)\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n            \n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  #'Score: {score.val:.4f}'\n                  'LR: {lr:.6f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          #score=scores,\n                          lr=scheduler.get_lr()[0]\n                         ))\n    TRAIN_LOSS = train_loss / len(train_loader)\n    #SCORE = score / len(train_loader)\n    metric_this_epoch=metric.value\n    return TRAIN_LOSS , metric_this_epoch\n    \ndef valid_fn(valid_loader, model, criterion,metric,epoch ,DEVICE):\n    metric.reset()\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    preds=[]\n    valid_loss=0.0\n    val_score=0.0\n    # switch to evaluation mode\n    model.eval()\n    start = end = time.time()\n    \n    for step, data in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        img, mask = data\n        img = img.to(DEVICE)\n        mask = mask.to(DEVICE)\n        batch_size = img.size(0)\n        # compute loss\n        with torch.no_grad():\n            outputs = model(img)\n        output_tuple = torch.unbind(outputs, dim=0)\n        preds.append(output_tuple)\n        loss = criterion(outputs, mask)\n        valid_loss += loss.item()\n        #val_score += metric(outputs,mask).item()\n        losses.update(loss.item(), batch_size)\n        #scores.update(metric(outputs,mask).item(), batch_size)\n        metric.accumulate(outputs.detach(), mask)\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                    'Elapsed {remain:s} '\n                    'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                    #'Score: {score.val:.4f}'\n                    .format(step, len(valid_loader),\n                            loss=losses,\n                            #score=scores,\n                            remain=timeSince(start, float(step+1)/len(valid_loader))))\n    VALID_LOSS = valid_loss / len(valid_loader)\n    #VALID_SCORE = val_score / len(valid_loader)\n    predictions = [item for t in preds for item in t]\n    metric_this_epoch=metric.value\n    return VALID_LOSS, metric_this_epoch , predictions","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.724820Z","iopub.execute_input":"2022-10-17T00:58:37.725159Z","iopub.status.idle":"2022-10-17T00:58:37.748283Z","shell.execute_reply.started":"2022-10-17T00:58:37.725126Z","shell.execute_reply":"2022-10-17T00:58:37.747252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop(folds, fold):\n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n    \n    # ====================================================\n    # loader\n    # ====================================================\n    if CFG.debug:\n        train_folds = folds[folds['fold'] != fold].sample(20)\n        valid_folds = folds[folds['fold'] == fold].sample(20)\n        \n    else:\n        train_folds = folds[folds['fold'] != fold]\n        valid_folds = folds[folds['fold'] == fold]\n        \n    train_folds = train_folds.reset_index(drop=True)\n    valid_folds = train_folds.reset_index(drop=True)\n        \n    best_loss = 999\n    best_score = 0\n    \n    ds_train = HuBMAPDataset(df=train_folds, tfms=get_aug())\n    ds_val = HuBMAPDataset(df=valid_folds)\n    \n    dataloader_train = torch.utils.data.DataLoader(ds_train,batch_size=CFG.BATCH_SIZE, shuffle=True,num_workers=CFG.NUM_WORKERS)\n    dataloader_val = torch.utils.data.DataLoader(ds_val,batch_size=CFG.BATCH_SIZE, shuffle=False,num_workers=CFG.NUM_WORKERS)\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, **CFG.reduce_params)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, **CFG.cosanneal_params)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, **CFG.cosanneal_res_params)\n        return scheduler\n    \n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    #model = define_model(encoder_name=CFG.ENCODER,decoder_name=CFG.DECODER).to(CFG.DEVICE)\n    model = UneXt50().to(CFG.DEVICE)\n    #model = get_model().to(CFG.DEVICE)\n    #model = init_model().to(CFG.DEVICE)\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 = torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n    #                                          max_lr=1e-2, epochs=CFG.EPOCHS, steps_per_epoch=len(dataloader_train))\n    optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n    #scheduler = get_scheduler(optimizer)\n    scheduler=torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.2,\n                                                  div_factor=1e2, max_lr=1e-4, epochs=CFG.EPOCHS,\n                                                  steps_per_epoch=len(dataloader_train))\n    #loss_func = CustomLoss()\n    loss_func=nn.BCEWithLogitsLoss()\n    #dice_coe = DiceCoef()\n    dice_coe=Dice_soft()\n    \n    for epoch in range(CFG.EPOCHS):\n        start_time = time.time()\n        \n        # train\n        train_loss, train_score = train_fn(dataloader_train, model, loss_func, dice_coe, optimizer, scheduler, epoch, CFG.DEVICE)\n        \n        # eval\n        valid_loss, val_score , predictions = valid_fn(dataloader_val, model, loss_func, dice_coe, epoch, CFG.DEVICE)\n        #RLES = PREDS_rle(predictions)\n        #print(f'RLES : {RLES}')\n        scheduler.step()\n        #if isinstance(scheduler, ReduceLROnPlateau):\n        #    scheduler.step(valid_loss)\n        #elif isinstance(scheduler, CosineAnnealingLR):\n        #    scheduler.step()\n        #elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n        #    scheduler.step()\n        \n        # scoring\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - train_loss: {train_loss:.4f}  train_score: {train_score:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - valid_loss: {valid_loss:.4f}  val_score: {val_score:.4f}')\n        \n        if val_score > best_score:\n            best_score = val_score\n            torch.save(model.state_dict(),\n                        f\"{OUTPUT_DIR}FOLD{fold}_best_score.pth\")\n            print(f\"Saved model for best score : FOLD{fold}_best_score.pth\")\n            LOGGER.info(f\"Saved model for best score : FOLD{fold}_best_score.pth\")\n        ","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.749833Z","iopub.execute_input":"2022-10-17T00:58:37.750236Z","iopub.status.idle":"2022-10-17T00:58:37.767598Z","shell.execute_reply.started":"2022-10-17T00:58:37.750181Z","shell.execute_reply":"2022-10-17T00:58:37.766556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Main","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# main\n# ====================================================\ndef main():\n\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.nfolds):\n            if fold in CFG.trn_folds:\n                _oof_df = train_loop(folds, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.770131Z","iopub.execute_input":"2022-10-17T00:58:37.770886Z","iopub.status.idle":"2022-10-17T00:58:37.782879Z","shell.execute_reply.started":"2022-10-17T00:58:37.770852Z","shell.execute_reply":"2022-10-17T00:58:37.781874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.785191Z","iopub.execute_input":"2022-10-17T00:58:37.785791Z","iopub.status.idle":"2022-10-17T00:58:37.943960Z","shell.execute_reply.started":"2022-10-17T00:58:37.785650Z","shell.execute_reply":"2022-10-17T00:58:37.942790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"out=[]\nfor i in range(5):    \n    tensor = torch.rand([8,3,256,256])\n    result = torch.unbind(tensor, dim=0)\n    # using list comprehension\n    out.append(result)\n\noutput = [item for t in out for item in t]\nlen(output)\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.945604Z","iopub.execute_input":"2022-10-17T00:58:37.946284Z","iopub.status.idle":"2022-10-17T00:58:37.958246Z","shell.execute_reply.started":"2022-10-17T00:58:37.946247Z","shell.execute_reply":"2022-10-17T00:58:37.957375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"execution":{"iopub.status.busy":"2022-10-17T00:58:37.959717Z","iopub.execute_input":"2022-10-17T00:58:37.960359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}