{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":11933012,"sourceType":"datasetVersion","datasetId":7502327},{"sourceId":11945583,"sourceType":"datasetVersion","datasetId":7509689},{"sourceId":11945589,"sourceType":"datasetVersion","datasetId":7509694},{"sourceId":12016306,"sourceType":"datasetVersion","datasetId":7559896},{"sourceId":12019924,"sourceType":"datasetVersion","datasetId":7444010},{"sourceId":12024955,"sourceType":"datasetVersion","datasetId":7565569},{"sourceId":12052322,"sourceType":"datasetVersion","datasetId":7553053}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport time\nimport argparse\nimport os\nfrom pathlib import Path\n#!pip install segmentation_models_pytorch==0.3.3\n#!pip install connected-components-3d\nimport segmentation_models_pytorch as smp\nimport cc3d\nimport matplotlib.pyplot as plt\nimport math\nimport subprocess\nimport sklearn.metrics\nimport cv2\nimport gc\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import DataLoader, default_collate\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms.functional as TTF\nimport torchvision\nimport torch.nn.functional as F\nfrom torch.nn import TransformerEncoder, TransformerEncoderLayer\nfrom torch.utils.data import Dataset\nfrom fastai.vision.all import *\nimport timm\nimport yaml\n\nimport multiprocessing as mp\nfrom queue import Empty\n\nVSR = 13.1\n\nINPUT_PATH = Path('/kaggle/input/byu-locating-bacterial-flagellar-motors-2025')\nRESNET34UNET512_PATH = [\n    '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_1',\n    '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_2',\n    '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_3',\n    '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_4',\n    '/kaggle/input/512byumyunet2dresnet34/resnet34_myUNet_5'\n]\nRESNET34UNET640_PATH = [\n    '/kaggle/input/byumyunet2d640/resnet34_myUNet_1',\n    '/kaggle/input/byumyunet2d640/resnet34_myUNet_2',\n    '/kaggle/input/byumyunet2d640/resnet34_myUNet_3',\n    '/kaggle/input/byumyunet2d640/resnet34_myUNet_4',\n    '/kaggle/input/byumyunet2d640/resnet34_myUNet_5'\n]\nRESNET18UNET512_PATH = [\n    '/kaggle/input/byumyunet2d/resnet18_5_myUNet_1',\n    '/kaggle/input/byumyunet2d/resnet18_5_myUNet_2',\n    '/kaggle/input/byumyunet2d/resnet18_5_myUNet_3',\n    '/kaggle/input/byumyunet2d/resnet18_5_myUNet_4',\n    '/kaggle/input/byumyunet2d/resnet18_5_myUNet_5'\n]\nDEPTHWISE_PATH = [\n    '/kaggle/input/byumydepthwisedisc32x64x64/models/resnet18_5_myDepthwiseDisc_1.pth',\n    '/kaggle/input/byumydepthwisedisc32x64x64/models/resnet18_5_myDepthwiseDisc_2.pth',\n    '/kaggle/input/byumydepthwisedisc32x64x64/models/resnet18_5_myDepthwiseDisc_3.pth',\n    '/kaggle/input/byumydepthwisedisc32x64x64/models/resnet18_5_myDepthwiseDisc_4.pth',\n    '/kaggle/input/byumydepthwisedisc32x64x64/models/resnet18_5_myDepthwiseDisc_5.pth'\n]\nViT_PATH =[\n    '/kaggle/input/byumyvitdisc64x128x128/models/resnet18_5_myViTDisc_1.pth',\n    '/kaggle/input/byumyvitdisc64x128x128/models/resnet18_5_myViTDisc_2.pth',\n    '/kaggle/input/byumyvitdisc64x128x128/models/resnet18_5_myViTDisc_3.pth',\n    '/kaggle/input/byumyvitdisc64x128x128/models/resnet18_5_myViTDisc_4.pth',\n    '/kaggle/input/byumyvitdisc64x128x128/models/resnet18_5_myViTDisc_5.pth'\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:12:20.641216Z","iopub.execute_input":"2025-05-15T08:12:20.641521Z","iopub.status.idle":"2025-05-15T08:12:33.46101Z","shell.execute_reply.started":"2025-05-15T08:12:20.641496Z","shell.execute_reply":"2025-05-15T08:12:33.46026Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"def get_tomos(input_path: Path,\n              data_type: str,\n              *, n=None) -> list[Path]:\n    \"\"\"\n    tomo..dachi...?\n    Args\n      input_path (Path): Kaggle input directory\n      data_type (str):   train or test\n      n (Optional[int]): Use only first n train tomos (not applicable to test)\n    \"\"\"\n    data_path = input_path / data_type\n    tomo_paths = sorted(data_path.glob('*'))\n\n    if (n is not None) and (data_type == 'train'):\n        tomo_paths = tomo_paths[:n]\n    \n    return tomo_paths\n    \n\ndef preprocess(\n    img: torch.Tensor\n) -> torch.Tensor:\n    \"\"\"\n    Resize and normalize\n\n    Arg:\n      img (Tensor[uint8]):   (batch_size, C, H, W)\n\n    Returns:\n      img (Tensor[float32]): (batch_size, C, size, size); size = 640\n    \"\"\"\n\n    img = img.to(dtype=torch.float32)\n    B,H,W = img.shape\n\n    if H > W:\n        w = 512\n        nh = np.rint(H*w/(32*W))\n        h = int(nh*32)\n    elif W > H:\n        h = 512\n        nw = np.rint(W*h/(H*32))\n        w = int(nw*32)\n    else:\n        h = w = 512\n    \n    img = F.interpolate(\n        img.view(1,B,H,W),\n        size=(h,w),\n        mode='bilinear'\n    ).view(B,1,h,w)\n\n    q = torch.Tensor([0.05, 0.95]).to(img.device)\n    x_min, x_max = torch.quantile(img.view(-1), q)\n\n    img = (img - x_min) / (x_max - x_min)\n\n    return img,x_max,x_min\n\n\nclass Dataset(torch.utils.data.Dataset):\n    \"\"\"\n    dataset = Dataset(tomo_path)\n    \n    Args:\n      tomo_path (Path): directory including jpg images    \n    \"\"\"\n    def __init__(self, tomo_path: Path):\n        self.filenames = sorted(tomo_path.glob('*'))  # list[Path]\n\n    def __len__(self) -> int:\n        return len(self.filenames)\n\n    def __getitem__(self, i: int) -> dict:\n        filename = self.filenames[i]  # Path\n        filebase = filename.stem\n        assert filebase[:6] == 'slice_'\n        slice_number = int(filebase[6:])  # slice_0000 -> int(0000)\n\n        # Load, resize and normalize image\n        img = Image.open(filename)\n        W, H = img.size\n        img = np.array(img)\n#       img = np.expand_dims(np.array(img), axis=0)  # array[uint8] (1, H, W)\n\n        ret = {\n            'img': img,\n            'slice_number': slice_number,\n            'shape': np.array((H, W), dtype=int),  # original shape H, W\n        }\n\n        return ret\n\n    def loader(self, batch_size: int, num_workers: int):\n        loader = DataLoader(self, batch_size=batch_size, num_workers=num_workers)\n        return loader\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:12:33.467527Z","iopub.execute_input":"2025-05-15T08:12:33.467752Z","iopub.status.idle":"2025-05-15T08:12:33.48028Z","shell.execute_reply.started":"2025-05-15T08:12:33.467732Z","shell.execute_reply":"2025-05-15T08:12:33.479573Z"}},"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\n\ndef 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\n\nclass 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)\n        self.disc = nn.Linear(512,classes)\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\n\nclass myViTDisc(nn.Module):\n    def __init__(\n        self,\n        myUNet,\n        emb_size=256,       # Embedding dimension\n        num_heads=8,        # Number of attention heads\n        depth=6,            # Number of transformer layers\n        dropout=0.1,        # Dropout rate\n        expansion=4,        # FFN expansion factor\n        max_seq_len=64,     # Maximum sequence allowed\n        pos_enc_scale=1.0   # Control position encoding strength (0.01 for weak init)\n    ):\n        super().__init__()\n        \n        # 1. UNet Encoder (processes each slice independently)\n        self.unet_encoder = myUNet.UNet.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. Learnable positional encoding initialized as sinusoidal\n        # First create the classic sinusoidal pattern\n        position = torch.arange(max_seq_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, emb_size, 2) * (-math.log(10000.0) / emb_size))\n        pe = torch.zeros(max_seq_len, emb_size)\n        pe[:, 0::2] = torch.sin(position * div_term)  # even dims\n        pe[:, 1::2] = torch.cos(position * div_term)  # odd dims\n    \n        # Make it learnable but initialize with sinusoids\n        self.positional_encoding = nn.Parameter(pe.unsqueeze(0) * pos_enc_scale) # [1, max_seq_len, emb_size]\n        \n        # 4. Transformer with multiple layers\n        encoder_layers = TransformerEncoderLayer(\n            d_model=emb_size,\n            nhead=num_heads,\n            dim_feedforward=emb_size * expansion,  # FFN hidden dim = emb_size * expansion\n            dropout=dropout,\n            batch_first=True,\n        )\n        self.transformer = TransformerEncoder(encoder_layers, num_layers=depth)\n        \n        # 5. 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        # Handle list input [images, mask]\n        if isinstance(x, list):\n            images, mask = x  # images: [B, 64, 1, 128, 128], mask: [B, 64]\n        else:\n            images, mask = x, None  # If no mask provided\n        \n        B, num_slices, C, H, W = images.shape\n        \n        # Process slices through UNet encoder\n        slices = images.view(B * num_slices, C, H, W)\n        encoded_slices = self.unet_encoder(slices)[-1]\n        \n        # Project to embeddings\n        embeddings = self.projection(encoded_slices).mean(dim=[2, 3])  # [B*64, emb_size]\n        embeddings = embeddings.view(B, num_slices, -1)  # [B, 64, emb_size]\n\n        # Add sinusoidal positional encoding\n        embeddings = embeddings + self.positional_encoding[:, :num_slices]\n        \n        # Transformer with mask (if provided)\n        if mask is not None:\n            # mask shape: [B, 64], dtype=bool (True for padding)\n            transformer_out = self.transformer(\n                embeddings, \n                src_key_padding_mask=mask\n            )\n        else:\n            transformer_out = self.transformer(embeddings)\n        \n        # Masked mean pooling (only average non-masked slices)\n        if mask is not None:\n            # Set masked embeddings to zero\n            transformer_out = transformer_out * (~mask).unsqueeze(-1)  # [B, 64, emb_size]\n            # Sum and divide by number of valid slices\n            pooled = transformer_out.sum(dim=1) / (~mask).sum(dim=1, keepdim=True)  # [B, emb_size]\n        else:\n            pooled = transformer_out.mean(dim=1)\n        \n        logits = self.classifier(pooled)\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:12:33.481521Z","iopub.execute_input":"2025-05-15T08:12:33.481815Z","iopub.status.idle":"2025-05-15T08:12:33.496863Z","shell.execute_reply.started":"2025-05-15T08:12:33.481786Z","shell.execute_reply":"2025-05-15T08:12:33.496198Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"def predict(\n    tomo_path: Path,\n    models_34_512: list[nn.Module],\n    models_34_640: list[nn.Module],\n    models_18_512: list[nn.Module],\n    depthwise: list[nn.Module],\n    ViT: list[nn.Module],\n    cfg: dict\n) -> dict:\n    \"\"\"\n    Predict moter coordinate for one tomo_id\n    At most one moter per tomo_id in test\n    \n    Args:\n      model (nn.Module): Pytorch model\n      dataset (Dataset): for one tomo_id\n    \"\"\"\n    tomo_id = tomo_path.name\n    batch_size = cfg['batch_size']\n    num_workers = cfg['num_workers']\n    use_amp = cfg['use_amp']\n    preprocess_device = cfg['preprocess_device']\n    assert preprocess_device in ['cuda', 'cpu']\n\n    dataset = Dataset(tomo_path)\n    loader = dataset.loader(batch_size=batch_size, num_workers=num_workers)\n    \n    device = next(models_18_512[0].parameters()).device\n#==================================================================================\n#   1st fast 2D scan over resized volume\n#==================================================================================\n    x_max_acc = 0\n    x_min_acc = 0\n    total = 0\n    v = []\n    MASK = []\n    for d in loader:\n        v.append(d['img'])\n        # Input image (batch_size, 1, H, W)\n        if preprocess_device == 'cuda':\n            img = d['img'].to(device) \n            img,x_max,x_min = preprocess(img)\n        elif preprocess_device == 'cpu':\n            img,x_max,x_min = preprocess(d['img'])\n            img = img.to(device)  # input image (batch_size, 1, H, W)\n        else:\n            raise ValueError(preprocess_device)\n\n        N,_,h,w = img.shape\n        x_max_acc += N*x_max\n        x_min_acc += N*x_min\n        total += N\n        \n        img_src = img[::2]\n        img_rot = torch.rot90(img[1::2],1,(-2,-1))\n        img_src[1::2] = torch.rot90(img_src[1::2],2,(-2,-1))\n        img_rot[1::2] = torch.rot90(img_rot[1::2],2,(-2,-1))\n        \n        y_pred_src_sum = None\n        y_pred_rot_sum = None\n        for model in models_34_512:\n            with torch.no_grad():\n                with torch.amp.autocast(\n                    device_type='cuda',\n                    enabled=use_amp,\n                    dtype=torch.float16\n                ):\n                    y_pred_src = model(img_src).softmax(1)[:,1]\n                    y_pred_rot = model(img_rot).softmax(1)[:,1]\n\n            if y_pred_src_sum is None:\n                y_pred_src_sum = y_pred_src\n            else:\n                y_pred_src_sum += y_pred_src\n\n            if y_pred_rot_sum is None:\n                y_pred_rot_sum = y_pred_rot\n            else:\n                y_pred_rot_sum += y_pred_rot\n\n        y_pred_src_sum[1::2] = torch.rot90(y_pred_src_sum[1::2],-2,(-2,-1))\n        y_pred_rot_sum[1::2] = torch.rot90(y_pred_rot_sum[1::2],-2,(-2,-1))\n        y_pred_sum = torch.zeros((N,h,w),device=device).float()\n        y_pred_sum[::2] = y_pred_src_sum\n        y_pred_sum[1::2] = torch.rot90(y_pred_rot_sum,-1,(-2,-1))\n        \n        MASK.append(y_pred_sum)\n\n    x_max = x_max_acc/total\n    x_min = x_min_acc/total\n    \n    MASK = torch.cat(MASK)/len(models_18_512)\n    binary_mask = (MASK > .5*MASK.max()).cpu().numpy()\n\n    cc = cc3d.connected_components(binary_mask)\n    stats = cc3d.statistics(cc)\n    if len(stats['centroids']) > 1:\n        centroids = stats['centroids'][1:]\n        kmax = stats['voxel_counts'][1:].argmax()\n        vcmax = stats['voxel_counts'][kmax+1]\n        m = stats['voxel_counts'][1:] > .5*vcmax\n        centroids = centroids[m].reshape(-1,3)\n        v = torch.cat(v)\n        D,H,W = v.shape\n        x = H*W/(512*512*math.pi)\n        zyx0 = centroids.copy()\n        zyx0[:,1] *= H/h\n        zyx0[:,2] *= W/w\n        dhw0 = np.rint(centroids).astype(int)\n        dhw1 = np.rint(zyx0).astype(int)\n        kmax = np.arange(len(m))[m]\n#==================================================================================\n#       2nd detailed 2D scan over centered square patches\n#==================================================================================\n        mask1 = []\n        vs0 = []\n        y0_pred = []\n        h_starts = []\n        w_starts = []\n        size = min([H,W])\n        size = size - size%32\n        for k in range(len(dhw0)):\n            d,h,w =  dhw0[k]\n            y0_pred.append(MASK[d,h,w].item())\n            glb = cc[\n                max([d-2,0]):min([D,d+3])\n            ] == kmax[k]+1\n            glb_mean = (glb.mean(0) > .5).sum()\n            if glb_mean > 0:\n                glb = glb_mean\n            else:\n                glb = glb.max(0).sum()\n            vs0.append(250/np.sqrt(\n                glb*x\n            ))\n            \n            d,h,w = dhw1[k]\n            start_h = max([h-size//2,0])\n            start_w = max([w-size//2,0])\n            start_h = min([start_h,H-size])\n            start_w = min([start_w,W-size])\n            h_starts.append(start_h)\n            w_starts.append(start_w)\n            img = (v[\n                d,\n                start_h:start_h+size,\n                start_w:start_w+size\n            ].to(device).view(1,1,size,size) - x_min)/(x_max - x_min)\n\n            img = torch.cat([\n                img,\n                torch.rot90(img,1,(-2,-1)),\n                torch.rot90(img,2,(-2,-1)),\n                torch.rot90(img,3,(-2,-1))\n            ])\n        \n            y_pred_sum = None\n            for model in models_34_512:\n                with torch.no_grad():\n                    with torch.amp.autocast(\n                        device_type='cuda',\n                        enabled=use_amp,\n                        dtype=torch.float16\n                    ):\n                        y_pred = model(img).softmax(1)[:,1]\n\n                if y_pred_sum is None:\n                    y_pred_sum = y_pred\n                else:\n                    y_pred_sum += y_pred\n\n            for model in models_18_512:\n                with torch.no_grad():\n                    with torch.amp.autocast(\n                        device_type='cuda',\n                        enabled=use_amp,\n                        dtype=torch.float16\n                    ):\n                        y_pred = model(img).softmax(1)[:,1]\n\n                y_pred_sum += y_pred\n\n            for model in models_34_640:\n                with torch.no_grad():\n                    with torch.amp.autocast(\n                        device_type='cuda',\n                        enabled=use_amp,\n                        dtype=torch.float16\n                    ):\n                        y_pred = model(img).softmax(1)[:,1]\n\n                y_pred_sum += y_pred\n\n            y_pred_sum[1] = torch.rot90(y_pred_sum[1],-1,(-2,-1))\n            y_pred_sum[2] = torch.rot90(y_pred_sum[2],-2,(-2,-1))\n            y_pred_sum[3] = torch.rot90(y_pred_sum[3],-3,(-2,-1))\n\n            mask1.append(y_pred_sum.sum(0))\n\n        mask1 = torch.stack(mask1)/(4*(len(models_34_512)+len(models_34_640)+len(models_18_512)))\n        binary_mask1 = (mask1 > .5*mask1.max()).cpu().numpy()\n\n        vs1 = []\n        zyx1 = []\n        y1_pred = []\n        for k in range(len(dhw0)):\n            dd,hh,ww =  dhw1[k]\n            cc = cc3d.connected_components(binary_mask1[k])\n    \n            value = cc[hh - h_starts[k], ww - w_starts[k]]\n\n            if value > 0:\n                glb = cc == value\n                vs1.append(250/np.sqrt(glb.sum()/(math.pi)))\n                h,w = np.where(glb)\n                h = h.mean()\n                w = w.mean()\n                zyx1.append([\n                    dd,\n                    h + h_starts[k],\n                    w + w_starts[k]\n                ])\n                h = np.rint(h).astype(int)\n                w = np.rint(w).astype(int)\n                y1_pred.append(mask1[k,h,w].item())\n                dhw1[k] = [\n                    dd,\n                    h + h_starts[k],\n                    w + w_starts[k]\n                ]\n            else:\n                vs1.append(np.nan)\n                d,h,w = dd,hh,ww\n                h -= h_starts[k]\n                w -= w_starts[k]\n                zyx1.append(zyx0[k])\n                y1_pred.append(mask1[k,h,w].item())\n                dhw1[k] = [\n                    dd,\n                    h + h_starts[k],\n                    w + w_starts[k]\n                ]\n                        \n        print(pd.DataFrame({\n            'vs0':vs0,\n            'vs1':vs1,\n            'y0_pred':y0_pred,\n            'y1_pred':y1_pred\n        }))\n        vs = np.array([vs0,vs1])\n#==================================================================================\n#       FP discrimination\n#==================================================================================\n        b = []\n        bm = []\n        for k in range(len(zyx1)):\n            d,h,w = dhw1[k]\n        \n            sc = VSR/np.nanmean(vs[:,k])\n            d_size = 64\n            p_size = max([1,np.rint(sc*128).astype(int)])\n        \n            d -= d_size//2\n            h -= p_size//2\n            w -= p_size//2\n\n            d_start = max(d, 0)\n            h_start = max(h, 0)\n            w_start = max(w, 0)\n        \n            d_end = min(d + d_size, D)\n            h_end = min(h + p_size, H)\n            w_end = min(w + p_size, W)\n\n            valid_d = d_end - d_start\n            valid_h = h_end - h_start\n            valid_w = w_end - w_start\n\n            target_d_start = d_start - d\n            target_h_start = h_start - h\n            target_w_start = w_start - w\n        \n            vv = torch.zeros((d_size,p_size,p_size),device=device).float()\n            mask = torch.zeros(d_size,device=device).bool()\n            vv[\n                target_d_start : target_d_start + valid_d,\n                target_h_start : target_h_start + valid_h,\n                target_w_start : target_w_start + valid_w\n            ] = (v[d_start:d_end, h_start:h_end, w_start:w_end].to(device) - x_min)/(x_max - x_min)\n            mask[:target_d_start] = True\n            mask[target_d_start + valid_d:] = True\n            b.append(F.interpolate(\n                vv.view(1,d_size,p_size,p_size),\n                size=(128, 128),\n                mode='bilinear'\n            )[0])\n            bm.append(mask)\n            \n        b = torch.stack(b)\n        b = torch.stack([\n            b,\n            torch.rot90(b,1,(-2,-1)),\n            torch.rot90(b,2,(-2,-1)),\n            torch.rot90(b,3,(-2,-1))\n        ]).view(-1,64,1,128,128)\n\n        bm = torch.stack(4*bm)\n\n        y_pred = 0\n        for model in ViT:\n            with torch.no_grad():\n                y_pred += model([b,bm]).softmax(-1)\n                y_pred += model([torch.rot90(b,2,(-4,-1)),bm.flip(-1)]).softmax(-1)\n                \n        b = b[:,16:-16,0,32:-32,32:-32].view(-1,1,32,64,64)\n        for model in depthwise:\n            with torch.no_grad():\n                y_pred += model(b).softmax(-1)\n                y_pred += model(torch.rot90(b,2,(-3,-1))).softmax(-1)\n\n        y_pred = (y_pred.view(4,-1,2).sum(0)/(8*(len(depthwise) + len(ViT))))[:,1]\n        y_pred = (y_pred/2 + (torch.tensor(y0_pred,device=device) + torch.tensor(y1_pred,device=device))/4)\n        kmax = y_pred.argmax()\n        y_pred = y_pred[kmax].item()\n        zyx = [i.item() for i in zyx1[kmax]]\n        vsf = np.nanmean(vs[:,kmax])\n\n        del b,bm,mask1,binary_mask1\n    \n    else:\n        y_pred = -1\n        zyx = (-1,-1,-1)\n        vsf = -1\n        D = -1\n        H = -1\n        W = -1\n\n    del v,MASK,binary_mask\n\n    return {\n        'tomo_id': tomo_id,\n        'y_pred':y_pred,\n        'zyx': zyx,\n        'vs':vsf,\n        'pmin':x_min.item(),\n        'pmax':x_max.item(),\n        'D':D,\n        'H':H,\n        'W':W\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:12:33.49771Z","iopub.execute_input":"2025-05-15T08:12:33.498006Z","iopub.status.idle":"2025-05-15T08:12:33.513496Z","shell.execute_reply.started":"2025-05-15T08:12:33.497977Z","shell.execute_reply":"2025-05-15T08:12:33.512705Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submit","metadata":{}},{"cell_type":"code","source":"def create_submission(preds: list, th: float, ofilename: str) -> pd.DataFrame:\n    \"\"\"\n    Args:\n      preds (list[dict]): predictions\n      th (float): threshold between no or one moter\n      ofilename (str): submission.csv\n    \"\"\"\n    rows = []\n    count_positive = 0\n    for pred in preds:\n        zyx = pred['zyx']\n        row = {\n            'tomo_id': pred['tomo_id'],\n            'Motor axis 0': zyx[0],\n            'Motor axis 1': zyx[1],\n            'Motor axis 2': zyx[2],\n            'score':pred['y_pred'],\n            'vs':pred['vs'],\n            'pmin':pred['pmin'],\n            'pmax':pred['pmax'],\n            'D':pred['D'],\n            'H':pred['H'],\n            'W':pred['W']\n        }\n        rows.append(row)\n\n    submit = pd.DataFrame(rows)\n    submission = submit.copy()\n    plt.hist(submit['score'],100)[2]\n    plt.show()\n    submission.loc[submit['score'] < submit['score'].mean(),['Motor axis 0','Motor axis 1','Motor axis 2']] = -1\n#   submit.loc[submit['score'] < th,['Motor axis 0','Motor axis 1','Motor axis 2']] = -1\n#   for filename in os.listdir('/kaggle/working/'):\n#       file_path = os.path.join('/kaggle/working/', filename)\n#       if os.path.isfile(file_path):\n#           os.remove(file_path)\n#           print(filename, \"is removed\")\n    submission[['tomo_id','Motor axis 0','Motor axis 1','Motor axis 2']].to_csv(ofilename, float_format='%.8e', index=False)\n\n    print('Submit %s: %d positives / %d tomo_ids' % (ofilename, count_positive, len(rows)))\n\n    return submit","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:12:33.514187Z","iopub.execute_input":"2025-05-15T08:12:33.514377Z","iopub.status.idle":"2025-05-15T08:12:33.523928Z","shell.execute_reply.started":"2025-05-15T08:12:33.51436Z","shell.execute_reply":"2025-05-15T08:12:33.523222Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main","metadata":{}},{"cell_type":"code","source":"def process_fn(process_id: int,\n               tomo_queue,\n               pred_queue,\n               cfg: dict):\n    \"\"\"\n    Prediction process for each GPU\n\n    Args:\n      process_id (int): 0 or 1 for 2 GPUs\n      tomo_queue (): tomo_ids\n      pred_queue (): output predictions\n      cfg (dict): config\n    \"\"\"\n    device = torch.device('cuda:%d' % process_id)\n    \n    models_34_512 = []\n    for p in RESNET34UNET512_PATH:\n        model = torch.load(\n            p,\n            map_location=device,\n            weights_only=False\n        ).eval()\n        models_34_512.append(model)\n\n        if process_id == 0:\n            print('Load model',p)\n\n    models_34_640 = []\n    for p in RESNET34UNET640_PATH:\n        model = torch.load(\n            p,\n            map_location=device,\n            weights_only=False\n        ).eval()\n        models_34_640.append(model)\n\n        if process_id == 0:\n            print('Load model',p)\n\n    models_18_512 = []\n    for p in RESNET18UNET512_PATH:\n        model = torch.load(\n            p,\n            map_location=device,\n            weights_only=False\n        ).eval()\n        models_18_512.append(model)\n\n        if process_id == 0:\n            print('Load model',p)\n\n    depthwise = []\n    for p in DEPTHWISE_PATH:\n        model = myDepthwiseDisc(torch.load(\n            '/kaggle/input/byumyunet2d/resnet18_5_myUNet_1',\n            map_location=device,\n            weights_only=False\n        )).eval().to(device)\n        model.load_state_dict(torch.load(p))\n        depthwise.append(model)\n\n        if process_id == 0:\n            print('Load model',p)\n        \n    ViT = []\n    for p in ViT_PATH:\n        model = myViTDisc(torch.load(\n            '/kaggle/input/byumyunet2d/resnet18_5_myUNet_1',\n            map_location=device,\n            weights_only=False\n        )).eval().to(device)\n        model.load_state_dict(torch.load(p))\n        ViT.append(model)\n\n        if process_id == 0:\n            print('Load model',p)\n        \n    # Loop over tomograms\n    while not tomo_queue.empty():\n        try:\n            tomo_path = tomo_queue.get(timeout=1)\n            pred = predict(\n                tomo_path,\n                models_34_512,\n                models_34_640,\n                models_18_512,\n                depthwise,\n                ViT,\n                cfg\n            )\n            pred_queue.put(pred)\n\n        except Empty:\n            break\n#\n# Main\n#\ntb = time.time()\n\n# Config\ncfg = {\n    'batch_size': 4,\n    'num_workers': 1,\n    'use_amp': True,\n    'preprocess_device': 'cuda'\n}\n#\n# List of tomograms\n#\ntomo_paths = get_tomos(INPUT_PATH, 'test')\n\nprint('Data %d' % len(tomo_paths))\n\nmanager = mp.Manager()\ntomo_queue = manager.Queue()\npred_queue = manager.Queue()\nfor tomo_path in tomo_paths:\n    tomo_queue.put(tomo_path)\n\ntime.sleep(1)\nassert not tomo_queue.empty()\n#\n# Launch process\n#\nnum_processes = 2\ntb = time.time()\n\nworkers = [mp.Process(target=process_fn,\n                      args=(i, tomo_queue, pred_queue, cfg))\n           for i in range(num_processes)]\n\nfor w in workers:\n    w.start()\n\nfor w in workers:\n    w.join()\n\n\ndt = time.time() - tb\nprint('%.2f sec for %d tomos' % (dt, len(tomo_paths)))\n\n# queue to list\npreds = []\ntry:\n    while not pred_queue.empty():\n        preds.append(pred_queue.get(timeout=1))\nexcept Empty:\n    pass\n\nassert len(preds) == len(tomo_paths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:12:33.524557Z","iopub.execute_input":"2025-05-15T08:12:33.524819Z","iopub.status.idle":"2025-05-15T08:13:57.347723Z","shell.execute_reply.started":"2025-05-15T08:12:33.524787Z","shell.execute_reply":"2025-05-15T08:13:57.346308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#\n# Submit\n#\nofilename = 'submission.csv'\nsubmit = create_submission(preds, .5, 'submission.csv')\nprint(ofilename, 'written')\nsubmit.tail()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:13:57.359687Z","iopub.execute_input":"2025-05-15T08:13:57.359954Z","iopub.status.idle":"2025-05-15T08:13:57.656319Z","shell.execute_reply.started":"2025-05-15T08:13:57.359932Z","shell.execute_reply":"2025-05-15T08:13:57.655548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# tomo_00e047\t169\t546\t603     15.6\n# tomo_01a877\t147\t638\t286     13.1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class myViTDisc(nn.Module):\n    def __init__(\n        self,\n        myUNet,\n        emb_size=256,       # Embedding dimension\n        num_heads=8,        # Number of attention heads\n        depth=6,            # Number of transformer layers\n        dropout=0.1,        # Dropout rate\n        expansion=4,        # FFN expansion factor\n        max_seq_len=64,     # Maximum sequence allowed\n        pos_enc_scale=1.0   # Control position encoding strength (0.01 for weak init)\n    ):\n        super().__init__()\n        \n        # 1. UNet Encoder (processes each slice independently)\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. Learnable positional encoding initialized as sinusoidal\n        # First create the classic sinusoidal pattern\n        position = torch.arange(max_seq_len).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, emb_size, 2) * (-math.log(10000.0) / emb_size))\n        pe = torch.zeros(max_seq_len, emb_size)\n        pe[:, 0::2] = torch.sin(position * div_term)  # even dims\n        pe[:, 1::2] = torch.cos(position * div_term)  # odd dims\n    \n        # Make it learnable but initialize with sinusoids\n        self.positional_encoding = nn.Parameter(pe.unsqueeze(0) * pos_enc_scale) # [1, max_seq_len, emb_size]\n        \n        # 4. Transformer with multiple layers\n        encoder_layers = TransformerEncoderLayer(\n            d_model=emb_size,\n            nhead=num_heads,\n            dim_feedforward=emb_size * expansion,  # FFN hidden dim = emb_size * expansion\n            dropout=dropout,\n            batch_first=True,\n        )\n        self.transformer = TransformerEncoder(encoder_layers, num_layers=depth)\n        \n        # 5. 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        # Handle list input [images, mask]\n        if isinstance(x, list):\n            images, mask = x  # images: [B, 64, 1, 128, 128], mask: [B, 64]\n        else:\n            images, mask = x, None  # If no mask provided\n        \n        B, num_slices, C, H, W = images.shape\n        \n        # Process slices through UNet encoder\n        slices = images.view(B * num_slices, C, H, W)\n        encoded_slices = self.unet_encoder(slices)[-1]\n        \n        # Project to embeddings\n        embeddings = self.projection(encoded_slices).mean(dim=[2, 3])  # [B*64, emb_size]\n        embeddings = embeddings.view(B, num_slices, -1)  # [B, 64, emb_size]\n\n        # Add sinusoidal positional encoding\n        embeddings = embeddings + self.positional_encoding[:, :num_slices]\n        \n        # Transformer with mask (if provided)\n        if mask is not None:\n            # mask shape: [B, 64], dtype=bool (True for padding)\n            transformer_out = self.transformer(\n                embeddings, \n                src_key_padding_mask=mask\n            )\n        else:\n            transformer_out = self.transformer(embeddings)\n        \n        # Masked mean pooling (only average non-masked slices)\n        if mask is not None:\n            # Set masked embeddings to zero\n            transformer_out = transformer_out * (~mask).unsqueeze(-1)  # [B, 64, emb_size]\n            # Sum and divide by number of valid slices\n            pooled = transformer_out.sum(dim=1) / (~mask).sum(dim=1, keepdim=True)  # [B, emb_size]\n        else:\n            pooled = transformer_out.mean(dim=1)\n        \n        logits = self.classifier(pooled)\n        return logits\n\nclass myDepthwiseDisc(nn.Module):\n    def __init__(\n        self,\n        myUNet,\n        classes=2\n    ):\n        super(myDepthwiseDisc, self).__init__()\n        self.encoder = myUNet.encoder\n        \n        # Convert initial layers\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)\n        \n        # Process all layers (layer1 to layer4)\n        for layer_name in ['layer1', 'layer2', 'layer3', 'layer4']:\n            layer = getattr(self.encoder, layer_name)\n            for block in layer:  # Iterate through all blocks in each layer\n                # Convert conv layers\n                block.conv1 = convert_conv_2dto3d(block.conv1)\n                block.conv2 = convert_conv_2dto3d(block.conv2)\n                \n                # Convert batch norms\n                block.bn1 = convert_bn_2dto3d(block.bn1)\n                block.bn2 = convert_bn_2dto3d(block.bn2)\n                \n                # Handle downsample if exists (first block only)\n                if hasattr(block, 'downsample') and block.downsample is not None:\n                    block.downsample[0] = convert_conv_2dto3d(block.downsample[0])\n                    block.downsample[1] = convert_bn_2dto3d(block.downsample[1])\n        \n        self.avgpool = nn.AdaptiveAvgPool3d(1)\n        self.disc = nn.Linear(512, classes)\n\n    def forward(self, X):\n        B = X.shape[0]\n        features = self.encoder(X)\n        x = self.avgpool(features[-1]).view(B, -1)\n        x = self.disc(x)\n        return x","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''submit.to_csv('submit.csv',index=False)\n#result = subprocess.run(['python', '/kaggle/input/byu-pseudo-labeling/byu-512-train.py'], capture_output=True, text=True)\nresult = subprocess.run(['python', '/kaggle/input/byu-pseudo-labeling/byu-train.py'], capture_output=True, text=True)\nprint(result.stdout)\nprint(result.stderr)\nos.remove('submit.csv')\nmodel2D = torch.load('myUNet_psdeudolabels')\nmodelViT = torch.load('myViT_psdeudolabels')\nmodelDepthwise = torch.load('myDepthwise_psdeudolabels')'''","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! head -n 5 submission.csv\n! wc -l submission.csv","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T08:47:44.060068Z","iopub.execute_input":"2025-05-15T08:47:44.060338Z","iopub.status.idle":"2025-05-15T08:47:44.362054Z","shell.execute_reply.started":"2025-05-15T08:47:44.060312Z","shell.execute_reply":"2025-05-15T08:47:44.360971Z"}},"outputs":[],"execution_count":null}]}