{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":117682,"databundleVersionId":15062069}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## PyTorch Custom Data pipeline notebook","metadata":{}},{"cell_type":"markdown","source":"This is the version that resolves the situation with standard PyTorch Dataloaders:  [Kaggle notebook for standard PyTorch data pipeline](https://www.kaggle.com/code/mohanarc/pytorch-dataloaders-demo/edit/run/301349909 )\n\nFeel free to run this notebook in one go. \n\n[GitHub link for notebook](https://github.com/MohanaRC/computefriendly_MLdemo/blob/4c6fdba5a77ffb9056443c306a2f817a3ed70541/pytorch-dataloaders-for-heavy-datasets-demo.ipynb)\n\n\n","metadata":{}},{"cell_type":"code","source":"!pip install imagecodecs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:00:46.406668Z","iopub.execute_input":"2026-03-05T13:00:46.407136Z","iopub.status.idle":"2026-03-05T13:00:51.772473Z","shell.execute_reply.started":"2026-03-05T13:00:46.407113Z","shell.execute_reply":"2026-03-05T13:00:51.771825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn.functional as F\nimport tifffile \nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport warnings\nfrom pathlib import Path\nfrom typing import Tuple, Optional, Dict, List, Callable,Union\nimport pytorch_lightning as pl\nimport random\n\n# Declare the paths\nDATA_DIR = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection\")\nTRAIN_IMAGES_DIR = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/train_images\")\nTRAIN_LABELS_DIR = Path(\"/kaggle/input/competitions/vesuvius-challenge-surface-detection/train_labels\")\nMODEL_INPUT_SIZE = (256, 256, 256)  # (depth, height, width) - standarize volume size to this\n# For training\nBATCH_SIZE = 1\nNUM_WORKERS = 2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:00:51.774390Z","iopub.execute_input":"2026-03-05T13:00:51.774725Z","iopub.status.idle":"2026-03-05T13:01:11.792497Z","shell.execute_reply.started":"2026-03-05T13:00:51.774700Z","shell.execute_reply":"2026-03-05T13:01:11.791685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#### Our Previous resize 3D class to replace pytorch transforms\nclass Resize3D:\n    \"\"\"\n    Resizing class separately written because Pytorch transforms doesn't directly handle 3d volumes :(\n    \"\"\"\n    def __init__(self, size=(256, 256, 256)):\n        self.size = size\n        \n    def __call__(self, volume_tensor):\n        # As per pytorch, the dataset has to be 5D arranged as [Batch, Channel, Depth, Height, Width]\n        # volume_tensor shape expected: [Channel, Depth, Height, Width]\n        volume_tensor = volume_tensor.unsqueeze(0) # Add Batch dim -> [1, C, D, H, W]       \n        resized_volume = F.interpolate(\n            volume_tensor, size=self.size, mode='trilinear', align_corners=False\n        )\n        return resized_volume # Remove Batch dim for stacking later","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:01:11.793504Z","iopub.execute_input":"2026-03-05T13:01:11.794454Z","iopub.status.idle":"2026-03-05T13:01:11.800210Z","shell.execute_reply.started":"2026-03-05T13:01:11.794416Z","shell.execute_reply":"2026-03-05T13:01:11.799519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Surface dataset class, the new addition\nclass SurfaceDataset3D(Dataset):\n    def __init__(self,images_dir: Path,labels_dir: Optional[Path],volume_files: Optional[List[str]] = None,\n        volume_shape: Tuple[int, int, int] = (256, 256, 256),):\n        \"\"\"Initialize the file paths, saves the volume shapes\"\"\"\n        super().__init__()\n        self.images_dir = images_dir\n        self.labels_dir = labels_dir \n        self.volume_files = volume_files\n        self.volume_shape = volume_shape\n        self._prepare_volume_files()\n    \n    def _prepare_volume_files(self):\n        \"\"\"Scan the path and get a list of all the tif files (since our data is tif)\"\"\"\n        if self.volume_files is None:\n             print(f\"No volume files specified. Scanning {self.images_dir}...\")\n             self.volume_files = sorted([p.name for p in self.images_dir.glob(f\"*{\".tif\"}\")])\n        valid_files = []\n        for filename in self.volume_files:\n            valid_files.append(filename)\n        self.volume_files = valid_files\n        print(f\"Found {len(self.volume_files)} volumes.\")\n    \n    def __len__(self) -> int:\n        return len(self.volume_files)  ### Just returning length of volume files\n    \n    def __getitem__(self, idx: int):\n        \"\"\"Similar to standard __getitem__ methods. Loads a single sample of image and labels\n        \"\"\"\n        filename = self.volume_files[idx]\n        image, mask = self._load_from_raw(filename)\n        \n        # Convert to Tensor (Channel dimension added here: [1, D, H, W])\n        image_t = torch.from_numpy(image).float().div_(255.0).unsqueeze(0)\n\n        if mask is not None:\n             mask_t = torch.from_numpy(mask).long().unsqueeze(0)\n        else:\n             mask_t = torch.full_like(image_t, 2, dtype=torch.long) ## The type casting is a topic for another day!\n\n        frag_id = Path(filename).stem\n        return image_t, mask_t, frag_id\n\n    def _load_file(self, path: Path) -> np.ndarray:\n        return tifffile.imread(str(path)) ### read the tiffile from the path\n\n    def _load_from_raw(self, volume_file: str) -> Tuple[np.ndarray, Optional[np.ndarray]]:\n        \"\"\"Method to read the volumes and labels and save it as numpy array\"\"\"\n        image_path = self.images_dir / volume_file\n        image_volume = self._load_file(image_path)\n        label_volume = None\n        return image_volume, label_volume","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:01:11.801029Z","iopub.execute_input":"2026-03-05T13:01:11.801253Z","iopub.status.idle":"2026-03-05T13:01:11.858122Z","shell.execute_reply.started":"2026-03-05T13:01:11.801230Z","shell.execute_reply":"2026-03-05T13:01:11.857471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def custom_collate(batch):\n    return batch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:01:11.859936Z","iopub.execute_input":"2026-03-05T13:01:11.860176Z","iopub.status.idle":"2026-03-05T13:01:11.876115Z","shell.execute_reply.started":"2026-03-05T13:01:11.860153Z","shell.execute_reply":"2026-03-05T13:01:11.875522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### SurfaceDataLoader class\nclass SurfaceDataLoader(pl.LightningDataModule):\n    def __init__(self,train_images_dir: Path,train_labels_dir: Path,volume_shape: Tuple[int, int, int] = (256, 256, 256), \n        val_split: float = 0.1,batch_size: int = BATCH_SIZE, num_workers: int = NUM_WORKERS):\n        \"\"\"Sets up the hyperparameters (batch size, workers etc.) and initializes the Resize3D tool, \n        which will be used later on the GPU.\"\"\"\n        super().__init__()\n        self.train_images_dir = train_images_dir\n        self.train_labels_dir = train_labels_dir\n        self.volume_shape = volume_shape\n        self.val_split = val_split\n        self.batch_size = batch_size\n        self.num_workers = num_workers\n        self.gpu_resizer = Resize3D(size=self.volume_shape) # Apply transform using pytorch resize3D\n\n    def setup(self, stage: Optional[str] = None):\n        #Shuffle and split data into train and val \n        print(f\"\\nSetting up training data...\")\n        all_files = sorted([f.name for f in self.train_images_dir.glob(f\"*{\".tif\"}\")])\n        random.seed(42)\n        random.shuffle(all_files)\n        split_idx = int(len(all_files) * (1 - self.val_split))\n        train_files = all_files[:split_idx]\n        val_files = all_files[split_idx:]\n\n        self.train_dataset = SurfaceDataset3D(\n            images_dir=self.train_images_dir,\n            labels_dir=self.train_labels_dir,\n            volume_files=train_files,\n            volume_shape=self.volume_shape\n        )\n        self.val_dataset = SurfaceDataset3D(\n            images_dir=self.train_images_dir,\n            labels_dir=self.train_labels_dir,\n            volume_files=val_files,\n            volume_shape=self.volume_shape\n        )\n\n    def train_dataloader(self) -> DataLoader:\n        # Wrap train dataset in a standard PyTorch dataloader\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            shuffle=True,\n            num_workers=self.num_workers,\n            pin_memory=True,\n            persistent_workers=bool(self.num_workers > 0),\n            collate_fn=custom_collate\n        )\n\n    def val_dataloader(self) -> DataLoader:\n        # Wrap val dataset in a standard PyTorch dataloader\n        return DataLoader(\n            self.val_dataset,\n            batch_size=self.batch_size,\n            shuffle=False,\n            num_workers=self.num_workers,\n            pin_memory=True,\n            persistent_workers=bool(self.num_workers > 0),\n            collate_fn=custom_collate\n        )\n\n    def on_after_batch_transfer(self, batch, dataloader_idx):\n        ## The main optimization hack\n        if not isinstance(batch, list):\n            return super().on_after_batch_transfer(batch, dataloader_idx)\n\n        x_list, y_list, frag_ids = [], [], []\n        device = self.trainer.strategy.root_device if self.trainer else torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n        for item in batch:\n            x, y, frag_id = item\n            # Move to GPU\n            x = x.to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n\n            # Apply Resize3D\n            x_resized = self.gpu_resizer(x)\n            y_resized = self.gpu_resizer(y.float()).long()\n            \n            x_list.append(x_resized)\n            y_list.append(y_resized)\n            frag_ids.append(frag_id)\n\n        # Final Batch Shape: [B, C, D, H, W]\n        return torch.stack(x_list), torch.stack(y_list), frag_ids","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:01:11.876832Z","iopub.execute_input":"2026-03-05T13:01:11.877065Z","iopub.status.idle":"2026-03-05T13:01:11.889970Z","shell.execute_reply.started":"2026-03-05T13:01:11.877038Z","shell.execute_reply":"2026-03-05T13:01:11.889208Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"datamodule = SurfaceDataLoader(\n    train_images_dir=TRAIN_IMAGES_DIR,\n    train_labels_dir=TRAIN_LABELS_DIR,\n    volume_shape=MODEL_INPUT_SIZE,\n)\ndatamodule.setup()\n\ndatamodule.volume_shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:01:11.890756Z","iopub.execute_input":"2026-03-05T13:01:11.891036Z","iopub.status.idle":"2026-03-05T13:01:12.007579Z","shell.execute_reply.started":"2026-03-05T13:01:11.891003Z","shell.execute_reply":"2026-03-05T13:01:12.007022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get a raw batch from the training loader\ntrain_loader = datamodule.train_dataloader()\nbatch = next(iter(train_loader))\n\n# Extract the first sample (raw_img: 320x320x320)\nraw_img, raw_mask, frag_id = batch[0]\nprint (raw_img.shape)\n\nimport matplotlib.pyplot as plt\nplt.imshow(raw_img.squeeze()[0], cmap=\"gray\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T13:01:12.008325Z","iopub.execute_input":"2026-03-05T13:01:12.008679Z","iopub.status.idle":"2026-03-05T13:01:15.055802Z","shell.execute_reply.started":"2026-03-05T13:01:12.008645Z","shell.execute_reply":"2026-03-05T13:01:15.053924Z"}},"outputs":[],"execution_count":null}]}