{"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":[{"sourceType":"competition","sourceId":84969,"databundleVersionId":10033515},{"sourceType":"datasetVersion","sourceId":10290944,"datasetId":6368954,"databundleVersionId":10593093},{"sourceType":"datasetVersion","sourceId":10358906,"datasetId":6415002,"databundleVersionId":10669503}],"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-02T23:22:43.596328Z","iopub.execute_input":"2025-01-02T23:22:43.596700Z","iopub.status.idle":"2025-01-02T23:23:10.323830Z","shell.execute_reply.started":"2025-01-02T23:22:43.596672Z","shell.execute_reply":"2025-01-02T23:23:10.322877Z"},"trusted":true,"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"VOLUMES_PATH = '/kaggle/input/czii-cryo-et-object-identification/train/static/ExperimentRuns/'\nLABELS_PATH = '/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns/'\nSAMPLES = [\n    'TS_73_6',\n    'TS_86_3',\n    'TS_69_2',\n    'TS_6_4',\n    'TS_99_9',\n    'TS_5_4',\n    'TS_6_6'\n]\nSYNTH_SAMPLES = [\n    'TS_0',\n    'TS_1',\n    'TS_10',\n    'TS_11',\n    'TS_12',\n    'TS_13',\n    'TS_14',\n    'TS_15',\n    'TS_16',\n    'TS_17',\n    'TS_18',\n    'TS_19',\n    'TS_2',\n    'TS_20',\n    'TS_21',\n    'TS_22',\n    'TS_23',\n    'TS_24',\n    'TS_25',\n    'TS_26',\n    'TS_3',\n    'TS_4',\n    'TS_5',\n    'TS_6',\n    'TS_7',\n    'TS_8',\n    'TS_9'\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\n#FOLDS = [1]#[1,2,3,4,5,6,7]\nENCODER_NAME = \"resnet18\"#\"efficientnet-b0\",\"resnet18\",\"timm-regnetx_002\",\"densenet121\",timm-resnest14d\nENCODER_DEPTH = 3\n#DROPOUT = .2\nPATCH_D = 64\nPATCH_H = 128\nPATCH_W = 128\nradius = {\n    'apo-ferritin':60,# easy\n    'beta-galactosidase':90,# hard\n    'ribosome':150,# easy\n    'thyroglobulin':130,# hard\n    'virus-like-particle':135# easy\n}\ncore = {\n    k:radius[k]*.05 for k in radius\n}\nshell = {\n    k:radius[k]*.11 for k in radius\n}\nBS = 8\nLR = 5e-4#5e-4\nEPOCHS = 1#10\nTH = .5\nCONTRAST_AUG = .25#.125\nBRIGTHNESS_AUG = .25#.125","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:23:10.325077Z","iopub.execute_input":"2025-01-02T23:23:10.325441Z","iopub.status.idle":"2025-01-02T23:23:10.332346Z","shell.execute_reply.started":"2025-01-02T23:23:10.325405Z","shell.execute_reply":"2025-01-02T23:23:10.331523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"V = {}\nL = {}\nfor sample in SYNTH_SAMPLES:\n        L[sample] = {}\n        file = zarr.open('/kaggle/input/czii-synth-data/'+sample+'.zarr', mode='r')\n        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        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            ),\n            mode='reflect'\n        )[0,4:-4]\n        for target in TARGETS:\n            L[sample][target] = []\n            for p in pd.read_json('/kaggle/input/czii-synth-data/synth_labels/'+sample+'/'+target+'.ndjson', lines=True)['location']:\n                L[sample][target].append([\n                    p['z']-4,\n                    p['y']+h//2,\n                    p['x']+w//2\n                ])\n            L[sample][target] = torch.tensor(L[sample][target]).float().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:23:49.227218Z","iopub.execute_input":"2025-01-02T23:23:49.227543Z","iopub.status.idle":"2025-01-02T23:25:32.665327Z","shell.execute_reply.started":"2025-01-02T23:23:49.227513Z","shell.execute_reply":"2025-01-02T23:25:32.664394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for sample in SAMPLES:\n        L[sample] = {}\n        file = zarr.open(VOLUMES_PATH + sample + '/VoxelSpacing10.000/denoised.zarr', mode='r')\n        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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:25:32.666605Z","iopub.execute_input":"2025-01-02T23:25:32.666966Z","iopub.status.idle":"2025-01-02T23:25:59.155184Z","shell.execute_reply.started":"2025-01-02T23:25:32.666939Z","shell.execute_reply":"2025-01-02T23:25:59.154320Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"z = torch.stack([\n    torch.stack([\n        torch.arange(PATCH_D)\n    ]*PATCH_H)\n]*PATCH_W).permute(2,1,0)\nx = torch.stack([\n    torch.stack([\n        torch.arange(PATCH_W)\n    ]*PATCH_H)\n]*PATCH_D)\ny = torch.stack([\n    torch.stack([\n        torch.arange(PATCH_H)\n    ]*PATCH_W)\n]*PATCH_D).permute(0,2,1)\nzyx = torch.stack([z,y,x]).float().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:25:59.157149Z","iopub.execute_input":"2025-01-02T23:25:59.157455Z","iopub.status.idle":"2025-01-02T23:25:59.233583Z","shell.execute_reply.started":"2025-01-02T23:25:59.157428Z","shell.execute_reply":"2025-01-02T23:25:59.232702Z"}},"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            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-02T23:25:59.235036Z","iopub.execute_input":"2025-01-02T23:25:59.235356Z","iopub.status.idle":"2025-01-02T23:25:59.253417Z","shell.execute_reply.started":"2025-01-02T23:25:59.235328Z","shell.execute_reply":"2025-01-02T23:25:59.252567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tds = CryoET_Dataset(V,L,SYNTH_SAMPLES)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:25:59.254295Z","iopub.execute_input":"2025-01-02T23:25:59.254580Z","iopub.status.idle":"2025-01-02T23:25:59.285770Z","shell.execute_reply.started":"2025-01-02T23:25:59.254543Z","shell.execute_reply":"2025-01-02T23:25:59.284966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(tds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:25:59.286647Z","iopub.execute_input":"2025-01-02T23:25:59.286879Z","iopub.status.idle":"2025-01-02T23:25:59.305496Z","shell.execute_reply.started":"2025-01-02T23:25:59.286859Z","shell.execute_reply":"2025-01-02T23:25:59.304627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image,heatmap = tds.__getitem__(np.random.randint(len(tds)))\nplt.imshow(image.sum(0).cpu() + (heatmap.sum(0).cpu()))\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:25:59.306304Z","iopub.execute_input":"2025-01-02T23:25:59.306560Z","iopub.status.idle":"2025-01-02T23:26:00.002031Z","shell.execute_reply.started":"2025-01-02T23:25:59.306532Z","shell.execute_reply":"2025-01-02T23:26:00.000990Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del tds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:26:18.581549Z","iopub.execute_input":"2025-01-02T23:26:18.581880Z","iopub.status.idle":"2025-01-02T23:26:18.768434Z","shell.execute_reply.started":"2025-01-02T23:26:18.581851Z","shell.execute_reply":"2025-01-02T23:26:18.767616Z"}},"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-02T23:26:22.367422Z","iopub.execute_input":"2025-01-02T23:26:22.367736Z","iopub.status.idle":"2025-01-02T23:26:22.372514Z","shell.execute_reply.started":"2025-01-02T23:26:22.367710Z","shell.execute_reply":"2025-01-02T23:26:22.371353Z"}},"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-02T23:26:25.746607Z","iopub.execute_input":"2025-01-02T23:26:25.746888Z","iopub.status.idle":"2025-01-02T23:26:25.751651Z","shell.execute_reply.started":"2025-01-02T23:26:25.746866Z","shell.execute_reply":"2025-01-02T23:26:25.750545Z"}},"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-02T23:26:28.172669Z","iopub.execute_input":"2025-01-02T23:26:28.173093Z","iopub.status.idle":"2025-01-02T23:26:28.178649Z","shell.execute_reply.started":"2025-01-02T23:26:28.173058Z","shell.execute_reply":"2025-01-02T23:26:28.177731Z"}},"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        \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#               nn.Dropout(DROPOUT),\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:26:30.478072Z","iopub.execute_input":"2025-01-02T23:26:30.478493Z","iopub.status.idle":"2025-01-02T23:26:30.490422Z","shell.execute_reply.started":"2025-01-02T23:26:30.478457Z","shell.execute_reply":"2025-01-02T23:26:30.489486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# https://arxiv.org/pdf/1708.02002\n# https://github.com/AdeelH/pytorch-multi-class-focal-loss/blob/master/README.md\nfrom typing import Optional, Sequence\n\nimport torch\nfrom torch import Tensor\nfrom torch import nn\nfrom torch.nn import functional as F\n\n\nclass FocalLoss(nn.Module):\n    \"\"\" Focal Loss, as described in https://arxiv.org/abs/1708.02002.\n\n    It is essentially an enhancement to cross entropy loss and is\n    useful for classification tasks when there is a large class imbalance.\n    x is expected to contain raw, unnormalized scores for each class.\n    y is expected to contain class labels.\n\n    Shape:\n        - x: (batch_size, C) or (batch_size, C, d1, d2, ..., dK), K > 0.\n        - y: (batch_size,) or (batch_size, d1, d2, ..., dK), K > 0.\n    \"\"\"\n\n    def __init__(self,\n                 alpha: Optional[Tensor] = None,\n                 gamma: float = 0.,\n                 reduction: str = 'mean',\n                 ignore_index: int = -100):\n        \"\"\"Constructor.\n\n        Args:\n            alpha (Tensor, optional): Weights for each class. Defaults to None.\n            gamma (float, optional): A constant, as described in the paper.\n                Defaults to 0.\n            reduction (str, optional): 'mean', 'sum' or 'none'.\n                Defaults to 'mean'.\n            ignore_index (int, optional): class label to ignore.\n                Defaults to -100.\n        \"\"\"\n        if reduction not in ('mean', 'sum', 'none', 'avg_mean'):\n            raise ValueError(\n                'Reduction must be one of: \"mean\", \"sum\", \"none\".')\n\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.ignore_index = ignore_index\n        self.reduction = reduction\n\n        self.nll_loss = nn.NLLLoss(\n            weight=alpha, reduction='none', ignore_index=ignore_index)\n\n    def __repr__(self):\n        arg_keys = ['alpha', 'gamma', 'ignore_index', 'reduction']\n        arg_vals = [self.__dict__[k] for k in arg_keys]\n        arg_strs = [f'{k}={v!r}' for k, v in zip(arg_keys, arg_vals)]\n        arg_str = ', '.join(arg_strs)\n        return f'{type(self).__name__}({arg_str})'\n\n    def forward(self, x: Tensor, y: Tensor) -> Tensor:\n        if x.ndim > 2:\n            # (N, C, d1, d2, ..., dK) --> (N * d1 * ... * dK, C)\n            c = x.shape[1]\n            x = x.permute(0, *range(2, x.ndim), 1).reshape(-1, c)\n            # (N, d1, d2, ..., dK) --> (N * d1 * ... * dK,)\n            y = y.view(-1)\n\n        unignored_mask = y != self.ignore_index\n        y = y[unignored_mask]\n        if len(y) == 0:\n            return torch.tensor(0.)\n        x = x[unignored_mask]\n\n        # compute weighted cross entropy term: -alpha * log(pt)\n        # (alpha is already part of self.nll_loss)\n        log_p = F.log_softmax(x, dim=-1)\n        ce = self.nll_loss(log_p, y)\n\n        # get true class column from each row\n        all_rows = torch.arange(len(x))\n        log_pt = log_p[all_rows, y]\n\n        # compute focal term: (1 - pt)^gamma\n        pt = log_pt.exp()\n        focal_term = (1 - pt)**self.gamma\n\n        # the full loss: -alpha * ((1 - pt)^gamma) * log(pt)\n        loss = focal_term * ce\n\n        if self.reduction == 'mean':\n            loss = loss.mean()\n        elif self.reduction == 'sum':\n            loss = loss.sum()\n#       Small addition, average mean\n        elif self.reduction == 'avg_mean':\n            loss = (loss.sum())/((self.alpha[y]*focal_term).sum())\n\n        return loss\n\n\ndef focal_loss(alpha: Optional[Sequence] = None,\n               gamma: float = 0.,\n               reduction: str = 'mean',\n               ignore_index: int = -100,\n               device='cpu',\n               dtype=torch.float32) -> FocalLoss:\n    \"\"\"Factory function for FocalLoss.\n\n    Args:\n        alpha (Sequence, optional): Weights for each class. Will be converted\n            to a Tensor if not None. Defaults to None.\n        gamma (float, optional): A constant, as described in the paper.\n            Defaults to 0.\n        reduction (str, optional): 'mean', 'sum' or 'none'.\n            Defaults to 'mean'.\n        ignore_index (int, optional): class label to ignore.\n            Defaults to -100.\n        device (str, optional): Device to move alpha to. Defaults to 'cpu'.\n        dtype (torch.dtype, optional): dtype to cast alpha to.\n            Defaults to torch.float32.\n\n    Returns:\n        A FocalLoss object\n    \"\"\"\n    if alpha is not None:\n        if not isinstance(alpha, Tensor):\n            alpha = torch.tensor(alpha)\n        alpha = alpha.to(device=device, dtype=dtype)\n\n    fl = FocalLoss(\n        alpha=alpha,\n        gamma=gamma,\n        reduction=reduction,\n        ignore_index=ignore_index)\n    return fl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:26:34.056573Z","iopub.execute_input":"2025-01-02T23:26:34.056892Z","iopub.status.idle":"2025-01-02T23:26:34.068455Z","shell.execute_reply.started":"2025-01-02T23:26:34.056862Z","shell.execute_reply":"2025-01-02T23:26:34.067590Z"}},"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-02T23:26:38.574940Z","iopub.execute_input":"2025-01-02T23:26:38.575374Z","iopub.status.idle":"2025-01-02T23:26:38.585510Z","shell.execute_reply.started":"2025-01-02T23:26:38.575338Z","shell.execute_reply":"2025-01-02T23:26:38.584577Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FL_DL(nn.Module):\n    \"\"\"https://academic.oup.com/bioinformaticsadvances/article/4/1/vbae169/7907198\n       The combined focal loss and dice loss function improves\n       the segmentation of beta-sheets in medium-resolution\n       cryo-electron-microscopy density maps\n    \"\"\"\n    def __init__(\n            self,\n            alpha: Optional[Sequence] = None,\n            gamma: float = 0.,\n            reduction: str = 'mean',\n            ignore_index: int = -100,\n            device='cpu',\n            dtype=torch.float32,\n            weight=None\n        ):\n        super(FL_DL, self).__init__()\n        if alpha is not None:\n            if not isinstance(alpha, Tensor):\n                alpha = torch.tensor(alpha)\n            alpha = alpha.to(device=device, dtype=dtype)\n        self.fl = FocalLoss(\n            alpha=alpha,\n            gamma=gamma,\n            reduction=reduction,\n            ignore_index=ignore_index\n        )\n\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        self.dl = DiceLoss(weight)\n\n    def forward(self, predict, target, alpha=.5):\n        fl = self.fl(predict, target)\n        dl = self.dl(predict, target)        \n\n        return alpha*fl + (1 - alpha)*dl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:26:42.317191Z","iopub.execute_input":"2025-01-02T23:26:42.317481Z","iopub.status.idle":"2025-01-02T23:26:42.323967Z","shell.execute_reply.started":"2025-01-02T23:26:42.317458Z","shell.execute_reply":"2025-01-02T23:26:42.322861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(SEED)\nmodel = myUNet2Dto3D(\n    len(TARGETS)+1\n)\n\ntds = CryoET_Dataset(V,L,SYNTH_SAMPLES)\nvds = CryoET_Dataset(V,L,SAMPLES,VALID=True)\n    \ntdl = torch.utils.data.DataLoader(tds, batch_size=BS, shuffle=True, drop_last=True)\nvdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n\ndls = DataLoaders(tdl,vdl)\n\nlearn = Learner(\n        dls,\n        model,\n        lr=LR,\n#       loss_func=focal_loss(alpha=weight,gamma=2,reduction='avg_mean',device=device),\n        loss_func=DiceLoss(),\n#       loss_func=FL_DL(alpha=weight,reduction='avg_mean',device=device),\n        cbs=[\n            ShowGraphCallback(),\n#           GradientAccumulation(n_acc=4)\n        ]\n)\nlearn.fit_one_cycle(EPOCHS)\ntorch.save(model,ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_synthetic_pretraining')\ndel tdl,vdl,dls,model,learn\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:26:45.811055Z","iopub.execute_input":"2025-01-02T23:26:45.811388Z","execution_failed":"2025-01-02T23:28:27.017Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:23:10.380873Z","iopub.status.idle":"2025-01-02T23:23:10.381161Z","shell.execute_reply":"2025-01-02T23:23:10.381028Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:23:10.382993Z","iopub.status.idle":"2025-01-02T23:23:10.383274Z","shell.execute_reply":"2025-01-02T23:23:10.383168Z"}},"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 = []\nmodel = torch.load(ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_segmentation_2Dto3D_synthetic_pretraining').eval()\nprint('TP: GREEN\\nFN: RED\\nFP: YELLOW')\nfor sample in SAMPLES:\n    print(sample)\n    volume = torch.as_tensor(V[sample]).to(device)\n    MASK = torch.zeros(192,640,640,len(TARGETS)+1)\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            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    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\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:23:10.383990Z","iopub.status.idle":"2025-01-02T23:23:10.384347Z","shell.execute_reply":"2025-01-02T23:23:10.384182Z"}},"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-02T23:23:10.385007Z","iopub.status.idle":"2025-01-02T23:23:10.385308Z","shell.execute_reply":"2025-01-02T23:23:10.385201Z"}},"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-02T23:23:10.386123Z","iopub.status.idle":"2025-01-02T23:23:10.386486Z","shell.execute_reply":"2025-01-02T23:23:10.386378Z"}},"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-02T23:23:10.387562Z","iopub.status.idle":"2025-01-02T23:23:10.387883Z","shell.execute_reply":"2025-01-02T23:23:10.387778Z"}},"outputs":[],"execution_count":null}]}