{"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":"!cp -r ../input/pytorch-segmentation-models-lib/ ./\n# !cp -r ../input/timm-pytorch-image-models/pytorch-image-models-master ./","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:29:33.996876Z","iopub.execute_input":"2022-08-31T19:29:33.998151Z","iopub.status.idle":"2022-08-31T19:29:35.509435Z","shell.execute_reply.started":"2022-08-31T19:29:33.998019Z","shell.execute_reply":"2022-08-31T19:29:35.507446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install -q ./timm-pytorch-image-models/pytorch-image-models-master","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:29:35.518041Z","iopub.execute_input":"2022-08-31T19:29:35.521299Z","iopub.status.idle":"2022-08-31T19:29:35.531691Z","shell.execute_reply.started":"2022-08-31T19:29:35.521217Z","shell.execute_reply":"2022-08-31T19:29:35.530090Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip config set global.disable-pip-version-check true","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:29:35.536571Z","iopub.execute_input":"2022-08-31T19:29:35.539812Z","iopub.status.idle":"2022-08-31T19:29:37.856549Z","shell.execute_reply.started":"2022-08-31T19:29:35.539755Z","shell.execute_reply":"2022-08-31T19:29:37.854427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q ./pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ./pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ./pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl\n!pip install -q ./pytorch-segmentation-models-lib/segmentation_models_pytorch-0.2.0-py3-none-any.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-31T19:29:37.870723Z","iopub.execute_input":"2022-08-31T19:29:37.878146Z","iopub.status.idle":"2022-08-31T19:30:32.883171Z","shell.execute_reply.started":"2022-08-31T19:29:37.878086Z","shell.execute_reply":"2022-08-31T19:30:32.881707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n\n!cp ../input/efficientnet-pytorch-b0-b7/efficientnet-b0-355c32eb.pth /root/.cache/torch/hub/checkpoints/efficientnet-b0-355c32eb.pth\n!cp ../input/efficientnet-pytorch-b0-b7/efficientnet-b7-dcc49843.pth /root/.cache/torch/hub/checkpoints/efficientnet-b7-dcc49843.pth","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:32.885347Z","iopub.execute_input":"2022-08-31T19:30:32.886117Z","iopub.status.idle":"2022-08-31T19:30:41.084367Z","shell.execute_reply.started":"2022-08-31T19:30:32.886065Z","shell.execute_reply":"2022-08-31T19:30:41.082740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\n# vizualizations\nimport cv2\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# PyTorch \nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\nimport torch.nn.functional as F\n\n# folds\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\n\nfrom tqdm import tqdm\nimport gc\nimport random\nimport os\n\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:41.087076Z","iopub.execute_input":"2022-08-31T19:30:41.088097Z","iopub.status.idle":"2022-08-31T19:30:47.022358Z","shell.execute_reply.started":"2022-08-31T19:30:41.088036Z","shell.execute_reply":"2022-08-31T19:30:47.020869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize(figsize=(30, 30), **images):\n    n = len(images)\n    plt.figure(figsize=figsize)\n    for i, (name, image) in enumerate(images.items()):\n        plt.subplot(1, n, i + 1)\n        plt.xticks([])\n        plt.yticks([])\n        plt.title(' '.join(name.split('_')).title().lower())\n        plt.imshow(image)\n    plt.show()\n    \n    \ndef plot_tiled_image(image, fig_size=[10, 10]):\n    fig = plt.figure(figsize=fig_size)\n    rows = int(np.sqrt(image.shape[0]))\n    columns = rows\n\n    for i in range(1, columns*rows+1):\n        img = image[i-1, 0, ...]\n        fig.add_subplot(rows, columns, i)\n        plt.imshow(img)\n        plt.axis('off')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.024823Z","iopub.execute_input":"2022-08-31T19:30:47.026017Z","iopub.status.idle":"2022-08-31T19:30:47.038619Z","shell.execute_reply.started":"2022-08-31T19:30:47.025962Z","shell.execute_reply":"2022-08-31T19:30:47.037164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ref: https://www.kaggle.com/inversion/run-length-decoding-quick-start\ndef rle_decode(mask_rle, shape, color=1):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height, width, channels) of array to return\n    color: color for the mask\n    Returns numpy array (mask)\n\n    '''\n    s = mask_rle.split()\n\n    starts = list(map(lambda x: int(x) - 1, s[0::2]))\n    lengths = list(map(int, s[1::2]))\n    ends = [x + y for x, y in zip(starts, lengths)]\n    if len(shape) == 3:\n        img = np.zeros((shape[0] * shape[1], shape[2]), dtype=np.float32)\n    else:\n        img = np.zeros(shape[0] * shape[1], dtype=np.float32)\n    for start, end in zip(starts, ends):\n        img[start : end] = color\n\n    return img.reshape(shape)\n\n#ref: https://www.kaggle.com/code/bguberfain/memory-aware-rle-encoding/notebook\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    This simplified method requires first and last pixel to be zero\n    '''\n    pixels = img.T.flatten()\n    \n    # This simplified method requires first and last pixel to be zero\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.040955Z","iopub.execute_input":"2022-08-31T19:30:47.041922Z","iopub.status.idle":"2022-08-31T19:30:47.057639Z","shell.execute_reply.started":"2022-08-31T19:30:47.041857Z","shell.execute_reply":"2022-08-31T19:30:47.056011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def nearest(n, x):\n    u = n % x > x // 2\n    r = n + (-1) ** (1 - u) * abs(x * u - n % x)\n    \n#     if r % n == 0:\n#         return r\n    return r + x","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.060107Z","iopub.execute_input":"2022-08-31T19:30:47.061136Z","iopub.status.idle":"2022-08-31T19:30:47.071722Z","shell.execute_reply.started":"2022-08-31T19:30:47.061090Z","shell.execute_reply":"2022-08-31T19:30:47.070444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalize_0_255 = lambda x: (x * 255) / x.max()\n\ndef normalize_mean_std(a, axis=None): \n    mean = np.mean(a, axis=axis, keepdims=True)\n    std = np.std(a, axis=axis, keepdims=True)\n    return (a - mean) / std\n\ndef normalize_max(a):\n    return a / a.max()\n\ndef normalize_255_mean_std(a):\n    mean = np.array([0.7720342, 0.74582646, 0.76392896])\n    std = np.array([0.24745085, 0.26182273, 0.25782376])\n    return (a / 255.0 - mean) / std\n\ndef normalize_255_mean_std_v2(a):\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    return (a / 255. - mean) / std","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.085359Z","iopub.execute_input":"2022-08-31T19:30:47.086044Z","iopub.status.idle":"2022-08-31T19:30:47.097468Z","shell.execute_reply.started":"2022-08-31T19:30:47.086005Z","shell.execute_reply":"2022-08-31T19:30:47.095670Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class cfg:\n    dataset_path        = '../input/hubmap-organ-segmentation'\n    \n    train_images        = f'{dataset_path}/train_images'\n    train_annotations   = f'{dataset_path}/train_annotations'\n    test_images         = f'{dataset_path}/test_images'\n    \n    train_csv           = f'{dataset_path}/train.csv'\n    test_csv            = f'{dataset_path}/test.csv'\n    submission_csv      = f'{dataset_path}/sample_submission.csv'\n    \n    name        = '[all] - 736 + [lung-t1] - 1024'\n    seed        = 3407\n    n_folds     = 3\n    \n    tiles       = False\n    img_size    = [736, 736]\n    tile_size   = 512\n    stride      = (nearest(img_size[0], tile_size) - img_size[0]) // 4\n\n    backbone    = 'convnext_base_22k_224'\n    num_classes = 1\n    tta         = True\n    \n    organ_mean_thresh_f = {\n        'HPA': {\n            'kidney': 1.75,\n            'largeintestine': 1.75,\n            'lung': 1.75,\n            'prostate': 1.75,\n            'spleen': 1.75,          # not much change from 2.5 -> 1.5 (stable preds)\n        },\n        'Hubmap': {\n            'kidney': 1.75,\n            'largeintestine': 1.75,\n            'lung': 1.75,\n            'prostate': 1.75,\n            'spleen': 1.75,          \n        },\n    }\n    \n    device      = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.100541Z","iopub.execute_input":"2022-08-31T19:30:47.102118Z","iopub.status.idle":"2022-08-31T19:30:47.187595Z","shell.execute_reply.started":"2022-08-31T19:30:47.101985Z","shell.execute_reply":"2022-08-31T19:30:47.186194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(cfg.img_size)\nprint(cfg.tile_size)\nprint(cfg.stride)","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.191353Z","iopub.execute_input":"2022-08-31T19:30:47.192107Z","iopub.status.idle":"2022-08-31T19:30:47.203961Z","shell.execute_reply.started":"2022-08-31T19:30:47.192047Z","shell.execute_reply":"2022-08-31T19:30:47.202675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 0.78\n# - lung: 0.10 LB\n# MODELS = {\n#     'fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold00-dice0.7249.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     'fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold01-dice0.8332.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     'fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold02-dice0.7435.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     'fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold03-dice0.7618.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     'fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold04-dice0.7421.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n# }\n\n\nMODELS = {\n    '[GUME]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold00-dice0.7249.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[GUME]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold01-dice0.8332.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[GUME]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold02-dice0.7435.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[GUME]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold03-dice0.7618.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[GUME]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v0/GUME/GUME-convnext_upernet-768-rgb-norm255_mean_std_v2-augheavy_-fold04-dice0.7421.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    \n    \n    \n    '[COOP]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/COOP/COOP-convnext_upernet_b-768-norm255_mean_std_v2-augheavy5cd-all-fold00-dice0.7423.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[COOP]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/COOP/COOP-convnext_upernet_b-768-norm255_mean_std_v2-augheavy5cd-all-fold01-dice0.8414.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[COOP]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/COOP/COOP-convnext_upernet_b-768-norm255_mean_std_v2-augheavy5cd-all-fold02-dice0.7587.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[COOP]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/COOP/COOP-convnext_upernet_b-768-norm255_mean_std_v2-augheavy5cd-all-fold03-dice0.7671.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[COOP]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/COOP/COOP-convnext_upernet_b-768-norm255_mean_std_v2-augheavy5cd-all-fold04-dice0.7624.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    \n    \n    \n#     '[JODO]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/JODO/JODO-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-fold00-dice0.7435.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[JODO]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/JODO/JODO-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-fold01-dice0.8396.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[JODO]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/JODO/JODO-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-fold02-dice0.7602.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n    '[JODO]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/JODO/JODO-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-fold03-dice0.7639.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n#     '[JODO]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/JODO/JODO-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-fold04-dice0.7616.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n    \n    \n    \n#     '[NEXT]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/NEXT/NEXT-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-lung_g1_f-fold00-dice0.7836.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[NEXT]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/NEXT/NEXT-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-lung_g1_f-fold01-dice0.8791.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[NEXT]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/NEXT/NEXT-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-lung_g1_f-fold02-dice0.7557.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n    '[NEXT]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v1/NEXT/NEXT-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-lung_g1_f-fold03-dice0.7937.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n#     '[NEXT]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v1/NEXT/NEXT-convnext_upernet_b-768-norm255_mean_std_v2-augheavy4-all-lung_g1_f-fold04-dice0.8005.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n    \n    \n#     '[A005]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A005/A005-hornet_upernet-b-768-norm255_mean_std_v2-augheavy4cds-all-fold00-dice0.7330.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n#     '[A005]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A005/A005-hornet_upernet-b-768-norm255_mean_std_v2-augheavy4cds-all-fold01-dice0.8489.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n#     '[A005]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A005/A005-hornet_upernet-b-768-norm255_mean_std_v2-augheavy4cds-all-fold02-dice0.7685.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n#     '[A005]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A005/A005-hornet_upernet-b-768-norm255_mean_std_v2-augheavy4cds-all-fold03-dice0.7695.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n#     '[A005]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A005/A005-hornet_upernet-b-768-norm255_mean_std_v2-augheavy4cds-all-fold04-dice0.7611.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n    \n    \n    \n    '[A007]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v2/A007/A007-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-fold00-dice0.7455.pth',\n        'num_classes': 1,\n        'arch': 'hornet_upernet',\n        'weight': 1\n    },\n    '[A007]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v2/A007/A007-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-fold01-dice0.8459.pth',\n        'num_classes': 1,\n        'arch': 'hornet_upernet',\n        'weight': 1\n    },\n    '[A007]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v2/A007/A007-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-fold02-dice0.7745.pth',\n        'num_classes': 1,\n        'arch': 'hornet_upernet',\n        'weight': 1\n    },\n    '[A007]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v2/A007/A007-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-fold03-dice0.7688.pth',\n        'num_classes': 1,\n        'arch': 'hornet_upernet',\n        'weight': 1\n    },\n    '[A007]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v2/A007/A007-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-fold04-dice0.7570.pth',\n        'num_classes': 1,\n        'arch': 'hornet_upernet',\n        'weight': 1\n    },\n    \n    \n#     '[A008]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A008/A008-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-xlung-fold00-dice0.8651.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n#     '[A008]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A008/A008-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-xlung-fold01-dice0.9110.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n#     '[A008]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A008/A008-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-xlung-fold02-dice0.8628.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n    '[A008]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-models-v2/A008/A008-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-xlung-fold03-dice0.8517.pth',\n        'num_classes': 1,\n        'arch': 'hornet_upernet',\n        'weight': 1\n    },\n#     '[A008]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy++~': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-models-v2/A008/A008-hornet_upernet-b-gf-768-norm255_mean_std_v2-augheavy4cds--all-xlung-fold04-dice0.8785.pth',\n#         'num_classes': 1,\n#         'arch': 'hornet_upernet',\n#         'weight': 1\n#     },\n}","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.206397Z","iopub.execute_input":"2022-08-31T19:30:47.207600Z","iopub.status.idle":"2022-08-31T19:30:47.234335Z","shell.execute_reply.started":"2022-08-31T19:30:47.207553Z","shell.execute_reply":"2022-08-31T19:30:47.231986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LUNG_MODELS = {\n#     '[lung_g1_f]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/MOTO/MOTO-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy4-lung-g1-f-fold00-dice0.2923.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/MOTO/MOTO-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy4-lung-g1-f-fold01-dice0.5040.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/MOTO/MOTO-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy4-lung-g1-f-fold02-dice0.3659.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/MOTO/MOTO-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy4-lung-g1-f-fold03-dice0.4879.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/MOTO/MOTO-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy4-lung-g1-f-fold04-dice0.3770.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n# }\n\n# LUNG_MODELS = {\n#     '[lung_t1]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold00-dice0.3542.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_t1]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold01-dice0.3732.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_t1]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold02-dice0.4176.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_t1]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold03-dice0.5012.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_t1]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold04-dice0.4236.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n# }\n\n\nLUNG_MODELS = {\n    '[lung_g1_f]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold00-dice0.3542.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[lung_g1_f]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-lung-models-v1/MOTO/MOTO-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy4-lung-g1-f-fold01-dice0.5040.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 2\n    },\n    '[lung_g1_f]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold02-dice0.4176.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    '[lung_g1_f]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold03-dice0.5012.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 3\n    },\n    '[lung_g1_f]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n        'backbone': None,\n        'weights': '../input/hubmap-hpa-lung-models-v1/GIGA/GIGA-convnext_upernet-1024x512x2-norm255_mean_std_v2-augheavy-fold04-dice0.4236.pth',\n        'num_classes': 1,\n        'arch': 'convnext_upernet',\n        'weight': 1\n    },\n    \n    \n#     '[XEOI]-[lung_g1_f]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold00-dice0.2870.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[XEOI]-[lung_g1_f]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold01-dice0.3065.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[XEOI]-[lung_g1_f]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold02-dice0.4827.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[XEOI]-[lung_g1_f]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold03-dice0.5141.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n#     '[XEOI]-[lung_g1_f]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold04-dice0.3544.pth',\n#         'num_classes': 1,\n#         'arch': 'convnext_upernet',\n#         'weight': 1\n#     },\n}\n\n\n# LUNG_MODELS = {\n#     '[lung_g1_f]-fold0-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold00-dice0.2870.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold1-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold01-dice0.3065.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold2-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold02-dice0.4827.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold3-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold03-dice0.5141.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n#     '[lung_g1_f]-fold4-upernet-convnext_base_22k_224-norm_255_mean_std_v2-aug_heavy+++': {\n#         'backbone': None,\n#         'weights': '../input/hubmap-hpa-lung-models-v1/XEOI/XEOI-convnext_upernet-1024-norm255_mean_std_v2-augheavy5cd-lung_g1-fold04-dice0.3544.pth',\n#         'num_classes': 1,\n#         'weight': 1\n#     },\n# }","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:47.237075Z","iopub.execute_input":"2022-08-31T19:30:47.238011Z","iopub.status.idle":"2022-08-31T19:30:47.254327Z","shell.execute_reply.started":"2022-08-31T19:30:47.237941Z","shell.execute_reply":"2022-08-31T19:30:47.252988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=42):\n    '''\n    Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY\n    '''\n    print(f'- SETTING SEED: {seed}')\n\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    \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    \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-08-31T19:30:47.257364Z","iopub.execute_input":"2022-08-31T19:30:47.258196Z","iopub.status.idle":"2022-08-31T19:30:47.277228Z","shell.execute_reply.started":"2022-08-31T19:30:47.258155Z","shell.execute_reply":"2022-08-31T19:30:47.273771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def conv3x3_bn_relu(in_planes, out_planes, stride=1):\n\t\"3x3 convolution + BN + relu\"\n\treturn nn.Sequential(\n\t\tnn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, padding=1, bias=False),\n\t\tnn.BatchNorm2d(out_planes),\n\t\tnn.ReLU(inplace=True),\n\t)\n\n \n# upernet\nclass UPerDecoder(nn.Module):\n\tdef __init__(self,\n\t    in_dim=[256, 512, 1024, 2048],\n\t    ppm_pool_scale=[1, 2, 3, 6],\n\t    ppm_dim=512,\n\t    fpn_out_dim=256\n\t):\n\t\tsuper(UPerDecoder, self).__init__()\n\n\t\t# PPM ----\n\t\tdim         = in_dim[-1]\n\t\tppm_pooling = []\n\t\tppm_conv    = []\n\t\t\n\t\tfor scale in ppm_pool_scale:\n\t\t\tppm_pooling.append(\n\t\t\t\tnn.AdaptiveAvgPool2d(scale)\n\t\t\t)\n\t\t\tppm_conv.append(\n\t\t\t\tnn.Sequential(\n\t\t\t\t\tnn.Conv2d(dim, ppm_dim, kernel_size=1, bias=False),\n\t\t\t\t\tnn.BatchNorm2d(ppm_dim),\n\t\t\t\t\tnn.ReLU(inplace=True)\n\t\t\t\t)\n\t\t\t)\n\t\tself.ppm_pooling    = nn.ModuleList(ppm_pooling)\n\t\tself.ppm_conv       = nn.ModuleList(ppm_conv)\n\t\tself.ppm_out        = conv3x3_bn_relu(dim + len(ppm_pool_scale)*ppm_dim, fpn_out_dim, 1)\n\t\t\n\t\t# FPN ----\n\t\tfpn_in = []\n\t\tfor i in range(0, len(in_dim)-1):  # skip the top layer\n\t\t\tfpn_in.append(\n\t\t\t\tnn.Sequential(\n\t\t\t\t\tnn.Conv2d(in_dim[i], fpn_out_dim, kernel_size=1, bias=False),\n\t\t\t\t\tnn.BatchNorm2d(fpn_out_dim),\n\t\t\t\t\tnn.ReLU(inplace=True)\n\t\t\t\t)\n\t\t\t)\n\t\tself.fpn_in = nn.ModuleList(fpn_in)\n\t\t\n\t\tfpn_out = []\n\t\tfor i in range(len(in_dim) - 1):  # skip the top layer\n\t\t\tfpn_out.append(\n\t\t\t\tconv3x3_bn_relu(fpn_out_dim, fpn_out_dim, 1),\n\t\t\t)\n\t\tself.fpn_out = nn.ModuleList(fpn_out)\n\t\tself.fpn_fuse = nn.Sequential(\n\t\t\tconv3x3_bn_relu(len(in_dim) * fpn_out_dim, fpn_out_dim, 1),\n\t\t)\n\t\n\tdef forward(self, feature):\n\t\tf = feature[-1]\n\t\tpool_shape = f.shape[2:]\n\t\t\n\t\tppm_out = [f]\n\t\tfor pool, conv in zip(self.ppm_pooling, self.ppm_conv):\n\t\t\tp = pool(f)\n\t\t\tp = F.interpolate(p, size=pool_shape, mode='bilinear', align_corners=False)\n\t\t\tp = conv(p)\n\t\t\tppm_out.append(p)\n\t\tppm_out = torch.cat(ppm_out, 1)\n\t\tdown = self.ppm_out(ppm_out)\n\t\t\n\t\t\n\t\t#--------------------------------------\n\t\tfpn_out = [down]\n\t\tfor i in reversed(range(len(feature) - 1)):\n\t\t\tlateral = feature[i]\n\t\t\tlateral = self.fpn_in[i](lateral) # lateral branch\n\t\t\tdown = F.interpolate(down, size=lateral.shape[2:], mode='bilinear', align_corners=False) # top-down branch\n\t\t\tdown = down + lateral\n\t\t\tfpn_out.append(self.fpn_out[i](down))\n\t\t\n\t\tfpn_out.reverse() # [P2 - P5]\n\t\tfusion_shape = fpn_out[0].shape[2:]\n\t\tfusion = [fpn_out[0]]\n\t\tfor i in range(1, len(fpn_out)):\n\t\t\tfusion.append(\n\t\t\t\tF.interpolate( fpn_out[i], fusion_shape, mode='bilinear', align_corners=False)\n\t\t\t)\n\t\tx = self.fpn_fuse( torch.cat(fusion, 1))\n\t\t\n\t\treturn x, fusion","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-08-31T19:30:47.279741Z","iopub.execute_input":"2022-08-31T19:30:47.280230Z","iopub.status.idle":"2022-08-31T19:30:47.328304Z","shell.execute_reply.started":"2022-08-31T19:30:47.280187Z","shell.execute_reply":"2022-08-31T19:30:47.326736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn, Tensor\nfrom torchvision.ops import StochasticDepth\n\n\nclass LayerNorm(nn.Module):\n    \"\"\"Channel first layer norm\n    \"\"\"\n    def __init__(self, normalized_shape, eps=1e-6) -> None:\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(normalized_shape))\n        self.bias = nn.Parameter(torch.zeros(normalized_shape))\n        self.eps = eps\n\n    def forward(self, x: Tensor) -> Tensor:\n        u = x.mean(1, keepdim=True)\n        s = (x - u).pow(2).mean(1, keepdim=True)\n        x = (x - u) / torch.sqrt(s + self.eps)\n        x = self.weight[:, None, None] * x + self.bias[:, None, None]\n        return x\n\n    \nclass Block(nn.Module):\n    def __init__(self, dim, dpr=0., init_value=1e-6):\n        super().__init__()\n        self.dwconv = nn.Conv2d(dim, dim, 7, 1, 3, groups=dim)\n        self.norm = nn.LayerNorm(dim, eps=1e-6)\n        self.pwconv1 = nn.Linear(dim, 4*dim)\n        self.act = nn.GELU()\n        self.pwconv2 = nn.Linear(4*dim, dim)\n        self.gamma = nn.Parameter(init_value * torch.ones((dim)), requires_grad=True) if init_value > 0 else None\n        self.drop_path = StochasticDepth(dpr, mode=\"batch\") if dpr > 0. else nn.Identity()\n\n    def forward(self, x: Tensor) -> Tensor:\n        input = x\n        x = self.dwconv(x)\n        x = x.permute(0, 2, 3, 1)   # NCHW to NHWC\n        x = self.norm(x)\n        x = self.pwconv1(x)\n        x = self.act(x)\n        x = self.pwconv2(x)\n\n        if self.gamma is not None:\n            x = self.gamma * x\n        \n        x = x.permute(0, 3, 1, 2)\n        x = input + self.drop_path(x)\n        return x\n\n\nclass Stem(nn.Sequential):\n    def __init__(self, c1, c2, k, s):\n        super().__init__(\n            nn.Conv2d(c1, c2, k, s),\n            LayerNorm(c2)\n        )\n\n\nclass Downsample(nn.Sequential):\n    def __init__(self, c1, c2, k, s):\n        super().__init__(\n            LayerNorm(c1),\n            nn.Conv2d(c1, c2, k, s)\n        )\n\n\nconvnext_settings = {\n    'T': [[3, 3, 9, 3], [96, 192, 384, 768], 0.0],       # [depths, dims, dpr]\n    'S': [[3, 3, 27, 3], [96, 192, 384, 768], 0.0],\n    'B': [[3, 3, 27, 3], [128, 256, 512, 1024], 0.0]\n}\n\n\nclass ConvNeXt(nn.Module):     \n    def __init__(self, model_name: str = 'T') -> None:\n        super().__init__()\n        assert model_name in convnext_settings.keys(), f\"ConvNeXt model name should be in {list(convnext_settings.keys())}\"\n        depths, embed_dims, drop_path_rate = convnext_settings[model_name]\n        self.channels = embed_dims\n    \n        self.downsample_layers = nn.ModuleList([\n            Stem(3, embed_dims[0], 4, 4),\n            *[Downsample(embed_dims[i], embed_dims[i+1], 2, 2) for i in range(3)]\n        ])\n\n        self.stages = nn.ModuleList()\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]\n        cur = 0\n\n        for i in range(4):\n            stage = nn.Sequential(*[\n                Block(embed_dims[i], dpr[cur+j])\n            for j in range(depths[i])])\n            self.stages.append(stage)\n            cur += depths[i]\n\n        for i in range(4):\n            self.add_module(f\"norm{i}\", LayerNorm(embed_dims[i]))\n\n    def forward(self, x: Tensor):\n        outs = []\n\n        for i in range(4):\n            x = self.downsample_layers[i](x)\n            x = self.stages[i](x)\n            norm_layer = getattr(self, f\"norm{i}\")\n            outs.append(norm_layer(x))\n        return outs\n\n\nif __name__ == '__main__':\n    model = ConvNeXt('B')\n\n    x = torch.randn(1, 3, 224, 224)\n    feats = model(x)\n    for y in feats:\n        print(y.shape)","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-08-31T19:30:47.330003Z","iopub.execute_input":"2022-08-31T19:30:47.330754Z","iopub.status.idle":"2022-08-31T19:30:50.796413Z","shell.execute_reply.started":"2022-08-31T19:30:47.330708Z","shell.execute_reply":"2022-08-31T19:30:50.794995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MixUpScaler(nn.Module):\n    def __init__(self, scale_factor=2):\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(x, scale_factor=self.scale_factor, mode='bilinear', align_corners=False) \\\n            + (1 - self.mixing) * F.interpolate(x, scale_factor=self.scale_factor, mode='bicubic')\n        return x\n\n\n\nclass ConvNextUperNet(nn.Module):\n    def load_pretrain(self):\n        checkpoint = '/content/convnext_base_22k_224.pth'\n\n        self.arch = 'convnext_large_22k_224'\n\n        print(f'loading [{checkpoint}]...')\n        checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)['model']\n        self.encoder.load_state_dict(checkpoint, strict=False)  # True\n        print(f'weights for [{self.arch}] loaded successfully!')\n\n\n    def __init__(self):\n        super(ConvNextUperNet, self).__init__()\n        # mixup scaling [https://www.ncbi.nlm.nih.gov/pmc/articles/PMC7924688/]\n        self.scale = MixUpScaler(scale_factor=4)  \n\n        # swin encoder\n        self.encoder = ConvNeXt('B')\n        encoder_dim = [128, 256, 512, 1024]\n        \n        # upernet decoder\n        self.decoder = UPerDecoder(\n            in_dim=encoder_dim,\n            ppm_pool_scale=[1, 2, 3, 6],\n            ppm_dim=512,\n            fpn_out_dim=256\n        )\n        \n        self.logit = nn.Sequential(\n            nn.Conv2d(256, 1, kernel_size=1)\n        )\n        self.aux = nn.ModuleList([\n            nn.Conv2d(256, 1, kernel_size=1, padding=0) for i in range(4)\n        ])\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n\n        # encoder branch\n        x = self.encoder(x)\n        \n        # decoder branch\n        x, _ = self.decoder(x)\n        \n        # out\n        logit = self.logit(x)\n        x = self.scale(logit)\n            \n        return x\n\n    \nconvnext_upernet = ConvNextUperNet()\n    \nif __name__ == '__main__':\n    model = ConvNextUperNet()\n\n    x = torch.randn(1, 3, 616, 616)\n    \n    out = model(x)\n    print(out.shape)","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:30:50.798399Z","iopub.execute_input":"2022-08-31T19:30:50.801823Z","iopub.status.idle":"2022-08-31T19:31:00.791196Z","shell.execute_reply.started":"2022-08-31T19:30:50.801772Z","shell.execute_reply":"2022-08-31T19:31:00.789880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from functools import partial\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom timm.models.layers import trunc_normal_, DropPath\nimport os\nimport sys\nimport torch.fft\nimport traceback\n\nimport torch.utils.checkpoint as checkpoint\n\n\n# if 'DWCONV_IMPL' in os.environ:\n#     try:\n#         sys.path.append(os.environ['DWCONV_IMPL'])\n#         from depthwise_conv2d_implicit_gemm import DepthWiseConv2dImplicitGEMM\n#\n#\n        # def get_dwconv(dim, kernel, bias):\n        #     return DepthWiseConv2dImplicitGEMM(dim, kernel, bias)\n#         # print('Using Megvii large kernel dw conv impl')\n#     except:\n#         print(traceback.format_exc())\n#\n#\ndef get_dwconv(dim, kernel, bias):\n    return nn.Conv2d(dim, dim, kernel_size=kernel, padding=(kernel - 1) // 2, bias=bias, groups=dim)\n\n\n#\n#         # print('[fail to use Megvii Large kernel] Using PyTorch large kernel dw conv impl')\n# else:\n#     def get_dwconv(dim, kernel, bias):\n#         return nn.Conv2d(dim, dim, kernel_size=kernel, padding=(kernel - 1) // 2, bias=bias, groups=dim)\n#\n#     # print('Using PyTorch large kernel dw conv impl')\n\n\nclass GlobalLocalFilter(nn.Module):\n    def __init__(self, dim, h=14, w=8):\n        super().__init__()\n        self.dw = nn.Conv2d(dim // 2, dim // 2, kernel_size=3, padding=1, bias=False, groups=dim // 2)\n        self.complex_weight = nn.Parameter(torch.randn(dim // 2, h, w, 2, dtype=torch.float32) * 0.02)\n        trunc_normal_(self.complex_weight, std=.02)\n        self.pre_norm = LayerNorm(dim, eps=1e-6, data_format='channels_first')\n        self.post_norm = LayerNorm(dim, eps=1e-6, data_format='channels_first')\n\n    def forward(self, x):\n        x = self.pre_norm(x)\n        x1, x2 = torch.chunk(x, 2, dim=1)\n        x1 = self.dw(x1)\n\n        x2 = x2.to(torch.float32)\n        B, C, a, b = x2.shape\n        x2 = torch.fft.rfft2(x2, dim=(2, 3), norm='ortho')\n\n        weight = self.complex_weight\n        if not weight.shape[1:3] == x2.shape[2:4]:\n            weight = F.interpolate(weight.permute(3, 0, 1, 2), size=x2.shape[2:4], mode='bilinear',\n                                   align_corners=True).permute(1, 2, 3, 0)\n\n        weight = torch.view_as_complex(weight.contiguous())\n\n        x2 = x2 * weight\n        x2 = torch.fft.irfft2(x2, s=(a, b), dim=(2, 3), norm='ortho')\n\n        x = torch.cat([x1.unsqueeze(2), x2.unsqueeze(2)], dim=2).reshape(B, 2 * C, a, b)\n        x = self.post_norm(x)\n        return x\n\n\nclass gnconv(nn.Module):\n    def __init__(self, dim, order=5, gflayer=None, h=14, w=8, s=1.0):\n        super().__init__()\n        self.order = order\n        self.dims = [dim // 2 ** i for i in range(order)]\n        self.dims.reverse()\n        self.proj_in = nn.Conv2d(dim, 2 * dim, 1)\n\n        if gflayer is None:\n            self.dwconv = get_dwconv(sum(self.dims), 7, True)\n        else:\n            self.dwconv = gflayer(sum(self.dims), h=h, w=w)\n\n        self.proj_out = nn.Conv2d(dim, dim, 1)\n\n        self.pws = nn.ModuleList(\n            [nn.Conv2d(self.dims[i], self.dims[i + 1], 1) for i in range(order - 1)]\n        )\n\n        self.scale = s\n\n        # print('[gconv]', order, 'order with dims=', self.dims, 'scale=%.4f' % self.scale)\n\n    def forward(self, x, mask=None, dummy=False):\n        B, C, H, W = x.shape\n\n        fused_x = self.proj_in(x)\n        pwa, abc = torch.split(fused_x, (self.dims[0], sum(self.dims)), dim=1)\n\n        dw_abc = self.dwconv(abc) * self.scale\n\n        dw_list = torch.split(dw_abc, self.dims, dim=1)\n        x = pwa * dw_list[0]\n\n        for i in range(self.order - 1):\n            x = self.pws[i](x) * dw_list[i + 1]\n\n        x = self.proj_out(x)\n\n        return x\n\n\nclass Block(nn.Module):\n    r\"\"\" HorNet block\n    \"\"\"\n\n    def __init__(self, dim, drop_path=0., layer_scale_init_value=1e-6, gnconv=gnconv):\n        super().__init__()\n\n        self.norm1 = LayerNorm(dim, eps=1e-6, data_format='channels_first')\n        self.gnconv = gnconv(dim)  # depthwise conv\n        self.norm2 = LayerNorm(dim, eps=1e-6)\n        self.pwconv1 = nn.Linear(dim, 4 * dim)  # pointwise/1x1 convs, implemented with linear layers\n        self.act = nn.GELU()\n        self.pwconv2 = nn.Linear(4 * dim, dim)\n\n        self.gamma1 = nn.Parameter(layer_scale_init_value * torch.ones(dim),\n                                   requires_grad=True) if layer_scale_init_value > 0 else None\n\n        self.gamma2 = nn.Parameter(layer_scale_init_value * torch.ones((dim)),\n                                   requires_grad=True) if layer_scale_init_value > 0 else None\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        if self.gamma1 is not None:\n            gamma1 = self.gamma1.view(C, 1, 1)\n        else:\n            gamma1 = 1\n        x = x + self.drop_path(gamma1 * self.gnconv(self.norm1(x)))\n\n        input = x\n        x = x.permute(0, 2, 3, 1)  # (N, C, H, W) -> (N, H, W, C)\n        x = self.norm2(x)\n        x = self.pwconv1(x)\n        x = self.act(x)\n        x = self.pwconv2(x)\n        if self.gamma2 is not None:\n            x = self.gamma2 * x\n        x = x.permute(0, 3, 1, 2)  # (N, H, W, C) -> (N, C, H, W)\n\n        x = input + self.drop_path(x)\n        return x\n\n\nclass HorNet(nn.Module):\n    r\"\"\" HorNet\n        A PyTorch impl of : `HorNet: Efficient High-Order Spatial Interactions with Recursive Gated Convolutions`\n    Args:\n        in_chans (int): Number of input image channels. Default: 3\n        num_classes (int): Number of classes for classification head. Default: 1000\n        depths (tuple(int)): Number of blocks at each stage. Default: [3, 3, 9, 3]\n        dims (int): Feature dimension at each stage. Default: [96, 192, 384, 768]\n        drop_path_rate (float): Stochastic depth rate. Default: 0.\n        layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.\n        head_init_scale (float): Init scaling value for classifier weights and biases. Default: 1.\n    \"\"\"\n\n    def __init__(self, in_chans=3, num_classes=1000,\n                 depths=[2, 3, 18, 2], base_dim=128, drop_path_rate=0.4,\n                 layer_scale_init_value=1e-6, head_init_scale=1.,\n                 gnconv=gnconv, block=Block, out_indices=[0, 1, 2, 3],\n                 pretrained=None,\n                 use_checkpoint=False,\n                 ):\n        super().__init__()\n        depths = [2, 3, 18, 2]\n        base_dim = 128\n#         gnconv = [\n#                      'partial(gnconv, order=2, s=1/3)',\n#                      'partial(gnconv, order=3, s=1/3)',\n#                      'partial(gnconv, order=4, s=1/3)',\n#                      'partial(gnconv, order=5, s=1/3)',\n#                  ]\n\n        gnconv = [\n            'partial(gnconv, order=2, s=1/3)',\n            'partial(gnconv, order=3, s=1/3)',\n            'partial(gnconv, order=4, s=1/3, h=14, w=8, gflayer=GlobalLocalFilter)',\n            'partial(gnconv, order=5, s=1/3, h=7, w=4, gflayer=GlobalLocalFilter)',\n        ]\n        drop_path_rate = 0.4\n        out_indices = [0, 1, 2, 3]\n\n        self.out_indices = out_indices\n        self.pretrained = pretrained\n        self.use_checkpoint = use_checkpoint\n\n        dims = [base_dim, base_dim * 2, base_dim * 4, base_dim * 8]\n        self.embed_dims = dims\n\n        self.downsample_layers = nn.ModuleList()  # stem and 3 intermediate downsampling conv layers\n        stem = nn.Sequential(\n            nn.Conv2d(in_chans, dims[0], kernel_size=4, stride=4),\n            LayerNorm(dims[0], eps=1e-6, data_format=\"channels_first\")\n        )\n        self.downsample_layers.append(stem)\n        for i in range(3):\n            downsample_layer = nn.Sequential(\n                LayerNorm(dims[i], eps=1e-6, data_format=\"channels_first\"),\n                nn.Conv2d(dims[i], dims[i + 1], kernel_size=2, stride=2),\n            )\n            self.downsample_layers.append(downsample_layer)\n\n        self.stages = nn.ModuleList()  # 4 feature resolution stages, each consisting of multiple residual blocks\n        dp_rates = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]\n\n        if not isinstance(gnconv, list):\n            gnconv = [gnconv, gnconv, gnconv, gnconv]\n        else:\n            gnconv = gnconv\n            assert len(gnconv) == 4\n\n        if isinstance(gnconv[0], str):\n            # print('[GConvNet]: convert str gconv to func')\n            gnconv = [eval(g) for g in gnconv]\n\n        if isinstance(block, str):\n            block = eval(block)\n\n        cur = 0\n        for i in range(4):\n            stage = nn.Sequential(\n                *[block(dim=dims[i], drop_path=dp_rates[cur + j],\n                        layer_scale_init_value=layer_scale_init_value, gnconv=gnconv[i]) for j in range(depths[i])]\n            )\n            self.stages.append(stage)\n            cur += depths[i]\n\n        norm_layer = partial(LayerNorm, eps=1e-6, data_format=\"channels_first\")\n        for i_layer in range(4):\n            layer = norm_layer(dims[i_layer])\n            layer_name = f'norm{i_layer}'\n            self.add_module(layer_name, layer)\n\n    def init_weights(self):\n        \"\"\"Initialize the weights in backbone.\n        Args:\n            pretrained (str, optional): Path to pre-trained weights.\n                Defaults to None.\n        \"\"\"\n        pretrained = self.pretrained\n\n        def _init_weights(m):\n            if isinstance(m, nn.Linear):\n                trunc_normal_(m.weight, std=.02)\n                if isinstance(m, nn.Linear) and m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.LayerNorm):\n                nn.init.constant_(m.bias, 0)\n                nn.init.constant_(m.weight, 1.0)\n\n        # if isinstance(pretrained, str):\n        #     self.apply(_init_weights)\n        #     logger = get_root_logger()\n        #     load_checkpoint(self, pretrained, strict=False, logger=logger)\n        # elif pretrained is None:\n        #     raise NotImplementedError()\n        #     self.apply(_init_weights)\n        # else:\n        #     raise TypeError('pretrained must be a str or None')\n\n    def forward_features(self, x):\n        outs = []\n        for i in range(4):\n            x = self.downsample_layers[i](x)\n            if self.use_checkpoint:\n                x = checkpoint.checkpoint_sequential(self.stages[i], 2, x)\n            else:\n                x = self.stages[i](x)\n            if i in self.out_indices:\n                norm_layer = getattr(self, f'norm{i}')\n                x_out = norm_layer(x)\n                outs.append(x_out)\n        return tuple(outs)\n\n    def forward(self, x):\n        x = self.forward_features(x)\n        return x\n\n\nclass LayerNorm(nn.Module):\n    r\"\"\" LayerNorm that supports two data formats: channels_last (default) or channels_first.\n    The ordering of the dimensions in the inputs. channels_last corresponds to inputs with\n    shape (batch_size, height, width, channels) while channels_first corresponds to inputs\n    with shape (batch_size, channels, height, width).\n    \"\"\"\n\n    def __init__(self, normalized_shape, eps=1e-6, data_format=\"channels_last\"):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(normalized_shape))\n        self.bias = nn.Parameter(torch.zeros(normalized_shape))\n        self.eps = eps\n        self.data_format = data_format\n        if self.data_format not in [\"channels_last\", \"channels_first\"]:\n            raise NotImplementedError\n        self.normalized_shape = (normalized_shape,)\n\n    def forward(self, x):\n        if self.data_format == \"channels_last\":\n            return F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps)\n        elif self.data_format == \"channels_first\":\n            u = x.mean(1, keepdim=True)\n            s = (x - u).pow(2).mean(1, keepdim=True)\n            x = (x - u) / torch.sqrt(s + self.eps)\n            x = self.weight[:, None, None] * x + self.bias[:, None, None]\n            return x\n\n\nimport numpy as np\nimport torch\nfrom torch import nn\n\n\nclass HorNetUperNet(nn.Module):\n    def load_pretrain(self, arch='L'):\n        checkpoints = {\n            'B': f'{cfg.work_dir}/models/HuBMAP_HPA_2022/pretrained/hornet/hornet_base_7x7.pth',\n            # 'L': f'{cfg.work_dir}/models/HuBMAP_HPA_2022/pretrained/convnext/convnext_large_22k_224.pth'\n        }\n\n        self.arch = 'hornet_base_7x7'\n\n        print(f'loading [{checkpoints[arch]}]...')\n        checkpoint = torch.load(checkpoints[arch], map_location=lambda storage, loc: storage)['model']\n        self.encoder.load_state_dict(checkpoint, strict=False)  # True\n        print(f'weights for [{self.arch}] loaded successfully!')\n\n    def __init__(self):\n        super(HorNetUperNet, self).__init__()\n        self.scale = MixUpScaler(scale_factor=4)\n\n        # convnext encoder\n        self.encoder = HorNet()\n        encoder_dim = self.encoder.embed_dims\n\n#         try:\n#             self.load_pretrain('B')\n#         except FileNotFoundError:\n#             print(\"Couldn't load pretrained weights!\")\n\n        # upernet decoder\n        self.decoder = UPerDecoder(\n            in_dim=encoder_dim,\n            ppm_pool_scale=[1, 2, 3, 6],\n            ppm_dim=512,\n            fpn_out_dim=256\n        )\n\n        self.logit = nn.Sequential(\n            nn.Conv2d(256, 1, kernel_size=1)\n        )\n\n        self.pre_scale_head = nn.ModuleList([\n            nn.Conv2d(256, 1, kernel_size=1, padding=0) for _ in range(4)\n        ])\n        self.norm = nn.LayerNorm(self.encoder.embed_dims[-1], eps=1e-6)  # final norm layer\n        self.cls_head = nn.Linear(self.encoder.embed_dims[-1], 5)\n\n    def forward(self, x):\n\n        B, C, H, W = x.shape\n\n        # encoder branch\n        x = self.encoder(x)\n\n        # decoder branch\n        x, decoder = self.decoder(x)\n\n        # out\n        logit = self.logit(x)\n        x = self.scale(logit)\n        return x\n\n\nhornet_upernet = HorNetUperNet()\n\n    \nif __name__ == '__main__':\n    model = HorNetUperNet()\n\n    x = torch.randn(1, 3, 224, 224)\n    out = model(x)\n    print(out.shape)","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-08-31T19:31:00.793371Z","iopub.execute_input":"2022-08-31T19:31:00.794079Z","iopub.status.idle":"2022-08-31T19:31:04.601738Z","shell.execute_reply.started":"2022-08-31T19:31:00.794036Z","shell.execute_reply":"2022-08-31T19:31:04.600258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from copy import deepcopy\n\ndef build_model(arch):\n    models = {\n        'convnext_upernet': deepcopy(convnext_upernet),\n        'hornet_upernet': deepcopy(hornet_upernet),\n    }\n\n    model = models[arch]\n    model.to(cfg.device)\n    \n    return model\n\ndef load_model(path, arch):\n    model = build_model(arch)\n    model.load_state_dict(torch.load(path, map_location=cfg.device))\n    model.eval()\n    print('- weights loaded!')\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:31:04.604146Z","iopub.execute_input":"2022-08-31T19:31:04.604952Z","iopub.status.idle":"2022-08-31T19:31:04.615283Z","shell.execute_reply.started":"2022-08-31T19:31:04.604892Z","shell.execute_reply":"2022-08-31T19:31:04.614129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_models = []\n\nfor model in MODELS:\n    print(f'Model: [{model}]')\n    \n    _model = load_model(\n        MODELS[model]['weights'],\n        MODELS[model]['arch'],\n    )\n    _weight = MODELS[model]['weight']\n    \n    all_models.append({\n        'model':  _model,\n        'weight': _weight\n    })\n    \n    print()","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:31:04.617735Z","iopub.execute_input":"2022-08-31T19:31:04.618313Z","iopub.status.idle":"2022-08-31T19:32:45.946571Z","shell.execute_reply.started":"2022-08-31T19:31:04.618271Z","shell.execute_reply":"2022-08-31T19:32:45.945215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LUNG_MODELS.update(MODELS)\n# len(LUNG_MODELS)","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:32:45.952500Z","iopub.execute_input":"2022-08-31T19:32:45.953225Z","iopub.status.idle":"2022-08-31T19:32:45.963933Z","shell.execute_reply.started":"2022-08-31T19:32:45.953178Z","shell.execute_reply":"2022-08-31T19:32:45.962527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lung_models = []\n\nfor model in LUNG_MODELS:\n    print(f'Model: [{model}]')\n    \n    _model = load_model(\n        LUNG_MODELS[model]['weights'],\n        LUNG_MODELS[model]['arch']\n    )\n    _weight = LUNG_MODELS[model]['weight']\n    \n    lung_models.append({\n        'model':  _model,\n        'weight': _weight\n    })\n    \n    print()","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:32:45.968751Z","iopub.execute_input":"2022-08-31T19:32:45.969814Z","iopub.status.idle":"2022-08-31T19:33:14.556515Z","shell.execute_reply.started":"2022-08-31T19:32:45.969765Z","shell.execute_reply":"2022-08-31T19:33:14.555139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def unfold(input, size=256, step=256):\n    c = input.shape[0]\n    patches = input.unfold(1, size, step).unfold(2, size, step)   \n    patches = patches.contiguous().view(c, -1, size, size)         \n    patches = patches.permute(1, 0, 2, 3) \n\n    return patches\n\ndef fold(patches, out_size, size=256, step=256):\n    c = patches.shape[1]\n\n    patches = patches.contiguous().transpose(1, 0).view(1, c, -1, size*size)  # [B, C, n_patches, kernel_size*kernel_size]\n    patches = patches.permute(1, 0, 3, 2)                                     # [B, C, kernel_size*kernel_size, n_patches]\n    patches = patches.contiguous().view(1, c*size*size, -1)                   # (B, C*kernel*kernel, n_patches)\n\n    weight_mask = F.fold(torch.ones_like(patches), \n                           output_size=out_size, \n                           kernel_size=size, stride=step)\n\n    out = F.fold(patches, out_size, \n                   kernel_size=size, stride=step)                             # [B, C, H, W]\n    out = out / weight_mask\n\n    return out","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:14.559069Z","iopub.execute_input":"2022-08-31T19:33:14.560216Z","iopub.status.idle":"2022-08-31T19:33:14.570652Z","shell.execute_reply.started":"2022-08-31T19:33:14.560171Z","shell.execute_reply":"2022-08-31T19:33:14.569344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HuBMAP_Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, normalization):\n        self.df = df\n        self.normalization = normalization\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        empty = False\n        data_source = self.df.iloc[idx]['data_source']\n        organ_type = self.df.iloc[idx]['organ']\n        \n        image_id = self.img_id(idx)\n        image = self.load_image(image_id)\n        \n        image_shape = image.shape[:2]\n\n        # reshape\n#         image = cv2.resize(image, tuple(self.img_size), cv2.INTER_AREA)    # no pixel normalization\n\n        # -- pixel size normalization v0 --\n        pixel_size = self.df.iloc[idx]['pixel_size']\n            \n        if cfg.tiles:\n            norm_size = int(image_shape[0] * pixel_size * (768 / (3000 * 0.4)))\n            \n            # decoder branch scale shifting alining\n            _shift = norm_size % 16\n            norm_size -= _shift\n        \n            cfg.img_size = [norm_size, norm_size]\n            cfg.tile_size = cfg.img_size[0] // 2\n            cfg.stride = (nearest(cfg.img_size[0], cfg.tile_size) - cfg.img_size[0]) // 4\n            \n        else:\n            if organ_type == 'lung':     # resize to lung train size\n                norm_size = int(image_shape[0] * pixel_size * (cfg.img_size[0] / (3000 * 0.4)))\n            else:\n                norm_size = int(image_shape[0] * pixel_size * (cfg.img_size[0] / (3000 * 0.4)))\n            \n            # decoder branch scale shifting alining\n            _shift = norm_size % 4\n            norm_size -= _shift\n\n        image = cv2.resize(image, (norm_size, norm_size))\n        \n        # check if image is empty\n        if self.empty_image(image):\n            empty = True\n        \n        if self.normalization:\n            image = self.normalization(image)\n        \n        # to tensor\n        image = np.transpose(image, (2, 0, 1))\n        image = torch.Tensor(image)\n        \n        # unfold\n        if cfg.tiles:\n            image = self.tiles(image, size=cfg.tile_size, step=cfg.stride)\n        \n        return image, image_shape, image_id, not empty, data_source, organ_type\n    \n    def tiles(self, image, size=256, step=256):\n        img_patches = unfold(image, size, step)\n        \n        return img_patches\n    \n    def load_image(self, image_id):\n        # get image\n        image = cv2.imread(f'{cfg.test_images}/{image_id}.tiff')\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n#         image = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)\n        \n        return image\n    \n    def img_id(self, idx):\n        return self.df.iloc[idx]['id']\n    \n    @staticmethod\n    def empty_image(image):\n        s_th = 40                                     # saturation blancking threshold\n        p_th = 1000 * (cfg.tile_size // 256) ** 2     # threshold for the minimum number of pixels\n        \n        # check for empty imges\n        hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)\n        h, s, v = cv2.split(hsv)\n        \n        if (s > s_th).sum() <= p_th or image.sum() <= p_th:\n            return True\n        return False","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:14.572496Z","iopub.execute_input":"2022-08-31T19:33:14.573124Z","iopub.status.idle":"2022-08-31T19:33:14.595548Z","shell.execute_reply.started":"2022-08-31T19:33:14.573081Z","shell.execute_reply":"2022-08-31T19:33:14.593334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sample = pd.read_csv('../input/hubmap-organ-segmentation/test.csv')","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:14.597773Z","iopub.execute_input":"2022-08-31T19:33:14.598672Z","iopub.status.idle":"2022-08-31T19:33:14.642733Z","shell.execute_reply.started":"2022-08-31T19:33:14.598624Z","shell.execute_reply":"2022-08-31T19:33:14.641566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# normalization = lambda x: normalize_mean_std(normalize_0_255(x))    # <- best?\n# normalization = lambda x: normalize_0_255(normalize_mean_std(x))\n# normalization = lambda x: normalize_mean_std(x)\n# normalization = lambda x: normalize_0_255(x) \n# normalization = lambda x: normalize_max(normalize_mean_std(normalize_0_255(x)))\n\n# normalization = lambda x: normalize_max(x)\n# normalization = lambda x: normalize_255_mean_std(x)\nnormalization = lambda x: normalize_255_mean_std_v2(x)\n\n\ntest_dataset = HuBMAP_Dataset(df_sample, normalization)\nimage, shape, _, empty, data_source, organ_type = test_dataset[0]\n\nprint(f'original image shape: \\t\\t{shape}')\nprint(f'image is not empty: \\t\\t{empty}')\nprint(f'tiled image shape: \\t\\t{image.shape}')\nprint(f'image data source: \\t\\t{data_source}')\nprint(f'image oragan type: \\t\\t{organ_type}')\n\nif cfg.tiles:\n    # merging the tiled mask\n    merged_image = fold(image, out_size=cfg.img_size, size=cfg.tile_size, step=cfg.stride)\n    print(f'reconstructed image shape: \\t{merged_image.shape}')\n    print()\n    print(image.min(), image.max())\n    print()\n\n    plot_tiled_image(image, [5, 5])\n\nif cfg.tiles:\n    visualize(\n        [5, 5],\n        recon_img=np.transpose(merged_image[0, ...], (1, 2, 0))\n    )\nelse:\n    visualize(\n        [5, 5],\n        image=np.transpose(image, (1, 2, 0)),\n    )","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:14.644770Z","iopub.execute_input":"2022-08-31T19:33:14.645217Z","iopub.status.idle":"2022-08-31T19:33:15.537377Z","shell.execute_reply.started":"2022-08-31T19:33:14.645176Z","shell.execute_reply":"2022-08-31T19:33:15.536123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataloader = DataLoader(test_dataset, batch_size=1, \n                             num_workers=2, shuffle=False, pin_memory=True, drop_last=False)\n\n# image, image_shape, image_id  = next(iter(test_dataloader))\n\n# if cfg.tiles: \n#     bs, n_tiles, c, h, w = image.size()\n#     image = image.contiguous().view(-1, c, h, w)\n            \n# print(image.shape, image_shape[0].cpu().detach().numpy()[0], image_shape[1].cpu().detach().numpy()[0])","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:15.544686Z","iopub.execute_input":"2022-08-31T19:33:15.545493Z","iopub.status.idle":"2022-08-31T19:33:15.554756Z","shell.execute_reply.started":"2022-08-31T19:33:15.545445Z","shell.execute_reply":"2022-08-31T19:33:15.553295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import scipy.ndimage as ndi\nimport skimage.morphology as morph\nfrom skimage.filters import threshold_otsu\n\ndef clean_mask(mask, kernel=(3, 3)):\n    structure = np.ones(kernel)\n    mask = ndi.binary_fill_holes(mask, structure=structure)\n    \n    return mask","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:15.557372Z","iopub.execute_input":"2022-08-31T19:33:15.558578Z","iopub.status.idle":"2022-08-31T19:33:15.928443Z","shell.execute_reply.started":"2022-08-31T19:33:15.558303Z","shell.execute_reply":"2022-08-31T19:33:15.927177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pred_single_mask(input, model, thresh=0.5):\n    pred = None \n    with torch.no_grad():\n        pred = model(input)\n        pred = (nn.Sigmoid()(pred)).double()    # fixed double thr problem\n\n    return pred\n\n\ndef pred_mask(input, models, m_thresh=0.5):\n    \"\"\"\n    predicting mask of each model (no tta)\n    \n    :input  : imput tensor of shape (b, c, h, w)\n    :models : input list of loaded models\n    :thresh : threshold used to binarize predicted prob map\n    ---------\n    :return : weighted mean of all masks\n    \"\"\"\n    preds = []\n\n    for model in models:\n        with torch.no_grad():\n            pred = pred_single_mask(input, model, thresh)\n            preds.append(pred)\n\n    preds = torch.mean(torch.stack(preds, dim=0), dim=0).cpu().detach()\n    \n    if cfg.tiles:\n        preds = fold(preds, out_size=cfg.img_size, size=cfg.tile_size, step=cfg.stride)\n    \n    preds = preds.numpy()\n    preds = (preds > m_thresh).astype(np.uint8)  # main fix with np.uint8\n\n    return preds\n\n\ndef pred_mask_tta(input, models, thresh_factor=2.5, avrg_type='weighted_mean'):\n    \"\"\"\n    predicting mask with tta of each model\n    \n    :tta    : flips = [[2], [3], [2, 3]]\n    :input  : imput tensor of shape (b, c, h, w)\n    :models : input list of loaded models\n    :thresh : threshold used to binarize predicted prob map\n    ---------\n    :return : weighted mean of all masks\n    \"\"\"\n    preds = []\n    weights = []\n\n    for model in models:\n        with torch.no_grad():\n            \n            model_weight = model['weight']\n            model = model['model']\n            \n            # sample prediction -- added v1.2\n            pred = model(input)\n            pred = (nn.Sigmoid()(pred)).double()    # fixed double thr problem\n\n            pred *= model_weight\n            preds.append(pred)\n            weights.append(model_weight)\n            \n            # tta\n            flips = [[2], [3], [2, 3]]\n            for d in flips:\n                t_image = torch.flip(input, dims=d)\n                \n                pred = model(t_image)\n                pred = (nn.Sigmoid()(pred)).double()    # fixed double thr problem\n\n                pred = torch.flip(pred, dims=d)\n\n                pred *= model_weight\n                preds.append(pred)\n                weights.append(model_weight)\n\n    # mask avareging\n#     preds = torch.mean(torch.stack(preds, dim=0), dim=0).cpu().detach()\n    \n    # mask weighted avareging\n    if avrg_type == 'weighted_mean':\n        preds = torch.sum(torch.stack(preds, dim=0), dim=0) / np.sum(weights)\n    elif avrg_type == 'sum':\n        preds = torch.sum(torch.stack(preds, dim=0), dim=0) \n        \n    preds = preds.cpu().detach()\n    \n    print(preds.shape)\n    \n    if cfg.tiles:\n        preds = fold(preds, out_size=cfg.img_size, size=cfg.tile_size, step=cfg.stride)\n        \n    preds = preds.numpy()\n    adaptive_thresh = np.mean(preds[preds > 0.02]) / thresh_factor\n    preds = (preds > adaptive_thresh).astype(np.uint8)  # main fix with np.uint8\n\n    return preds","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:15.930411Z","iopub.execute_input":"2022-08-31T19:33:15.931109Z","iopub.status.idle":"2022-08-31T19:33:15.951404Z","shell.execute_reply.started":"2022-08-31T19:33:15.931064Z","shell.execute_reply":"2022-08-31T19:33:15.948962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def disk_struct(radius):\n    x = np.arange(-radius, radius+1)\n    x, y = np.meshgrid(x, x)\n    r = x**2 + y**2\n    struct = r < radius**2\n    return struct\n\ndef refine_mask(pred_mask):\n    struct = disk_struct(5)\n    closing = ndi.binary_opening(pred_mask, struct)\n\n    struct = disk_struct(10)\n    opening = ndi.binary_opening(closing, struct)\n    \n    refined_mask = opening\n    \n    return refined_mask","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:15.953860Z","iopub.execute_input":"2022-08-31T19:33:15.954547Z","iopub.status.idle":"2022-08-31T19:33:15.969736Z","shell.execute_reply.started":"2022-08-31T19:33:15.954501Z","shell.execute_reply":"2022-08-31T19:33:15.968331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"names, preds = [], []\n\nfor i in tqdm(range(len(test_dataset))):\n    # get image\n    image, image_shape, image_id, not_empty, data_source, organ_type = test_dataset[i]\n    # add batch dimention\n    image = image.view(1, *image.shape)  # [C, H, W] -> [B, C, H, W]\n    image = image.to(cfg.device, dtype=torch.float)\n\n    # handle empty images\n#     if not not_empty or organ_type != 'spleen':   # include only spleen\n#     if not not_empty or data_source != 'Hubmap':  # + include Hubmap\n#     if not not_empty or data_source != 'HPA':     # + include HPA\n\n#     if organ_type != 'lung':                        # predict on lung only\n#         rle = ''\n#         names.append(image_id)\n#         preds.append(rle)\n#         continue\n    \n    \n    # get prediction\n    thr_f = cfg.organ_mean_thresh_f[data_source][organ_type]\n    print(f'using thresh_factor: {thr_f}')\n    \n    if organ_type == 'lung':\n        print(f'using lung models')\n    \n        # predict\n        if cfg.tta:\n            pred = pred_mask_tta(image, lung_models, thresh_factor=thr_f, avrg_type='weighted_mean')\n        else:\n            pred = pred_mask(image, lung_models, thresh_factor=thr_f)\n            \n        if np.all(pred == 0):\n            for ii in range(5):\n                thr_f -= 0.2\n                # predict\n                if cfg.tta:\n                    pred = pred_mask_tta(image, lung_models, thresh_factor=thr_f, avrg_type='weighted_mean')\n                else:\n                    pred = pred_mask(image, lung_models, thresh_factor=thr_f)\n            \n    else:\n        print(f'using all models')\n\n        # predict\n        if cfg.tta:\n            pred = pred_mask_tta(image, all_models, thresh_factor=thr_f, avrg_type='weighted_mean')\n        else:\n            pred = pred_mask(image, all_models, thresh_factor=thr_f)\n            \n        if np.all(pred == 0):\n            for ii in range(5):\n                thr_f -= 0.2\n                # predict\n                if cfg.tta:\n                    pred = pred_mask_tta(image, all_models, thresh_factor=thr_f, avrg_type='weighted_mean')\n                else:\n                    pred = pred_mask(image, all_models, thresh_factor=thr_f)\n                    \n#     if organ_type in ['kidney']:\n#         print('applying postprocessing')\n#         pred = refine_mask(pred.squeeze()).astype(np.uint8)\n#         pred = np.expand_dims(pred, 0)\n#         pred = np.expand_dims(pred, 0)\n            \n    # resize back\n    pred = cv2.resize(pred.squeeze(), image_shape)\n    \n    # fill the mask\n#     pred = clean_mask(pred, kernel=(3, 3))\n\n    visualize(\n        [10, 10],\n        pred=pred,\n    )\n    \n    rle = rle_encode(pred)\n    names.append(image_id)\n    preds.append(rle)\n    \n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:15.972603Z","iopub.execute_input":"2022-08-31T19:33:15.973213Z","iopub.status.idle":"2022-08-31T19:33:30.539225Z","shell.execute_reply.started":"2022-08-31T19:33:15.973166Z","shell.execute_reply":"2022-08-31T19:33:30.537937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame({'id': names, 'rle': preds})\ndf.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:30.541127Z","iopub.execute_input":"2022-08-31T19:33:30.541873Z","iopub.status.idle":"2022-08-31T19:33:30.555247Z","shell.execute_reply.started":"2022-08-31T19:33:30.541829Z","shell.execute_reply":"2022-08-31T19:33:30.554025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:30.557786Z","iopub.execute_input":"2022-08-31T19:33:30.558761Z","iopub.status.idle":"2022-08-31T19:33:30.575281Z","shell.execute_reply.started":"2022-08-31T19:33:30.558717Z","shell.execute_reply":"2022-08-31T19:33:30.573640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://stackoverflow.com/questions/62995726/pytorch-sliding-window-with-unfold-fold\n\n\n# v2.0:\n# - added BGR->RGB\n# - fixed double thr problem \n# - norm_255_mean_std\n\n# v2.1\n# - added sample tta\n\n# v2.2\n# - upernet-swin \n# - larger pred size 224 -> 512 (tiles - 1024x512x2) - 0.63\n# v2.2.1\n# - added INTER_AREA to image rescaling before unfold + INTER_CUBIC for final resize to original shape\n\n# v3.1.3\n# - added INTER_LANCZOS4 to image rescaling for final resize to original shape\n\n# - infer size 768 -> lb 0.76\n# - infer size 736 -> lb 0.76+\n# - infer size 704 -> lb 0.77-\n# - infer size 672 -> lb 0.77+\n# - infer size 640 -> \n\n# - infer size 768 + [normalize image shape by its pixel size] -> ","metadata":{"execution":{"iopub.status.busy":"2022-08-31T19:33:30.577951Z","iopub.execute_input":"2022-08-31T19:33:30.579222Z","iopub.status.idle":"2022-08-31T19:33:30.585616Z","shell.execute_reply.started":"2022-08-31T19:33:30.579174Z","shell.execute_reply":"2022-08-31T19:33:30.584461Z"},"trusted":true},"execution_count":null,"outputs":[]}]}