{"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":10358906,"sourceType":"datasetVersion","datasetId":6415002},{"sourceId":10558265,"sourceType":"datasetVersion","datasetId":6532180}],"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!pip3 install topaz-em\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\nimport topaz.denoise\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":"2025-01-25T23:16:06.081969Z","iopub.execute_input":"2025-01-25T23:16:06.082202Z","iopub.status.idle":"2025-01-25T23:19:11.368382Z","shell.execute_reply.started":"2025-01-25T23:16:06.082174Z","shell.execute_reply":"2025-01-25T23:19:11.367513Z"},"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/'\nSAMPLES = [\n    'TS_73_6',\n    'TS_86_3',\n    'TS_69_2',\n    'TS_6_4',\n    'TS_99_9',\n    'TS_5_4',\n    'TS_6_6'\n]\nSYNTH_SAMPLES = [\n    'TS_19', 'TS_17', 'TS_10',\n    'TS_13', 'TS_25', 'TS_26', 'TS_15',\n    'TS_5', 'TS_8', 'TS_20', 'TS_4',\n    'TS_0', 'TS_12', 'TS_2', 'TS_24',\n    'TS_22', 'TS_18', 'TS_16', 'TS_3',\n    'TS_6', 'TS_11', 'TS_9', 'TS_23',\n    'TS_7', 'TS_14', 'TS_21', 'TS_1'\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 = 4\nLR = 2.5e-5\nEPOCHS = 1\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:19:11.369396Z","iopub.execute_input":"2025-01-25T23:19:11.369703Z","iopub.status.idle":"2025-01-25T23:19:11.376609Z","shell.execute_reply.started":"2025-01-25T23:19:11.369674Z","shell.execute_reply":"2025-01-25T23:19:11.375760Z"}},"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-25T23:19:11.377549Z","iopub.execute_input":"2025-01-25T23:19:11.377874Z","iopub.status.idle":"2025-01-25T23:19:11.398038Z","shell.execute_reply.started":"2025-01-25T23:19:11.377843Z","shell.execute_reply":"2025-01-25T23:19:11.397312Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"V = {}\nL = {}\nfor sample in SYNTH_SAMPLES:\n        L[sample] = {}\n        file = zarr.open('/kaggle/input/czii-synth-data/'+sample+'.zarr', mode='r')\n        scale = file.attrs['multiscales'][0]['datasets'][0]['coordinateTransformations'][0]['scale']\n        vol = np.array(file[0])\n        pmin,pmax = np.percentile(vol,(1,99))\n        V[sample] = (vol - pmin)/(pmax - pmin)\n        V[sample] = torch.as_tensor(V[sample])\n        V[sample] = torch.nn.functional.pad(\n            V[sample].unsqueeze(0),\n            (\n                5,5,\n                5,5\n            ),\n            mode='reflect'\n        )[0]\n        for target in TARGETS:\n            L[sample][target] = []\n            for p in pd.read_json('/kaggle/input/czii-synth-data/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-25T23:19:11.399943Z","iopub.execute_input":"2025-01-25T23:19:11.400189Z","iopub.status.idle":"2025-01-25T23:21:21.347691Z","shell.execute_reply.started":"2025-01-25T23:19:11.400170Z","shell.execute_reply":"2025-01-25T23:21:21.346986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for sample in SAMPLES:\n    L[sample] = {}\n    file = zarr.open(VOLUMES_PATH + sample + '/VoxelSpacing10.000/denoised.zarr', mode='r')\n    for target in TARGETS:\n        L[sample][target] = []\n        f = open(LABELS_PATH + sample + '/Picks/' + target + '.json')\n        for p in json.loads(f.read())['points']:\n            L[sample][target].append([\n                p['location']['z']/scale[0],\n                p['location']['y']/scale[1],\n                p['location']['x']/scale[2]\n            ])\n        L[sample][target] = torch.tensor(L[sample][target]).float().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:21.348931Z","iopub.execute_input":"2025-01-25T23:21:21.349184Z","iopub.status.idle":"2025-01-25T23:21:21.541965Z","shell.execute_reply.started":"2025-01-25T23:21:21.349163Z","shell.execute_reply":"2025-01-25T23:21:21.541386Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for sample in SAMPLES:\n    L[sample]['distribution'] = []\n    for target in TARGETS:\n        L[sample]['distribution'].append(len(L[sample][target]))\n    L[sample]['distribution'] = torch.tensor(L[sample]['distribution']).float()\n    L[sample]['distribution'] /= L[sample]['distribution'].sum()\n    plt.plot(L[sample]['distribution'],'lightgreen')\nfor sample in SYNTH_SAMPLES:\n    L[sample]['distribution'] = []\n    for target in TARGETS:\n        L[sample]['distribution'].append(len(L[sample][target]))\n    L[sample]['distribution'] = torch.tensor(L[sample]['distribution']).float()\n    L[sample]['distribution'] /= L[sample]['distribution'].sum()\n    plt.plot(L[sample]['distribution'],'r')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:21.542729Z","iopub.execute_input":"2025-01-25T23:21:21.543026Z","iopub.status.idle":"2025-01-25T23:21:21.825834Z","shell.execute_reply.started":"2025-01-25T23:21:21.542995Z","shell.execute_reply":"2025-01-25T23:21:21.825094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"t = []\nfor sample in SYNTH_SAMPLES:\n    total = 0\n    for target in TARGETS:\n        total += len(L[sample][target])\n    t.append(total)\n    L[sample]['t'] = total\nplt.bar(np.arange(len(t)),t)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:21.826818Z","iopub.execute_input":"2025-01-25T23:21:21.827196Z","iopub.status.idle":"2025-01-25T23:21:22.001366Z","shell.execute_reply.started":"2025-01-25T23:21:21.827162Z","shell.execute_reply":"2025-01-25T23:21:22.000648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1 FOLD x3 and 6 FOLD x4\nSYNTH_SAMPLES.sort(key=lambda k:L[k]['t'],reverse=True)\nt = []\nfor sample in SYNTH_SAMPLES:\n    total = 0\n    for target in TARGETS:\n        total += len(L[sample][target])\n    t.append(total)\n    L[sample]['t'] = total\nplt.bar(np.arange(len(t)),t)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.002086Z","iopub.execute_input":"2025-01-25T23:21:22.002310Z","iopub.status.idle":"2025-01-25T23:21:22.175759Z","shell.execute_reply.started":"2025-01-25T23:21:22.002292Z","shell.execute_reply":"2025-01-25T23:21:22.175042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FOLDS = [SYNTH_SAMPLES[:3]]\nFOLDS[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.176669Z","iopub.execute_input":"2025-01-25T23:21:22.176992Z","iopub.status.idle":"2025-01-25T23:21:22.181879Z","shell.execute_reply.started":"2025-01-25T23:21:22.176959Z","shell.execute_reply":"2025-01-25T23:21:22.181114Z"}},"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-25T23:21:22.182857Z","iopub.execute_input":"2025-01-25T23:21:22.183092Z","iopub.status.idle":"2025-01-25T23:21:22.197108Z","shell.execute_reply.started":"2025-01-25T23:21:22.183049Z","shell.execute_reply":"2025-01-25T23:21:22.196403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"REMAINING = SYNTH_SAMPLES[3:]\nseed_everything(SEED)\nrandom.shuffle(REMAINING)\nfor k in range(6):\n    FOLDS.append(REMAINING[k*4:k*4+4])\n    print(FOLDS[-1])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.197968Z","iopub.execute_input":"2025-01-25T23:21:22.198242Z","iopub.status.idle":"2025-01-25T23:21:22.215845Z","shell.execute_reply.started":"2025-01-25T23:21:22.198211Z","shell.execute_reply":"2025-01-25T23:21:22.215128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#2nd sample has the higher amount of particles\nt = []\nfor sample in SAMPLES:\n    total = 0\n    for target in TARGETS:\n        total += len(L[sample][target])\n    t.append(total)\n    L[sample]['t'] = total\nplt.bar(np.arange(len(t)),t)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.216621Z","iopub.execute_input":"2025-01-25T23:21:22.216954Z","iopub.status.idle":"2025-01-25T23:21:22.368873Z","shell.execute_reply.started":"2025-01-25T23:21:22.216923Z","shell.execute_reply":"2025-01-25T23:21:22.368047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"z = torch.stack([\n    torch.stack([\n        torch.arange(PATCH_D)\n    ]*PATCH_H)\n]*PATCH_W).permute(2,1,0)\nx = torch.stack([\n    torch.stack([\n        torch.arange(PATCH_W)\n    ]*PATCH_H)\n]*PATCH_D)\ny = torch.stack([\n    torch.stack([\n        torch.arange(PATCH_H)\n    ]*PATCH_W)\n]*PATCH_D).permute(0,2,1)\nzyx = torch.stack([z,y,x]).float().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.371491Z","iopub.execute_input":"2025-01-25T23:21:22.371695Z","iopub.status.idle":"2025-01-25T23:21:22.418691Z","shell.execute_reply.started":"2025-01-25T23:21:22.371678Z","shell.execute_reply":"2025-01-25T23:21:22.418024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CryoET_Dataset(Dataset):\n    def __init__(self, V, L, samples, VALID=False):\n        self.data = V\n        self.L = L\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        o = torch.tensor([d,h,w]).to(device)\n        hm = torch.zeros(\n            PATCH_D,\n            PATCH_H,\n            PATCH_W\n        ).long().to(device)\n        searching = True\n        for k in range(len(TARGETS)):\n#           Find and label particles within this voxel\n            t = TARGETS[k]\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                r = zyx.view(1,3,PATCH_D,PATCH_H,PATCH_W) - (self.L[sample][t][inside] - o).view(-1,3,1,1,1)\n                cm = ((r*r).sum(1).sqrt() < core[t]).sum(0) > 0\n                hm[cm] = k + 1\n#           Find and mask partially present particles\n            pinside = ~inside* \\\n                      (self.L[sample][t][:,0] > d - shell[t])* \\\n                      (self.L[sample][t][:,1] > h - shell[t])* \\\n                      (self.L[sample][t][:,2] > w - shell[t])* \\\n                      (self.L[sample][t][:,0] < d + PATCH_D + shell[t])* \\\n                      (self.L[sample][t][:,1] < h + PATCH_H + shell[t])* \\\n                      (self.L[sample][t][:,2] < w + PATCH_W + shell[t])\n            if pinside.sum() > 0:\n                r = zyx.view(1,3,PATCH_D,PATCH_H,PATCH_W) - (self.L[sample][t][pinside] - o).view(-1,3,1,1,1)\n                sm = ((r*r).sum(1).sqrt() < shell[t]).sum(0) > 0\n                image[sm] = image[sm][torch.randperm(sm.sum())]\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#           To flip or not to flip, that is the question\n            if np.random.rand() < .5:\n                axis = np.random.randint(2) - 2\n                image = image.flip(axis)\n                hm = hm.flip(axis)\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-25T23:21:22.419910Z","iopub.execute_input":"2025-01-25T23:21:22.420235Z","iopub.status.idle":"2025-01-25T23:21:22.434957Z","shell.execute_reply.started":"2025-01-25T23:21:22.420204Z","shell.execute_reply":"2025-01-25T23:21:22.434286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = CryoET_Dataset(V,L,SYNTH_SAMPLES)\nlen(ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.435640Z","iopub.execute_input":"2025-01-25T23:21:22.435829Z","iopub.status.idle":"2025-01-25T23:21:22.458350Z","shell.execute_reply.started":"2025-01-25T23:21:22.435812Z","shell.execute_reply":"2025-01-25T23:21:22.457674Z"}},"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-25T23:21:22.458964Z","iopub.execute_input":"2025-01-25T23:21:22.459196Z","iopub.status.idle":"2025-01-25T23:21:22.868041Z","shell.execute_reply.started":"2025-01-25T23:21:22.459178Z","shell.execute_reply":"2025-01-25T23:21:22.867132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:22.868820Z","iopub.execute_input":"2025-01-25T23:21:22.869040Z","iopub.status.idle":"2025-01-25T23:21:23.075688Z","shell.execute_reply.started":"2025-01-25T23:21:22.869022Z","shell.execute_reply":"2025-01-25T23:21:23.074975Z"}},"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)\n    \nclass 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)\n    \nclass myUNet2Dto3D(nn.Module):\n    def __init__(\n        self,\n        classes\n        ):\n        super(myUNet2Dto3D, self).__init__()\n        \n        decoder_channels = (256, 128, 64, 32, 16)[-ENCODER_DEPTH:]\n        decoder_in_channels = (768, 384, 192, 128 , 32)[-ENCODER_DEPTH:]    \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                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                classes,\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-25T23:21:23.076665Z","iopub.execute_input":"2025-01-25T23:21:23.076996Z","iopub.status.idle":"2025-01-25T23:21:23.095360Z","shell.execute_reply.started":"2025-01-25T23:21:23.076960Z","shell.execute_reply":"2025-01-25T23:21:23.094555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myDenoisET(nn.Module):\n    def __init__(\n        self,\n        ):\n        super(myDenoisET, self).__init__()    \n        \n        self.model = topaz.denoise.Denoise3D('unet').model.to(device)\n        \n    def forward(self,X):\n        H,W = X.shape[-2:]\n        x = self.model(X.view(-1,1,H,W))\n        \n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:23.096180Z","iopub.execute_input":"2025-01-25T23:21:23.096413Z","iopub.status.idle":"2025-01-25T23:21:23.114707Z","shell.execute_reply.started":"2025-01-25T23:21:23.096394Z","shell.execute_reply":"2025-01-25T23:21:23.114139Z"}},"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\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-25T23:21:23.115437Z","iopub.execute_input":"2025-01-25T23:21:23.115666Z","iopub.status.idle":"2025-01-25T23:21:23.130808Z","shell.execute_reply.started":"2025-01-25T23:21:23.115636Z","shell.execute_reply":"2025-01-25T23:21:23.130142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myLoss(nn.Module):\n    def __init__(self, model, weight=None):\n        super(myLoss, self).__init__()\n        self.Loss = DiceLoss(weight)\n        self.model = model.eval()\n\n    def forward(self, predict, target):\n        predict = self.model(predict.view(-1,PATCH_D,PATCH_H,PATCH_W))\n\n        return self.Loss(predict, target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:23.131504Z","iopub.execute_input":"2025-01-25T23:21:23.131772Z","iopub.status.idle":"2025-01-25T23:21:23.151721Z","shell.execute_reply.started":"2025-01-25T23:21:23.131745Z","shell.execute_reply":"2025-01-25T23:21:23.151121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in range(len(FOLDS)):\n    seed_everything(SEED)\n    model = torch.nn.DataParallel(myDenoisET(), device_ids=[0,1])\n\n    train_samples = concat(*FOLDS[:f],*FOLDS[f+1:])\n    valid_samples = FOLDS[f]\n\n    tds = CryoET_Dataset(V,L,train_samples)\n    vds = CryoET_Dataset(V,L,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(\n        dls,\n        model,\n        lr=LR,\n        loss_func=myLoss(torch.nn.DataParallel(torch.load(\n            '/kaggle/input/resnet18-3-segmentation-2dto3d/resnet18_3_segmentation_2Dto3D_'+str([2,1,3,4,5,6,7][f])#2nd sample has the higher amount of particles\n        ), device_ids=[0,1])),\n        cbs=[ShowGraphCallback()]\n    )\n    learn.fit_one_cycle(EPOCHS)\n    torch.save(model,'myDenoisET_'+str([2,1,3,4,5,6,7][f]))#2nd sample has the higher amount of particles\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-25T23:21:23.152496Z","iopub.execute_input":"2025-01-25T23:21:23.152680Z","iopub.status.idle":"2025-01-26T00:39:51.363703Z","shell.execute_reply.started":"2025-01-25T23:21:23.152665Z","shell.execute_reply":"2025-01-26T00:39:51.362299Z"}},"outputs":[],"execution_count":null}]}