{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# This notebook serves as an example to show how we can load the train/validation/test data directly in a pytorch dataloader. \n\n# Warning: I am not sure if this is an efficient way to load the data.\n","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport matplotlib.pyplot as plt\n\nGLOBAL_PATH = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'","metadata":{"execution":{"iopub.status.busy":"2023-05-11T12:03:35.826643Z","iopub.execute_input":"2023-05-11T12:03:35.827062Z","iopub.status.idle":"2023-05-11T12:03:39.941985Z","shell.execute_reply.started":"2023-05-11T12:03:35.827025Z","shell.execute_reply":"2023-05-11T12:03:39.940636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailDataset(Dataset):\n    def __init__(self, base_dir, data_type='train', transform=None):\n        assert data_type in ['train', 'validation', 'test'], \\\n            \"'data_type' should be one of 'train', 'validation', or 'test'\"\n\n        self.base_dir = base_dir\n        self.data_type = data_type\n        self.transform = transform\n        self.record = os.listdir(self.base_dir +'/'+ self.data_type)\n\n    def __len__(self):\n        return len(self.record)\n\n    def __getitem__(self, idx):\n        record_id = self.record[idx]\n        record_dir = os.path.join(self.base_dir, self.data_type, record_id)\n\n        # Load the necessary .npy files\n        bands_data = []\n        for i in range(8, 17):\n            band_file = os.path.join(record_dir, f'band_{str(i).zfill(2)}.npy')\n            band_data = np.load(band_file)\n            bands_data.append(band_data)\n\n        # Stack band data along the channel axis\n        bands_data = np.stack(bands_data, axis=-1)\n\n        # If the data type is 'train' or 'validation', load the masks\n        if self.data_type in ['train', 'validation']:\n            pixel_masks_file = os.path.join(record_dir, 'human_pixel_masks.npy')\n            pixel_masks = np.load(pixel_masks_file)\n        else:\n            pixel_masks = None  # No masks for 'test' data\n\n        sample = {'bands': bands_data, 'masks': pixel_masks}\n\n        if self.transform:\n            sample = self.transform(sample)\n\n        return sample\n\n\ndef collate_fn(batch):\n    \"\"\"Collate function to use with DataLoader for batching\"\"\"\n    # Here we assume that each element in batch is a dict {'bands': bands, 'masks': masks}\n    bands = torch.stack([torch.from_numpy(item['bands']) for item in batch])\n    \n    # Check if masks are present\n    if batch[0]['masks'] is not None:\n        masks = torch.stack([torch.from_numpy(item['masks']) for item in batch])\n    else:\n        masks = None\n\n    return bands, masks\n\ndef get_dataloader(base_dir, data_type, batch_size, transform=None):\n    dataset = ContrailDataset(base_dir, data_type=data_type, transform=transform)\n    dataloader = DataLoader(dataset, batch_size=batch_size, collate_fn=collate_fn)\n\n    return dataloader\n\n","metadata":{"execution":{"iopub.status.busy":"2023-05-10T23:36:28.789112Z","iopub.execute_input":"2023-05-10T23:36:28.78952Z","iopub.status.idle":"2023-05-10T23:36:28.807839Z","shell.execute_reply.started":"2023-05-10T23:36:28.789489Z","shell.execute_reply":"2023-05-10T23:36:28.806286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Instantiate the dataloader by defining the global path and the \ntrain_dataloader = get_dataloader(GLOBAL_PATH, 'train', batch_size=16)\nvalidation_dataloader = get_dataloader(GLOBAL_PATH, 'validation', batch_size=16)\n#test_dataloader = get_dataloader(GLOBAL_PATH, 'test', batch_size=16)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T23:36:28.810149Z","iopub.execute_input":"2023-05-10T23:36:28.81051Z","iopub.status.idle":"2023-05-10T23:36:28.838463Z","shell.execute_reply.started":"2023-05-10T23:36:28.810481Z","shell.execute_reply":"2023-05-10T23:36:28.8372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the first batch\nbands, masks = next(iter(train_dataloader))\nprint(bands.shape)  # Batch size, H, W, Timeframe, Bands\nprint(masks.shape)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T23:36:28.83997Z","iopub.execute_input":"2023-05-10T23:36:28.840341Z","iopub.status.idle":"2023-05-10T23:36:29.670029Z","shell.execute_reply.started":"2023-05-10T23:36:28.840308Z","shell.execute_reply":"2023-05-10T23:36:29.668841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the first band and first time frame for all 16 samples\nplt.figure(figsize=(16, 16))\n\n# Loop through all bands and plot them\nfor i in range(16):\n    # Extract the first band and first time frame\n    band = bands[i, :, :, 0, 0]\n\n    # Convert tensor to numpy array if necessary\n    if isinstance(band, torch.Tensor):\n        band = band.cpu().numpy()\n\n    # Create a subplot for each band\n    plt.subplot(4, 4, i + 1)\n    plt.imshow(band)\n    plt.axis('off')\n\n# Show the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-10T23:36:29.673082Z","iopub.execute_input":"2023-05-10T23:36:29.67358Z","iopub.status.idle":"2023-05-10T23:36:31.200876Z","shell.execute_reply.started":"2023-05-10T23:36:29.673533Z","shell.execute_reply":"2023-05-10T23:36:31.199439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Plot the ground truth masks for all 16 samples:\nplt.figure(figsize=(16, 16))\n\n# Loop through all masks and plot them\nfor i in range(16):\n    mask = masks[i].numpy()\n\n    # Convert tensor to numpy array if necessary\n\n    # Create a subplot for each mask\n    plt.subplot(4, 4, i + 1)\n    plt.imshow(mask, cmap='gray')\n    plt.axis('off')\n\n# Show the plot\nplt.show()\n\n# Use the funct","metadata":{"execution":{"iopub.status.busy":"2023-05-10T23:36:31.202659Z","iopub.execute_input":"2023-05-10T23:36:31.203086Z","iopub.status.idle":"2023-05-10T23:36:32.100371Z","shell.execute_reply.started":"2023-05-10T23:36:31.203049Z","shell.execute_reply":"2023-05-10T23:36:32.099015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}