{"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":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11607257,"sourceType":"datasetVersion","datasetId":7280452},{"sourceId":11607757,"sourceType":"datasetVersion","datasetId":7280742}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q \"/kaggle/input/monai-1-3-0/monai-1.3.0-202310121228-py3-none-any.whl\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-29T04:56:34.026047Z","iopub.execute_input":"2025-04-29T04:56:34.026479Z","iopub.status.idle":"2025-04-29T04:56:38.488085Z","shell.execute_reply.started":"2025-04-29T04:56:34.026447Z","shell.execute_reply":"2025-04-29T04:56:38.486737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport os\nimport sys\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom typing import Literal, Tuple\nfrom monai.data.utils import dense_patch_slices\nfrom monai.inferers.utils import _get_scan_interval\nimport timm\nprint('timm.__version__',timm.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-29T04:56:38.489772Z","iopub.execute_input":"2025-04-29T04:56:38.490123Z","iopub.status.idle":"2025-04-29T04:56:38.497004Z","shell.execute_reply.started":"2025-04-29T04:56:38.490079Z","shell.execute_reply":"2025-04-29T04:56:38.496130Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_slice(slice_data: np.ndarray) -> np.ndarray:\n    \"\"\"Нормализация слайда по 2-му и 98-му перцентилям.\"\"\"\n    p2, p98 = np.percentile(slice_data, [2, 98])\n    clipped_data = np.clip(slice_data, p2, p98)\n    normalized = 255 * (clipped_data - p2) / (p98 - p2)\n    return normalized.astype(np.uint8)\n\nclass FlagellarMotorsFullDataset(Dataset):\n    def __init__(self, \n                 tomo_ids: list[str], \n                 annotations_file: str, \n                 img_dir: str, \n                 transform=None,\n                 mode: Literal[\"train\", \"val\", \"inference\"] = \"train\",\n                 patch_size: tuple = (192, 384, 384),\n                 overlap: tuple = (0, 0, 0)):\n        \n        self.img_dir = img_dir\n        self.transform = transform\n        self.mode = mode\n        \n        self.annotations = pd.read_csv(annotations_file)\n        self.tomo_ids = tomo_ids\n        self.patch_size = patch_size\n        self.overlap = overlap\n\n    def __len__(self):\n        return len(self.tomo_ids)\n    \n    def __getitem__(self, idx):\n        tomo_id = self.tomo_ids[idx]\n        tomo_path = os.path.join(self.img_dir, tomo_id)\n        slice_files = sorted([f for f in os.listdir(tomo_path) if f.endswith(\".jpg\")])\n\n        # Загружаем все слайды томограммы\n        volume_list = []\n        for slice_file in slice_files:\n            img = cv2.imread(os.path.join(tomo_path, slice_file), cv2.IMREAD_GRAYSCALE)\n            if len(volume_list) == 0:\n                orig_h, orig_w = img.shape  # store original resolution\n            img = cv2.resize(img, (384, 384))\n            img = normalize_slice(img)\n            volume_list.append(img)\n\n        # Собираем в один numpy-массив и превращаем в тензор\n        volume_np = np.stack(volume_list, axis=0)  # (D, H, W)\n        volume = torch.from_numpy(volume_np).float() / 255.0\n        volume = volume.unsqueeze(0)  # -> (1, D, H, W)\n\n        # Логика разбиения на патчи (обобщённая для всех режимов)\n        image_size = volume.shape[1:]  # (D, H, W)\n        scan_interval = _get_scan_interval(image_size, self.patch_size, 3, self.overlap)\n        patches = dense_patch_slices(image_size, self.patch_size, scan_interval, return_slice=True)\n        vol_patches = [volume[:, sl[0], sl[1], sl[2]] for sl in patches]\n        vol_patches = torch.stack(vol_patches)  # (N, 1, patch_D, patch_H, patch_W)\n\n        scale_h = orig_h / 384.0\n        scale_w = orig_w / 384.0\n\n        # Оформляем meta\n        # converted_patches = [tuple((sl.start, sl.stop) for sl in s) for s in patches]\n        patch_slices = [\n            (\n                (int(sl_z.start), int(sl_z.stop)),\n                (int(sl_y.start), int(sl_y.stop)),\n                (int(sl_x.start), int(sl_x.stop)),\n            )\n            for (sl_z, sl_y, sl_x) in patches\n        ]\n        meta = {\n            \"tomo_id\": tomo_id,\n            \"orig_shape\": (len(slice_files), orig_h, orig_w),  # original D,H,W\n            \"patch_slices\": patch_slices,                      # list[(z0,z1),(y0,y1),(x0,x1)]\n            \"scale_hw\": (scale_h, scale_w),                    # scale factors H, W\n        }\n\n\n        if self.mode in (\"train\", \"val\"):\n            # Для train/val формируем тепловую карту\n            motor_data = self.annotations[self.annotations['tomo_id'] == tomo_id]\n            if motor_data.empty:\n                heatmap_np = np.zeros_like(volume_np, dtype=np.float32)\n            else:\n                z = int(motor_data['Motor axis 0'].iloc[0])\n                y = int(motor_data['Motor axis 1'].iloc[0])\n                x = int(motor_data['Motor axis 2'].iloc[0])\n                # volume_shape = volume_np.shape  # (D, H, W)\n                volume_shape: Tuple[int, int, int] = (volume_np.shape[0], volume_np.shape[1], volume_np.shape[2])\n                heatmap_np = gaussian_heatmap_RSNA_9place(volume_shape, (z, y, x), distance_size=(45,45,45))\n\n            heatmap = torch.from_numpy(heatmap_np).float()\n\n            # Разбиваем тепловую карту на патчи\n            hm_patches = [heatmap[sl[0], sl[1], sl[2]] for sl in patches]\n            hm_patches = torch.stack(hm_patches)  # (N, patch_D, patch_H, patch_W)\n\n            \n            return vol_patches, hm_patches, meta\n        else:\n            # Режим inference: возвращаем только патчи объёма и метаданные\n            return vol_patches, meta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-29T04:56:38.499129Z","iopub.execute_input":"2025-04-29T04:56:38.499474Z","iopub.status.idle":"2025-04-29T04:56:38.520662Z","shell.execute_reply.started":"2025-04-29T04:56:38.499434Z","shell.execute_reply":"2025-04-29T04:56:38.519823Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def encode_for_pvtv2(e, x, B, depth_scaling=[2,2,2,2,]):\n\t#poor man's attention = avg + max pool\n    def pool_in_depth(x, depth_scaling):\n        bd, c, h, w = x.shape\n        x1 = x.reshape(B, -1, c, h, w).permute(0, 2, 1, 3, 4)\n        x1 = F.avg_pool3d(x1, kernel_size=(depth_scaling, 1, 1), stride=(depth_scaling, 1, 1), padding=0) \\\n                + F.gelu(F.max_pool3d(x1, kernel_size=(depth_scaling, 1, 1), stride=(depth_scaling, 1, 1), padding=0))\n        x = x1.permute(0, 2, 1, 3, 4).reshape(-1, c, h, w)\n        return x, x1\n\n    encode=[]  #x = 1,256,512,512\n    x = e.patch_embed(x) # x seq: 1024, 128, 128, 64\n\n    x = e.stages_0(x)   #4, 64, 128, 128\n    x, x1 = pool_in_depth(x, depth_scaling[0])\n    #encode.append(x1)\n    x = e.stages_1(x)   #4, 128, 64, 64\n    x, x1 = pool_in_depth(x, depth_scaling[1])\n    #encode.append(x1)\n    x = e.stages_2(x)   #4, 320, 32, 32\n    x, x1 = pool_in_depth(x, depth_scaling[2])\n    #encode.append(x1)\n    x = e.stages_3(x)   #4, 512, 16, 16\n    x, x1 = pool_in_depth(x, depth_scaling[3])\n    encode.append(x1)\n\n    return encode\n\n\nclass Net(nn.Module):\n    def __init__(self, pretrained=False, cfg=None):\n        super(Net, self).__init__()\n        self.output_type = ['infer', 'loss', ]\n        self.register_buffer('D', torch.tensor(0))\n\n        self.arch = 'pvt_v2_b1'\n        encoder_dim = {\n            'resnet34d': [64, 64, 128, 256, 512, ],\n            'resnet50d': [64, 256, 512, 1024, 2048, ],\n            'seresnext26d_32x4d': [64, 256, 512, 1024, 2048, ],\n            'convnext_small.fb_in22k': [96, 192, 384, 768],\n            'pvt_v2_b1': [64, 128, 320, 512],\n            'pvt_v2_b2': [64, 128, 320, 512],\n        }.get(self.arch, [1024])\n\n        self.encoder = timm.create_model(\n            model_name=self.arch, pretrained=pretrained, in_chans=3, num_classes=0, global_pool='', features_only=True,\n        )\n        self.mask = nn.Conv3d(encoder_dim[-1],1, kernel_size=1)\n\n    def forward(self, batch):\n        device = self.D.device\n\n        image = batch['image'].to(device)\n        B, D, H, W = image.shape\n        image = image.reshape(B*D, 1,H, W)\n\n        x = (image.half() - 128) / 128\n        x = x.expand(-1, 3, -1, -1)\n\n        encode = encode_for_pvtv2(self.encoder, x, B)\n        last = encode[-1] #this is the feature map !!!!\n        logit = self.mask(last) .squeeze(1)\n\n        \n        print(f'last', last.shape)\n        [print(f'encode_{i}', e.shape) for i,e in enumerate(encode)]\n        print('logit', logit.shape)\n\n        output = {} \n\n        #loss for pretraining 2d-3d encoder\n        if 'loss' in self.output_type:\n            truth = batch['truth'].to(device)\n            output['mask_loss'] = F.binary_cross_entropy_with_logits(logit,truth)\n\n        if 'infer' in self.output_type:\n            output['motor'] = torch.sigmoid(logit)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-29T04:56:38.521628Z","iopub.execute_input":"2025-04-29T04:56:38.521897Z","iopub.status.idle":"2025-04-29T04:56:38.538565Z","shell.execute_reply.started":"2025-04-29T04:56:38.521877Z","shell.execute_reply":"2025-04-29T04:56:38.537670Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Настройки ===\nannotations_file = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv'\nimg_dir = '/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/test/'\noutput_path = 'submission.csv'\nmodel_path = '/kaggle/input/by-hengck/00003808.pth'\nbatch_size = 1\nTHRESHOLD = 0.25   # минимальная вероятность, при которой считаем, что мотор найден\n\n# Загружаем шаблон\nsubmission_df = pd.DataFrame(columns=[\"tomo_id\", \"Motor axis 0\", \"Motor axis 1\", \"Motor axis 2\"])\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n\n# Загружаем модель\nnet = Net(pretrained=False).to(device)\nnet.output_type = ['infer']\nf = torch.load(model_path, map_location=lambda storage, loc: storage, weights_only=False)\nstate_dict = f.get('state_dict', f) \nnet.load_state_dict(state_dict, strict=False)\nnet.eval()\n\ntomo_ids = sorted(os.listdir(img_dir))\n\ntrain_dataset = FlagellarMotorsFullDataset(tomo_ids, annotations_file, img_dir, mode='inference')\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=False, num_workers=4)\n\n\n# batch = {\n#     'image': torch.from_numpy(image).unsqueeze(0).byte(), must be [1, D, H, W]\n# }\t\t\nwith torch.amp.autocast('cuda', dtype=torch.float16):\n    with torch.no_grad():\n        for vol_patches, meta in tqdm(train_loader):\n            \n            # meta is a list (len = batch_size) of dicts when default collate is used\n            # meta = meta[0]      # because batch_size = 1\n            tomo_id      = meta[\"tomo_id\"][0]\n            patch_slices = meta[\"patch_slices\"]\n            original_shape = meta[\"orig_shape\"]\n            scale_h, scale_w = meta[\"scale_hw\"]\n            \n            print('tomo_id', tomo_id)\n            print('original_shape', original_shape)\n            print('patch_slices', patch_slices)\n\n            print('vol_patches', vol_patches.shape)\n            vol_patches = vol_patches.squeeze(0)\n            print('vol_patches', vol_patches.shape)\n\n            best_prob_val = -1.0\n            best_coord_global = (None, None, None)\n\n            for patch_idx, patch in enumerate(vol_patches):\n                # patch_dict = {'image': patch}\n                # print('patch', patch.shape)\n                patch_uint8 = (patch.squeeze(0) * 255).clamp_(0, 255).to(torch.uint8)\n                patch_dict  = {'image': patch_uint8.unsqueeze(0).to(device)}\n                output = net(patch_dict)\n\n                prob = F.interpolate(\n                    output['motor'].unsqueeze(1),\n                    scale_factor=(16, 32, 32), mode='nearest',\n                )[0, 0].float().cpu().numpy()\n\n                # --- найти максимум внутри текущего патча ---\n                local_max_val = prob.max()\n                # учитываем только патчи, где максимум ≥ THRESHOLD\n                if local_max_val >= THRESHOLD and local_max_val > best_prob_val:\n                    # локальные координаты внутри патча\n                    local_z, local_y, local_x = np.unravel_index(prob.argmax(), prob.shape)\n\n                    # координаты патча в исходном объёме\n                    z_slice, y_slice, x_slice = patch_slices[patch_idx]   # (start, stop) for each axis\n                    global_coord = (\n                        int(z_slice[0] + local_z),\n                        int(y_slice[0] + local_y),\n                        int(x_slice[0] + local_x),\n                    )\n\n                    best_prob_val = float(local_max_val)\n                    best_coord_global = global_coord\n\n            if best_prob_val < THRESHOLD:\n                best_coord_global = (None, None, None)\n\n            if best_coord_global[0] is not None:\n                scaled_coord = (\n                    best_coord_global[0],\n                    int(best_coord_global[1] * scale_h),\n                    int(best_coord_global[2] * scale_w),\n                )\n                print(f\"[{tomo_id}] best_prob: {best_prob_val:.4f} coord(resized): {best_coord_global}  ->  coord(orig): {scaled_coord}\")\n            else:\n                scaled_coord = (-1, -1, -1)\n                print(f\"[{tomo_id}] no motor detected (max prob {best_prob_val:.4f} < {THRESHOLD})\")\n            # при необходимости можно добавить в submission_df:\n            submission_df.loc[len(submission_df)] = [\n                tomo_id,\n                scaled_coord[0],\n                scaled_coord[1],\n                scaled_coord[2],\n            ]\n\n    submission_df.to_csv(output_path, index=False)\n    print(f\"Saved submission to {output_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-29T04:56:38.539429Z","iopub.execute_input":"2025-04-29T04:56:38.539750Z","iopub.status.idle":"2025-04-29T04:56:58.208779Z","shell.execute_reply.started":"2025-04-29T04:56:38.539728Z","shell.execute_reply":"2025-04-29T04:56:58.207646Z"}},"outputs":[],"execution_count":null}]}