{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":11871779,"sourceType":"datasetVersion","datasetId":7444010}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# 1st 2D fast candidates finder trained on positives only","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:27:03.701312Z","iopub.execute_input":"2025-05-15T21:27:03.701665Z","iopub.status.idle":"2025-05-15T21:27:03.705722Z","shell.execute_reply.started":"2025-05-15T21:27:03.701643Z","shell.execute_reply":"2025-05-15T21:27:03.704919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport sklearn.metrics\nimport matplotlib.pyplot as plt\nimport cv2\nimport gc\nfrom tqdm import tqdm\n\nimport torchvision\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:27:03.707287Z","iopub.execute_input":"2025-05-15T21:27:03.707571Z","iopub.status.idle":"2025-05-15T21:28:47.865707Z","shell.execute_reply.started":"2025-05-15T21:27:03.707547Z","shell.execute_reply":"2025-05-15T21:28:47.864883Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ENCODER_NAME = \"resnet18\"\nENCODER_DEPTH = 5\nNB = 150 # Å\nRADIUS = 250 # Å\nMIN_SIZE = 2500 # Å\nPATCH = 512\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25\nSEED = 1337\nBS = 32\nLR = 1e-4\nEPOCHS = 16\nFOLDS = [1,2,3,4,5]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:47.866617Z","iopub.execute_input":"2025-05-15T21:28:47.867076Z","iopub.status.idle":"2025-05-15T21:28:47.873274Z","shell.execute_reply.started":"2025-05-15T21:28:47.867052Z","shell.execute_reply":"2025-05-15T21:28:47.87236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv')\ntrain.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:47.874079Z","iopub.execute_input":"2025-05-15T21:28:47.874506Z","iopub.status.idle":"2025-05-15T21:28:47.928407Z","shell.execute_reply.started":"2025-05-15T21:28:47.87448Z","shell.execute_reply":"2025-05-15T21:28:47.927682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train['count'] = 1\ntrain.groupby('Voxel spacing').count()['count']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:47.930585Z","iopub.execute_input":"2025-05-15T21:28:47.930777Z","iopub.status.idle":"2025-05-15T21:28:47.946166Z","shell.execute_reply.started":"2025-05-15T21:28:47.930761Z","shell.execute_reply":"2025-05-15T21:28:47.945367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_unique = train.groupby('tomo_id')[['tomo_id','Voxel spacing']].max().reset_index(drop=True).sort_values('tomo_id').reset_index(drop=True)\ntrain_unique.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:47.947057Z","iopub.execute_input":"2025-05-15T21:28:47.947266Z","iopub.status.idle":"2025-05-15T21:28:48.001475Z","shell.execute_reply.started":"2025-05-15T21:28:47.94725Z","shell.execute_reply":"2025-05-15T21:28:48.000788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives = train[train['Motor axis 0'] > -1].reset_index(drop=True)\npositives.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:48.002103Z","iopub.execute_input":"2025-05-15T21:28:48.002479Z","iopub.status.idle":"2025-05-15T21:28:48.016004Z","shell.execute_reply.started":"2025-05-15T21:28:48.00246Z","shell.execute_reply":"2025-05-15T21:28:48.015322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives.groupby('Voxel spacing').count()['count']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:48.016726Z","iopub.execute_input":"2025-05-15T21:28:48.016934Z","iopub.status.idle":"2025-05-15T21:28:48.033535Z","shell.execute_reply.started":"2025-05-15T21:28:48.016918Z","shell.execute_reply":"2025-05-15T21:28:48.032795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"L = {}\nS = {}\nVS = {}\nfor tomo_id in tqdm(train.tomo_id.unique()):\n    df = train[train.tomo_id == tomo_id]\n    L[tomo_id] = torch.as_tensor(df[['Motor axis 0','Motor axis 1','Motor axis 2']].values).float().to(device)\n    S[tomo_id] = df[['Array shape (axis 0)','Array shape (axis 1)','Array shape (axis 2)']].max(0).values\n    VS[tomo_id] = df['Voxel spacing'].max()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:48.034477Z","iopub.execute_input":"2025-05-15T21:28:48.034717Z","iopub.status.idle":"2025-05-15T21:28:49.08276Z","shell.execute_reply.started":"2025-05-15T21:28:48.034697Z","shell.execute_reply":"2025-05-15T21:28:49.082004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives = positives.groupby('tomo_id')[['Voxel spacing','Number of motors']].max().reset_index()\npositives.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.083509Z","iopub.execute_input":"2025-05-15T21:28:49.083735Z","iopub.status.idle":"2025-05-15T21:28:49.09554Z","shell.execute_reply.started":"2025-05-15T21:28:49.083719Z","shell.execute_reply":"2025-05-15T21:28:49.094749Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.hist(positives['Voxel spacing'], bins=50, edgecolor='black')\nplt.xlabel('Voxel spacing')\nplt.ylabel('Frequency')\nplt.title('Distribution of positives')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.09658Z","iopub.execute_input":"2025-05-15T21:28:49.096865Z","iopub.status.idle":"2025-05-15T21:28:49.374724Z","shell.execute_reply.started":"2025-05-15T21:28:49.09681Z","shell.execute_reply":"2025-05-15T21:28:49.373994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives['Group'] = 'NaN'\npositives.loc[positives['Voxel spacing'] < 14,'Group'] = 'A'\npositives.loc[(positives['Voxel spacing'] > 14) & (positives['Voxel spacing'] < 16),'Group'] = 'B'\npositives.loc[(positives['Voxel spacing'] > 16) & (positives['Voxel spacing'] < 18),'Group'] = 'C'\npositives.loc[positives['Voxel spacing'] > 18,'Group'] = 'D'\npositives.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.375397Z","iopub.execute_input":"2025-05-15T21:28:49.375625Z","iopub.status.idle":"2025-05-15T21:28:49.390025Z","shell.execute_reply.started":"2025-05-15T21:28:49.375608Z","shell.execute_reply":"2025-05-15T21:28:49.389179Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\n\nn_folds = 5\nskf = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=SEED)\nfor fold, (train_idx, test_idx) in enumerate(skf.split(list(range(len(positives))), positives['Group'])):\n    print(f\"Fold {fold + 1}:\")\n    print(f\"  Train Indices: {train_idx[:10]}...\")  # Print first 10 for brevity\n    print(f\"  Test Indices: {test_idx[:10]}...\")\n    print(f\"  Train Labels Distribution:\\n{pd.Series([positives['Group'][i] for i in train_idx]).value_counts()}\")\n    print(f\"  Test Labels Distribution:\\n{pd.Series([positives['Group'][i] for i in test_idx]).value_counts()}\")\n    print()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.390901Z","iopub.execute_input":"2025-05-15T21:28:49.391632Z","iopub.status.idle":"2025-05-15T21:28:49.430402Z","shell.execute_reply.started":"2025-05-15T21:28:49.39161Z","shell.execute_reply":"2025-05-15T21:28:49.429538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"negatives = train[train['Motor axis 0'] < 0].groupby('tomo_id')[['Voxel spacing','Number of motors']].max().reset_index()\nnegatives.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.433188Z","iopub.execute_input":"2025-05-15T21:28:49.433389Z","iopub.status.idle":"2025-05-15T21:28:49.444553Z","shell.execute_reply.started":"2025-05-15T21:28:49.433373Z","shell.execute_reply":"2025-05-15T21:28:49.443659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.hist(negatives['Voxel spacing'], bins=50, edgecolor='black')\nplt.xlabel('Voxel spacing')\nplt.ylabel('Frequency')\nplt.title('Distribution of negatives')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.445601Z","iopub.execute_input":"2025-05-15T21:28:49.445909Z","iopub.status.idle":"2025-05-15T21:28:49.630872Z","shell.execute_reply.started":"2025-05-15T21:28:49.445886Z","shell.execute_reply":"2025-05-15T21:28:49.630053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"negatives['Group'] = 'NaN'\nnegatives.loc[negatives['Voxel spacing'] < 10,'Group'] = 'A'\nnegatives.loc[(negatives['Voxel spacing'] > 10) & (negatives['Voxel spacing'] < 14),'Group'] = 'B'\nnegatives.loc[(negatives['Voxel spacing'] > 14) & (negatives['Voxel spacing'] < 16),'Group'] = 'C'\nnegatives.loc[(negatives['Voxel spacing'] > 16) & (negatives['Voxel spacing'] < 16.5),'Group'] = 'D'\nnegatives.loc[(negatives['Voxel spacing'] > 16.5) & (negatives['Voxel spacing'] < 18),'Group'] = 'E'\nnegatives.loc[negatives['Voxel spacing'] > 18,'Group'] = 'F'\nnegatives.groupby('Group').count()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.631637Z","iopub.execute_input":"2025-05-15T21:28:49.631853Z","iopub.status.idle":"2025-05-15T21:28:49.648182Z","shell.execute_reply.started":"2025-05-15T21:28:49.631837Z","shell.execute_reply":"2025-05-15T21:28:49.647336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for fold, (train_idx, test_idx) in enumerate(skf.split(list(range(len(negatives))), negatives['Group'])):\n    print(f\"Fold {fold + 1}:\")\n    print(f\"  Train Indices: {train_idx[:10]}...\")  # Print first 10 for brevity\n    print(f\"  Test Indices: {test_idx[:10]}...\")\n    print(f\"  Train Labels Distribution:\\n{pd.Series([negatives['Group'][i] for i in train_idx]).value_counts()}\")\n    print(f\"  Test Labels Distribution:\\n{pd.Series([negatives['Group'][i] for i in test_idx]).value_counts()}\")\n    print()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.649503Z","iopub.execute_input":"2025-05-15T21:28:49.649786Z","iopub.status.idle":"2025-05-15T21:28:49.682916Z","shell.execute_reply.started":"2025-05-15T21:28:49.649768Z","shell.execute_reply":"2025-05-15T21:28:49.682179Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open('/kaggle/input/byu-dicts/p.pkl', 'rb') as f:\n    p = pickle.load(f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.683791Z","iopub.execute_input":"2025-05-15T21:28:49.684061Z","iopub.status.idle":"2025-05-15T21:28:49.700346Z","shell.execute_reply.started":"2025-05-15T21:28:49.684042Z","shell.execute_reply":"2025-05-15T21:28:49.699486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rot_centers(yx,angle,o):\n    yx -= o\n    theta_rad = math.radians(angle)\n    c = math.cos(theta_rad)\n    s =  math.sin(theta_rad)\n    rot = torch.tensor([\n        [c,s],\n        [-s,c]\n    ]).to(device)\n    return yx@rot + o","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.701123Z","iopub.execute_input":"2025-05-15T21:28:49.701354Z","iopub.status.idle":"2025-05-15T21:28:49.705846Z","shell.execute_reply.started":"2025-05-15T21:28:49.701335Z","shell.execute_reply":"2025-05-15T21:28:49.704924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BYU_Dataset(Dataset):\n#   work hard, code card\n    def __init__(self, L, S, VS, samples, VALID=False):\n        self.sample = list(samples)\n        self.VALID = VALID\n        \n        self.path_cache = {\n            tomo_id: [f'/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/{tomo_id}/slice_{num:04d}.jpg' \n                     for num in range(S[tomo_id][0])]\n            for tomo_id in self.sample\n        }\n\n        self.d = []\n        self.t = []\n        self.idx = []\n        for i in range(len(samples)):\n            tomo_id = list(samples)[i]\n            D = S[tomo_id][0]\n            for j in range(len(L[tomo_id])):\n                d = np.rint(L[tomo_id][j,0].cpu()).long().item()\n                if self.VALID:\n                    self.d = self.d + [d]\n                    self.t = self.t + [j]\n                    self.idx = self.idx + [i]\n                else:\n                    N = int(NB/VS[tomo_id])\n                    start_d = max([0,d-N])\n                    end_d = min([D,d+N+1])\n                    self.d = self.d + torch.arange(start_d,end_d).tolist()\n                    self.t = self.t + (end_d - start_d)*[j]\n                    self.idx = self.idx + (end_d - start_d)*[i]\n\n        self.indices = torch.tensor(np.indices((PATCH,PATCH))).to(device)\n\n    def __len__(self):\n        return len(self.idx)\n    \n    def augmentations(self,img,h,w,H,W,selected_size,rot90,angle):\n        if rot90:\n#           Rot90\n            img = torch.rot90(img[:,h:h+selected_size,w:w+selected_size], k=angle, dims=(-2, -1))\n                    \n        else:\n#           Free rotation\n            diag = (H**2 + W**2)**0.5\n            pad_h = int((diag - H)/2)\n            pad_w = int((diag - W)/2)\n            img = F.pad(\n                img,\n                (pad_w, pad_w, pad_h, pad_h),\n                mode='reflect'\n            )\n            img = torchvision.transforms.functional.rotate(\n                img,\n                angle,\n                interpolation=torchvision.transforms.InterpolationMode.BILINEAR,\n                center = (pad_w + w + selected_size//2,pad_h + h + selected_size//2)\n            )[:,pad_h+h:pad_h+h+selected_size,pad_w+w:pad_w+w+selected_size]\n\n#       Resize\n        img = F.interpolate(\n            img.unsqueeze(0),\n            size=(PATCH, PATCH),\n            mode='bilinear',\n            align_corners=False\n        )\n        \n        return img[0]\n\n    def __getitem__(self, idx):\n        \n        tomo_id = self.sample[self.idx[idx]]\n        d = self.d[idx]\n        t = self.t[idx]\n        dhw = L[tomo_id].clone()\n        _,H,W = S[tomo_id]\n        th = torch.as_tensor(RADIUS/VS[tomo_id]).float().to(device)\n        pmin,pmax = p[tomo_id]\n\n        _,hh,ww = dhw[t].tolist()\n        _,h,w = np.rint(dhw[t].cpu()).long().tolist()\n        dhw[:,0] -= d\n        if self.VALID:\n            h -= PATCH//2\n            w -= PATCH//2\n            if h < 0: h = 0\n            if w < 0: w = 0\n            if h > H - PATCH: h = H - PATCH\n            if w > W - PATCH: w = W - PATCH\n\n            dhw[:,1] -= h\n            dhw[:,2] -= w\n\n            img = cv2.imread(self.path_cache[tomo_id][d], cv2.IMREAD_ANYDEPTH)[h:h+PATCH,w:w+PATCH]\n            img = torch.as_tensor((img - pmin)/(pmax - pmin),device=device).float().view(1,PATCH,PATCH)\n            \n        else:\n            min_size = int(MIN_SIZE/VS[tomo_id])\n            max_size = min([H,W])\n            selected_size = min_size + np.random.randint(max_size - min_size)\n            h -= int(th) + np.random.randint(selected_size - int(2*th))\n            w -= int(th) + np.random.randint(selected_size - int(2*th))\n            if h < 0: h = 0\n            if w < 0: w = 0\n            if h > H - selected_size: h = H - selected_size\n            if w > W - selected_size: w = W - selected_size\n            dhw[:,1] -= h\n            dhw[:,2] -= w\n            dh = hh - (h + selected_size//2)\n            dw = ww - (w + selected_size//2)\n            rot90 = np.sqrt(dh*dh + dw*dw) > selected_size//2 - th\n\n            if rot90:\n                angle = np.random.randint(4)\n                dhw[:,1:] = rot_centers(dhw[:,1:],angle*90,torch.tensor([[selected_size,selected_size]]).to(device)/2)\n            else:\n                angle = torch.as_tensor(random.uniform(-180, 180)).item()\n                dhw[:,1:] = rot_centers(dhw[:,1:],angle,torch.tensor([[selected_size,selected_size]]).to(device)/2)\n            \n            zoom = PATCH/selected_size\n            dhw[:,1:] *= zoom\n            th *= zoom\n\n            img = (torch.as_tensor(cv2.imread(self.path_cache[tomo_id][d], cv2.IMREAD_ANYDEPTH),device=device).float() - pmin)/(pmax - pmin)\n            img = self.augmentations(img.view(1,H,W),h,w,H,W,selected_size,rot90,angle)\n\n#           GAUSSIAN NOISE\n            img += torch.normal(0,.01,(1,PATCH,PATCH),device=device)\n#           CONTRAST\n            img *= np.random.normal(1,CONTRAST_AUG)\n#           BRIGTHNESS\n            img += np.random.normal(0,BRIGTHNESS_AUG)\n#           YX FLIP\n            if np.random.rand() < .5:\n                axis = np.random.randint(2) - 2\n                img = img.flip(axis)\n                if axis == -1:\n                    dhw[:,2] = PATCH - 1 - dhw[:,2]\n                else:\n                    dhw[:,1] = PATCH - 1 - dhw[:,1]\n\n        c = dhw[t,1:].view(2,1,1)\n        r = c - self.indices\n        r = (r*r).sum(0).sqrt()\n        label = r < th\n        \n        xt = torch.arange(len(dhw)) != t\n        if xt.sum() > 0:\n            cxt = dhw[xt,1:].view(-1,2)\n            dxt = c.view(1,2) - cxt\n            cxt = cxt[(dxt*dxt).sum(1).sqrt() > 1.8*th].view(-1,2,1,1)\n            r = cxt - self.indices.unsqueeze(0)\n            r = (r*r).sum(1).sqrt()\n            m = (r < th).sum(0) > 0\n            img[0,m] = 0\n\n        return img,label.long()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.706681Z","iopub.execute_input":"2025-05-15T21:28:49.707085Z","iopub.status.idle":"2025-05-15T21:28:49.728757Z","shell.execute_reply.started":"2025-05-15T21:28:49.707061Z","shell.execute_reply":"2025-05-15T21:28:49.728017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = BYU_Dataset(L, S, VS, positives.tomo_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.729681Z","iopub.execute_input":"2025-05-15T21:28:49.729928Z","iopub.status.idle":"2025-05-15T21:28:49.877173Z","shell.execute_reply.started":"2025-05-15T21:28:49.729894Z","shell.execute_reply":"2025-05-15T21:28:49.876435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"v,l = ds.__getitem__(np.random.randint(len(ds)))\nimg = (v[0] + l)\nplt.imshow(img.cpu())\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:49.877807Z","iopub.execute_input":"2025-05-15T21:28:49.878001Z","iopub.status.idle":"2025-05-15T21:28:50.494886Z","shell.execute_reply.started":"2025-05-15T21:28:49.877986Z","shell.execute_reply":"2025-05-15T21:28:50.494101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:50.495801Z","iopub.execute_input":"2025-05-15T21:28:50.496106Z","iopub.status.idle":"2025-05-15T21:28:50.722425Z","shell.execute_reply.started":"2025-05-15T21:28:50.496077Z","shell.execute_reply":"2025-05-15T21:28:50.721584Z"}},"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-05-15T21:28:50.723292Z","iopub.execute_input":"2025-05-15T21:28:50.723552Z","iopub.status.idle":"2025-05-15T21:28:50.73446Z","shell.execute_reply.started":"2025-05-15T21:28:50.72353Z","shell.execute_reply":"2025-05-15T21:28:50.733735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myUNet(nn.Module):\n    def __init__(\n        self,\n        classes=2\n        ):\n        super(myUNet, self).__init__()\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    def forward(self,X):\n#       EDIT from\n#       https://github.com/qubvel-org/segmentation_models.pytorch/blob/main/segmentation_models_pytorch/decoders/unet/decoder.py\n        features = self.UNet.encoder(X)\n        \n        features = features[1:]  # remove first skip with same spatial resolution\n        features = features[::-1]  # reverse channels to start from head of encoder\n\n        head = features[0]\n        skip_connections = features[1:]\n\n        x = self.UNet.decoder.center(head)\n\n        for i, decoder_block in enumerate(self.UNet.decoder.blocks):\n            # upsample to the next spatial shape\n            skip_connection = skip_connections[i] if i < len(skip_connections) else None\n            x = decoder_block(x, skip_connection)\n\n        x = self.UNet.segmentation_head(x)\n        \n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:50.734913Z","iopub.execute_input":"2025-05-15T21:28:50.735147Z","iopub.status.idle":"2025-05-15T21:28:50.747576Z","shell.execute_reply.started":"2025-05-15T21:28:50.73513Z","shell.execute_reply":"2025-05-15T21:28:50.74696Z"}},"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\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-05-15T21:28:50.748363Z","iopub.execute_input":"2025-05-15T21:28:50.748587Z","iopub.status.idle":"2025-05-15T21:28:50.760407Z","shell.execute_reply.started":"2025-05-15T21:28:50.748568Z","shell.execute_reply":"2025-05-15T21:28:50.759603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for f in FOLDS:\n    ptidx, pvidx = list(skf.split(list(range(len(positives))), positives['Group']))[f-1]\n    print(pvidx)\n    tds = BYU_Dataset(L,S,VS,positives.tomo_id[ptidx])\n    vds = BYU_Dataset(L,S,VS,positives.tomo_id[pvidx],VALID=True)\n\n    seed_everything(SEED)\n    model = torch.nn.DataParallel(myUNet(), device_ids=[0,1])\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.module,ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_myUNet_'+str(f))\n    del tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:28:50.761116Z","iopub.execute_input":"2025-05-15T21:28:50.761378Z","iopub.status.idle":"2025-05-15T21:36:29.789433Z","shell.execute_reply.started":"2025-05-15T21:28:50.761353Z","shell.execute_reply":"2025-05-15T21:36:29.788227Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = torch.load(ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_myUNet_'+str(1),weights_only=False).eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:42:15.434186Z","iopub.execute_input":"2025-05-15T21:42:15.43444Z","iopub.status.idle":"2025-05-15T21:42:15.497239Z","shell.execute_reply.started":"2025-05-15T21:42:15.43442Z","shell.execute_reply":"2025-05-15T21:42:15.49648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ptidx, pvidx = list(skf.split(list(range(len(positives))), positives['Group']))[0]\nprint(pvidx)\nvds = BYU_Dataset(L,S,VS,positives.tomo_id[pvidx],VALID=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:42:27.105691Z","iopub.execute_input":"2025-05-15T21:42:27.10636Z","iopub.status.idle":"2025-05-15T21:42:27.139524Z","shell.execute_reply.started":"2025-05-15T21:42:27.106328Z","shell.execute_reply":"2025-05-15T21:42:27.138632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"total = 0\ncatchs = 0\nfor i in np.random.permutation(len(vds))[:50]:\n    total += 1\n    v,l = vds.__getitem__(i)\n    out = 0\n    v = torch.stack([\n        v,\n        torch.rot90(v,1,(-2,-1)),\n        torch.rot90(v,2,(-2,-1)),\n        torch.rot90(v,3,(-2,-1))\n    ])\n    with torch.no_grad():\n        out = model(v).softmax(1)[:,1]\n        out = out[0] + torch.rot90(out[1],-1,(-2,-1)) + torch.rot90(out[2],-2,(-2,-1)) + torch.rot90(out[3],-3,(-2,-1))\n        out = (out/4 > .5).float().cpu()\n\n    if out[l.bool().cpu()].sum() > 0: catchs += 1\n    fig, axes = plt.subplots(1, 2, figsize=(10,10))\n    axes[0].imshow(v[0,0].cpu() + l.cpu())\n    axes[1].imshow(v[0,0].cpu() + out)\n    plt.show()\n\nprint(100*catchs/total)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:42:30.791159Z","iopub.execute_input":"2025-05-15T21:42:30.791434Z","iopub.status.idle":"2025-05-15T21:42:56.920588Z","shell.execute_reply.started":"2025-05-15T21:42:30.791414Z","shell.execute_reply":"2025-05-15T21:42:56.919851Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del model,vds\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T21:44:28.82212Z","iopub.execute_input":"2025-05-15T21:44:28.822792Z","iopub.status.idle":"2025-05-15T21:44:29.173458Z","shell.execute_reply.started":"2025-05-15T21:44:28.822761Z","shell.execute_reply":"2025-05-15T21:44:29.172735Z"}},"outputs":[],"execution_count":null}]}