{"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-09-22T13:04:19.038214Z","iopub.execute_input":"2022-09-22T13:04:19.038702Z","iopub.status.idle":"2022-09-22T13:04:20.144465Z","shell.execute_reply.started":"2022-09-22T13:04:19.038658Z","shell.execute_reply":"2022-09-22T13:04:20.142954Z"},"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-09-22T13:04:20.147621Z","iopub.execute_input":"2022-09-22T13:04:20.148453Z","iopub.status.idle":"2022-09-22T13:04:20.153300Z","shell.execute_reply.started":"2022-09-22T13:04:20.148419Z","shell.execute_reply":"2022-09-22T13:04:20.151936Z"},"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-09-22T13:04:20.155083Z","iopub.execute_input":"2022-09-22T13:04:20.156179Z","iopub.status.idle":"2022-09-22T13:04:21.508060Z","shell.execute_reply.started":"2022-09-22T13:04:20.156140Z","shell.execute_reply":"2022-09-22T13:04:21.506903Z"},"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-09-22T13:04:21.511837Z","iopub.execute_input":"2022-09-22T13:04:21.512198Z","iopub.status.idle":"2022-09-22T13:05:02.206008Z","shell.execute_reply.started":"2022-09-22T13:04:21.512164Z","shell.execute_reply":"2022-09-22T13:05:02.204762Z"},"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-09-22T12:55:46.247599Z","iopub.execute_input":"2022-09-22T12:55:46.248462Z","iopub.status.idle":"2022-09-22T12:55:46.253212Z","shell.execute_reply.started":"2022-09-22T12:55:46.248427Z","shell.execute_reply":"2022-09-22T12:55:46.252114Z"},"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-09-22T13:06:29.435281Z","iopub.execute_input":"2022-09-22T13:06:29.435688Z","iopub.status.idle":"2022-09-22T13:06:32.418342Z","shell.execute_reply.started":"2022-09-22T13:06:29.435656Z","shell.execute_reply":"2022-09-22T13:06:32.417225Z"},"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-09-22T13:06:34.667345Z","iopub.execute_input":"2022-09-22T13:06:34.667748Z","iopub.status.idle":"2022-09-22T13:06:34.677994Z","shell.execute_reply.started":"2022-09-22T13:06:34.667702Z","shell.execute_reply":"2022-09-22T13:06:34.676799Z"},"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-09-22T13:06:34.816454Z","iopub.execute_input":"2022-09-22T13:06:34.817018Z","iopub.status.idle":"2022-09-22T13:06:34.828095Z","shell.execute_reply.started":"2022-09-22T13:06:34.816990Z","shell.execute_reply":"2022-09-22T13:06:34.827096Z"},"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-09-22T13:06:35.190722Z","iopub.execute_input":"2022-09-22T13:06:35.191338Z","iopub.status.idle":"2022-09-22T13:06:35.196787Z","shell.execute_reply.started":"2022-09-22T13:06:35.191299Z","shell.execute_reply":"2022-09-22T13:06:35.195917Z"},"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-09-22T13:06:35.344348Z","iopub.execute_input":"2022-09-22T13:06:35.344612Z","iopub.status.idle":"2022-09-22T13:06:35.352018Z","shell.execute_reply.started":"2022-09-22T13:06:35.344588Z","shell.execute_reply":"2022-09-22T13:06:35.350790Z"},"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-09-22T13:06:35.664496Z","iopub.execute_input":"2022-09-22T13:06:35.665477Z","iopub.status.idle":"2022-09-22T13:06:35.724808Z","shell.execute_reply.started":"2022-09-22T13:06:35.665438Z","shell.execute_reply":"2022-09-22T13:06:35.723788Z"},"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-09-22T13:06:35.839198Z","iopub.execute_input":"2022-09-22T13:06:35.839489Z","iopub.status.idle":"2022-09-22T13:06:35.847912Z","shell.execute_reply.started":"2022-09-22T13:06:35.839463Z","shell.execute_reply":"2022-09-22T13:06:35.846757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODELS = {\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-09-22T13:06:36.361823Z","iopub.execute_input":"2022-09-22T13:06:36.362999Z","iopub.status.idle":"2022-09-22T13:06:36.383979Z","shell.execute_reply.started":"2022-09-22T13:06:36.362959Z","shell.execute_reply":"2022-09-22T13:06:36.382991Z"},"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/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': 2\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}","metadata":{"execution":{"iopub.status.busy":"2022-09-22T13:06:36.631170Z","iopub.execute_input":"2022-09-22T13:06:36.631456Z","iopub.status.idle":"2022-09-22T13:06:36.638799Z","shell.execute_reply.started":"2022-09-22T13:06:36.631431Z","shell.execute_reply":"2022-09-22T13:06:36.637462Z"},"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-09-22T13:06:36.824370Z","iopub.execute_input":"2022-09-22T13:06:36.824665Z","iopub.status.idle":"2022-09-22T13:06:36.834847Z","shell.execute_reply.started":"2022-09-22T13:06:36.824640Z","shell.execute_reply":"2022-09-22T13:06:36.833521Z"},"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-09-22T13:06:36.969402Z","iopub.execute_input":"2022-09-22T13:06:36.969659Z","iopub.status.idle":"2022-09-22T13:06:36.986820Z","shell.execute_reply.started":"2022-09-22T13:06:36.969636Z","shell.execute_reply":"2022-09-22T13:06:36.985899Z"},"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-09-22T13:06:37.118709Z","iopub.execute_input":"2022-09-22T13:06:37.119019Z","iopub.status.idle":"2022-09-22T13:06:38.704032Z","shell.execute_reply.started":"2022-09-22T13:06:37.118994Z","shell.execute_reply":"2022-09-22T13:06:38.702877Z"},"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-09-22T13:06:38.707353Z","iopub.execute_input":"2022-09-22T13:06:38.708166Z","iopub.status.idle":"2022-09-22T13:06:46.345324Z","shell.execute_reply.started":"2022-09-22T13:06:38.708125Z","shell.execute_reply":"2022-09-22T13:06:46.344202Z"},"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\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\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":{"execution":{"iopub.status.busy":"2022-09-22T13:06:46.347422Z","iopub.execute_input":"2022-09-22T13:06:46.348075Z","iopub.status.idle":"2022-09-22T13:06:49.411550Z","shell.execute_reply.started":"2022-09-22T13:06:46.348038Z","shell.execute_reply":"2022-09-22T13:06:49.410389Z"},"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-09-22T13:06:49.414542Z","iopub.execute_input":"2022-09-22T13:06:49.414932Z","iopub.status.idle":"2022-09-22T13:06:49.421132Z","shell.execute_reply.started":"2022-09-22T13:06:49.414903Z","shell.execute_reply":"2022-09-22T13:06:49.420005Z"},"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-09-22T13:06:49.422658Z","iopub.execute_input":"2022-09-22T13:06:49.423653Z","iopub.status.idle":"2022-09-22T13:08:18.304154Z","shell.execute_reply.started":"2022-09-22T13:06:49.423595Z","shell.execute_reply":"2022-09-22T13:08:18.302980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LUNG_MODELS.update(MODELS)\n# len(LUNG_MODELS)","metadata":{"execution":{"iopub.status.busy":"2022-09-22T13:08:18.305991Z","iopub.execute_input":"2022-09-22T13:08:18.306368Z","iopub.status.idle":"2022-09-22T13:08:18.311059Z","shell.execute_reply.started":"2022-09-22T13:08:18.306331Z","shell.execute_reply":"2022-09-22T13:08:18.309770Z"},"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-09-22T13:08:18.312795Z","iopub.execute_input":"2022-09-22T13:08:18.313143Z","iopub.status.idle":"2022-09-22T13:08:36.709869Z","shell.execute_reply.started":"2022-09-22T13:08:18.313109Z","shell.execute_reply":"2022-09-22T13:08:36.708712Z"},"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-09-22T13:08:36.711426Z","iopub.execute_input":"2022-09-22T13:08:36.711806Z","iopub.status.idle":"2022-09-22T13:08:36.720902Z","shell.execute_reply.started":"2022-09-22T13:08:36.711771Z","shell.execute_reply":"2022-09-22T13:08:36.719085Z"},"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 shift 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            elif organ_type == 'prostate':\n                norm_size = int(image_shape[0] * pixel_size * (cfg.img_size[0] / (3000 * 0.4)))\n            else:\n                norm_size = cfg.img_size[0]\n            \n            # decoder branch scale shift 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-09-22T13:08:36.722596Z","iopub.execute_input":"2022-09-22T13:08:36.723058Z","iopub.status.idle":"2022-09-22T13:08:36.740036Z","shell.execute_reply.started":"2022-09-22T13:08:36.723014Z","shell.execute_reply":"2022-09-22T13:08:36.739059Z"},"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-09-22T13:08:36.745203Z","iopub.execute_input":"2022-09-22T13:08:36.745460Z","iopub.status.idle":"2022-09-22T13:08:36.779805Z","shell.execute_reply.started":"2022-09-22T13:08:36.745437Z","shell.execute_reply":"2022-09-22T13:08:36.778939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalization = 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-09-22T13:08:36.781090Z","iopub.execute_input":"2022-09-22T13:08:36.781868Z","iopub.status.idle":"2022-09-22T13:08:37.269848Z","shell.execute_reply.started":"2022-09-22T13:08:36.781835Z","shell.execute_reply":"2022-09-22T13:08:37.268852Z"},"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)","metadata":{"execution":{"iopub.status.busy":"2022-09-22T13:08:37.271014Z","iopub.execute_input":"2022-09-22T13:08:37.274070Z","iopub.status.idle":"2022-09-22T13:08:37.280350Z","shell.execute_reply.started":"2022-09-22T13:08:37.274033Z","shell.execute_reply":"2022-09-22T13:08:37.279267Z"},"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-09-22T13:08:37.282192Z","iopub.execute_input":"2022-09-22T13:08:37.283136Z","iopub.status.idle":"2022-09-22T13:08:37.553343Z","shell.execute_reply.started":"2022-09-22T13:08:37.283098Z","shell.execute_reply":"2022-09-22T13:08:37.552409Z"},"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):\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 weighted avareging\n    preds = torch.sum(torch.stack(preds, dim=0), dim=0) / np.sum(weights)\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.04]) / 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-09-22T13:08:37.554564Z","iopub.execute_input":"2022-09-22T13:08:37.554925Z","iopub.status.idle":"2022-09-22T13:08:37.569850Z","shell.execute_reply.started":"2022-09-22T13:08:37.554891Z","shell.execute_reply":"2022-09-22T13:08:37.568827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from skimage.segmentation import watershed\n\n\ndef clean_mask(mask, kernel):\n    mask = ndi.binary_fill_holes(mask, structure=kernel)\n    return mask\n\n\ndef 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\n\ndef refine_mask(pred_mask):\n    # instacnce segmentation via watershedding\n    D = ndi.distance_transform_edt(pred_mask)\n    ret, sure_fg = cv2.threshold(D, 0.1 * D.max(), 255, 0)\n    sure_fg = np.uint8(sure_fg)\n    _, markers = cv2.connectedComponents(sure_fg)\n\n    labels = watershed(-D, markers, mask=pred_mask)\n\n    filled = np.zeros((labels.shape))\n    for i in np.unique(labels):\n        instance_mask = labels.copy()\n        instance_mask[instance_mask != i] = 0\n\n        # convert to 8-bit format\n        instance_mask = instance_mask.astype(np.uint8)\n\n        # remove connections\n        struct = disk_struct(5)\n        closing = ndi.binary_opening(instance_mask, struct)\n        struct = disk_struct(10)\n        instance_mask = ndi.binary_opening(closing, struct)\n        \n        filled += instance_mask\n\n    refined_mask = filled\n\n    return refined_mask","metadata":{"execution":{"iopub.status.busy":"2022-09-22T13:47:38.793042Z","iopub.execute_input":"2022-09-22T13:47:38.793809Z","iopub.status.idle":"2022-09-22T13:47:38.805950Z","shell.execute_reply.started":"2022-09-22T13:47:38.793771Z","shell.execute_reply":"2022-09-22T13:47:38.804984Z"},"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    \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    # 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)\n        else:\n            pred = pred_mask(image, lung_models, thresh_factor=thr_f)\n            \n        if np.all(pred == 0):\n            # adaptive threshold\n            for ii in range(7):\n                thr_f -= 0.2\n                # predict\n                if cfg.tta:\n                    pred = pred_mask_tta(image, lung_models, thresh_factor=thr_f)\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)\n        else:\n            pred = pred_mask(image, all_models, thresh_factor=thr_f)\n            \n        if np.all(pred == 0):\n            # adaptive threshold\n            for ii in range(7):\n                thr_f -= 0.2\n                # predict\n                if cfg.tta:\n                    pred = pred_mask_tta(image, all_models, thresh_factor=thr_f)\n                else:\n                    pred = pred_mask(image, all_models, thresh_factor=thr_f)\n    \n    # resize back\n    pred = cv2.resize(pred.squeeze(), image_shape)\n    \n    # refine the mask\n    if organ_type in ['kidney']:\n        print('doing post-processing')\n        pred = refine_mask(pred)\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-09-22T13:48:25.192920Z","iopub.execute_input":"2022-09-22T13:48:25.193371Z","iopub.status.idle":"2022-09-22T13:48:50.458464Z","shell.execute_reply.started":"2022-09-22T13:48:25.193332Z","shell.execute_reply":"2022-09-22T13:48:50.457312Z"},"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-09-22T13:48:54.917924Z","iopub.execute_input":"2022-09-22T13:48:54.918596Z","iopub.status.idle":"2022-09-22T13:48:54.926025Z","shell.execute_reply.started":"2022-09-22T13:48:54.918561Z","shell.execute_reply":"2022-09-22T13:48:54.924804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2022-09-22T13:48:55.444884Z","iopub.execute_input":"2022-09-22T13:48:55.445902Z","iopub.status.idle":"2022-09-22T13:48:55.456643Z","shell.execute_reply.started":"2022-09-22T13:48:55.445864Z","shell.execute_reply":"2022-09-22T13:48:55.455351Z"},"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-09-22T12:55:49.265442Z","iopub.status.idle":"2022-09-22T12:55:49.266227Z","shell.execute_reply.started":"2022-09-22T12:55:49.265970Z","shell.execute_reply":"2022-09-22T12:55:49.265994Z"},"trusted":true},"execution_count":null,"outputs":[]}]}