{"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},{"sourceId":11334039,"sourceType":"datasetVersion","datasetId":7089855},{"sourceId":11334586,"sourceType":"datasetVersion","datasetId":7090188},{"sourceId":11334967,"sourceType":"datasetVersion","datasetId":7090457},{"sourceId":11335172,"sourceType":"datasetVersion","datasetId":7090608},{"sourceId":11335176,"sourceType":"datasetVersion","datasetId":7090612},{"sourceId":11335178,"sourceType":"datasetVersion","datasetId":7090613},{"sourceId":11336205,"sourceType":"datasetVersion","datasetId":7091379},{"sourceId":11336283,"sourceType":"datasetVersion","datasetId":7091443},{"sourceId":11336287,"sourceType":"datasetVersion","datasetId":7091446},{"sourceId":11336292,"sourceType":"datasetVersion","datasetId":7091451},{"sourceId":11337379,"sourceType":"datasetVersion","datasetId":7092331},{"sourceId":11337393,"sourceType":"datasetVersion","datasetId":7092341},{"sourceId":11337394,"sourceType":"datasetVersion","datasetId":7092342},{"sourceId":11337395,"sourceType":"datasetVersion","datasetId":7092343},{"sourceId":11341432,"sourceType":"datasetVersion","datasetId":7095622},{"sourceId":11341451,"sourceType":"datasetVersion","datasetId":7095636},{"sourceId":11341472,"sourceType":"datasetVersion","datasetId":7095649},{"sourceId":11341492,"sourceType":"datasetVersion","datasetId":7095666},{"sourceId":11341503,"sourceType":"datasetVersion","datasetId":7095675},{"sourceId":11341517,"sourceType":"datasetVersion","datasetId":7095687},{"sourceId":11341518,"sourceType":"datasetVersion","datasetId":7095688},{"sourceId":11341532,"sourceType":"datasetVersion","datasetId":7095699},{"sourceId":11341549,"sourceType":"datasetVersion","datasetId":7095713},{"sourceId":11341554,"sourceType":"datasetVersion","datasetId":7095716},{"sourceId":11341585,"sourceType":"datasetVersion","datasetId":7095745},{"sourceId":11341600,"sourceType":"datasetVersion","datasetId":7095756},{"sourceId":11341603,"sourceType":"datasetVersion","datasetId":7095758},{"sourceId":11341611,"sourceType":"datasetVersion","datasetId":7095764},{"sourceId":11341614,"sourceType":"datasetVersion","datasetId":7095766},{"sourceId":11341615,"sourceType":"datasetVersion","datasetId":7095767},{"sourceId":11341618,"sourceType":"datasetVersion","datasetId":7095770},{"sourceId":11525984,"sourceType":"datasetVersion","datasetId":7228703},{"sourceId":11526000,"sourceType":"datasetVersion","datasetId":7228717}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-24T07:35:41.602670Z","iopub.execute_input":"2025-04-24T07:35:41.602922Z","iopub.status.idle":"2025-04-24T07:35:42.612332Z","shell.execute_reply.started":"2025-04-24T07:35:41.602902Z","shell.execute_reply":"2025-04-24T07:35:42.611599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom torch.utils.data import Dataset\n\nclass LazyLoadingDataset(Dataset):\n    def __init__(self, patch_dirs):\n        self.file_paths = []\n        for d in patch_dirs:\n            self.file_paths += [\n                os.path.join(d, f) for f in os.listdir(d) if f.endswith(\".npz\")\n            ]\n\n    def __len__(self):\n        return len(self.file_paths)\n\n    def __getitem__(self, idx):\n        patch = np.load(self.file_paths[idx])  # 또는 torch.load\n        return patch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T07:35:46.899786Z","iopub.execute_input":"2025-04-24T07:35:46.900098Z","iopub.status.idle":"2025-04-24T07:35:50.314431Z","shell.execute_reply.started":"2025-04-24T07:35:46.900071Z","shell.execute_reply":"2025-04-24T07:35:50.313562Z"}},"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-24T07:35:56.773775Z","iopub.execute_input":"2025-04-24T07:35:56.774213Z","iopub.status.idle":"2025-04-24T07:36:00.908336Z","shell.execute_reply.started":"2025-04-24T07:35:56.774191Z","shell.execute_reply":"2025-04-24T07:36:00.907210Z"}},"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T07:36:02.902438Z","iopub.execute_input":"2025-04-24T07:36:02.902807Z","iopub.status.idle":"2025-04-24T07:36:07.180873Z","shell.execute_reply.started":"2025-04-24T07:36:02.902778Z","shell.execute_reply":"2025-04-24T07:36:07.179797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom torch.utils.data import Dataset\n\nPATCH_D, PATCH_H, PATCH_W = 64, 256, 256\n\nclass OnTheFlyPatchDataset(Dataset):\n    def __init__(self, npz_files):\n        self.vol_paths = npz_files\n        self.records = []             # (file_idx, z, y, x)\n        self._vol_cache = (None, None)  # (file_idx, (volume, mask))\n\n        for fi, path in enumerate(npz_files):\n            with np.load(path) as d:\n                infos = d[\"patch_infos\"].astype(int)\n            self.records.extend([(fi, z, y, x) for z, y, x in infos])\n\n    def _get_volume(self, fi):\n        if self._vol_cache[0] != fi:\n            with np.load(self.vol_paths[fi]) as d:\n                vol = d[\"volume\"].astype(np.float32)\n                msk = d[\"mask\"].astype(np.float32)\n            self._vol_cache = (fi, (vol, msk))\n        return self._vol_cache[1]\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, idx):\n        fi, z, y, x = self.records[idx]\n        vol, msk = self._get_volume(fi)\n\n        patch  = vol[z:z+PATCH_D, y:y+PATCH_H, x:x+PATCH_W]\n        target = msk[z:z+PATCH_D, y:y+PATCH_H, x:x+PATCH_W]\n\n        patch_tensor  = torch.from_numpy(patch).unsqueeze(0).contiguous()   # (1,D,H,W)\n        target_tensor = torch.from_numpy(target).unsqueeze(0).contiguous()\n        return patch_tensor, target_tensor\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T07:36:10.874516Z","iopub.execute_input":"2025-04-24T07:36:10.874971Z","iopub.status.idle":"2025-04-24T07:36:10.887279Z","shell.execute_reply.started":"2025-04-24T07:36:10.874925Z","shell.execute_reply":"2025-04-24T07:36:10.886417Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LazyMultiNPZDataset(Dataset):\n    def __init__(self, npz_paths):\n        self.npz_paths = npz_paths\n        self.index_map = []  # (file_idx, local_patch_idx)\n\n        self.vol_infos = []\n        for fi, path in enumerate(npz_paths):\n            with np.load(path) as d:\n                infos = d['patch_infos'].astype(int)\n            self.vol_infos.append((path, len(infos)))\n            self.index_map.extend([(fi, i) for i in range(len(infos))])\n\n        self._vol_cache = (None, None)  # (fi, (volume, mask, infos))\n\n    def _load_volume(self, fi):\n        if self._vol_cache[0] != fi:\n            path, _ = self.vol_infos[fi]\n            d = np.load(path)\n            vol = d['volume'].astype(np.float32)\n            msk = d['mask'].astype(np.float32)\n            infos = d['patch_infos'].astype(int)\n            self._vol_cache = (fi, (vol, msk, infos))\n        return self._vol_cache[1]\n\n    def __len__(self):\n        return len(self.index_map)\n\n    def __getitem__(self, idx):\n        fi, patch_idx = self.index_map[idx]\n        vol, msk, infos = self._load_volume(fi)\n        z, y, x = infos[patch_idx]\n        patch = vol[z:z+64, y:y+256, x:x+256]\n        target = msk[z:z+64, y:y+256, x:x+256]\n        return torch.from_numpy(patch).unsqueeze(0), torch.from_numpy(target).unsqueeze(0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T07:36:14.707629Z","iopub.execute_input":"2025-04-24T07:36:14.707961Z","iopub.status.idle":"2025-04-24T07:36:14.716199Z","shell.execute_reply.started":"2025-04-24T07:36:14.707937Z","shell.execute_reply":"2025-04-24T07:36:14.715449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. Dataset 정의\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader\n# 1. Dataset 정의\npatch_dirs = [\n    \"/kaggle/input/byu-flagellar-part-01\",\n    *[f\"/kaggle/input/byu-preprocessed-part-{i:02d}\" for i in range(2, 34)]\n]\n\nnpz_paths = [\n    os.path.join(d, f)\n    for d in patch_dirs\n    for f in os.listdir(d)\n    if f.endswith(\".npz\")\n]\n\ndataset = LazyMultiNPZDataset(npz_paths)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-24T07:36:17.339515Z","iopub.execute_input":"2025-04-24T07:36:17.339850Z","iopub.status.idle":"2025-04-24T07:36:28.518113Z","shell.execute_reply.started":"2025-04-24T07:36:17.339825Z","shell.execute_reply":"2025-04-24T07:36:28.517432Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Subset\nfrom monai.networks.nets import UNet\nfrom monai.losses import DiceLoss\nfrom torch.optim import Adam\nfrom torch.amp import autocast, GradScaler\nfrom tqdm import tqdm\n\n# 학습 파라미터\nPATCH_SIZE = (64, 256, 256)\nSTRIDE = (32, 128, 128)\nBATCH_SIZE = 1            # ✅ 안정 우선\nNUM_WORKERS = 1           # ✅ worker 없음 (RAM 안전)\nCHUNK_SIZE = 10000\nEPOCHS = 10\nLEARNING_RATE = 1e-4\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\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\nloss_fn = DiceLoss(sigmoid=True)\noptimizer = Adam(model.parameters(), lr=LEARNING_RATE)\nscaler = GradScaler()\n\n# 전체 데이터 크기\nTOTAL_SIZE = len(dataset)\n\n# Chunk 단위로 학습\nfor chunk_start in range(0, TOTAL_SIZE, CHUNK_SIZE):\n    chunk_end = min(chunk_start + CHUNK_SIZE, TOTAL_SIZE)\n    subset = Subset(dataset, list(range(chunk_start, chunk_end)))\n\n    dataloader = DataLoader(\n        subset,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=NUM_WORKERS,\n        pin_memory=True,\n        persistent_workers=True,\n        prefetch_factor = 2\n    )\n\n    print(f\"\\n🔥 Training patch chunk [{chunk_start} ~ {chunk_end}]\")\n\n    for epoch in range(EPOCHS):\n        model.train()\n        running_loss = 0.0\n        print(f\"\\n🌀 Epoch {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'):\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\"[Chunk {chunk_start}-{chunk_end}] Epoch {epoch+1} Avg Loss: {avg_loss:.4f}\")\n\n    # 청크별 저장\n    filename = f\"unet3d_chunk_{chunk_start}_{chunk_end}.pth\"\n    torch.save(model.state_dict(), filename)\n    print(f\"✅ 모델 저장 완료: {filename}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patch_dirs = [\n    \"/kaggle/input/byu-flagellar-part-01\",\n    *[f\"/kaggle/input/byu-preprocessed-part-{i:02d}\" for i in range(2, 34)]\n]\n\nnpz_paths = [\n    os.path.join(d, f)\n    for d in patch_dirs\n    for f in os.listdir(d)\n    if f.endswith(\".npz\")\n]\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}}]}