{"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":"code","source":"import os\nimport sys\nimport random\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import animation\nimport matplotlib.pyplot as plt\nfrom IPython import display\n\n\n# choice of ML tool\nimport torch\nfrom tqdm import tqdm\n\n# for input of the model\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\n\n\nimport torch.nn as nn\nfrom torchvision import models\nfrom torch.nn.functional import relu\n\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nimport copy\nimport time\nfrom collections import defaultdict\nimport torch.nn.functional as F\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-07T13:14:19.426021Z","iopub.execute_input":"2023-08-07T13:14:19.426761Z","iopub.status.idle":"2023-08-07T13:14:19.436202Z","shell.execute_reply.started":"2023-08-07T13:14:19.426722Z","shell.execute_reply":"2023-08-07T13:14:19.435217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\nTRAIN_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/train'\nTEST_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/test'","metadata":{"execution":{"iopub.status.busy":"2023-08-07T13:14:21.898417Z","iopub.execute_input":"2023-08-07T13:14:21.898848Z","iopub.status.idle":"2023-08-07T13:14:21.904699Z","shell.execute_reply.started":"2023-08-07T13:14:21.898816Z","shell.execute_reply":"2023-08-07T13:14:21.903318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#A custom Dataset class must implement three functions: __init__, __len__, and __getitem__\n_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\nclass ContrailDataset(Dataset):\n    \n\n    \n    # params :  all params required for len and get item method\n    def __init__(self , base_dir :str, dataset_type: str , image_size :int):\n        self.base_dir = base_dir \n        self.dataset_type = dataset_type\n        self.dataset_dir = os.path.join(self.base_dir, self.dataset_type)\n        self.records = os.listdir(self.dataset_dir)\n        self.image_size = image_size\n        print(f\"dataset dir = {self.dataset_dir} -> _exists {os._exists(self.dataset_dir)}\")\n        \n    \n                \n    def __len__(self):\n        return len(self.records)\n    \n    def normalize_range(self , data, bounds):\n        \"\"\"Maps data to the range [0, 1].\"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\n    \n    def preprocess_items(self):\n        # waring : needs too much memory and 20 minutes of time\n        # observation : Uses only cpu and ~max memory\n        # todo : optimize, vectorize, utilizer parallel processing\n        self.input_data =  []\n        self.output_labels = []\n#         self.input_data = np.zeros((self.image_size,self.image_size, 3 , self.__len__()))\n#         self.output_labels = np.zeros((self.image_size, self.image_size, self.__len__()))\n\n# observation : as this loop runs, memory increases incrementally and time taken by each loop oncreases too\n        for index , record_id in tqdm(enumerate(self.records)) : \n        \n            input_data ,output_labels = preprocess_item(record_id)\n                \n            # save in numpy array\n            self.input_data.append(input_data)\n            self.output_labels.append(output_labels)\n            \n\n    def preprocess_item(self , record_id):\n        band11 = np.load(os.path.join(self.dataset_dir, record_id, 'band_11.npy'))\n        band14 = np.load(os.path.join(self.dataset_dir, record_id, 'band_14.npy'))\n        band15 = np.load(os.path.join(self.dataset_dir, record_id, 'band_15.npy'))\n\n        r = self.normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        g = self.normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        b = self.normalize_range(band14, _T11_BOUNDS)\n        \n        false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n#         print(f'false_color {false_color.shape}')\n        human_pixel_mask = np.load(os.path.join(self.dataset_dir, record_id, 'human_pixel_masks.npy'))\n\n        return false_color , human_pixel_mask\n    \n    def __getitem__(self , idx):\n        \n        input_data ,output_labels = self.preprocess_item(self.records[idx])        \n\n        # start data augmentation\n        v = random.random()\n        hflip = v < 0.25\n        vflip = v > 0.5 and v < 0.75\n        if vflip: \n#             print(\"fliped V\")\n            input_data = np.flip(input_data,0) #0 vertical\n            output_labels = np.flip(output_labels,0)\n            \n        if hflip:\n#             print(\"fliped H\")\n            input_data = np.flip(input_data,1)# 1 horizontal\n            output_labels = np.flip(output_labels,1)\n        # end data augmentation\n        \n#         print(f'{input_data.shape=}, {output_labels.shape}')\n\n        # start data prep for torch training\n        # re arranges channels required for traning\n        input_data = torch.tensor( np.array(input_data).copy()).permute(2,0,1,3)\n        output_labels = torch.tensor(np.array(output_labels).copy(),  dtype=torch.float)\n        # end  data prep for torch training\n        \n        return input_data ,output_labels","metadata":{"execution":{"iopub.status.busy":"2023-08-07T13:14:22.386308Z","iopub.execute_input":"2023-08-07T13:14:22.386919Z","iopub.status.idle":"2023-08-07T13:14:22.406947Z","shell.execute_reply.started":"2023-08-07T13:14:22.386886Z","shell.execute_reply":"2023-08-07T13:14:22.405568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = ContrailDataset(base_dir=BASE_DIR, dataset_type='validation' ,image_size =  256)\ntrain_loader = DataLoader(train_data, batch_size=5, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T18:17:32.115407Z","iopub.execute_input":"2023-08-06T18:17:32.116187Z","iopub.status.idle":"2023-08-06T18:17:32.130734Z","shell.execute_reply.started":"2023-08-06T18:17:32.11615Z","shell.execute_reply":"2023-08-06T18:17:32.128936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img =  train_data.__getitem__(6)[0]\nimg.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-06T18:17:32.501725Z","iopub.execute_input":"2023-08-06T18:17:32.502189Z","iopub.status.idle":"2023-08-06T18:17:32.610246Z","shell.execute_reply.started":"2023-08-06T18:17:32.502153Z","shell.execute_reply":"2023-08-06T18:17:32.608967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for inputs, label in train_loader:\n    frame1 = inputs[...,7]\n    print(f'{frame1.shape=} {inputs.shape=} , {label.shape=}')","metadata":{"execution":{"iopub.status.busy":"2023-08-06T18:17:32.864766Z","iopub.execute_input":"2023-08-06T18:17:32.86522Z","iopub.status.idle":"2023-08-06T18:17:40.228332Z","shell.execute_reply.started":"2023-08-06T18:17:32.865184Z","shell.execute_reply":"2023-08-06T18:17:40.225229Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img8 , human_pixel_mask =train_data.__getitem__(1)\nimg = img8[...,0]\nprint(img.shape)\nplt.figure(figsize=(18, 6))\nax = plt.subplot(1, 3, 1)\nax.imshow(img)\nax.set_title('False color image/satelite recorded image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(human_pixel_mask, interpolation='none')\nax.set_title('Ground truth contrail mask/labeled data')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(img)\nax.imshow(human_pixel_mask, cmap='Reds', alpha=.4, interpolation='none')\nax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-08-06T18:17:40.230607Z","iopub.status.idle":"2023-08-06T18:17:40.231272Z","shell.execute_reply.started":"2023-08-06T18:17:40.230961Z","shell.execute_reply":"2023-08-06T18:17:40.230991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Animation\nfig = plt.figure(figsize=(6, 6))\nim = plt.imshow(img8[..., 0])\ndef draw(i):\n    im.set_array(img8[..., i])\n    return [im]\nanim = animation.FuncAnimation(\n    fig, draw, frames=img8.shape[-1], interval=500, blit=True\n)\nplt.close()\ndisplay.HTML(anim.to_jshtml())","metadata":{"execution":{"iopub.status.busy":"2023-08-06T18:09:52.7065Z","iopub.execute_input":"2023-08-06T18:09:52.706972Z","iopub.status.idle":"2023-08-06T18:09:55.388653Z","shell.execute_reply.started":"2023-08-06T18:09:52.706933Z","shell.execute_reply":"2023-08-06T18:09:55.386963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## ASH dataset","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ash_dataset = '/kaggle/input/contrails-images-ash-color/contrails'\ntrain_path = '/kaggle/input/contrails-images-ash-color/train_df.csv'\nval_path = '/kaggle/input/contrails-images-ash-color/valid_df.csv'","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:07:58.31872Z","iopub.execute_input":"2023-08-06T17:07:58.319116Z","iopub.status.idle":"2023-08-06T17:07:58.325009Z","shell.execute_reply.started":"2023-08-06T17:07:58.319086Z","shell.execute_reply":"2023-08-06T17:07:58.323655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"record_id = \"5728069425727341010\"\n# record_id = contrail_records[10]\nTRAIN_DIR = ash_dataset\n\nwith open(os.path.join(TRAIN_DIR,  f'{record_id}.npy'), 'rb') as f:\n    data = np.load(f)\n# with open(os.path.join(TRAIN_DIR, record_id, 'band_14.npy'), 'rb') as f:\n#     band14 = np.load(f)\n# with open(os.path.join(TRAIN_DIR, record_id, 'band_15.npy'), 'rb') as f:\n#     band15 = np.load(f)\n# with open(os.path.join(TRAIN_DIR, record_id, 'human_pixel_masks.npy'), 'rb') as f:\n#     human_pixel_mask = np.load(f)\n# with open(os.path.join(TRAIN_DIR, record_id, 'human_individual_masks.npy'), 'rb') as f:\n#     human_individual_mask = np.load(f)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:17:21.219965Z","iopub.execute_input":"2023-08-06T17:17:21.22044Z","iopub.status.idle":"2023-08-06T17:17:21.229733Z","shell.execute_reply.started":"2023-08-06T17:17:21.220403Z","shell.execute_reply":"2023-08-06T17:17:21.228607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:17:21.565611Z","iopub.execute_input":"2023-08-06T17:17:21.566719Z","iopub.status.idle":"2023-08-06T17:17:21.577413Z","shell.execute_reply.started":"2023-08-06T17:17:21.566675Z","shell.execute_reply":"2023-08-06T17:17:21.575854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_TIMES_BEFORE = 4\nimg = data[:,:,2]\nhuman_pixel_mask = data[...,3]\nprint(img.shape)\nplt.figure(figsize=(18, 6))\nax = plt.subplot(1, 3, 1)\nax.imshow(img)\nax.set_title('False color image/satelite recorded image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(human_pixel_mask, interpolation='none')\nax.set_title('Ground truth contrail mask/labeled data')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(img)\nax.imshow(human_pixel_mask, cmap='Reds', alpha=.4, interpolation='none')\nax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:17:22.664752Z","iopub.execute_input":"2023-08-06T17:17:22.665193Z","iopub.status.idle":"2023-08-06T17:17:23.890733Z","shell.execute_reply.started":"2023-08-06T17:17:22.66516Z","shell.execute_reply":"2023-08-06T17:17:23.889777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Animation\nfig = plt.figure(figsize=(6, 6))\nim = plt.imshow(img8[..., 0])\ndef draw(i):\n    im.set_array(img8[..., i])\n    return [im]\nanim = animation.FuncAnimation(\n    fig, draw, frames=img8.shape[-1], interval=500, blit=True\n)\nplt.close()\ndisplay.HTML(anim.to_jshtml())","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#A custom Dataset class must implement three functions: __init__, __len__, and __getitem__\nclass ContrailAshDataset(Dataset):\n    \n    ash_dataset = '/kaggle/input/contrails-images-ash-color/contrails'\n\n    train_path = '/kaggle/input/contrails-images-ash-color/train_df.csv'\n    val_path = '/kaggle/input/contrails-images-ash-color/valid_df.csv'\n\n    \n    # params :  all params required for len and get item method\n    def __init__(self , dataset_type: str , image_size :int):\n        self.dataset_type = dataset_type\n        self.train_ids = pd.read_csv(train_path)['record_id']\n        self.val_ids = pd.read_csv(val_path)['record_id']\n        self.image_size = image_size\n        \n    \n                \n    def __len__(self):\n        if(self.dataset_type == 'train'):\n            return len(self.train_ids)\n        if(self.dataset_type == 'val'):\n            return len(self.val_ids)\n        \n    \n            \n\n    def preprocess_item(self , record_id):\n        print(record_id)\n        data = np.load(os.path.join(ash_dataset, f'{record_id}.npy'))\n\n        print(f'data {data.shape}')\n\n        return data[:,:,:3] , data[:, :, 3]\n    \n    def __getitem__(self , idx):\n        if(self.dataset_type == 'train'):\n            input_data ,output_labels = self.preprocess_item(self.train_ids[idx])\n            \n        elif(self.dataset_type == 'val'):\n            input_data ,output_labels = self.preprocess_item(self.val_ids[idx])\n            \n            \n        input_data = torch.tensor(input_data)\n        output_labels = torch.tensor(output_labels)\n        return input_data ,output_labels","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:39:46.179316Z","iopub.execute_input":"2023-08-06T17:39:46.17979Z","iopub.status.idle":"2023-08-06T17:39:46.19409Z","shell.execute_reply.started":"2023-08-06T17:39:46.179753Z","shell.execute_reply":"2023-08-06T17:39:46.192523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ash_train_data = ContrailAshDataset( dataset_type='val' ,image_size =  256)\n# train_loader = DataLoader(train_data, batch_size=batch_size, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:39:46.487993Z","iopub.execute_input":"2023-08-06T17:39:46.488536Z","iopub.status.idle":"2023-08-06T17:39:46.51426Z","shell.execute_reply.started":"2023-08-06T17:39:46.488498Z","shell.execute_reply":"2023-08-06T17:39:46.512939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\n# Assuming you have the NumPy array with shape (256, 256, 4)\n# Let's call the array \"data_array\"\ndata_array = np.random.rand(256, 256, 4)  # Replace this with your actual data\n\n# Get the first three channels (slices along the third dimension)\nfirst_three_channels = data_array[:, :, :3]\n\n# Get the last channel (slice along the third dimension)\nlast_channel = data_array[:, :, 3]\n\n# Print the shapes of the extracted arrays for verification\nprint(\"First three channels shape:\", first_three_channels.shape)\nprint(\"Last channel shape:\", last_channel.shape)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:39:46.792572Z","iopub.execute_input":"2023-08-06T17:39:46.79365Z","iopub.status.idle":"2023-08-06T17:39:46.803962Z","shell.execute_reply.started":"2023-08-06T17:39:46.793606Z","shell.execute_reply":"2023-08-06T17:39:46.802939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img8 , human_pixel_mask = ash_train_data.__getitem__(6)\nimg = img8\nprint(img.shape)\n# plt.figure(figsize=(18, 6))\n# ax = plt.subplot(1, 3, 1)\n# ax.imshow(img)\n# ax.set_title('False color image/satelite recorded image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(human_pixel_mask, interpolation='none')\nax.set_title('Ground truth contrail mask/labeled data')\n\n# ax = plt.subplot(1, 3, 3)\n# ax.imshow(img)\n# ax.imshow(human_pixel_mask, cmap='Reds', alpha=.4, interpolation='none')\n# ax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:40:12.900686Z","iopub.execute_input":"2023-08-06T17:40:12.901163Z","iopub.status.idle":"2023-08-06T17:40:13.167047Z","shell.execute_reply.started":"2023-08-06T17:40:12.901126Z","shell.execute_reply":"2023-08-06T17:40:13.16566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"human_pixel_mask.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-06T17:40:13.246457Z","iopub.execute_input":"2023-08-06T17:40:13.247131Z","iopub.status.idle":"2023-08-06T17:40:13.253808Z","shell.execute_reply.started":"2023-08-06T17:40:13.247092Z","shell.execute_reply":"2023-08-06T17:40:13.252376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}