{"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":"none","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import json\nimport matplotlib.pyplot as plt\nimport seaborn as sns\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":"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]\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_H = 512\nPATCH_W = 512\nradius = {\n    'apo-ferritin':60,# easy\n    'beta-galactosidase':90,# hard\n    'ribosome':150,# easy\n    'thyroglobulin':130,# hard\n    'virus-like-particle':135\n}\nscale = 5/100\nr2 = {\n    k:(radius[k]*scale)*(radius[k]*scale) for k in radius\n}\nBS = 16\nLR = 1e-3\nEPOCHS = 10\nTH = .5\nFLIP_AUG = .5\nCONTRAST_AUG = .25#.125\nBRIGTHNESS_AUG = .25#.125","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        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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"HM = {}\nptotal = []\nfor sample in SAMPLES:\n        print(sample)\n        D,H,W = V[sample].shape\n        z = torch.stack([\n            torch.stack([\n                torch.arange(D)\n            ]*H)\n        ]*W).permute(2,1,0)\n        x = torch.stack([\n                torch.stack([\n                    torch.arange(W)\n                ]*H)\n            ]*D\n        )\n        y = torch.stack([\n            torch.stack([\n                torch.arange(H)\n            ]*W)\n        ]*D).permute(0,2,1)\n        zyx = torch.stack([z,y,x]).float().to(device)\n        HM[sample] = [torch.zeros_like(zyx[0]).to(device)]\n        for target in TARGETS:\n            print(target)\n            mask = 0\n            for p in tqdm(L[sample][target]):\n                r = zyx - p.reshape(-1,1,1,1)\n                r = r*r\n                r = r.sum(0)\n                mask += r < r2[target]\n            HM[sample].append(mask > 0)\n\n        HM[sample] = torch.stack(HM[sample]).argmax(0).cpu()\n    \n        total = []\n        for k in range(len(TARGETS)+1):\n            total.append((HM[sample] == k).sum())\n        ptotal.append(torch.stack(total))\nptotal = torch.stack(ptotal)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ptotal","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CryoET_Dataset(Dataset):\n    def __init__(self, V, L, HM, samples, VALID=False):\n        self.data = V\n        self.L = L\n        self.hm = HM\n        self.samples = samples\n        self.sample = []\n        self.d = []\n        for sample in samples:\n            self.d = self.d + torch.arange(len(V[sample])).tolist()\n            self.sample = self.sample + [sample]*len(V[sample])\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        H,W = self.data[sample].shape[-2:]\n        d = self.d[idx]\n\n        if not self.VALID:\n            h = np.random.randint(H - PATCH_H)\n            w = np.random.randint(W - PATCH_W)\n        else :\n            h = (H - PATCH_H)//2\n            w = (W - PATCH_W)//2\n\n        hm = torch.as_tensor(self.hm[sample][\n            d,\n            h:h+PATCH_H,\n            w:w+PATCH_W\n        ]).to(device)\n\n        image = torch.as_tensor(self.data[sample][\n            d,\n            h:h+PATCH_H,\n            w:w+PATCH_W\n        ]).to(device)\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#           FreeRot\n            '''image = torch.nn.functional.pad(\n                image.unsqueeze(0),\n                (\n                    64,64,\n                    64,64\n                ),\n                mode='reflect'\n            )[0]\n            hm = torch.nn.functional.pad(\n                hm.unsqueeze(0).float(),\n                (\n                    64,64,\n                    64,64\n                ),\n                mode='reflect'\n            )[0].long()\n            angle = torch.as_tensor(random.uniform(-180, 180))\n            image = torchvision.transforms.functional.rotate(\n                image,angle.item(),\n#               interpolation=torchvision.transforms.InterpolationMode.BILINEAR\n            )[:,64:-64,64:-64]\n            hm = torchvision.transforms.functional.rotate(\n                hm,angle.item(),\n#               interpolation=torchvision.transforms.InterpolationMode.BILINEAR\n            )[:,64:-64,64:-64]'''\n#           Rot180\n            '''if np.random.rand() < .5:\n                axis = np.random.randint(2) - 2\n                image = torch.rot90(image, k=2, dims=(-3, axis))\n                hm = torch.rot90(hm, k=2, dims=(-3, axis))'''\n#           Flipz\n            '''if np.random.rand() < .5:\n                image = image.flip(0)\n                hm = hm.flip(0)'''\n#           CONTRAST\n            x = np.random.normal(1,CONTRAST_AUG)\n            image *= x\n#           BRIGTHNESS\n            x = np.random.normal(0,BRIGTHNESS_AUG)\n            image += x\n\n        return image,hm\n    \n    def __getitem__(self, idx):\n        \n        image,hm = self.__rawgetitem__(idx)\n        if not self.VALID:\n            while hm.sum() == 0:\n                image,hm = self.__rawgetitem__(np.random.randint(self.__len__()))\n\n        return image,hm","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tds = CryoET_Dataset(V,L,HM,SAMPLES)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(tds)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image,heatmap = tds.__getitem__(np.random.randint(len(tds)))\nplt.imshow(image.cpu() + (heatmap.cpu()))\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del tds\ngc.collect()","metadata":{},"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 myUNet2D(nn.Module):\n    def __init__(\n        self,\n        classes\n        ):\n        super(myUNet2D, 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    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":"max_radius_cube = max(radius.values())\nmax_radius_cube = max_radius_cube*max_radius_cube*max_radius_cube\nmax_radius_cube","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weight = torch.tensor([.1]+[max_radius_cube/(radius[k]*radius[k]*radius[k]) for k in radius])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#weight[2] *= 2\n#weight[4] *= 2\nweight","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n    model = myUNet2D(\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    weight = ptotal[:,1:].clone()\n    weight[f-1] = 0\n    weight = weight.sum(0)\n    weight = weight/(weight.max())\n    weight = 1/weight\n    weight = torch.cat([torch.tensor([.1]),weight])\n\n    print(weight)\n\n    tds = CryoET_Dataset(V,L,HM,train_samples)\n    vds = CryoET_Dataset(V,L,HM,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=nn.CrossEntropyLoss(\n            weight=weight\n        ),\n        cbs=[\n            ShowGraphCallback(),\n#           GradientAccumulation(n_acc=4)\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS)\n    torch.save(model,'CryoET_segmentation_2D_'+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,628),np.clip(j,1,628)\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":"BS = 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('CryoET_segmentation_2D_'+str(f))\n    sample = SAMPLES[f-1]\n    print(sample)\n    volume = torch.as_tensor(V[sample]).to(device)\n    volume = torch.nn.functional.pad(\n            volume.unsqueeze(0),\n            (\n                5,5,\n                5,5\n            ),\n            mode='reflect'\n        )[0]\n    MASK = []\n    STEPS = len(volume)//BS\n    if len(volume) - STEPS*BS > 0: STEPS += 1\n    with torch.no_grad():\n            for k in tqdm(range(STEPS)):\n                mask = model(\n                        volume[\n                            k*BS:(k+1)*BS\n                        ]\n                ).argmax(1).cpu()\n                \n                MASK.append(mask)\n\n    MASK = torch.cat(MASK)[:,5:-5,5:-5].numpy()\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[sample] == k+1).numpy()).sum(0)\n        FN = ((MASK != k+1)*(HM[sample] == k+1).numpy()).sum(0)\n        FP = ((MASK == k+1)*(HM[sample] != k+1).numpy()).sum(0)\n        TP = TP/TP.max()\n        FN = FN/FN.max()\n        FP = FP/FP.max()\n        img = np.ones((630,630,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\n    del volume,MASK,model","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\nsubmission.tail()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"solution = 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\nsolution.tail()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"score(\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}]}