{"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"},{"sourceId":10290944,"sourceType":"datasetVersion","datasetId":6368954}],"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 sys\nsys.path.insert(0,'/kaggle/input/czii-metric')\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":"2024-11-08T08:59:28.355532Z","iopub.status.busy":"2024-11-08T08:59:28.355101Z","iopub.status.idle":"2024-11-08T08:59:29.675343Z","shell.execute_reply":"2024-11-08T08:59:29.674142Z","shell.execute_reply.started":"2024-11-08T08:59:28.35549Z"},"trusted":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]\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 = 8\nLR = 2e-4\nEPOCHS = 10\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25","metadata":{},"outputs":[],"execution_count":null},{"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":{},"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":{},"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 *= 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\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#           Flip or not to flip, that's 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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = CryoET_Dataset(V,L,SAMPLES)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(ds)","metadata":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true},"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":{},"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":{},"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":{},"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        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":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = myUNet2Dto3D(\n        len(TARGETS)+1\n    )\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(),\n        cbs=[\n            ShowGraphCallback()\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS)\n    torch.save(model,ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def add_contour(mask,N):\n    for _ in range(N):\n        i,j = np.where(mask)\n        i,j=np.clip(i,1,638),np.clip(j,1,638)\n        mask[i+1,j] = True\n        mask[i-1,j] = True\n        mask[i,j+1] = True\n        mask[i,j-1] = True\n        mask[i+1,j+1] = True\n        mask[i+1,j-1] = True\n        mask[i-1,j+1] = True\n        mask[i-1,j-1] = True\n\n    return mask","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"D,H,W = 192,640,640\nz = torch.stack([\n    torch.stack([\n        torch.arange(D)\n    ]*H)\n]*W).permute(2,1,0)\nx = torch.stack([\n    torch.stack([\n        torch.arange(W)\n    ]*H)\n]*D)\ny = torch.stack([\n    torch.stack([\n        torch.arange(H)\n    ]*W)\n]*D).permute(0,2,1)\nzyx = torch.stack([z,y,x]).float().to(device)\nHM = torch.zeros(192,640,640).long()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_D = 32\nexperiment = []\nparticle_type = []\nx = []\ny = []\nz = []\nGT_experiment = []\nGT_particle_type = []\nGT_x = []\nGT_y = []\nGT_z = []\nprint('TP: GREEN\\nFN: RED\\nFP: YELLOW')\nfor f in FOLDS:\n    model = torch.load(ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_'+str(f))\n    sample = SAMPLES[f-1]\n    print(sample)\n    volume = torch.as_tensor(V[sample]).to(device)\n    MASK = []\n    y_pred = torch.zeros(6,32,640,640).float().to(device)\n    for k in range(len(TARGETS)):\n        for p in L[sample][TARGETS[k]]:\n            r = zyx - p.view(3,1,1,1)\n            cm = (r*r).sum(0).sqrt() < core[TARGETS[k]]\n            HM[cm] = k + 1\n    STEPS = 11\n    with torch.no_grad():\n        vv = volume[: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 = volume[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    efig, e_axes = plt.subplots(1, 3, figsize=(10,10))\n    hfig, h_axes = plt.subplots(1, 2, figsize=(10,10))\n    for k in range(len(TARGETS)):\n        TP = ((MASK == k+1)*(HM == k+1).numpy()).sum(0)\n        FN = ((MASK != k+1)*(HM == k+1).numpy()).sum(0)\n        FP = ((MASK == k+1)*(HM != k+1).numpy()).sum(0)\n        TP = TP/TP.max()\n        FN = FN/FN.max()\n        FP = FP/FP.max()\n        img = np.ones((640,640,3))\n        img[add_contour(FP>0,1)] = 0,0,0\n        img[FP>0] = 1,1,0\n        img[add_contour(FN>0,1)] = 0,0,0\n        img[FN>0] = 1,0,0\n        img[add_contour(TP>0,1)] = 0,0,0\n        img[TP>0] = 0,1,0\n        if k in [0,2,4]:\n            e_axes[k//2].imshow(img)\n            e_axes[k//2].set_title(TARGETS[k])\n        else:\n            h_axes[k//2].imshow(img)\n            h_axes[k//2].set_title(TARGETS[k])\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    plt.show()\n    HM[:] = 0\n    del volume,MASK,model\n    gc.collect()","metadata":{},"outputs":[],"execution_count":null},{"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":{},"outputs":[],"execution_count":null}]}