{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10504270,"sourceType":"datasetVersion","datasetId":6502943},{"sourceId":10587858,"sourceType":"datasetVersion","datasetId":6552690},{"sourceId":10588584,"sourceType":"datasetVersion","datasetId":6553182}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install segmentation_models_pytorch==0.3.3\n!pip install connected-components-3d\n!pip install zarr\n\nimport json\nimport matplotlib.pyplot as plt\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-26T22:41:28.202669Z","iopub.execute_input":"2025-01-26T22:41:28.203043Z","iopub.status.idle":"2025-01-26T22:42:11.749782Z","shell.execute_reply.started":"2025-01-26T22:41:28.202981Z","shell.execute_reply":"2025-01-26T22:42:11.748558Z"},"trusted":true,"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"VOLUMES_PATH = '/kaggle/input/czii-cryo-et-object-identification/train/static/ExperimentRuns/'\nLABELS_PATH = '/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns/'\nSYNTH_SAMPLES = [\n    ['TS_13', 'TS_25', 'TS_0', 'TS_15'],\n    ['TS_19', 'TS_17', 'TS_10'],\n    ['TS_5', 'TS_8', 'TS_20', 'TS_24'],\n    ['TS_26', 'TS_12', 'TS_2', 'TS_4'],\n    ['TS_22', 'TS_18', 'TS_16', 'TS_3'],\n    ['TS_6', 'TS_11', 'TS_9', 'TS_1'],\n    ['TS_7', 'TS_14', 'TS_21', 'TS_23']\n]\nTARGETS = [\n    'apo-ferritin',# easy\n    'beta-galactosidase',# hard\n    'ribosome',# easy\n    'thyroglobulin',# hard\n    'virus-like-particle'# easy\n]\nSEED = 1337\nENCODER_NAME = \"resnet18\"\nENCODER_DEPTH = 3\nPATCH_D = 64\nPATCH_H = 128\nPATCH_W = 128\nradius = {\n    'apo-ferritin':60,# easy\n    'beta-galactosidase':90,# hard\n    'ribosome':150,# easy\n    'thyroglobulin':130,# hard\n    'virus-like-particle':135# easy\n}\ncore = {\n    k:radius[k]*.05 for k in radius\n}\nshell = {\n    k:radius[k]*.11 for k in radius\n}\nBS = 16\nLR = 1e-4\nEPOCHS = 2\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T22:42:11.758357Z","iopub.execute_input":"2025-01-26T22:42:11.758740Z","iopub.status.idle":"2025-01-26T22:42:11.767276Z","shell.execute_reply.started":"2025-01-26T22:42:11.758700Z","shell.execute_reply":"2025-01-26T22:42:11.765914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_to_path = {\n    'apo-ferritin':'101/ferritin_complex-1.0_orientedpoint.ndjson',\n    'beta-galactosidase':'103/beta_galactosidase-1.0_orientedpoint.ndjson',\n    'ribosome':'104/cytosolic_ribosome-1.0_orientedpoint.ndjson',\n    'thyroglobulin':'105/thyroglobulin-1.0_orientedpoint.ndjson',\n    'virus-like-particle':'106/pp7_vlp-1.0_orientedpoint.ndjson'    \n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T22:42:24.019841Z","iopub.execute_input":"2025-01-26T22:42:24.020346Z","iopub.status.idle":"2025-01-26T22:42:24.027518Z","shell.execute_reply.started":"2025-01-26T22:42:24.020304Z","shell.execute_reply":"2025-01-26T22:42:24.025247Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"V = {}\nM = {}\nL = {}\nfor f in SYNTH_SAMPLES:\n    for sample in tqdm(f):\n        V[sample] = np.load('/kaggle/input/synth-denoised/'+sample+'.npy')\n        \n        file = zarr.open('/kaggle/input/cryoet-10441-masks/'+sample+target_to_path[TARGETS[0]][3:-20]+'segmentationmask.zarr', mode='r')\n        m = np.array(file[0])\n        m = torch.as_tensor(m)\n        M[sample] = torch.nn.functional.pad(\n            m.unsqueeze(0),\n            (\n                5,5,\n                5,5\n            )\n        )[0]\n        for k in range(4):\n            file = zarr.open('/kaggle/input/cryoet-10441-masks/'+sample+target_to_path[TARGETS[k+1]][3:-20]+'segmentationmask.zarr', mode='r')\n            m = np.array(file[0])\n            m = torch.as_tensor(m)\n            m = torch.nn.functional.pad(\n                m.unsqueeze(0),\n                (\n                    5,5,\n                    5,5\n                )\n            )[0]\n            M[sample][m>0] = k + 2\n\n        L[sample] = {}\n        for target in TARGETS:\n            L[sample][target] = []\n            for p in pd.read_json('/kaggle/input/synth-labels/synth_labels/'+sample+'/'+target+'.ndjson', lines=True)['location']:\n                L[sample][target].append([\n                    p['z'],\n                    p['y']+5,\n                    p['x']+5\n                ])\n            L[sample][target] = torch.tensor(L[sample][target]).float().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T22:59:10.768725Z","iopub.execute_input":"2025-01-26T22:59:10.769223Z","iopub.status.idle":"2025-01-26T23:01:14.283832Z","shell.execute_reply.started":"2025-01-26T22:59:10.769183Z","shell.execute_reply":"2025-01-26T23:01:14.281529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for axis in [0,1,2]:\n    plt.imshow(M[sample].sum(axis))\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:01:26.661849Z","iopub.execute_input":"2025-01-26T23:01:26.662317Z","iopub.status.idle":"2025-01-26T23:01:29.245381Z","shell.execute_reply.started":"2025-01-26T23:01:26.662282Z","shell.execute_reply":"2025-01-26T23:01:29.244104Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T22:49:09.336910Z","iopub.execute_input":"2025-01-26T22:49:09.337443Z","iopub.status.idle":"2025-01-26T22:49:09.344324Z","shell.execute_reply.started":"2025-01-26T22:49:09.337396Z","shell.execute_reply":"2025-01-26T22:49:09.343059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CryoET_Dataset(Dataset):\n    def __init__(self, V, L, M, samples, VALID=False):\n        self.data = V\n        self.L = L\n        self.M = M\n        self.samples = samples\n        self.sample = []\n        self.d = []\n        self.h = []\n        self.w = []\n        for sample in samples:\n            d,h,w = np.where(np.ones((5,9,9)))\n            d = d*32 + 4\n            h *= 64\n            w *= 64\n            self.sample = self.sample + [sample]*len(d)\n            self.d = self.d + list(d)\n            self.h = self.h + list(h)\n            self.w = self.w + list(w)\n        self.VALID = VALID\n\n    def __len__(self):\n        return len(self.d)\n\n    def __rawgetitem__(self, idx):\n        \n        sample = self.sample[idx]\n        d = self.d[idx]\n        h = self.h[idx]\n        w = self.w[idx]\n\n        if not self.VALID:\n            d += np.random.randint(32) - 16\n            h += np.random.randint(64) - 32\n            w += np.random.randint(64) - 32\n\n        if d < 0: d = 0\n        if h < 0: h = 0\n        if w < 0: w = 0\n\n        if d > 200 - 64: d = 200 - 64\n        if h > 512: h = 512\n        if w > 512: w = 512\n\n        image = torch.as_tensor(self.data[sample][\n            d:d+PATCH_D,\n            h:h+PATCH_H,\n            w:w+PATCH_W\n        ]).to(device)\n\n        hm = torch.as_tensor(self.M[sample][\n            d:d+PATCH_D,\n            h:h+PATCH_H,\n            w:w+PATCH_W\n        ]).long().to(device)\n\n        searching = True\n        for k in range(5):\n            t = TARGETS[k] \n#           Find particles within this voxel\n            inside = (self.L[sample][t][:,0] > d)* \\\n                     (self.L[sample][t][:,1] > h + core[t])* \\\n                     (self.L[sample][t][:,2] > w + core[t])* \\\n                     (self.L[sample][t][:,0] < d + PATCH_D)* \\\n                     (self.L[sample][t][:,1] < h + PATCH_H - core[t])* \\\n                     (self.L[sample][t][:,2] < w + PATCH_W - core[t])\n            if inside.sum() > 0:\n                searching = False\n\n        if not self.VALID:\n#           Rot90\n            angle = np.random.randint(4)\n            image = torch.rot90(image, k=angle, dims=(-2, -1))\n            hm = torch.rot90(hm, k=angle, dims=(-2, -1))\n#           CONTRAST\n            image *= np.random.normal(1,CONTRAST_AUG)\n#           BRIGTHNESS\n            image += np.random.normal(0,BRIGTHNESS_AUG)\n\n        return image,hm,searching\n    \n    def __getitem__(self, idx):\n        \n        image,hm,searching = self.__rawgetitem__(idx)\n        if not self.VALID:\n            while searching:\n                image,hm,searching = self.__rawgetitem__(np.random.randint(self.__len__()))\n\n        return image,hm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:02:46.820717Z","iopub.execute_input":"2025-01-26T23:02:46.821149Z","iopub.status.idle":"2025-01-26T23:02:46.837982Z","shell.execute_reply.started":"2025-01-26T23:02:46.821113Z","shell.execute_reply":"2025-01-26T23:02:46.836689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = CryoET_Dataset(V,L,M,concat(*SYNTH_SAMPLES))\nlen(ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:02:53.836206Z","iopub.execute_input":"2025-01-26T23:02:53.836699Z","iopub.status.idle":"2025-01-26T23:02:54.153154Z","shell.execute_reply.started":"2025-01-26T23:02:53.836656Z","shell.execute_reply":"2025-01-26T23:02:54.151920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image,heatmap = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(image.sum(0).cpu() + heatmap.sum(0).cpu())\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:06:38.970827Z","iopub.execute_input":"2025-01-26T23:06:38.971335Z","iopub.status.idle":"2025-01-26T23:06:39.253666Z","shell.execute_reply.started":"2025-01-26T23:06:38.971288Z","shell.execute_reply":"2025-01-26T23:06:39.251992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:08:08.610334Z","iopub.execute_input":"2025-01-26T23:08:08.610891Z","iopub.status.idle":"2025-01-26T23:08:08.907855Z","shell.execute_reply.started":"2025-01-26T23:08:08.610846Z","shell.execute_reply":"2025-01-26T23:08:08.906719Z"}},"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        C,H,W = X.shape[-3:]\n        return X.view(-1,PATCH_D,C,H,W).permute(0,2,1,3,4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:08:11.521335Z","iopub.execute_input":"2025-01-26T23:08:11.521692Z","iopub.status.idle":"2025-01-26T23:08:11.527640Z","shell.execute_reply.started":"2025-01-26T23:08:11.521663Z","shell.execute_reply":"2025-01-26T23:08:11.526225Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:08:14.231732Z","iopub.execute_input":"2025-01-26T23:08:14.232146Z","iopub.status.idle":"2025-01-26T23:08:14.238547Z","shell.execute_reply.started":"2025-01-26T23:08:14.232111Z","shell.execute_reply":"2025-01-26T23:08:14.237051Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:08:16.561886Z","iopub.execute_input":"2025-01-26T23:08:16.562379Z","iopub.status.idle":"2025-01-26T23:08:16.577730Z","shell.execute_reply.started":"2025-01-26T23:08:16.562341Z","shell.execute_reply":"2025-01-26T23:08:16.576181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# https://github.com/shuaizzZ/Dice-Loss-PyTorch/blob/master/dice_loss.py\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\n\n\nclass DiceLoss(nn.Module):\n    \"\"\"Dice Loss PyTorch\n        Created by: Zhang Shuai\n        Email: shuaizzz666@gmail.com\n        dice_loss = 1 - 2*p*t / (p^2 + t^2). p and t represent predict and target.\n    Args:\n        weight: An array of shape [C,]\n        predict: A float32 tensor of shape [N, C, *], for Semantic segmentation task is [N, C, H, W]\n        target: A int64 tensor of shape [N, *], for Semantic segmentation task is [N, H, W]\n    Return:\n        diceloss\n    \"\"\"\n    def __init__(self, weight=None):\n        super(DiceLoss, self).__init__()\n        if weight is not None:\n            weight = torch.Tensor(weight)\n            self.weight = weight / torch.sum(weight) # Normalized weight\n        self.smooth = 1e-5\n\n    def forward(self, predict, target):\n        N, C = predict.size()[:2]\n        predict = predict.view(N, C, -1) # (N, C, *)\n        target = target.view(N, 1, -1) # (N, 1, *)\n\n        predict = F.softmax(predict, dim=1) # (N, C, *) ==> (N, C, *)\n        ## convert target(N, 1, *) into one hot vector (N, C, *)\n        target_onehot = torch.zeros(predict.size()).cuda()  # (N, 1, *) ==> (N, C, *)\n        target_onehot.scatter_(1, target, 1)  # (N, C, *)\n\n        intersection = torch.sum(predict * target_onehot, dim=2)  # (N, C)\n        union = torch.sum(predict.pow(2), dim=2) + torch.sum(target_onehot, dim=2)  # (N, C)\n        ## p^2 + t^2 >= 2*p*t, target_onehot^2 == target_onehot\n        dice_coef = (2 * intersection + self.smooth) / (union + self.smooth)  # (N, C)\n\n        if hasattr(self, 'weight'):\n            if self.weight.type() != predict.type():\n                self.weight = self.weight.type_as(predict)\n                dice_coef = dice_coef * self.weight * C  # (N, C)\n        dice_loss = 1 - torch.mean(dice_coef)  # 1\n\n        return dice_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:08:23.672745Z","iopub.execute_input":"2025-01-26T23:08:23.673262Z","iopub.status.idle":"2025-01-26T23:08:23.718370Z","shell.execute_reply.started":"2025-01-26T23:08:23.673216Z","shell.execute_reply":"2025-01-26T23:08:23.716959Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in range(len(SYNTH_SAMPLES)):\n    seed_everything(SEED)\n    model = torch.nn.DataParallel(myUNet2Dto3D(6), device_ids=[0,1])\n\n    train_samples = concat(*SYNTH_SAMPLES[:f],*SYNTH_SAMPLES[f+1:])\n    valid_samples = SYNTH_SAMPLES[f]\n\n    tds = CryoET_Dataset(V,L,M,train_samples)\n    vds = CryoET_Dataset(V,L,M,valid_samples,VALID=True)\n\n    tdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n\n    dls = DataLoaders(tdl,vdl)\n\n    learn = Learner(dls,model,lr=LR,loss_func=DiceLoss(),cbs=[ShowGraphCallback()])\n    learn.fit_one_cycle(EPOCHS)\n    torch.save(model,ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_pretraining_2Dto3D_'+str(f+1))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-26T23:08:26.982180Z","iopub.execute_input":"2025-01-26T23:08:26.982656Z","execution_failed":"2025-01-26T23:09:13.928Z"}},"outputs":[],"execution_count":null}]}