{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":9869730,"sourceType":"datasetVersion","datasetId":6058495},{"sourceId":10106143,"sourceType":"datasetVersion","datasetId":6233973},{"sourceId":10364147,"sourceType":"datasetVersion","datasetId":6383310},{"sourceId":10522003,"sourceType":"datasetVersion","datasetId":6512054},{"sourceId":10524633,"sourceType":"datasetVersion","datasetId":6513752},{"sourceId":10551753,"sourceType":"datasetVersion","datasetId":6528712},{"sourceId":10573055,"sourceType":"datasetVersion","datasetId":6509390},{"sourceId":10587481,"sourceType":"datasetVersion","datasetId":6552409},{"sourceId":10591029,"sourceType":"datasetVersion","datasetId":6554845},{"sourceId":10595206,"sourceType":"datasetVersion","datasetId":6557851},{"sourceId":10595259,"sourceType":"datasetVersion","datasetId":6557884},{"sourceId":10616474,"sourceType":"datasetVersion","datasetId":6572779},{"sourceId":10622345,"sourceType":"datasetVersion","datasetId":6576941},{"sourceId":10662702,"sourceType":"datasetVersion","datasetId":6603433},{"sourceId":206640467,"sourceType":"kernelVersion"},{"sourceId":213055282,"sourceType":"kernelVersion"},{"sourceId":215061124,"sourceType":"kernelVersion"},{"sourceId":208471,"sourceType":"modelInstanceVersion","modelInstanceId":176271,"modelId":198599},{"sourceId":208598,"sourceType":"modelInstanceVersion","modelInstanceId":175864,"modelId":198200},{"sourceId":208866,"sourceType":"modelInstanceVersion","modelInstanceId":163890,"modelId":186240},{"sourceId":209184,"sourceType":"modelInstanceVersion","modelInstanceId":175864,"modelId":198200},{"sourceId":209261,"sourceType":"modelInstanceVersion","modelInstanceId":176271,"modelId":198599},{"sourceId":211515,"sourceType":"modelInstanceVersion","modelInstanceId":180333,"modelId":202595},{"sourceId":221414,"sourceType":"modelInstanceVersion","modelInstanceId":180111,"modelId":186240},{"sourceId":237845,"sourceType":"modelInstanceVersion","modelInstanceId":203128,"modelId":224862},{"sourceId":238613,"sourceType":"modelInstanceVersion","modelInstanceId":203787,"modelId":225518},{"sourceId":240586,"sourceType":"modelInstanceVersion","modelInstanceId":205581,"modelId":227328},{"sourceId":241798,"sourceType":"modelInstanceVersion","modelInstanceId":206545,"modelId":228293},{"sourceId":242194,"sourceType":"modelInstanceVersion","modelInstanceId":206545,"modelId":228293},{"sourceId":242790,"sourceType":"modelInstanceVersion","modelInstanceId":207362,"modelId":229084},{"sourceId":243128,"sourceType":"modelInstanceVersion","modelInstanceId":207362,"modelId":229084},{"sourceId":250544,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":214151,"modelId":235824}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"deps_path = '/kaggle/input/czii-cryoet-dependencies'","metadata":{"execution":{"iopub.status.busy":"2025-02-04T16:42:19.133824Z","iopub.execute_input":"2025-02-04T16:42:19.134134Z","iopub.status.idle":"2025-02-04T16:42:19.143128Z","shell.execute_reply.started":"2025-02-04T16:42:19.134094Z","shell.execute_reply":"2025-02-04T16:42:19.142283Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! cp -r /kaggle/input/czii-cryoet-dependencies/asciitree-0.3.3/ asciitree-0.3.3/","metadata":{"execution":{"iopub.status.busy":"2025-02-04T16:42:19.145367Z","iopub.execute_input":"2025-02-04T16:42:19.14579Z","iopub.status.idle":"2025-02-04T16:42:20.259635Z","shell.execute_reply.started":"2025-02-04T16:42:19.145743Z","shell.execute_reply":"2025-02-04T16:42:20.25836Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip wheel asciitree-0.3.3/asciitree-0.3.3/\n! cp /kaggle/input/best-weights-code/* ./ -r\n! cp /kaggle/input/best-weights ./ -r\n! cp /kaggle/input/voxhrnet-v2/voxhrnetV2.py ./ -r\n! cp /kaggle/input/load-vox-net/load_model.py ./ -r\n! cp /kaggle/input/hrnet-v1/* ./ -r","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-02-04T16:42:20.261269Z","iopub.execute_input":"2025-02-04T16:42:20.261555Z","iopub.status.idle":"2025-02-04T16:43:02.379154Z","shell.execute_reply.started":"2025-02-04T16:42:20.26151Z","shell.execute_reply":"2025-02-04T16:43:02.377938Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! cp /kaggle/input/d/luoziqian/unet2e3d-6c/* ./\n! cp /kaggle/input/voxhrnet/voxhrnet.py ./ -r\n! cp /kaggle/input/vox-networks-dataset/voxhrnet.py ./\n! cp /kaggle/input/vox-networks-dataset/config_small.yaml ./\n! cp /kaggle/input/vox-networks-dataset/model10.py ./\n! cp /kaggle/input/voxresnet-v0/voxresnetV0.py ./","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:43:02.38067Z","iopub.execute_input":"2025-02-04T16:43:02.380978Z","iopub.status.idle":"2025-02-04T16:43:08.44662Z","shell.execute_reply.started":"2025-02-04T16:43:02.380949Z","shell.execute_reply":"2025-02-04T16:43:08.445355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install asciitree-0.3.3-py3-none-any.whl\n! pip install /kaggle/input/einops-0-8-none-any/einops-0.8.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2025-02-04T16:43:08.448456Z","iopub.execute_input":"2025-02-04T16:43:08.448878Z","iopub.status.idle":"2025-02-04T16:44:29.108446Z","shell.execute_reply.started":"2025-02-04T16:43:08.448838Z","shell.execute_reply":"2025-02-04T16:44:29.107557Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-02-04T16:44:29.109944Z","iopub.execute_input":"2025-02-04T16:44:29.110295Z","iopub.status.idle":"2025-02-04T16:44:47.401417Z","shell.execute_reply.started":"2025-02-04T16:44:29.110268Z","shell.execute_reply":"2025-02-04T16:44:47.400577Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install /kaggle/input/vox-networks-dataset/yacs-0.1.8-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:44:47.404239Z","iopub.execute_input":"2025-02-04T16:44:47.404559Z","iopub.status.idle":"2025-02-04T16:45:28.339718Z","shell.execute_reply.started":"2025-02-04T16:44:47.404511Z","shell.execute_reply":"2025-02-04T16:45:28.338849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List, Tuple, Union\nimport numpy as np\nimport torch\nfrom monai.data import DataLoader, Dataset, CacheDataset, decollate_batch\nfrom monai.transforms import (\n    Compose, \n    EnsureChannelFirstd, \n    Orientationd,  \n    AsDiscrete,  \n    RandFlipd, \n    RandRotate90d, \n    NormalizeIntensityd,\n    RandCropByLabelClassesd,\n)","metadata":{"execution":{"iopub.status.busy":"2025-02-04T16:45:28.340983Z","iopub.execute_input":"2025-02-04T16:45:28.341264Z","iopub.status.idle":"2025-02-04T16:46:01.603626Z","shell.execute_reply.started":"2025-02-04T16:45:28.341238Z","shell.execute_reply":"2025-02-04T16:46:01.60288Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport shutil\nimport tempfile\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nfrom monai.transforms import (\n    Compose,\n    Orientationd,\n    RandFlipd,\n    RandShiftIntensityd,\n    RandRotate90d,\n)\nfrom monai.data import (\n    ThreadDataLoader,\n    CacheDataset,\n    load_decathlon_datalist,\n    decollate_batch,\n    set_track_meta,\n)\nfrom monai.inferers import sliding_window_inference\nfrom monai.networks.nets import SwinUNETR\nfrom monai.metrics import DiceMetric\nfrom monai.losses import DiceCELoss\nimport torch\nimport einops\nimport warnings\n\n\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"CUDA_DEVICE_ORDER\"] = \"PCI_BUS_ID\"\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:01.60461Z","iopub.execute_input":"2025-02-04T16:46:01.605258Z","iopub.status.idle":"2025-02-04T16:46:01.662254Z","shell.execute_reply.started":"2025-02-04T16:46:01.60523Z","shell.execute_reply":"2025-02-04T16:46:01.661008Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define some helper functions\n\n\n### Patching helper functions\n\nThese are mostly used to split large volumes into smaller ones and stitch them back together. ","metadata":{}},{"cell_type":"code","source":"def calculate_patch_starts_with_overlap(\n    dimension_size: int, patch_size: int, overlap: int\n) -> List[int]:\n    if dimension_size <= patch_size:\n        return [0]\n\n    num_patches = np.ceil(\n        (dimension_size - overlap) / (patch_size - overlap) + 1\n    ).astype(int)\n    patch_starts = []\n    for i in range(num_patches):\n        pos = int(i * (patch_size - overlap))\n        if pos + patch_size > dimension_size:\n            pos = dimension_size - patch_size\n        if pos not in patch_starts:\n            patch_starts.append(pos)\n    return patch_starts\n\n\ndef extract_3d_patches_overlap(\n    arrays: List[np.ndarray],\n    patch_sizes: Tuple[int, int, int],\n    overlap_sizes: Tuple[int, int, int],\n) -> Tuple[List[np.ndarray], List[Tuple[int, int, int]]]:\n\n    patch_starts_x = calculate_patch_starts_with_overlap(\n        arrays[0].shape[0], patch_sizes[0], overlap_sizes[0]\n    )\n    patch_starts_y = calculate_patch_starts_with_overlap(\n        arrays[0].shape[1], patch_sizes[1], overlap_sizes[1]\n    )\n    patch_starts_z = calculate_patch_starts_with_overlap(\n        arrays[0].shape[2], patch_sizes[2], overlap_sizes[2]\n    )\n    patch_size_d, patch_size_h, patch_size_w = patch_sizes\n    patches = []\n    coordinates = []\n    for arr in arrays:\n        for x in patch_starts_x:\n            for y in patch_starts_y:\n                for z in patch_starts_z:\n                    patch = arr[\n                        x : x + patch_size_d, y : y + patch_size_h, z : z + patch_size_w\n                    ]\n                    patches.append(patch)\n                    coordinates.append((x, y, z))\n\n    return patches, coordinates","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:01.663644Z","iopub.execute_input":"2025-02-04T16:46:01.664134Z","iopub.status.idle":"2025-02-04T16:46:01.694749Z","shell.execute_reply.started":"2025-02-04T16:46:01.664103Z","shell.execute_reply":"2025-02-04T16:46:01.694027Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Reading in the data","metadata":{}},{"cell_type":"code","source":"TRAIN_DATA_DIR = \"/kaggle/input/create-numpy-dataset-exp-name\"\nTEST_DATA_DIR = \"/kaggle/input/czii-cryo-et-object-identification\"","metadata":{"execution":{"iopub.status.busy":"2025-02-04T16:46:01.695683Z","iopub.execute_input":"2025-02-04T16:46:01.696016Z","iopub.status.idle":"2025-02-04T16:46:01.708722Z","shell.execute_reply.started":"2025-02-04T16:46:01.695975Z","shell.execute_reply":"2025-02-04T16:46:01.708011Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Initialize the model\n\nThis model is pretty much directly copied from [3D U-Net PyTorch Lightning distributed training](https://www.kaggle.com/code/zhuowenzhao11/3d-u-net-pytorch-lightning-distributed-training)","metadata":{}},{"cell_type":"code","source":"from torch.nn.modules import Module\nfrom torch import nn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:01.709611Z","iopub.execute_input":"2025-02-04T16:46:01.709883Z","iopub.status.idle":"2025-02-04T16:46:01.717886Z","shell.execute_reply.started":"2025-02-04T16:46:01.709858Z","shell.execute_reply":"2025-02-04T16:46:01.717262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/unet2d3e/* ./","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:01.71887Z","iopub.execute_input":"2025-02-04T16:46:01.719195Z","iopub.status.idle":"2025-02-04T16:46:02.77268Z","shell.execute_reply.started":"2025-02-04T16:46:01.719159Z","shell.execute_reply":"2025-02-04T16:46:02.771419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import lightning.pytorch as pl\n\nfrom monai.networks.nets import UNet\nfrom monai.losses import TverskyLoss\nfrom monai.metrics import DiceMetric\n\n# use warmup lr scheduler\nfrom torch.optim.lr_scheduler import (\n    CosineAnnealingLR,\n    CosineAnnealingWarmRestarts,\n    StepLR,\n)\nimport lightning.pytorch as pl\n\nfrom monai.networks.nets import UNet, AttentionUnet\nfrom monai.losses import TverskyLoss\nfrom monai.metrics import DiceMetric\n\n# use warmup lr scheduler\nfrom torch.optim.lr_scheduler import (\n    CosineAnnealingLR,\n    CosineAnnealingWarmRestarts,\n    StepLR,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:02.774398Z","iopub.execute_input":"2025-02-04T16:46:02.77538Z","iopub.status.idle":"2025-02-04T16:46:03.233264Z","shell.execute_reply.started":"2025-02-04T16:46:02.775335Z","shell.execute_reply":"2025-02-04T16:46:03.232341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config_v2 = {\n    \"MODEL\": {\n        \"NAME\": \"voxhrnet\",\n        \"EXTRA\": {\n            \"STAGE2\": {\n                \"NUM_MODULES\": 1,\n                \"NUM_BRANCHES\": 2,\n                \"BLOCK\": \"BASIC\",\n                \"NUM_BLOCKS\": [3, 3],\n                \"NUM_CHANNELS\": [16, 32],\n            },\n            \"STAGE3\": {\n                \"NUM_MODULES\": 1,\n                \"NUM_BRANCHES\": 3,\n                \"BLOCK\": \"BASIC\",\n                \"NUM_BLOCKS\": [3, 3, 3],\n                \"NUM_CHANNELS\": [16, 32, 64],\n            },\n        },\n    }\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:03.23467Z","iopub.execute_input":"2025-02-04T16:46:03.235368Z","iopub.status.idle":"2025-02-04T16:46:03.240529Z","shell.execute_reply.started":"2025-02-04T16:46:03.235328Z","shell.execute_reply":"2025-02-04T16:46:03.239492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from monai.networks.nets import UNet,SegResNet,DynUNet\nfrom custom_vnet import CustomVNet\nfrom model2_6c import Net\n# from voxhrnet import build_model\nfrom voxhrnetV2 import build_model as build_model_v2\nfrom load_model import build_model as build_model_v3\nfrom model10_v1 import build_model as build_model_v4\nfrom voxresnetV0 import VoxResNet as VoxResNetV0\n\nbasic_unet = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=7,\n    channels=(48, 64, 80, 80),\n    strides=(2, 2, 1),\n    num_res_units=1,\n)\nbasic_unet_1 = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=7,\n    channels=(48, 64, 80, 80),\n    strides=(2, 2, 1),\n    num_res_units=2,\n)\nbasic_unet_6c = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=6,\n    channels=(48, 64, 80, 80),\n    strides=(2, 2, 1),\n    num_res_units=1,\n)\n\nsegresnet_6c_v1 = SegResNet(\n    in_channels=1,\n    out_channels=6,\n    dropout_prob=0.1,\n    upsample_mode=\"deconv\",\n)\n\nbasic_unet2e3d = Net(\n    out_channels=7,\n    arch=\"resnet18d\",\n    decoder_dim=[80, 80, 64, 32, 16],\n    pretrained=False,\n)\n\nbasic_dynunet_v1 = DynUNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=6,\n    kernel_size=(3, 3, 3, 3),\n    strides=((1, 1, 1), 2, 2, 1),\n    upsample_kernel_size=(2, 2, 1),\n    filters=[16, 24, 48, 80, 80],\n    norm_name=\"instance\",\n    act_name=\"PRELU\",\n    deep_supervision=True,\n)\n\nbasic_dynunet_v1 = DynUNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=6,\n    kernel_size=(3, 3, 3, 3),\n    strides=((1, 1, 1), 2, 2, 1),\n    upsample_kernel_size=(2, 2, 1),\n    filters=[16, 24, 48, 80, 80],\n    norm_name=\"instance\",\n    act_name=\"PRELU\",\n    deep_supervision=True,\n)\n\nbasic_dynunet_v2 = DynUNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=6,\n    kernel_size=(3, 3, 3, 3, 3),\n    strides=((1, 1, 1), 2, 2, 2, 1),\n    upsample_kernel_size=(2, 2, 2, 1),\n    filters=[16, 24, 48, 80, 80],\n    norm_name=\"instance\",\n    act_name=\"PRELU\",\n    deep_supervision=True,\n)\n\n# basic_voxhrnet_v0 = build_model(1,6)\nbasic_voxhrnet_v2 = build_model_v2(1,6,config_v2)\nbasic_voxhrnet_v3 = build_model_v3()\nbisic_voxhrnet_v4 = build_model_v4()\n\nbasic_voxresnet_v0 = VoxResNetV0(in_channels=1, n_classes=7)\n\nbasic_dense_vnet = CustomVNet(in_channels=1, classes=7)","metadata":{"execution":{"iopub.status.busy":"2025-02-04T16:46:03.241641Z","iopub.execute_input":"2025-02-04T16:46:03.24191Z","iopub.status.idle":"2025-02-04T16:46:05.142583Z","shell.execute_reply.started":"2025-02-04T16:46:03.241885Z","shell.execute_reply":"2025-02-04T16:46:05.141467Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nbest_weights_dir = \"best-weights\"\nos.listdir(best_weights_dir)\n# best_weights_path_list = [\n#  'best-weights/epoch122-step2952-valid_loss0.3625-val_metric0.8367.ckpt',\n#  'best-weights/epoch148-step3576-valid_loss1.1154-val_metric0.7722.ckpt',\n#  'best-weights/epoch153-step3696-valid_loss0.3021-val_metric0.8900.ckpt',\n#  'best-weights/epoch194-step4680-valid_loss1.0213-val_metric0.8788.ckpt',\n#  '/kaggle/input/voxresnet-v0/epoch195-step4704-valid_loss0.4258-val_metric0.7914.ckpt',\n# ] \\\n#  + ['/kaggle/input/hrnet-v1/epoch188-step4536-valid_loss0.4231-val_metric0.8659.ckpt'] \\\n#  + ['/kaggle/input/unet2e3d-6c/pytorch/default/1/unet2E3D-v1-epoch114-val_loss0.55-val_metric0.53-step2760.ckpt'] \\\n#  + ['/kaggle/input/my-weights/unet3D-epoch173-val_loss0.53-val_metric0.54-step4176.ckpt']  \\\n#  + ['best-weights/epoch152-step3672-valid_loss0.4333-val_metric0.7929.ckpt'] \\\n#  +['/kaggle/input/segresnet-6c/pytorch/default/1/epoch314-val_loss0.54-val_metric0.54-step7560.ckpt']\n\n#  # + ['/kaggle/input/hrnet-v1/epoch188-step4536-valid_loss0.4231-val_metric0.8659.ckpt'] \n\n\n # +['/kaggle/input/my-weights/unet3D-epoch173-val_loss0.53-val_metric0.54-step4176.ckpt']\n\n\n\n # +['/kaggle/input/segresnet-6c/pytorch/default/1/epoch314-val_loss0.54-val_metric0.54-step7560.ckpt']\n\n#  + ['/kaggle/input/dynunet-6c/pytorch/default/1/DynUnet-epoch189-val_loss0.51-val_metric0.55-step4560.ckpt'] \\\n#  +['/kaggle/input/segresnet-6c/pytorch/default/1/epoch314-val_loss0.54-val_metric0.54-step7560.ckpt']\n\n#  +['/kaggle/input/my-weights/unet3D-epoch173-val_loss0.53-val_metric0.54-step4176.ckpt']\n # +['/kaggle/input/unet2e3d-6c/pytorch/default/1/unet2E3D-v1-epoch114-val_loss0.55-val_metric0.53-step2760.ckpt'] \\\n# best_weights_path_list=['/kaggle/input/dynunet-6c/pytorch/default/1/DynUnet-epoch189-val_loss0.51-val_metric0.55-step4560.ckpt']\n# best_weights_path_list=['/kaggle/input/dynunet-6c/pytorch/default/2/DynUnet-v2-epoch164-val_loss0.53-val_metric0.53-step3960.ckpt']\n# best_weights_path_list=['/kaggle/input/vox-networks-dataset/epoch133-step3216-valid_loss0.4223-val_metric0.8612.ckpt']\nbest_weights_path_list = ['/kaggle/input/unet2e3d-7c/pytorch/default/1/unet2e3d-epoch135-step3264-valid_loss0.4315-val_metric0.7775.ckpt']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:05.144141Z","iopub.execute_input":"2025-02-04T16:46:05.144521Z","iopub.status.idle":"2025-02-04T16:46:05.154177Z","shell.execute_reply.started":"2025-02-04T16:46:05.144483Z","shell.execute_reply":"2025-02-04T16:46:05.1532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy\nimport torch\ndummy_input = torch.rand((1, 1, 128, 384, 384)).half().cuda()\n\nmodel_list = []\nfor idx,path in enumerate(best_weights_path_list):\n    ckpt = torch.load(path)\n    state_dict_ = ckpt[\"state_dict\"]\n    state_dict = {}\n    for k in state_dict_.keys():\n        if \"model.\" in k:\n            state_dict[k[6:]] = state_dict_[k]\n    try:\n        model = copy.deepcopy(basic_unet)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load unet\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_unet.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        pass\n    try:\n        model = copy.deepcopy(basic_unet_1)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load unet_1\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_unet_1.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(basic_dense_vnet)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load vnet\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_vnet.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(basic_unet_6c)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load unet_6c\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_unet_6c.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(basic_unet2e3d)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load unet2e3d\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_unet2e3d.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(segresnet_6c_v1)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load segresnet_6c_v1\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_segresnet_6c_v1.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(basic_dynunet_v1)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load dynunet\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_dynunet.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(basic_dynunet_v2)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load basic_dynunet_v2\")  \n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_basic_dynunet_v2.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        # print(\"load failed\")\n        pass\n    try:\n        model = copy.deepcopy(basic_voxhrnet_v0)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load basic_voxhrnet_v0\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_basic_voxhrnet_v0.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        pass\n    try:\n        model = copy.deepcopy(basic_voxhrnet_v3)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load basic_voxhrnet_v3\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_basic_voxhrnet_v3.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        pass\n    try:\n        model = copy.deepcopy(bisic_voxhrnet_v4)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load bisic_voxhrnet_v4\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_bisic_voxhrnet_v4.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        pass\n    try:\n        model = copy.deepcopy(basic_voxhrnet_v2)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load basic_voxhrnet_v2\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_basic_voxhrnet_v2.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        pass\n    try:\n        model = copy.deepcopy(basic_voxresnet_v0)\n        model.load_state_dict(state_dict)\n        model_list.append(model.to(\"cpu\"))\n        print(\"load basic_voxresnet_v0\")\n        model=model.eval().half().cuda()\n        dynamic_axes = {\"images\": {0: \"batch\"}, \"output\": {0: \"batch\"}}\n        output_name=f\"{idx}_basic_voxresnet_v0.onnx\"\n        torch.onnx.export(\n            model,\n            dummy_input,\n            output_name,\n            verbose=False,\n            input_names=['images'],\n            output_names=['output'],\n            dynamic_axes=dynamic_axes,\n            do_constant_folding=True,\n            opset_version=12        \n        )\n        continue\n    except:\n        pass\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-04T16:46:05.155634Z","iopub.execute_input":"2025-02-04T16:46:05.156008Z","iopub.status.idle":"2025-02-04T16:46:07.749782Z","shell.execute_reply.started":"2025-02-04T16:46:05.155965Z","shell.execute_reply":"2025-02-04T16:46:07.748736Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}