{"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10381424,"sourceType":"datasetVersion","datasetId":6430991},{"sourceId":10290944,"sourceType":"datasetVersion","datasetId":6368954}],"dockerImageVersionId":30823,"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.status.busy":"2025-01-06T00:52:52.256736Z","iopub.execute_input":"2025-01-06T00:52:52.257180Z","iopub.status.idle":"2025-01-06T00:53:19.093392Z","shell.execute_reply.started":"2025-01-06T00:52:52.257126Z","shell.execute_reply":"2025-01-06T00:53:19.092501Z"},"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\"#\"efficientnet-b0\",\"resnet18\",\"timm-regnetx_002\",\"densenet121\",timm-resnest14d\nENCODER_DEPTH = 3\nPATCH_D = 64\nPATCH_H = 256\nPATCH_W = 256\nradius = {\n    'apo-ferritin':60,# easy\n    'beta-galactosidase':90,# hard\n    'ribosome':150,# easy\n    'thyroglobulin':130,# hard\n    'virus-like-particle':135# hard\n}\ncore = {\n    k:radius[k]*.05 for k in radius\n}\nshell = {\n    k:radius[k]*.12 for k in radius\n}\nLR = 1e-4#5e-4\nEPOCHS = 10#30#50#10\nTH = .5\nCONTRAST_AUG = .25#.125\nBRIGTHNESS_AUG = .25#.125","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:19.094481Z","iopub.execute_input":"2025-01-06T00:53:19.094731Z","iopub.status.idle":"2025-01-06T00:53:19.100760Z","shell.execute_reply.started":"2025-01-06T00:53:19.094709Z","shell.execute_reply":"2025-01-06T00:53:19.099992Z"}},"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        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.as_tensor(V[sample])\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:19.103246Z","iopub.execute_input":"2025-01-06T00:53:19.103590Z","iopub.status.idle":"2025-01-06T00:53:45.799392Z","shell.execute_reply.started":"2025-01-06T00:53:19.103559Z","shell.execute_reply":"2025-01-06T00:53:45.798705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CryoET_Mosaic_Dataset(Dataset):\n#   Simple homemade data Dataset\n#   We will fold 16 64x64x64 particle voxels in one single 64x256x256 \"synthetic sample\" and see what happens\n#   The idea is to learn as much as possible from hard to remember synthetic data for train with real data after\n    def __init__(self, V, L, samples):\n        z = torch.stack([\n            torch.stack([\n                torch.arange(64)\n            ]*64)\n        ]*64).permute(2,1,0)\n        x = torch.stack([\n            torch.stack([\n                torch.arange(64)\n            ]*64)\n        ]*64)\n        y = torch.stack([\n            torch.stack([\n                torch.arange(64)\n            ]*64)\n        ]*64).permute(0,2,1)\n        self.zyx = torch.stack([z,y,x]).float().to(device)\n        self.data = V\n        self.L = L\n        self.samples = samples\n        self.s = []\n        self.t = []\n        self.p = []\n        for s in samples:\n            for t in TARGETS:\n                n = len(L[s][t])\n                self.s = self.s + [s]*n\n                self.t = self.t + [t]*n\n                self.p = self.p + torch.arange(n).tolist()\n\n        r = len(self.s)%16\n        if r > 0:\n            self.s = self.s + list(np.array(samples)[np.random.randint(0,len(samples),16-r)])\n            self.t = self.t + ['background']*(16-r)\n            self.p = self.p + [None]*(16-r)\n\n        zipped = list(zip(self.s,self.t,self.p))\n        random.shuffle(zipped)\n        self.s,self.t,self.p = zip(*zipped)\n\n    def __len__(self):\n        return len(self.s)//16\n    \n    def __shuffle__(self):\n        zipped = list(zip(self.s,self.t,self.p))\n        random.shuffle(zipped)\n        self.s,self.t,self.p = zip(*zipped)\n\n    def __getitem__(self, idx):\n        \n        samples = self.s[idx*16:16*(idx+1)]\n        targets = self.t[idx*16:16*(idx+1)]\n        ipoints = self.p[idx*16:16*(idx+1)]\n\n        image = []\n        mask = []\n        for k in range(16):\n            if targets[k] != 'background':\n                img = torch.zeros(64,64*3,64*3).to(device) + .5\n                m = torch.zeros(64,64*3,64*3).long().to(device) - 100\n\n                p = z,y,x = np.rint(self.L[samples[k]][targets[k]][ipoints[k]].cpu()).long().to(device)\n                d = z - 32\n                h = y - 96\n                w = x - 96\n                \n                freedom = 32 - int(shell[targets[k]])\n                d += np.random.randint(2*freedom) - freedom\n                h += np.random.randint(2*freedom) - freedom\n                w += np.random.randint(2*freedom) - freedom\n                o = torch.tensor([d,h,w]).to(device)\n\n                end_d = d + 64\n                end_h = h + 192\n                end_w = w + 192\n                if d < 0:\n                    start_d = - d\n                    d = 0\n                else:\n                    start_d = 0\n                if h < 0:\n                    start_h = - h\n                    h = 0\n                else:\n                    start_h = 0\n                if w < 0:\n                    start_w = - w\n                    w = 0\n                else:\n                    start_w = 0\n\n                crop = torch.as_tensor(self.data[samples[k]][\n                    d:end_d,\n                    h:end_h,\n                    w:end_w\n                ]).to(device)\n                D,H,W = crop.shape\n                \n                img[\n                    start_d:start_d+D,\n                    start_h:start_h+H,\n                    start_w:start_w+W\n                ] = crop\n                m[\n                    start_d:start_d+D,\n                    start_h:start_h+H,\n                    start_w:start_w+W\n                ] = 0\n                center = (p - o)[[2,1]].tolist()\n#               Free rotation around chosen particle\n                angle = torch.as_tensor(random.uniform(-180, 180))\n                img = torchvision.transforms.functional.rotate(\n                    img,angle.item(),\n#                   interpolation=torchvision.transforms.InterpolationMode.BILINEAR,\n                    center=center\n                )\n                m = torchvision.transforms.functional.rotate(\n                    m,angle.item(),\n#                   interpolation=torchvision.transforms.InterpolationMode.BILINEAR,\n                    center=center\n                )\n                angle = -angle*math.pi/180\n                s = torch.sin(angle)\n                c = torch.cos(angle)\n                rot = torch.stack([\n                    torch.stack([c, s]),\n                    torch.stack([-s, c])\n                ]).to(device)\n                img = img[:,64:128,64:128]\n                m = m[:,64:128,64:128]\n                missing = m == - 100\n                img[missing] = img[~missing].mean()\n                m[:] = 0\n#               What particles are completely within this voxel?\n                for kk in range(5):\n                    t = TARGETS[kk]\n                    particles = (self.L[samples[k]][t] - o)\n                    center = torch.as_tensor(center).float().to(device)\n                    particles[:,[[2,1]]] = ((particles[:,[[2,1]]] - center) @ rot) + center - 64\n                    inside = (particles[:,0] > 0)* \\\n                             (particles[:,0] < 64)* \\\n                             (particles[:,1] > shell[t])* \\\n                             (particles[:,2] > shell[t])* \\\n                             (particles[:,1] < 64 - shell[t])* \\\n                             (particles[:,2] < 64 - shell[t])\n                    if inside.sum() > 0:\n                        r = self.zyx.view(1,3,64,64,64) - particles[inside].view(-1,3,1,1,1)\n                        r = (r*r).sum(1).sqrt()\n                        cm = (r < core[t]).sum(0) > 0\n                        m[cm] = kk + 1\n#                   Find and mask partially present particles\n                    pinside = ~inside* \\\n                              (particles[:,0] > - shell[t])* \\\n                              (particles[:,1] > - shell[t])* \\\n                              (particles[:,2] > - shell[t])* \\\n                              (particles[:,0] < 64 + shell[t])* \\\n                              (particles[:,1] < 64 + shell[t])* \\\n                              (particles[:,2] < 64 + shell[t])\n                    if pinside.sum() > 0:\n                        r = self.zyx.view(1,3,64,64,64) - particles[pinside].view(-1,3,1,1,1)\n                        r = (r*r).sum(1).sqrt()\n                        sm = (r < shell[t]).sum(0) > 0\n#                       Confetti\n                        img[sm] = img[sm][torch.randperm(sm.sum())]\n\n#               Rot90\n                angle = np.random.randint(4)\n                img = torch.rot90(img, k=angle, dims=(-2, -1))\n                m = torch.rot90(m, k=angle, dims=(-2, -1))\n#               Flip party\n                for axis in [0,1,2]:\n                    if np.random.rand() < .5:\n                        img = img.flip(axis)\n                        m = m.flip(axis)\n\n            else:\n                searching = True\n                while searching:\n                    s = self.samples[np.random.randint(len(self.samples))]\n                    D,H,W = self.data[s].shape\n                    d = np.random.randint(D - 64)\n                    h = np.random.randint(H - 64)\n                    w = np.random.randint(W - 64)\n#                   What particles are partially within this voxel?\n                    searching =  False\n                    for kk in range(5):\n                        t = TARGETS[kk]\n                        inside = (self.L[samples[k]][t][:,0] > d - shell[t])* \\\n                                 (self.L[samples[k]][t][:,1] > h - shell[t])* \\\n                                 (self.L[samples[k]][t][:,2] > w - shell[t])* \\\n                                 (self.L[samples[k]][t][:,0] < d + 64 + shell[t])* \\\n                                 (self.L[samples[k]][t][:,1] < h + 64 + shell[t])* \\\n                                 (self.L[samples[k]][t][:,2] < w + 64 + shell[t])\n                        if inside.sum() > 0: searching = True\n\n                img = torch.as_tensor(self.data[s][\n                    d:d+64,\n                    h:h+64,\n                    w:w+64\n                ]).to(device)\n                m = torch.zeros(64,64,64).long().to(device)\n                \n            image.append(img)\n            mask.append(m)\n\n        image = torch.cat([\n            torch.cat(image[:4],-1),\n            torch.cat(image[4:8],-1),\n            torch.cat(image[8:12],-1),\n            torch.cat(image[12:16],-1)\n        ],-2)\n        mask = torch.cat([\n            torch.cat(mask[:4],-1),\n            torch.cat(mask[4:8],-1),\n            torch.cat(mask[8:12],-1),\n            torch.cat(mask[12:16],-1)\n        ],-2)\n#       Wrap\n        for k in range(4):\n                drift = np.random.randint(64)\n                image[:,k*64:(k+1)*64] = torch.nn.functional.pad(\n                    image[:,k*64:(k+1)*64],\n                    (64,64),\n                    mode='circular'\n                )[:,:,drift:drift+256]\n                mask[:,k*64:(k+1)*64] = torch.nn.functional.pad(\n                    mask[:,k*64:(k+1)*64],\n                    (64,64),\n                    mode='circular'\n                )[:,:,drift:drift+256]\n#       Rot90\n        angle = np.random.randint(4)\n        image = torch.rot90(image, k=angle, dims=(-2, -1))\n        mask = torch.rot90(mask, 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 party\n        for axis in [0,1,2]:\n            if np.random.rand() < .5:\n                image = image.flip(axis)\n                mask = mask.flip(axis)\n\n        return image,mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:45.800608Z","iopub.execute_input":"2025-01-06T00:53:45.800837Z","iopub.status.idle":"2025-01-06T00:53:45.828133Z","shell.execute_reply.started":"2025-01-06T00:53:45.800816Z","shell.execute_reply":"2025-01-06T00:53:45.827224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = CryoET_Mosaic_Dataset(V,L,SAMPLES)\nlen(ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:45.886702Z","iopub.execute_input":"2025-01-06T00:53:45.886992Z","iopub.status.idle":"2025-01-06T00:53:45.930095Z","shell.execute_reply.started":"2025-01-06T00:53:45.886972Z","shell.execute_reply":"2025-01-06T00:53:45.929303Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image,mask = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(image.sum(0).cpu() + mask.sum(0).cpu())\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:45.930997Z","iopub.execute_input":"2025-01-06T00:53:45.931227Z","iopub.status.idle":"2025-01-06T00:53:47.096321Z","shell.execute_reply.started":"2025-01-06T00:53:45.931208Z","shell.execute_reply":"2025-01-06T00:53:47.095550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.097187Z","iopub.execute_input":"2025-01-06T00:53:47.097407Z","iopub.status.idle":"2025-01-06T00:53:47.283096Z","shell.execute_reply.started":"2025-01-06T00:53:47.097386Z","shell.execute_reply":"2025-01-06T00:53:47.282228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Callback to shuffle tds\ndef cb(self):\n    learn.dls.train_ds.__shuffle__()\nshuffle_cb = Callback(before_epoch=cb)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.286350Z","iopub.execute_input":"2025-01-06T00:53:47.286602Z","iopub.status.idle":"2025-01-06T00:53:47.294984Z","shell.execute_reply.started":"2025-01-06T00:53:47.286582Z","shell.execute_reply":"2025-01-06T00:53:47.294158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CryoET_Dataset(Dataset):\n    def __init__(self, V, L, samples, VALID=False, D=64, H=128, W=128):\n        self.DHW = (D,H,W)\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        y = torch.stack([\n            torch.stack([\n                torch.arange(H)\n            ]*W)\n        ]*D).permute(0,2,1)\n        self.zyx = torch.stack([z,y,x]).float().to(device)\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        PATCH_D,PATCH_H,PATCH_W = self.DHW\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 = self.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 = self.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#               Confetti mask\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#           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#           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,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-06T00:53:47.296323Z","iopub.execute_input":"2025-01-06T00:53:47.296560Z","iopub.status.idle":"2025-01-06T00:53:47.314595Z","shell.execute_reply.started":"2025-01-06T00:53:47.296529Z","shell.execute_reply":"2025-01-06T00:53:47.313974Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = CryoET_Dataset(V,L,SAMPLES)\nlen(ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.315261Z","iopub.execute_input":"2025-01-06T00:53:47.315511Z","iopub.status.idle":"2025-01-06T00:53:47.345365Z","shell.execute_reply.started":"2025-01-06T00:53:47.315491Z","shell.execute_reply":"2025-01-06T00:53:47.344607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image,mask = ds.__getitem__(np.random.randint(len(ds)))\nplt.imshow(image.sum(0).cpu() + mask.sum(0).cpu())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.346278Z","iopub.execute_input":"2025-01-06T00:53:47.346591Z","iopub.status.idle":"2025-01-06T00:53:47.578604Z","shell.execute_reply.started":"2025-01-06T00:53:47.346562Z","shell.execute_reply":"2025-01-06T00:53:47.577731Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.579402Z","iopub.execute_input":"2025-01-06T00:53:47.579613Z","iopub.status.idle":"2025-01-06T00:53:47.766410Z","shell.execute_reply.started":"2025-01-06T00:53:47.579595Z","shell.execute_reply":"2025-01-06T00:53:47.765541Z"}},"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-06T00:53:47.767198Z","iopub.execute_input":"2025-01-06T00:53:47.767451Z","iopub.status.idle":"2025-01-06T00:53:47.777644Z","shell.execute_reply.started":"2025-01-06T00:53:47.767430Z","shell.execute_reply":"2025-01-06T00:53:47.776839Z"}},"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-06T00:53:47.778349Z","iopub.execute_input":"2025-01-06T00:53:47.778586Z","iopub.status.idle":"2025-01-06T00:53:47.790423Z","shell.execute_reply.started":"2025-01-06T00:53:47.778566Z","shell.execute_reply":"2025-01-06T00:53:47.789663Z"}},"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-06T00:53:47.791122Z","iopub.execute_input":"2025-01-06T00:53:47.791355Z","iopub.status.idle":"2025-01-06T00:53:47.803318Z","shell.execute_reply.started":"2025-01-06T00:53:47.791336Z","shell.execute_reply":"2025-01-06T00:53:47.802673Z"}},"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#       decoder_channels = (256, 128, 64, 32, 16)[-ENCODER_DEPTH:]\n#       decoder_in_channels = (432, 296 , 64, 96, 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                self.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-06T00:53:47.803982Z","iopub.execute_input":"2025-01-06T00:53:47.804171Z","iopub.status.idle":"2025-01-06T00:53:47.819961Z","shell.execute_reply.started":"2025-01-06T00:53:47.804154Z","shell.execute_reply":"2025-01-06T00:53:47.819282Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# https://math-projects.elte.hu/media/works/187/report/tversky_loss_and_variants.pdf\n# https://scikit-learn.org/1.5/modules/generated/sklearn.metrics.fbeta_score.html\nclass myLoss(nn.Module):\n    def __init__(\n            self,\n            beta=4\n    ):\n        super(myLoss, self).__init__()\n        self.beta = beta\n        \n    def forward(\n            self,\n            pred,\n            label,\n            weight=torch.tensor([\n                1.,\n                2.,\n                1.,\n                2.,\n                1.\n            ]).to(device),\n            smooth=1e-5\n        ):\n        pred = pred.softmax(1)\n        FB = W = torch.tensor([smooth]).to(device)\n        for k in range(len(TARGETS)):\n            m = label == k + 1\n            if m.sum() > 0:\n                bm = m.view(len(m),-1).sum(-1) > 0\n                y_true = m[bm].view(bm.sum(),-1).float()\n                y_pred = pred[bm,k+1,:,:,:].view(bm.sum(),-1)\n#               We'll check the max value other than ith, doesn't needs to evercome all of them to be positive\n                y_pred_max = torch.cat([\n                    pred[bm,:k+1],\n                    pred[bm,k+2:]\n                ],1).max(1)[0].detach().view(bm.sum(),-1)\n                y_pred = y_pred/(y_pred+y_pred_max)\n\n                TP = (y_pred*y_true).sum(-1)\n                FN = ((1-y_pred)*y_true).sum(-1)\n                FP = (y_pred*(1-y_true)).sum(-1)\n                \n                FB = FB + (weight[k]*((1 + self.beta*self.beta)*TP + smooth) / ((1 + self.beta*self.beta)*TP + FP + self.beta*self.beta*FN + smooth)).mean()\n                W = W + weight[k]\n\n        return 1 - FB/W","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.820706Z","iopub.execute_input":"2025-01-06T00:53:47.820914Z","iopub.status.idle":"2025-01-06T00:53:47.836624Z","shell.execute_reply.started":"2025-01-06T00:53:47.820873Z","shell.execute_reply":"2025-01-06T00:53:47.835983Z"}},"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-06T00:53:47.837210Z","iopub.execute_input":"2025-01-06T00:53:47.837425Z","iopub.status.idle":"2025-01-06T00:53:47.848307Z","shell.execute_reply.started":"2025-01-06T00:53:47.837408Z","shell.execute_reply":"2025-01-06T00:53:47.847700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    seed_everything(SEED)\n#   model = myUNet2Dto3D(len(TARGETS)+1)\n    model = torch.load('/kaggle/input/czii-synthetic-data-pretraining/resnet18_3_segmentation_2Dto3D_synthetic_pretraining')\n        \n    train_samples = []\n    valid_samples = []\n    weight = torch.zeros(5)\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_Mosaic_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=2, shuffle=True, drop_last=True)\n    vdl = torch.utils.data.DataLoader(vds, batch_size=4, 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            shuffle_cb,\n#           GradientAccumulation(n_acc=4)\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS)\n    torch.save(model,ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_pre_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T00:53:47.849222Z","iopub.execute_input":"2025-01-06T00:53:47.849438Z","iopub.status.idle":"2025-01-06T01:10:07.366246Z","shell.execute_reply.started":"2025-01-06T00:53:47.849420Z","shell.execute_reply":"2025-01-06T01:10:07.364785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_D = 32\nBS = 32\nexperiment = []\nparticle_type = []\nx = []\ny = []\nz = []\nGT_experiment = []\nGT_particle_type = []\nGT_x = []\nGT_y = []\nGT_z = []\nfor f in FOLDS:\n    model = torch.load(ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_pre_'+str(f))\n    sample = SAMPLES[f-1]\n    print(sample)\n    volume = torch.as_tensor(V[sample]).to(device)\n    MASK = torch.zeros(192,640,640,len(TARGETS)+1)\n    STEPS = 11\n    with torch.no_grad():\n            for k in tqdm(range(STEPS)):\n                mask = model(\n                        volume[\n                            k*PATCH_D//2:k*PATCH_D//2+PATCH_D\n                        ]\n                )[0].argmax(0).cpu()\n                \n                patch = MASK[\n                    k*PATCH_D//2:k*PATCH_D//2+PATCH_D\n                ].reshape(-1,6)\n                    \n                patch[torch.arange(mask.numel()),mask.view(-1)] += 1\n\n                MASK[\n                    k*PATCH_D//2:k*PATCH_D//2+PATCH_D\n                ] = patch.reshape(32,640,640,len(TARGETS)+1)\n\n    MASK = MASK.argmax(-1).numpy()\n    for k in range(len(TARGETS)):\n                    \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        try:\n            preds = torch.tensor(preds).float().to(device)\n            dis = (preds.unsqueeze(1) - L[sample][TARGETS[k]].unsqueeze(0))\n            dis = (dis*dis).sum(-1).sqrt()\n            hits = (dis < .05*radius[TARGETS[k]]).max(0)[0].sum().item()\n            miss = (dis > .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        except:\n            None\n\n    del volume,MASK,model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T01:10:16.065865Z","iopub.execute_input":"2025-01-06T01:10:16.066164Z","iopub.status.idle":"2025-01-06T01:10:32.362046Z","shell.execute_reply.started":"2025-01-06T01:10:16.066143Z","shell.execute_reply":"2025-01-06T01:10:32.360963Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T01:10:50.916780Z","iopub.execute_input":"2025-01-06T01:10:50.917132Z","iopub.status.idle":"2025-01-06T01:10:50.944150Z","shell.execute_reply.started":"2025-01-06T01:10:50.917104Z","shell.execute_reply":"2025-01-06T01:10:50.943334Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T01:10:53.628496Z","iopub.execute_input":"2025-01-06T01:10:53.628820Z","iopub.status.idle":"2025-01-06T01:10:53.641320Z","shell.execute_reply.started":"2025-01-06T01:10:53.628793Z","shell.execute_reply":"2025-01-06T01:10:53.640663Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-06T01:10:57.165319Z","iopub.execute_input":"2025-01-06T01:10:57.165628Z","iopub.status.idle":"2025-01-06T01:10:57.190638Z","shell.execute_reply.started":"2025-01-06T01:10:57.165602Z","shell.execute_reply":"2025-01-06T01:10:57.189970Z"}},"outputs":[],"execution_count":null}]}