{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.14"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10558265,"sourceType":"datasetVersion","datasetId":6532180},{"sourceId":10649514,"sourceType":"datasetVersion","datasetId":6594173},{"sourceId":10665229,"sourceType":"datasetVersion","datasetId":6605153}],"dockerImageVersionId":30805,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":263.718805,"end_time":"2024-12-15T20:33:00.918064","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-12-15T20:28:37.199259","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import json\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport zarr\nimport pandas as pd\nimport numpy as np\nimport os\nfrom tqdm import tqdm\nimport gc\nimport cc3d\n\nimport torchvision\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom fastai.vision.all import *\nimport segmentation_models_pytorch as smp\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2025-01-21T15:28:01.750248Z","iopub.execute_input":"2025-01-21T15:28:01.750517Z","iopub.status.idle":"2025-01-21T15:28:10.152333Z","shell.execute_reply.started":"2025-01-21T15:28:01.750491Z","shell.execute_reply":"2025-01-21T15:28:10.151627Z"},"papermill":{"duration":9.794406,"end_time":"2024-12-15T20:28:49.380257","exception":false,"start_time":"2024-12-15T20:28:39.585851","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATH = '/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns/'\nSAMPLES = [\n    x for x in os.listdir(PATH)\n]\nTARGETS = [\n    'apo-ferritin',# easy\n    'beta-galactosidase',# hard\n    'ribosome',# easy\n    'thyroglobulin',# hard\n    'virus-like-particle'# easy\n]\nMODELS = [\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_1',\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_2',\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_3',\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_4',\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_5',\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_6',\n    '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_7'\n]","metadata":{"execution":{"iopub.status.busy":"2025-01-21T15:28:10.153392Z","iopub.execute_input":"2025-01-21T15:28:10.153659Z","iopub.status.idle":"2025-01-21T15:28:10.164672Z","shell.execute_reply.started":"2025-01-21T15:28:10.153634Z","shell.execute_reply":"2025-01-21T15:28:10.164073Z"},"papermill":{"duration":0.016185,"end_time":"2024-12-15T20:28:49.39961","exception":false,"start_time":"2024-12-15T20:28:49.383425","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class repack_3D(nn.Module):\n    def __init__(\n        self\n        ):\n        super(repack_3D, self).__init__()\n\n    def forward(self,X):\n        D,C,H,W = X.shape[-4:]\n        return X.view(-1,D,C,H,W).permute(0,2,1,3,4)","metadata":{"execution":{"iopub.status.busy":"2025-01-21T15:28:10.166052Z","iopub.execute_input":"2025-01-21T15:28:10.166356Z","iopub.status.idle":"2025-01-21T15:28:10.23865Z","shell.execute_reply.started":"2025-01-21T15:28:10.166331Z","shell.execute_reply":"2025-01-21T15:28:10.23779Z"},"papermill":{"duration":0.009604,"end_time":"2024-12-15T20:28:49.439364","exception":false,"start_time":"2024-12-15T20:28:49.42976","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class repack_2D(nn.Module):\n    def __init__(\n        self\n        ):\n        super(repack_2D, self).__init__()\n\n    def forward(self,X):\n        C,D,H,W = X.shape[-4:]\n        return X.permute(0,2,1,3,4).reshape(-1,C,H,W)","metadata":{"execution":{"iopub.status.busy":"2025-01-21T15:28:10.239784Z","iopub.execute_input":"2025-01-21T15:28:10.240089Z","iopub.status.idle":"2025-01-21T15:28:10.249276Z","shell.execute_reply.started":"2025-01-21T15:28:10.240062Z","shell.execute_reply":"2025-01-21T15:28:10.248554Z"},"papermill":{"duration":0.01023,"end_time":"2024-12-15T20:28:49.452379","exception":false,"start_time":"2024-12-15T20:28:49.442149","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myUNet2Dto3D(nn.Module):\n    def __init__(\n        self,\n        classes\n        ):\n        super(myUNet2Dto3D, self).__init__()\n\n        self.classes = classes\n        \n        decoder_channels = (256, 128, 64, 32, 16)[-ENCODER_DEPTH:]\n        decoder_in_channels = (768, 384, 192, 128 , 32)[-ENCODER_DEPTH:]\n    \n        \n        self.UNet = smp.Unet(\n            encoder_name=ENCODER_NAME,\n            encoder_depth=ENCODER_DEPTH,\n            decoder_channels=decoder_channels,\n            classes=classes,\n            in_channels=1\n        ).to(device)\n\n        self.UNet.encoder.layer3 = nn.Identity()\n        self.UNet.encoder.layer4 = nn.Identity()        \n\n        for k in range(ENCODER_DEPTH):\n            self.UNet.decoder.blocks[k].conv1[0] = nn.Sequential(\n                nn.Dropout(DROPOUT),\n                repack_3D(),\n                nn.Conv3d(\n                    decoder_in_channels[k],\n                    decoder_channels[k],\n                    kernel_size=3,\n                    stride=1,\n                    padding=1,\n                    bias=False\n                )\n            )\n            self.UNet.decoder.blocks[k].conv1[1] = nn.BatchNorm3d(\n                decoder_channels[k],\n                eps=1e-05, momentum=0.1,\n                affine=True,\n                track_running_stats=True\n            )\n            self.UNet.decoder.blocks[k].conv2[0] = nn.Conv3d(\n                decoder_channels[k],\n                decoder_channels[k],\n                kernel_size=3,\n                stride=1,\n                padding=1,\n                bias=False\n                )\n            self.UNet.decoder.blocks[k].conv2[1] = nn.Sequential(\n                nn.BatchNorm3d(\n                    decoder_channels[k],\n                    eps=1e-05, momentum=0.1,\n                    affine=True,\n                    track_running_stats=True\n                ),\n                repack_2D()\n            )\n\n        self.UNet.segmentation_head[0] = nn.Sequential(\n            repack_3D(),\n            nn.Conv3d(\n                16,\n                6,\n                kernel_size=3,\n                stride=1,\n                padding=1\n            )\n        )\n\n    def forward(self,X):\n        H,W = X.shape[-2:]\n        x = self.UNet(X.view(-1,1,H,W))\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2025-01-21T15:28:10.250429Z","iopub.execute_input":"2025-01-21T15:28:10.250669Z","iopub.status.idle":"2025-01-21T15:28:10.26198Z","shell.execute_reply.started":"2025-01-21T15:28:10.250627Z","shell.execute_reply":"2025-01-21T15:28:10.261279Z"},"papermill":{"duration":0.013107,"end_time":"2024-12-15T20:28:49.468154","exception":false,"start_time":"2024-12-15T20:28:49.455047","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if 1:#len(SAMPLES) > 3:\n    models = []\n    for path in MODELS:\n#       models.append(torch.load(path,map_location=device).eval())\n        models.append(torch.nn.DataParallel(torch.load(path).eval(), device_ids=[0, 1]))\n\n    experiment = []\n    particle_type = []\n    x = []\n    y = []\n    z = []\n    p = (1,99)\n    rot = [0,1,2,3,0,1,2,3,0]\n    y_pred = torch.zeros(2,6,36,640,640).float().to(device)\n    with torch.no_grad():\n        for sample in tqdm(SAMPLES):\n            MASK = []\n            file = zarr.open(PATH + sample + '/VoxelSpacing10.000/denoised.zarr', mode='r')\n            scale = np.array([file.attrs['multiscales'][0]['datasets'][0]['coordinateTransformations'][0]['scale']])\n            volume =np.array(file[0])\n            pmin,pmax = np.percentile(volume,p)\n            volume = (volume - pmin)/(pmax - pmin)\n            volume = torch.as_tensor(volume)\n            volume = torch.nn.functional.pad(\n                volume.unsqueeze(0),\n                (\n                    5,5,\n                    5,5\n                ),\n                mode='reflect'\n            )[0]\n\n            v = torch.stack(\n                volume[:36],\n                torch.rot90(volume[:36],1,(-2,-1)\n            ).to(device)\n\n            y_pred.zero_()\n            for model in models:\n                y_pred += model(v).softmax(1)\n\n            y_pred[0] += torch.rot90(y_pred[1],-1,(-2,-1))\n            \n            MASK.append(y_pred[0,:,:20].argmax(1).cpu())\n            mask = y_pred[0,:,20:].clone()\n            \n            for k in range(1,7):\n#               Drill Scan\n                v = torch.stack(\n                    torch.rot90(volume[k*16+4:k*16+36],k=rot[k],dims=(-2,-1)),\n                    torch.rot90(volume[k*16+4:k*16+36],k=rot[k]+1,dims=(-2,-1))\n                ).to(device)\n\n                y_pred.zero_()\n                for model in models:\n                    y_pred[:,:,:32] += model(v).softmax(1)\n                \n                y_pred[0,:,:32] += torch.rot90(y_pred[1,:,:32],-1,(-2,-1))\n\n                y_pred[0,:,:32] = torch.rot90(y_pred[0,:,:32],k=-rot[k],dims=(-2,-1))\n                y_pred[0,:,:16] += mask\n                MASK.append((y_pred[0,:,:16]).argmax(1).cpu())\n                mask = y_pred[0,:,16:32].clone()\n\n            v = torch.stack(\n                torch.rot90(volume[-36:],k=rot[7],dims=(-2,-1)),\n                torch.rot90(volume[-36:],k=rot[7]+1,dims=(-2,-1))\n            ).to(device)\n\n            y_pred.zero_()\n            for model in models:\n                y_pred += model(v).softmax(1)\n\n            y_pred[0] += torch.rot90(y_pred[1],-1,(-2,-1))\n\n            y_pred[0] = torch.rot90(y_pred[0],k=-rot[7],dims=(-2,-1))\n            y_pred[0,:,:16] += mask\n            MASK.append(y_pred[0].argmax(1).cpu())\n        \n            mask = torch.concat(MASK)[:,5:-5,5:-5].numpy()\n            for k in range(5):\n                stats = cc3d.statistics(cc3d.connected_components(mask == k+1))\n                preds = stats['centroids'][1:]*scale\n                experiment = experiment + [sample]*len(preds)\n                particle_type = particle_type + [TARGETS[k]]*len(preds)\n                x = x + list(preds[:,2])\n                y = y + list(preds[:,1])\n                z = z + list(preds[:,0])\n\nelse:\n    experiment = SAMPLES\n    particle_type = [TARGETS[0]]*len(SAMPLES)\n    x = [0]*len(SAMPLES)\n    y = [0]*len(SAMPLES)\n    z = [0]*len(SAMPLES)\n    \n\npd.DataFrame({\n    'id':np.arange(len(experiment)),\n    'experiment':experiment,\n    'particle_type':particle_type,\n    'x':x,\n    'y':y,\n    'z':z\n}).to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2025-01-21T15:31:27.03161Z","iopub.execute_input":"2025-01-21T15:31:27.032589Z","execution_failed":"2025-01-21T15:32:57.311Z"},"papermill":{"duration":246.50845,"end_time":"2024-12-15T20:32:57.962947","exception":false,"start_time":"2024-12-15T20:28:51.454497","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}