{"metadata":{"kernelspec":{"display_name":"pytorch_gpu","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.7.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":30401.907005,"end_time":"2025-02-04T11:09:18.317995","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-02-04T02:42:36.410990","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"6c23033d","cell_type":"code","source":"import 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\nfrom Metric import compute_metrics,score\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.execute_input":"2025-02-04T02:42:38.974920Z","iopub.status.busy":"2025-02-04T02:42:38.974622Z","iopub.status.idle":"2025-02-04T02:43:12.082595Z","shell.execute_reply":"2025-02-04T02:43:12.081584Z"},"papermill":{"duration":33.114489,"end_time":"2025-02-04T02:43:12.084232","exception":false,"start_time":"2025-02-04T02:42:38.969743","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"f9b46276","cell_type":"code","source":"VOLUMES_PATH = 'C:/Users/Angel/Kaggle/czii-cryo-et-object-identification/train/static/ExperimentRuns/'\nLABELS_PATH = 'C:/Users/Angel/Kaggle/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]\nTARGETS = [\n    'apo-ferritin',# easy\n    'beta-galactosidase',# hard\n    'ribosome',# easy\n    'thyroglobulin',# hard\n    'virus-like-particle'# easy\n]\nSEED = 1337\nFOLDS = [1,2,3,4,5,6,7]\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]*.1 for k in radius\n}\nBS = 4\nLR = 1e-4\nEPOCHS = 1\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25","metadata":{"execution":{"iopub.execute_input":"2025-02-04T02:43:12.096498Z","iopub.status.busy":"2025-02-04T02:43:12.096217Z","iopub.status.idle":"2025-02-04T02:43:12.101877Z","shell.execute_reply":"2025-02-04T02:43:12.101057Z"},"papermill":{"duration":0.012869,"end_time":"2025-02-04T02:43:12.103193","exception":false,"start_time":"2025-02-04T02:43:12.090324","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d3556519","cell_type":"code","source":"V = {}\nL = {}\nfor sample in SAMPLES:\n        L[sample] = {}\n        file = zarr.open(VOLUMES_PATH + sample + '/VoxelSpacing10.000/denoised.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        D,H,W =V[sample].shape   \n        h = 128 - H%128\n        w = 128 - W%128\n        d = 64 - D%64\n        V[sample] = torch.nn.functional.pad(\n            V[sample].unsqueeze(0),\n            (\n                w//2,w - w//2,\n                h//2,h - h//2,\n                d//2,d - d//2\n            ),\n            mode='reflect'\n        )[0]\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]+d//2,\n                    p['location']['y']/scale[1]+h//2,\n                    p['location']['x']/scale[2]+w//2\n                ])\n            L[sample][target] = torch.tensor(L[sample][target]).float().to(device)","metadata":{"execution":{"iopub.execute_input":"2025-02-04T02:43:12.114564Z","iopub.status.busy":"2025-02-04T02:43:12.114308Z","iopub.status.idle":"2025-02-04T02:43:40.903345Z","shell.execute_reply":"2025-02-04T02:43:40.902393Z"},"papermill":{"duration":28.796457,"end_time":"2025-02-04T02:43:40.904913","exception":false,"start_time":"2025-02-04T02:43:12.108456","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e7ff5b93","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":{"execution":{"iopub.execute_input":"2025-02-04T02:43:41.025486Z","iopub.status.busy":"2025-02-04T02:43:41.025238Z","iopub.status.idle":"2025-02-04T02:43:41.091489Z","shell.execute_reply":"2025-02-04T02:43:41.090540Z"},"papermill":{"duration":0.074109,"end_time":"2025-02-04T02:43:41.093233","exception":false,"start_time":"2025-02-04T02:43:41.019124","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"007f3169","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 *= 32\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 > 128: d = 128\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":{"execution":{"iopub.execute_input":"2025-02-04T02:43:41.106711Z","iopub.status.busy":"2025-02-04T02:43:41.106448Z","iopub.status.idle":"2025-02-04T02:43:41.124848Z","shell.execute_reply":"2025-02-04T02:43:41.124033Z"},"papermill":{"duration":0.026221,"end_time":"2025-02-04T02:43:41.126165","exception":false,"start_time":"2025-02-04T02:43:41.099944","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a54201a5","cell_type":"code","source":"ds = CryoET_Dataset(V,L,SAMPLES)","metadata":{"execution":{"iopub.execute_input":"2025-02-04T02:43:41.138127Z","iopub.status.busy":"2025-02-04T02:43:41.137896Z","iopub.status.idle":"2025-02-04T02:43:41.141783Z","shell.execute_reply":"2025-02-04T02:43:41.141193Z"},"papermill":{"duration":0.011095,"end_time":"2025-02-04T02:43:41.142935","exception":false,"start_time":"2025-02-04T02:43:41.131840","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"603c4f39","cell_type":"code","source":"len(ds)","metadata":{"execution":{"iopub.execute_input":"2025-02-04T02:43:41.154934Z","iopub.status.busy":"2025-02-04T02:43:41.154718Z","iopub.status.idle":"2025-02-04T02:43:41.158955Z","shell.execute_reply":"2025-02-04T02:43:41.158322Z"},"papermill":{"duration":0.011602,"end_time":"2025-02-04T02:43:41.160201","exception":false,"start_time":"2025-02-04T02:43:41.148599","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"48060623","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":{"execution":{"iopub.execute_input":"2025-02-04T02:43:41.172507Z","iopub.status.busy":"2025-02-04T02:43:41.172284Z","iopub.status.idle":"2025-02-04T02:43:41.782891Z","shell.execute_reply":"2025-02-04T02:43:41.781845Z"},"papermill":{"duration":0.620048,"end_time":"2025-02-04T02:43:41.785889","exception":false,"start_time":"2025-02-04T02:43:41.165841","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"78a2e020","cell_type":"code","source":"del ds\ngc.collect()","metadata":{"execution":{"iopub.execute_input":"2025-02-04T02:43:41.806698Z","iopub.status.busy":"2025-02-04T02:43:41.806450Z","iopub.status.idle":"2025-02-04T02:43:42.022822Z","shell.execute_reply":"2025-02-04T02:43:42.022042Z"},"papermill":{"duration":0.227709,"end_time":"2025-02-04T02:43:42.024089","exception":false,"start_time":"2025-02-04T02:43:41.796380","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"394c95b5","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":{"execution":{"iopub.execute_input":"2025-02-04T02:43:42.042299Z","iopub.status.busy":"2025-02-04T02:43:42.041982Z","iopub.status.idle":"2025-02-04T02:43:42.045888Z","shell.execute_reply":"2025-02-04T02:43:42.045182Z"},"papermill":{"duration":0.014076,"end_time":"2025-02-04T02:43:42.047012","exception":false,"start_time":"2025-02-04T02:43:42.032936","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"802286d7","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":{"execution":{"iopub.execute_input":"2025-02-04T02:43:42.064843Z","iopub.status.busy":"2025-02-04T02:43:42.064592Z","iopub.status.idle":"2025-02-04T02:43:42.068812Z","shell.execute_reply":"2025-02-04T02:43:42.067998Z"},"papermill":{"duration":0.014423,"end_time":"2025-02-04T02:43:42.069974","exception":false,"start_time":"2025-02-04T02:43:42.055551","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d99ff3df","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.execute_input":"2025-02-04T02:43:42.088050Z","iopub.status.busy":"2025-02-04T02:43:42.087808Z","iopub.status.idle":"2025-02-04T02:43:42.091884Z","shell.execute_reply":"2025-02-04T02:43:42.091073Z"},"papermill":{"duration":0.014356,"end_time":"2025-02-04T02:43:42.093026","exception":false,"start_time":"2025-02-04T02:43:42.078670","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ae5c0e53","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        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                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.execute_input":"2025-02-04T02:43:42.111901Z","iopub.status.busy":"2025-02-04T02:43:42.111639Z","iopub.status.idle":"2025-02-04T02:43:42.119278Z","shell.execute_reply":"2025-02-04T02:43:42.118547Z"},"papermill":{"duration":0.018378,"end_time":"2025-02-04T02:43:42.120624","exception":false,"start_time":"2025-02-04T02:43:42.102246","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"3ff2af01","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":{"execution":{"iopub.execute_input":"2025-02-04T02:43:42.139337Z","iopub.status.busy":"2025-02-04T02:43:42.139017Z","iopub.status.idle":"2025-02-04T02:43:42.145710Z","shell.execute_reply":"2025-02-04T02:43:42.145063Z"},"papermill":{"duration":0.017435,"end_time":"2025-02-04T02:43:42.146937","exception":false,"start_time":"2025-02-04T02:43:42.129502","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ee3cb372","cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n#   model = myUNet2Dto3D(6)\n    model = torch.load('resnet18_3_segmentation_2Dto3D_'+str(f))\n    \n    train_samples = []\n    valid_samples = []\n    for sample in SAMPLES:\n        if sample == SAMPLES[f-1]:\n            valid_samples = valid_samples + [sample]\n        else:\n            train_samples = train_samples + [sample]\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=DiceLoss([2,1,2,1,2,1]),\n        cbs=[\n            ShowGraphCallback()\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS)\n    torch.save(model,ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_FT_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"execution":{"iopub.execute_input":"2025-02-04T02:43:42.165017Z","iopub.status.busy":"2025-02-04T02:43:42.164797Z","iopub.status.idle":"2025-02-04T11:03:14.279792Z","shell.execute_reply":"2025-02-04T11:03:14.279030Z"},"papermill":{"duration":29972.125635,"end_time":"2025-02-04T11:03:14.281321","exception":false,"start_time":"2025-02-04T02:43:42.155686","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"43201b1a","cell_type":"code","source":"del zyx\ngc.collect()","metadata":{"execution":{"iopub.execute_input":"2025-02-04T11:03:14.308852Z","iopub.status.busy":"2025-02-04T11:03:14.308606Z","iopub.status.idle":"2025-02-04T11:03:14.525802Z","shell.execute_reply":"2025-02-04T11:03:14.524899Z"},"papermill":{"duration":0.232024,"end_time":"2025-02-04T11:03:14.526975","exception":false,"start_time":"2025-02-04T11:03:14.294951","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0cb12367","cell_type":"code","source":"PATCH_D = 32\nexperiment = []\nparticle_type = []\nx = []\ny = []\nz = []\nGT_experiment = []\nGT_particle_type = []\nGT_x = []\nGT_y = []\nGT_z = []\ny_pred = torch.zeros(6,32,640,640).float().to(device)\nSTEPS = 11\nfor f in FOLDS:\n    model = torch.load(ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_FT_'+str(f))\n    sample = SAMPLES[f-1]\n    print(sample)\n    MASK = []\n    with torch.no_grad():\n        vv = torch.as_tensor(V[sample][:32]).to(device)\n        y_pred.zero_()\n        for rot in [0,1,2,3]:\n            v = torch.rot90(vv,rot,(-2,-1))\n            y_pred += torch.rot90(model(v)[0].softmax(0),-rot,(-2,-1))\n            \n        MASK.append(y_pred[:,:16].argmax(0).cpu())\n        mask = y_pred[:,16:].clone()\n            \n        for k in tqdm(range(1,STEPS)):\n            vv = torch.as_tensor(V[sample][k*16:k*16+32]).to(device)\n            y_pred.zero_()\n            for rot in [0,1,2,3]:\n                v = torch.rot90(vv,k=rot,dims=(-2,-1))\n                y_pred += torch.rot90(model(v)[0].softmax(0),k=-rot,dims=(-2,-1))\n\n            y_pred[:,:-16] += mask\n            MASK.append((y_pred[:,:16]).argmax(0).cpu())\n            mask = y_pred[:,16:].clone()\n\n        MASK.append(mask.argmax(0).cpu())\n        MASK = torch.cat(MASK).numpy()\n\n    for k in range(len(TARGETS)):\n        labels_out = cc3d.connected_components(MASK == k+1)\n        stats = cc3d.statistics(labels_out)\n        preds = stats['centroids'][1:]\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\n        GT_experiment = GT_experiment + [sample]*len(L[sample][TARGETS[k]])\n        GT_particle_type = GT_particle_type + [TARGETS[k]]*len(L[sample][TARGETS[k]])\n        GT_x = GT_x + list(L[sample][TARGETS[k]][:,2].tolist())\n        GT_y = GT_y + list(L[sample][TARGETS[k]][:,1].tolist())\n        GT_z = GT_z + list(L[sample][TARGETS[k]][:,0].tolist())\n\n        preds = torch.tensor(preds).float().to(device)\n        d = (preds.unsqueeze(1) - L[sample][TARGETS[k]].unsqueeze(0))\n        d = (d*d).sum(-1).sqrt()\n        hits = (d < .05*radius[TARGETS[k]]).max(0)[0].sum().item()\n        miss = (d > .05*radius[TARGETS[k]]).min(1)[0].sum().item()\n\n        print(TARGETS[k])\n        print('hits: ',hits/len(L[sample][TARGETS[k]]))\n        print('miss: ',miss/len(preds))\n        print()\n\n    del MASK,model\n    gc.collect()","metadata":{"execution":{"iopub.execute_input":"2025-02-04T11:03:14.554438Z","iopub.status.busy":"2025-02-04T11:03:14.554085Z","iopub.status.idle":"2025-02-04T11:09:14.941055Z","shell.execute_reply":"2025-02-04T11:09:14.940084Z"},"papermill":{"duration":360.402031,"end_time":"2025-02-04T11:09:14.942650","exception":false,"start_time":"2025-02-04T11:03:14.540619","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0084943d","cell_type":"code","source":"submission = 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})\nsubmission[['x','y','z']] = submission[['x','y','z']]*10\n\nsolution = pd.DataFrame({\n    'id':np.arange(len(GT_experiment)),\n    'experiment':GT_experiment,\n    'particle_type':GT_particle_type,\n    'x':GT_x,\n    'y':GT_y,\n    'z':GT_z\n})\nsolution[['x','y','z']] = solution[['x','y','z']]*10\n\nscore(\n    solution = solution,\n    submission = submission,\n    row_id_column_name = 'id',\n    distance_multiplier = .5,\n    beta = 4\n)","metadata":{"execution":{"iopub.execute_input":"2025-02-04T11:09:14.980824Z","iopub.status.busy":"2025-02-04T11:09:14.980550Z","iopub.status.idle":"2025-02-04T11:09:15.089033Z","shell.execute_reply":"2025-02-04T11:09:15.088182Z"},"papermill":{"duration":0.128301,"end_time":"2025-02-04T11:09:15.090336","exception":false,"start_time":"2025-02-04T11:09:14.962035","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}