{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30635,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport math\nimport glob\nimport gc\nimport tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nimport torchvision.transforms as T\nimport torchvision.transforms.functional as TF\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import Dataset, DataLoader\nfrom fastai.vision.all import *\nfrom typing import Optional\nfrom torch.nn.functional import one_hot\nfrom sklearn.model_selection import KFold\nimport random\n!pip install segmentation_models_pytorch\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:05:55.147337Z","iopub.execute_input":"2024-01-22T12:05:55.148368Z","iopub.status.idle":"2024-01-22T12:06:20.739001Z","shell.execute_reply.started":"2024-01-22T12:05:55.148322Z","shell.execute_reply":"2024-01-22T12:06:20.737864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ENCODER = \"resnet50\"\nWEIGHTS = \"imagenet\"","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:20.741526Z","iopub.execute_input":"2024-01-22T12:06:20.741889Z","iopub.status.idle":"2024-01-22T12:06:20.747794Z","shell.execute_reply.started":"2024-01-22T12:06:20.741857Z","shell.execute_reply":"2024-01-22T12:06:20.746997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We start from a 2D SMP Unet\nmodel = smp.Unet(\n    encoder_name=ENCODER,\n    encoder_weights=WEIGHTS,\n    in_channels=1,\n    classes=2\n)","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:20.749191Z","iopub.execute_input":"2024-01-22T12:06:20.749457Z","iopub.status.idle":"2024-01-22T12:06:22.004070Z","shell.execute_reply.started":"2024-01-22T12:06:20.749435Z","shell.execute_reply":"2024-01-22T12:06:22.003100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# original 2D layers and values\n# I manually copied the output to a string\nkeys = []\nvalues = []\nfor name, value in model.named_parameters():\n    layer = 'model'\n    for sub_string in name.split('.')[:-1]:\n        \n        if sub_string in ['0','1','2','3','4','5','6','7','8','9']:\n            layer = layer + '['+sub_string+']'\n        else:\n            layer = layer + '.' + sub_string\n\n    exec('print('+layer+')')\n    keys.append(name)\n    values.append(value)","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:22.005722Z","iopub.execute_input":"2024-01-22T12:06:22.006079Z","iopub.status.idle":"2024-01-22T12:06:22.022045Z","shell.execute_reply.started":"2024-01-22T12:06:22.006028Z","shell.execute_reply":"2024-01-22T12:06:22.021065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"layers = '''Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 64, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 128, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 512, kernel_size=(1, 1), stride=(2, 2), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 128, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 1024, kernel_size=(1, 1), stride=(2, 2), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 256, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 1024, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(1024, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 512, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(1024, 2048, kernel_size=(1, 1), stride=(2, 2), bias=False)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(2048, 512, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(512, 2048, kernel_size=(1, 1), stride=(1, 1), bias=False)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(3072, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(768, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(384, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(128, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(32, 16, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(16, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(16, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(16, 16, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\nBatchNorm2d(16, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nBatchNorm2d(16, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\nConv2d(16, 2, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))\nConv2d(16, 2, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))'''","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2024-01-22T12:06:22.024744Z","iopub.execute_input":"2024-01-22T12:06:22.025817Z","iopub.status.idle":"2024-01-22T12:06:22.038661Z","shell.execute_reply.started":"2024-01-22T12:06:22.025778Z","shell.execute_reply":"2024-01-22T12:06:22.037617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"layers =  layers.split('\\n')","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:22.039810Z","iopub.execute_input":"2024-01-22T12:06:22.040165Z","iopub.status.idle":"2024-01-22T12:06:22.052894Z","shell.execute_reply.started":"2024-01-22T12:06:22.040133Z","shell.execute_reply":"2024-01-22T12:06:22.052124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For each layer if needed I declared the corresponding 3D one\nlayers_3D = []\nfor layer in layers:\n    layer = layer.replace('2d','3d')\n    if 'kernel_size=(' in layer:\n        START = layer.find('kernel_size=(')\n        END = START + layer[START:].find(')')\n        layer = layer[:START] + layer[START:END] + ',' + layer[START:END].split(',')[-1] + layer[END:]\n    if 'stride=(' in layer:\n        START = layer.find('stride=(')\n        END = START + layer[START:].find(')')\n        layer = layer[:START] + layer[START:END] + ',' + layer[START:END].split(',')[-1] + layer[END:]\n    if 'padding=(' in layer:\n        START = layer.find('padding=(')\n        END = START + layer[START:].find(')')\n        layer = layer[:START] + layer[START:END] + ',' + layer[START:END].split(',')[-1] + layer[END:]\n    layers_3D.append(layer)","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:22.054556Z","iopub.execute_input":"2024-01-22T12:06:22.054928Z","iopub.status.idle":"2024-01-22T12:06:22.064489Z","shell.execute_reply.started":"2024-01-22T12:06:22.054894Z","shell.execute_reply":"2024-01-22T12:06:22.063533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# I reassign the layers and tiled weights to obtain the equivalent 3D Unet\nfor i in range(len(keys)):\n    key = keys[i]\n    layer = 'model'\n    for sub_string in key.split('.')[:-1]:\n        \n        if sub_string in ['0','1','2','3','4','5','6','7','8','9']:\n            layer = layer + '['+sub_string+']'\n        else:\n            layer = layer + '.' + sub_string\n\n    exec(layer+'=nn.'+layers_3D[i])  \n    if 'kernel_size=(' in layers_3D[i]:\n        value = values[i]\n        if value.dim() > 2: value = torch.tile(value.unsqueeze(-3),(value.shape[-1],1,1))\n        exec(layer+'.weights=value')\n          ","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:22.065720Z","iopub.execute_input":"2024-01-22T12:06:22.066013Z","iopub.status.idle":"2024-01-22T12:06:22.918876Z","shell.execute_reply.started":"2024-01-22T12:06:22.065989Z","shell.execute_reply":"2024-01-22T12:06:22.917677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# One last change\n# I found that layer after some debugging attempts with\n#for name, child in model.named_children():\n#        print('name',name)\n#        for x, y in child.named_children():\n#            print('x',x)\n#            for i, j in y.named_children():\n#                print('i',i)\n#                for k, l in j.named_children():\n#                     print('k',k)\nmodel.encoder.maxpool = nn.MaxPool3d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:22.919914Z","iopub.execute_input":"2024-01-22T12:06:22.920264Z","iopub.status.idle":"2024-01-22T12:06:22.926017Z","shell.execute_reply.started":"2024-01-22T12:06:22.920233Z","shell.execute_reply":"2024-01-22T12:06:22.924805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 3D SMP Unet working\nmodel(torch.zeros((16,1,32,32,32))).shape","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:22.927137Z","iopub.execute_input":"2024-01-22T12:06:22.927438Z","iopub.status.idle":"2024-01-22T12:06:24.000271Z","shell.execute_reply.started":"2024-01-22T12:06:22.927408Z","shell.execute_reply":"2024-01-22T12:06:23.999366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We can save and reload it\n# Everything looks good\ntorch.save(model,'SMP_resnet50_3D')\ndel model,layers,layers_3D,values,keys\nmodel = torch.load('SMP_resnet50_3D')\nprint(model)\nmodel(torch.zeros((16,1,32,32,32))).shape\n","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:06:24.001309Z","iopub.execute_input":"2024-01-22T12:06:24.001608Z","iopub.status.idle":"2024-01-22T12:06:25.628120Z","shell.execute_reply.started":"2024-01-22T12:06:24.001579Z","shell.execute_reply":"2024-01-22T12:06:25.627315Z"},"trusted":true},"execution_count":null,"outputs":[]}]}