{"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":11933012,"sourceType":"datasetVersion","datasetId":7502327},{"sourceId":11935894,"sourceType":"datasetVersion","datasetId":7444010}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Depthwise 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 = \"resnet18\"\nENCODER_DEPTH = 5\nRADIUS = 250 # Å\nVSR = 13.1 # Å\nPATCH = 64\nDEPTH = 32\nCONTRAST_AUG = .25\nBRIGTHNESS_AUG = .25\nSEED = 1337\nBS = 32\nLR = 1e-4\nEPOCHS = 20\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.ones(DEPTH,device=device).bool()\n        img = torch.empty((1,DEPTH,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[0,i] = 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[0,i] = self.augmentations(img0.view(1,H,W),h,w,H,W,selected_size,rot90,angle)\n\n            img[0,sm] = img.flip(1)[0,sm]\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[0,i] = 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[0,i] = self.augmentations(img0.view(1,H,W),h,w,H,W,selected_size,rot90,angle)\n\n            img[0,sm] = img.flip(1)[0,sm]\n#           GAUSSIAN NOISE\n            img += torch.normal(0,.01,(1,DEPTH,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#           Z FLIP\n            if np.random.rand() < .5:\n                img = img.flip(-3)\n                dhw[:,0] = DEPTH - 1 - dhw[:,0]\n#           Rot180\n            if np.random.rand() < .5:\n                angle = 2*np.random.randint(2)\n                axis = np.random.randint(2) - 2\n                img = torch.rot90(img, k=angle, dims=(-3, axis))\n                dhw[:,[0,axis]] = rot_centers(dhw[:,[0,axis]],angle*90,torch.tensor([[DEPTH,PATCH]]).to(device)/2)\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 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":"def convert_conv_2dto3d(conv2d):\n    \"\"\"Convert 2D conv to 3D conv operating on last two dimensions\"\"\"\n    new_conv = nn.Conv3d(\n        conv2d.in_channels,\n        conv2d.out_channels,\n        kernel_size=(1, *conv2d.kernel_size),\n        stride=(1, *conv2d.stride),\n        padding=(0, *conv2d.padding),\n        bias=conv2d.bias is not None\n    )\n\n    depthwise = nn.Conv3d(\n        conv2d.out_channels,\n        conv2d.out_channels,\n        groups=conv2d.out_channels,\n        kernel_size=(conv2d.kernel_size[0], 1, 1),\n        stride=(conv2d.stride[0], 1, 1),\n        padding=(conv2d.padding[0], 0, 0),\n    )\n    \n    with torch.no_grad():\n        # Handle weights\n        new_conv.weight.data = conv2d.weight.data.unsqueeze(2)\n\n        depthwise.weight.data = nn.Parameter(torch.ones_like(depthwise.weight.data)/conv2d.kernel_size[0])\n        \n        # Handle bias if exists\n        if conv2d.bias is not None:\n            new_conv.bias.data = conv2d.bias.data\n            \n    return nn.Sequential(\n        new_conv,\n        depthwise\n    )\n\ndef convert_bn_2dto3d(bn2d):\n    \"\"\"Convert BatchNorm2d to BatchNorm3d with preserved parameters\"\"\"\n    bn3d = nn.BatchNorm3d(\n        bn2d.num_features,\n        eps=bn2d.eps,\n        momentum=bn2d.momentum,\n        affine=bn2d.affine,\n        track_running_stats=bn2d.track_running_stats\n    )\n    \n    with torch.no_grad():\n        if bn2d.affine:\n            bn3d.weight.data = bn2d.weight.data.clone()\n            bn3d.bias.data = bn2d.bias.data.clone()\n        \n        if bn2d.track_running_stats:\n            bn3d.running_mean.data = bn2d.running_mean.data.clone()\n            bn3d.running_var.data = bn2d.running_var.data.clone()\n            bn3d.num_batches_tracked.data = bn2d.num_batches_tracked.data.clone()\n    \n    return bn3d","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myDepthwiseDisc(nn.Module):\n    def __init__(\n        self,\n        myUNet,\n        classes=2\n        ):\n        super(myDepthwiseDisc, self).__init__()\n        self.encoder = myUNet.UNet.encoder\n        self.encoder.conv1 = convert_conv_2dto3d(self.encoder.conv1)\n        self.encoder.bn1 = convert_bn_2dto3d(self.encoder.bn1)\n        self.encoder.maxpool = nn.MaxPool3d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)\n        for n,m in self.encoder.named_children():\n            if 'layer' in n:\n                m[0].conv1 = convert_conv_2dto3d(m[0].conv1)\n                m[0].conv2 = convert_conv_2dto3d(m[0].conv2)\n                m[1].conv1 = convert_conv_2dto3d(m[1].conv1)\n                m[1].conv2 = convert_conv_2dto3d(m[1].conv2)\n\n                m[0].bn1 = convert_bn_2dto3d(m[0].bn1)\n                m[0].bn2 = convert_bn_2dto3d(m[0].bn2)\n                m[1].bn1 = convert_bn_2dto3d(m[1].bn1)\n                m[1].bn2 = convert_bn_2dto3d(m[1].bn2)\n\n                if int(n[5]) > 1:\n                    m[0].downsample[0] = convert_conv_2dto3d(m[0].downsample[0])\n                    m[0].downsample[1] = convert_bn_2dto3d(m[0].downsample[1])\n\n        self.avgpool = nn.AdaptiveAvgPool3d(1).to(device)\n        self.disc = nn.Linear(512,classes).to(device)\n\n    def forward(self,X):\n        B = X.shape[0]\n        x = self.encoder(X)\n        x = self.avgpool(x[-1]).view(B,-1)\n        x = self.disc(x)\n        \n        return x","metadata":{},"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 = myDepthwiseDisc(torch.load(\n        '/kaggle/input/byumyunet2d/'+ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_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=nn.CrossEntropyLoss(weight=weight),\n        cbs=[\n            ShowGraphCallback(),\n            SaveModelCallback(\n                fname=ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_myDepthwiseDisc_'+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 = myDepthwiseDisc(myUNet()).eval().to(device)\nfor f in FOLDS:\n    model.load_state_dict(torch.load('models/'+ENCODER_NAME+'_'+str(ENCODER_DEPTH)+'_myDepthwiseDisc_'+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                ensemble += model(torch.rot90(b[0],rot,(-2,-1))).softmax(-1)\n                ensemble += model(torch.rot90(torch.rot90(b[0],2,(-3,-1)),rot,(-2,-1))).softmax(-1)\n            y_pred = y_pred + (ensemble[:,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('Depthwise.csv',index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}