{"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":11257843,"sourceType":"datasetVersion","datasetId":7035785}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\nfrom torch.amp import autocast, GradScaler\nfrom tqdm import tqdm\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:26:36.632315Z","iopub.execute_input":"2025-04-04T05:26:36.632612Z","iopub.status.idle":"2025-04-04T05:26:40.856622Z","shell.execute_reply.started":"2025-04-04T05:26:36.632589Z","shell.execute_reply":"2025-04-04T05:26:40.855877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch==2.5.1 torchvision==0.20.1 torchaudio==2.5.1\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:26:40.857672Z","iopub.execute_input":"2025-04-04T05:26:40.858032Z","iopub.status.idle":"2025-04-04T05:26:44.500843Z","shell.execute_reply.started":"2025-04-04T05:26:40.858012Z","shell.execute_reply":"2025-04-04T05:26:44.499929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    import monai\n    print(\"✅ MONAI is already installed.\")\nexcept ImportError:\n    print(\"📦 Installing MONAI from local .whl file...\")\n    !pip install /kaggle/input/monai-1-3/monai-1.3.0-202310121228-py3-none-any.whl\n\nDATA_DIR = \"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025\"\nTRAIN_IMG_DIR = os.path.join(DATA_DIR, \"train\")\ntest_DIR = os.path.join(DATA_DIR,\"test\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:26:44.502656Z","iopub.execute_input":"2025-04-04T05:26:44.502889Z","iopub.status.idle":"2025-04-04T05:27:09.430467Z","shell.execute_reply.started":"2025-04-04T05:26:44.502869Z","shell.execute_reply":"2025-04-04T05:27:09.429614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from collections import OrderedDict\nimport os\nimport numpy as np\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset\nfrom tqdm import tqdm\n\n\nclass PatchBasedTrainDataset(Dataset):\n    def __init__(self, tomo_root_dir, label_df, patch_size=(64, 256, 256), stride=(32, 128, 128), max_cache=4):\n        self.tomo_root_dir = tomo_root_dir\n        self.label_df = label_df\n        self.patch_size = patch_size\n        self.stride = stride\n        self.patch_infos = []  # (tomo_id, z_start, y_start, x_start)\n        self.volume_cache = OrderedDict()  # soft cache\n        self.max_cache = max_cache\n\n        self._build_index()\n\n    def _build_index(self):\n        unique_ids = self.label_df['tomo_id'].unique()\n        print(f\"Indexing patches from {len(unique_ids)} volumes...\")\n\n        for tomo_id in tqdm(unique_ids, desc=\"Indexing\"):\n            tomo_path = os.path.join(self.tomo_root_dir, tomo_id)\n            slices = sorted(os.listdir(tomo_path))\n            if not slices:\n                continue\n\n            sample = cv2.imread(os.path.join(tomo_path, slices[0]), cv2.IMREAD_GRAYSCALE)\n            z = len(slices)\n            y, x = sample.shape\n\n            if z < self.patch_size[0] or y < self.patch_size[1] or x < self.patch_size[2]:\n                continue\n\n            for z0 in range(0, z - self.patch_size[0] + 1, self.stride[0]):\n                for y0 in range(0, y - self.patch_size[1] + 1, self.stride[1]):\n                    for x0 in range(0, x - self.patch_size[2] + 1, self.stride[2]):\n                        self.patch_infos.append((tomo_id, z0, y0, x0))\n\n        print(f\"[DEBUG] Total patches created: {len(self.patch_infos)}\")\n\n    def __len__(self):\n        return len(self.patch_infos)\n\n    def _load_volume(self, tomo_id):\n        tomo_path = os.path.join(self.tomo_root_dir, tomo_id)\n        slices = sorted(os.listdir(tomo_path))\n        volume = np.stack([\n            cv2.imread(os.path.join(tomo_path, sl), cv2.IMREAD_GRAYSCALE)\n            for sl in slices\n        ]).astype(np.float32) / 255.0\n        return volume\n\n    def __getitem__(self, idx):\n        tomo_id, z0, y0, x0 = self.patch_infos[idx]\n\n        # soft cache with LRU strategy\n        if tomo_id not in self.volume_cache:\n            volume = self._load_volume(tomo_id)\n            self.volume_cache[tomo_id] = volume\n            if len(self.volume_cache) > self.max_cache:\n                self.volume_cache.popitem(last=False)\n        else:\n            # move to end to mark as recently used\n            self.volume_cache.move_to_end(tomo_id)\n\n        volume = self.volume_cache[tomo_id]\n\n        patch = volume[z0:z0+self.patch_size[0], y0:y0+self.patch_size[1], x0:x0+self.patch_size[2]]\n        patch = np.expand_dims(patch, axis=0)  # (1, Z, Y, X)\n\n        coords = self.label_df[self.label_df['tomo_id'] == tomo_id][['Motor axis 2', 'Motor axis 1', 'Motor axis 0']].values\n        mask = np.zeros_like(patch[0], dtype=np.uint8)\n        for x, y, z in coords:\n            if z0 <= z < z0+self.patch_size[0] and y0 <= y < y0+self.patch_size[1] and x0 <= x < x0+self.patch_size[2]:\n                mask[int(z - z0), int(y - y0), int(x - x0)] = 1\n        mask = np.expand_dims(mask, axis=0)\n\n        return torch.tensor(patch, dtype=torch.float32), torch.tensor(mask, dtype=torch.float32)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:27:09.431469Z","iopub.execute_input":"2025-04-04T05:27:09.432179Z","iopub.status.idle":"2025-04-04T05:27:09.446984Z","shell.execute_reply.started":"2025-04-04T05:27:09.432152Z","shell.execute_reply":"2025-04-04T05:27:09.445840Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader\nfrom monai.networks.nets import UNet\nfrom monai.losses import DiceLoss\nfrom torch.optim import Adam\nfrom torch.amp import autocast, GradScaler\n\n# 학습에 필요한 파라미터 설정\nPATCH_SIZE = (64, 256, 256)\nSTRIDE = (32, 128, 128)\nBATCH_SIZE = 4\nEPOCHS = 10\nLEARNING_RATE = 1e-4\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nmax_cache_num = 2\n\n# 데이터셋 생성\ndataset = PatchBasedTrainDataset(\n    tomo_root_dir=\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train\",\n    label_df=pd.read_csv(\"/kaggle/input/byu-locating-bacterial-flagellar-motors-2025/train_labels.csv\"),\n    patch_size=PATCH_SIZE,\n    stride=STRIDE,\n    max_cache = max_cache_num\n)\ndataloader = DataLoader(\n    dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=4,     # 💡 시스템 상황 따라 2~8 사이 실험 가능\n    pin_memory=True,    # 💡 GPU로 tensor 전송시 속도 향상\n    persistent_workers=False,\n    prefetch_factor = 2\n)\n# 모델 정의\nmodel = UNet(\n    spatial_dims=3,\n    in_channels=1,\n    out_channels=1,\n    channels=(16, 32, 64, 128, 256),\n    strides=(2, 2, 2, 2),\n    num_res_units=2,\n    norm='instance'\n).to(DEVICE)\n\n# 손실 함수 및 옵티마이저\nloss_fn = DiceLoss(sigmoid=True)\noptimizer = Adam(model.parameters(), lr=LEARNING_RATE)\nscaler = GradScaler()\n\n\n# 학습 루프\nfor epoch in range(EPOCHS):\n    model.train()\n    running_loss = 0.0\n    print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n    \n    for i, (inputs, targets) in tqdm(enumerate(dataloader), total=len(dataloader), desc=\"Training\"):\n        inputs, targets = inputs.to(DEVICE), targets.to(DEVICE)\n        optimizer.zero_grad()\n        \n        with autocast(device_type='cuda'):  # deprecation warning 대응\n            outputs = model(inputs)\n            loss = loss_fn(outputs, targets)\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        running_loss += loss.item()\n\n    avg_loss = running_loss / len(dataloader)\n    print(f\"[Epoch {epoch+1}] Avg Loss: {avg_loss:.4f}\")\n\n\n# 모델 저장\ntorch.save(model.state_dict(), \"unet3d_patch_based.pth\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:30:37.114740Z","iopub.execute_input":"2025-04-04T05:30:37.115102Z","iopub.status.idle":"2025-04-04T05:32:24.083650Z","shell.execute_reply.started":"2025-04-04T05:30:37.115072Z","shell.execute_reply":"2025-04-04T05:32:24.079436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(next(model.parameters()).device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:32:26.529842Z","iopub.execute_input":"2025-04-04T05:32:26.530157Z","iopub.status.idle":"2025-04-04T05:32:26.535236Z","shell.execute_reply.started":"2025-04-04T05:32:26.530133Z","shell.execute_reply":"2025-04-04T05:32:26.534329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(inputs.device, targets.device)\n# 둘 다 cuda:0 이어야 GPU가 작동해\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-04T05:32:33.402631Z","iopub.execute_input":"2025-04-04T05:32:33.403006Z","iopub.status.idle":"2025-04-04T05:32:33.407867Z","shell.execute_reply.started":"2025-04-04T05:32:33.402976Z","shell.execute_reply":"2025-04-04T05:32:33.406935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}