{"metadata":{"kernelspec":{"display_name":"pytorch_gpu","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":11871779,"sourceType":"datasetVersion","datasetId":7444010},{"sourceId":11933012,"sourceType":"datasetVersion","datasetId":7502327}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Obtaining FP as out-of-threshold predictions","metadata":{},"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\nimport cc3d\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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ENCODER_NAME = \"resnet18\"\nENCODER_DEPTH = 5\nNTH = .75 # Negatives th\nPTH = .5 # Positives th\nVCTH = .5 # Voxel counts th\nNB = 150 # Å\nRADIUS = 250 # Å\nMIN_SIZE = 2500 # Å\nPATCH = 512\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25\nSEED = 1337\nBS = 36\nLR = 1e-4\nEPOCHS = 10\nFOLDS = [1,2,3,4,5]","metadata":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train['count'] = 1\ntrain.groupby('Voxel spacing').count()['count']","metadata":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives = train[train['Motor axis 0'] > -1].reset_index(drop=True)\npositives.tail()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives.groupby('Voxel spacing').count()['count']","metadata":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"positives = positives.groupby('tomo_id')[['Voxel spacing','Number of motors']].max().reset_index()\npositives.tail()","metadata":{},"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":{},"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":{},"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":{},"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":{},"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":{},"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":{},"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":{},"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":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"avgpool = nn.AvgPool3d(2).to(device)\nmodels = [\n    torch.load(\n        '/kaggle/input/byumyunet2d/'+ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_myUNet_'+str(f),\n        weights_only=False\n    ).eval() for f in FOLDS\n]","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bs = 8\nfp = {}\nfor tomo_id in tqdm(negatives.tomo_id):\n        D,H,W = S[tomo_id]\n        pmin,pmax = p[tomo_id]\n        paths = [f'/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/{tomo_id}/slice_{num:04d}.jpg' for num in range(D)]\n        volume = np.empty((D,H,W), dtype=np.float32)\n\n        h,w = H//2,W//2\n        REMAINING_H = h%32\n        REMAINING_W = w%32\n\n        with torch.no_grad():\n            MASK = []\n            for i in range(0,D,bs):\n                for k, path in enumerate(paths[i:i+bs]):\n                    volume[i+k] = cv2.imread(path, cv2.IMREAD_ANYDEPTH)\n\n                v = (torch.as_tensor(volume[\n                    i:i+bs\n                ],device=device).view(1,-1,H,W) - pmin)/(pmax - pmin)\n                v = avgpool(v).view(-1,1,h,w)[\n                    ...,\n                    REMAINING_H//2:h-REMAINING_H+REMAINING_H//2,\n                    REMAINING_W//2:w-REMAINING_W+REMAINING_W//2\n                ]\n                \n                N =  len(v)\n\n                v = torch.cat([\n                    v,\n                    torch.rot90(v,2,(-2,-1))\n                ])\n\n                vv = torch.rot90(v,1,(-2,-1))\n\n                y_pred,yy_pred = 0,0\n                for model in models:\n                    y_pred = model(v).softmax(1)[:,1]\n                    yy_pred = model(vv).softmax(1)[:,1]\n                y_pred = y_pred + torch.rot90(yy_pred,-1,(-2,-1))\n                y_pred = y_pred[:N] + torch.rot90(y_pred[N:],-2,(-2,-1)) \n                \n                MASK.append(y_pred/(4*len(models)))\n\n            MASK = torch.cat(MASK)\n            y_pred_max = MASK.max()\n            cc = cc3d.connected_components((MASK > NTH*y_pred_max).cpu().numpy())\n            stats = cc3d.statistics(cc)\n            dhw = torch.as_tensor(stats['centroids'][1:],device=device).float().view(-1,3)\n            dhw[:,0] *= 2\n            dhw[:,1] = 2*(dhw[:,1] + REMAINING_H//2)\n            dhw[:,2] = 2*(dhw[:,2] + REMAINING_W//2)\n            vcmax = stats['voxel_counts'][1:].max()\n            m = torch.as_tensor(\n                stats['voxel_counts'][1:] > VCTH*vcmax,\n                device=device,\n                dtype=torch.bool\n            )\n            fp[tomo_id] = dhw[m].view(-1,3).cpu()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open('fp.pkl', 'wb') as handle:\n    pickle.dump(fp, handle)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bs = 4\ncatchs = 0\ntotal = 0\nfor f in FOLDS:\n    model = models[f-1]\n    _,pvidx = list(skf.split(list(range(len(positives))), positives['Group']))[f-1]\n    for tomo_id in tqdm(positives.tomo_id[pvidx]):\n        th = 1000/VS[tomo_id]\n        total += len(L[tomo_id])\n        D,H,W = S[tomo_id]\n        pmin,pmax = p[tomo_id]\n        paths = [f'/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train/{tomo_id}/slice_{num:04d}.jpg' for num in range(D)]\n        volume = np.empty((D,H,W), dtype=np.float32)\n\n        REMAINING_H = H%32\n        REMAINING_W = W%32\n\n        with torch.no_grad():\n            MASK = []\n            for i in range(0,D,bs):\n                for k, path in enumerate(paths[i:i+bs]):\n                    volume[i+k] = cv2.imread(path, cv2.IMREAD_ANYDEPTH)\n\n                v = (torch.as_tensor(volume[\n                    i:i+bs,\n                    REMAINING_H//2:H-REMAINING_H+REMAINING_H//2,\n                    REMAINING_W//2:W-REMAINING_W+REMAINING_W//2\n                ],device=device).view(-1,1,H-REMAINING_H,W-REMAINING_W) - pmin)/(pmax - pmin)\n                \n                N =  len(v)//2\n                v = v[:2*N]\n\n                v = torch.cat([\n                    v[::2],\n                    torch.rot90(v[1::2],2,(-2,-1))\n                ])\n\n                vv = torch.rot90(v,1,(-2,-1))\n\n                y_pred = model(v).softmax(1)[:,1]\n                yy_pred = model(vv).softmax(1)[:,1]\n                y_pred = y_pred + torch.rot90(yy_pred,-1,(-2,-1))\n                y_pred = y_pred[:N] + torch.rot90(y_pred[N:],-2,(-2,-1)) \n                \n                MASK.append(y_pred/4)\n\n            MASK = torch.cat(MASK)\n            cc = cc3d.connected_components((MASK > PTH).cpu().numpy())\n            stats = cc3d.statistics(cc)\n            if len(stats['centroids'][1:]) > 0:\n                dhw = torch.as_tensor(stats['centroids'][1:],device=device).float().view(-1,3)\n                dhw[:,0] *= 2\n                dhw[:,1] += REMAINING_H//2\n                dhw[:,2] += REMAINING_W//2\n                vcmax = stats['voxel_counts'][1:].max()\n                m = torch.as_tensor(\n                    stats['voxel_counts'][1:] > VCTH*vcmax,\n                    device=device,\n                    dtype=torch.bool\n                )\n                dhw = dhw[m].view(-1,3)\n                error = L[tomo_id].view(-1,1,3) - dhw.view(1,-1,3)\n                error = (error*error).sum(-1).sqrt()\n                catchs += ((error < th).sum(1) > 0).sum()\n                m = (error < th).sum(0) == 0\n                if m.sum() > 0: fp[tomo_id] = dhw[m].view(-1,3).cpu()\n\n    print(100*catchs/total)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open('fp.pkl', 'wb') as handle:\n    pickle.dump(fp, handle)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(fp)","metadata":{},"outputs":[],"execution_count":null}]}