{"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9942824,"sourceType":"datasetVersion","datasetId":6113517}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"DEBUG = False\nif DEBUG:\n    !pip install zarr\n    !pip install segmentation_models_pytorch==0.3.3\n    !pip install connected-components-3d","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T14:14:49.098830Z","iopub.execute_input":"2024-11-19T14:14:49.099782Z","iopub.status.idle":"2024-11-19T14:15:13.483324Z","shell.execute_reply.started":"2024-11-19T14:14:49.099743Z","shell.execute_reply":"2024-11-19T14:15:13.482198Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport zarr\nimport pandas as pd\nimport matplotlib.pyplot as plt\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":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T14:15:13.485288Z","iopub.execute_input":"2024-11-19T14:15:13.485608Z","iopub.status.idle":"2024-11-19T14:15:18.828881Z","shell.execute_reply.started":"2024-11-19T14:15:13.485578Z","shell.execute_reply":"2024-11-19T14:15:18.827887Z"}},"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]\nFOLDS = [1,2,3,4,5,6]#[1,2,3,4,5,6,7]\nBS = 32\nTH = .3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T14:15:18.830234Z","iopub.execute_input":"2024-11-19T14:15:18.831237Z","iopub.status.idle":"2024-11-19T14:15:18.838620Z","shell.execute_reply.started":"2024-11-19T14:15:18.831205Z","shell.execute_reply":"2024-11-19T14:15:18.837873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myUNet(nn.Module):\n    def __init__(\n        self,\n        classes\n        ):\n        super(myUNet, self).__init__()\n\n        self.classes = classes\n        self.UNet = smp.Unet(\n            encoder_name=ENCODER_NAME,\n            classes=classes,\n            in_channels=1\n        ).to(device)\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.view(-1,H,W)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T14:15:21.716344Z","iopub.execute_input":"2024-11-19T14:15:21.716642Z","iopub.status.idle":"2024-11-19T14:15:21.721972Z","shell.execute_reply.started":"2024-11-19T14:15:21.716615Z","shell.execute_reply":"2024-11-19T14:15:21.721086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = {}\nfor target in TARGETS:\n    model[target] = []\n    for f in FOLDS:\n        model[target].append(torch.load('/kaggle/input/czii2d/CryoET_segmentation_'+target+'_'+str(f),map_location=device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T14:15:21.723208Z","iopub.execute_input":"2024-11-19T14:15:21.723560Z","iopub.status.idle":"2024-11-19T14:15:24.569248Z","shell.execute_reply.started":"2024-11-19T14:15:21.723500Z","shell.execute_reply":"2024-11-19T14:15:24.568561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"experiment = []\nparticle_type = []\nx = []\ny = []\nz = []\nfor sample in SAMPLES:\n    print(sample)\n    V = torch.as_tensor(np.array(zarr.open(PATH + sample + '/VoxelSpacing10.000/denoised.zarr', mode='r')[0])).to(device).float()\n    V /= V.std()\n    D,H,W = V.shape\n    h = 32 - H%32\n    w = 32 - W%32\n    d = D%BS\n    STEPS = D//BS\n    if d > 0: STEPS += 1\n    for target in TARGETS:\n        print(target)\n        MASK = 0\n        with torch.no_grad():\n            for f in tqdm(FOLDS):\n                for rot in [0]:#[0,1,2,3]:\n                    mask = []\n                    for k in range(STEPS):\n                        mask.append(\n                            torch.rot90(\n                                model[target][f-1](                \n                                    torch.nn.functional.pad(\n                                        torch.rot90(V[k*BS:(k+1)*BS],rot,(-2,-1)),\n                                        (w//2,w - w//2,h//2,h - h//2),\n                                        mode='reflect'\n                                    )\n                                )[:,h//2:h//2-h,w//2:w//2-w],\n                                -rot,\n                                (-2,-1)\n                            )\n                        )\n\n                    mask = torch.concat(mask).cpu()\n                    max_mask = mask.max()\n                    min_mask = mask.min()\n                    mask = (mask - min_mask)/(max_mask - min_mask)\n                    MASK += mask\n                    del mask\n                    gc.collect()\n\n        max_mask = MASK.max()    \n        min_mask = MASK.min()\n        MASK = (MASK - min_mask)/(max_mask - min_mask) > .5\n        labels_out = cc3d.connected_components(MASK.numpy())\n        stats = cc3d.statistics(labels_out)\n        preds = 10*stats['centroids'][1:]\n        experiment = experiment + [sample]*len(preds)\n        particle_type = particle_type + [target]*len(preds)\n        x = x + list(preds[:,2])\n        y = y + list(preds[:,1])\n        z = z + list(preds[:,0])\n\n        del MASK\n        gc.collect()\n\n    del V\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T14:15:24.570356Z","iopub.execute_input":"2024-11-19T14:15:24.570644Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pd.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":{"trusted":true},"outputs":[],"execution_count":null}]}