{"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":12019924,"sourceType":"datasetVersion","datasetId":7444010},{"sourceId":12024955,"sourceType":"datasetVersion","datasetId":7565569}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# GRU Discriminator trained on positives and precomputed FP","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\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 = \"resnet34\"\nRADIUS = 250 # Å\nVSR = 13.1 # Å\nPATCH = 128\nDEPTH = 64\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25\nSEED = 1337\nBS = 16\nLR = 2e-5\nEPOCHS = 16\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":"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":"with open('/kaggle/input/byu-dicts/fp.pkl', 'rb') as f:\n    fp = pickle.load(f)","metadata":{},"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":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BYU_Dataset(Dataset):\n#   work hard, code card\n    def __init__(self, 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.t = []\n        self.idx = []\n        self.label = []\n        for i in range(len(samples)):\n            tomo_id = list(samples)[i]\n            if L[tomo_id][0,0] > -1:\n                if self.VALID:\n                    self.t =self.t + torch.arange(len(L[tomo_id])).tolist()\n                    self.idx = self.idx + len(L[tomo_id])*[i]\n                    self.label = self.label + len(L[tomo_id])*[1]\n                else:\n                    self.t =self.t + [0]\n                    self.idx = self.idx + [i]\n                    self.label = self.label + [1]\n            if tomo_id in fp.keys():\n                if self.VALID:\n                    self.t =self.t + torch.arange(len(fp[tomo_id])).tolist()\n                    self.idx = self.idx + len(fp[tomo_id])*[i]\n                    self.label = self.label + len(fp[tomo_id])*[0]\n                else:\n                    self.y = self.t + [0]\n                    self.idx = self.idx + [i]\n                    self.label = self.label + [0]\n\n        self.indices = torch.tensor(np.indices((DEPTH,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        if self.label[idx]:\n            if self.VALID:\n                t = self.t[idx]\n            else:\n                t = np.random.randint(len(L[tomo_id]))\n            dhw = L[tomo_id].clone()\n        else:\n            if self.VALID:\n                t = self.t[idx]\n            else:\n                t = np.random.randint(len(fp[tomo_id]))\n            if L[tomo_id][0,0] > -1:\n                dhw = torch.cat([\n                    fp[tomo_id][t].view(1,3).to(device),\n                    L[tomo_id]\n                ])\n            else:\n                dhw = fp[tomo_id][t].view(1,3).to(device).clone()\n            t = 0\n\n        D,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        d,h,w = np.rint(dhw[t].cpu()).long().tolist()\n        sm = torch.zeros(DEPTH,device=device).bool()\n        img = torch.empty((DEPTH,1,PATCH,PATCH),device=device).float()\n        if self.VALID:\n            sc = VSR/VS[tomo_id]\n            selected_size = np.rint(sc*PATCH).astype(int)\n\n            d -= DEPTH//2\n            h -= selected_size//2\n            w -= selected_size//2\n\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[:,0] -= d\n            dhw[:,1] -= h\n            dhw[:,2] -= w\n            rot90 = False\n            angle = 0\n\n            dhw /= sc\n            th /= sc\n            \n            for i,dd in enumerate(range(d,d+DEPTH)):\n                if dd < 0 or dd > D - 1:\n                    sm[i] = True\n                    img[i,0] = torch.zeros((PATCH,PATCH),device=device).float()\n                else:\n                    img0 = (torch.as_tensor(cv2.imread(self.path_cache[tomo_id][dd], cv2.IMREAD_ANYDEPTH),device=device).float() - pmin)/(pmax - pmin)\n                    img[i,0] = self.augmentations(img0.view(1,H,W),h,w,H,W,selected_size,rot90,angle)\n            \n        else:\n            sc = (.8 + .4*np.random.rand())*VSR/VS[tomo_id]\n            selected_size = np.rint(sc*PATCH).astype(int)\n            d -= DEPTH//2\n            if selected_size - int(2*th) > 0:\n                h -= int(th)\n                w -= int(th)\n                h -= np.random.randint(selected_size - int(2*th))\n                w -= np.random.randint(selected_size - int(2*th))\n            else:\n                h -= selected_size//2\n                w -= selected_size//2\n\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[:,0] -= d\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            dhw /= sc\n            th /= sc\n\n            for i,dd in enumerate(range(d,d+DEPTH)):\n                if dd < 0 or dd > D - 1:\n                    sm[i] = True\n                    img[i,0] = torch.zeros((PATCH,PATCH),device=device).float()\n                else:\n                    img0 = (torch.as_tensor(cv2.imread(self.path_cache[tomo_id][dd], cv2.IMREAD_ANYDEPTH),device=device).float() - pmin)/(pmax - pmin)\n                    img[i,0] = self.augmentations(img0.view(1,H,W),h,w,H,W,selected_size,rot90,angle)\n\n#           GAUSSIAN NOISE\n            img += torch.normal(0,.01,(DEPTH,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        img[sm] = img.flip(0)[sm]\n\n        label = torch.tensor(self.label[idx]).to(device).long()\n\n        return img,label","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = BYU_Dataset(train.tomo_id.unique(),VALID=True)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"v,l = ds.__getitem__(np.random.randint(len(ds)))\nprint(l.item())\nplt.imshow(v[:,0].sum(0).cpu())\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"N = len(ds.label)\npos = sum(ds.label)\nneg = N - pos\nweight = torch.tensor([N/neg,N/pos],device=device)\nweight","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del ds\ngc.collect()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myGRUDisc(nn.Module):\n    def __init__(\n        self,\n        myUNet,\n        emb_size=256,  # Embedding dimension\n        depth=6,       # Number of GRU layers\n        dropout=0.1,   # Dropout rate\n    ):\n        super().__init__()\n        \n        # 1. UNet Encoder\n        self.unet_encoder = myUNet.encoder\n        \n        # 2. Projection to embedding size\n        encoder_out_channels = self.unet_encoder.out_channels[-1]\n        if encoder_out_channels != emb_size:\n            self.projection = nn.Conv2d(encoder_out_channels, emb_size, kernel_size=1)\n        else:\n            self.projection = nn.Identity()\n\n        # 3. GRU (non-bidirectional)\n        self.gru = nn.GRU(\n            input_size=emb_size,\n            hidden_size=emb_size,  # Hidden size = emb_size (for last-step pooling)\n            num_layers=depth,\n            batch_first=True,\n            dropout=dropout if depth > 1 else 0,\n        )\n        \n        # 4. Classification head\n        self.classifier = nn.Sequential(\n            nn.LayerNorm(emb_size),\n            nn.Linear(emb_size, 2)\n        )\n    \n    def forward(self, x):\n        # Input shape: [B, 128, C, H, W] (no mask needed)\n        B, num_slices, C, H, W = x.shape\n        \n        # Process slices through UNet encoder\n        slices = x.view(B * num_slices, C, H, W)\n        encoded_slices = self.unet_encoder(slices)[-1]\n        \n        # Project to embeddings: [B, 128, emb_size]\n        embeddings = self.projection(encoded_slices).mean(dim=[2, 3])\n        embeddings = embeddings.view(B, num_slices, -1)\n        \n        # Split into first half (0:64) and second half (64:128, flipped)\n        first_half = embeddings[:, :num_slices//2]       # [B, 64, emb_size]\n        second_half = embeddings[:,num_slices//2:].flip(dims=[1])  # [B, 64, emb_size] (reversed)\n        \n        # Process first half with GRU (normal order)\n        _, first_last_hidden = self.gru(first_half)  # last_hidden: [depth, B, emb_size]\n        first_features = first_last_hidden[-1]        # [B, emb_size] (last layer's output)\n        \n        # Process second half with GRU (reversed order)\n        _, second_last_hidden = self.gru(second_half)\n        second_features = second_last_hidden[-1]      # [B, emb_size]\n        \n        first_logits = self.classifier(first_features)\n        second_logits = self.classifier(second_features)\n        return torch.stack([first_logits,second_logits])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myLoss(nn.Module):\n    def __init__(self):\n        super(myLoss, self).__init__()\n        self.CE = nn.CrossEntropyLoss(weight=weight)\n\n    def forward(self, y_pred, y_true):\n        return (self.CE(y_pred[0],y_true) + self.CE(y_pred[1],y_true))/2","metadata":{"trusted":true},"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    ntidx, nvidx = list(skf.split(list(range(len(negatives))), negatives['Group']))[f-1]\n    print(nvidx)\n    tds = BYU_Dataset(list(positives.tomo_id[ptidx]) + list(negatives.tomo_id[ntidx]))\n    vds = BYU_Dataset(list(positives.tomo_id[pvidx]) + list(negatives.tomo_id[nvidx]),VALID=True)\n\n    seed_everything(SEED)\n    model = myGRUDisc(torch.load(\n        '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_'+str(f),\n        weights_only=False\n    ))\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=myLoss(),\n        cbs=[\n            ShowGraphCallback(),\n            SaveModelCallback(\n                fname=ENCODER_NAME+'_myGRUDisc_'+str(f)\n            )\n        ]\n    )\n    learn.fit_one_cycle(EPOCHS)\n    del tds,vds,tdl,vdl,dls,model,learn\n    gc.collect()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix,fbeta_score\n\ny_pred = []\ny_true = []\nmodel = myGRUDisc(torch.load(\n    '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_1',\n    weights_only=False\n)).eval().to(device)\nfor f in FOLDS:\n    model.load_state_dict(torch.load('models/'+ENCODER_NAME+'_myGRUDisc_'+str(f)+'.pth'))\n    _, pvidx = list(skf.split(list(range(len(positives))), positives['Group']))[f-1]\n    _, nvidx = list(skf.split(list(range(len(negatives))), negatives['Group']))[f-1]\n    vds = BYU_Dataset(list(positives.tomo_id[pvidx]) + list(negatives.tomo_id[nvidx]))\n    vdl = torch.utils.data.DataLoader(vds, batch_size=BS, shuffle=False)\n    with torch.no_grad():\n        for b in tqdm(vdl):\n            ensemble = 0\n            for rot in [0,1,2,3]:\n                v = torch.rot90(b[0],rot,(-2,-1))\n                ensemble += model(v).softmax(-1)\n            y_pred = y_pred + (ensemble.sum(0)[:,1]/8).tolist()\n            y_true = y_true + b[1].tolist()\n\ny_true = np.array(y_true)\ny_pred = np.array(y_pred)\nprint(((y_pred > .5) == y_true).sum()/len(y_pred))\nprint(fbeta_score(y_true, y_pred > .5, beta=2))\nprint(confusion_matrix(y_true, y_pred > .5))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pd.DataFrame({\n    'y_true':y_true,\n    'y_pred':y_pred\n}).to_csv('GRU.csv',index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}