{"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":51753,"databundleVersionId":5692552,"sourceType":"competition"},{"sourceId":7144812,"sourceType":"datasetVersion","datasetId":4119014}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 0. Import Libraries","metadata":{}},{"cell_type":"code","source":"!pip install torch-summary -q\n!pip install onnx -q\n!pip install onnxscript -q","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:14:34.807122Z","iopub.execute_input":"2023-12-07T11:14:34.807396Z","iopub.status.idle":"2023-12-07T11:15:10.749585Z","shell.execute_reply.started":"2023-12-07T11:14:34.807373Z","shell.execute_reply":"2023-12-07T11:15:10.748474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nimport os\nimport shutil\nimport random\nimport time\nimport gc\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import TensorDataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchmetrics.functional as F_metrics\nfrom torchsummary import summary","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-07T11:15:10.751804Z","iopub.execute_input":"2023-12-07T11:15:10.752182Z","iopub.status.idle":"2023-12-07T11:15:16.788133Z","shell.execute_reply.started":"2023-12-07T11:15:10.752148Z","shell.execute_reply":"2023-12-07T11:15:16.787352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Choose device type\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:15:16.789194Z","iopub.execute_input":"2023-12-07T11:15:16.789629Z","iopub.status.idle":"2023-12-07T11:15:16.82635Z","shell.execute_reply.started":"2023-12-07T11:15:16.789602Z","shell.execute_reply":"2023-12-07T11:15:16.825419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Hyperparameters\n\nLEARNING_RATE = 0.005\nNUM_EPOCHS = 100\nBATCH_SIZE = 32\n\nnum_classes=1\nweight_decay = 1e-5","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:15:16.828957Z","iopub.execute_input":"2023-12-07T11:15:16.829402Z","iopub.status.idle":"2023-12-07T11:15:16.83587Z","shell.execute_reply.started":"2023-12-07T11:15:16.829371Z","shell.execute_reply":"2023-12-07T11:15:16.835104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create directories for training and testing \n\nos.mkdir(\"/kaggle/working/dataset\")\nos.mkdir(\"/kaggle/working/dataset/train\")\nos.mkdir(\"/kaggle/working/dataset/train/images\")\nos.mkdir(\"/kaggle/working/dataset/train/labels\")\nos.mkdir(\"/kaggle/working/dataset/validation\")\nos.mkdir(\"/kaggle/working/dataset/validation/images\")\nos.mkdir(\"/kaggle/working/dataset/validation/labels\")\nos.mkdir(\"/kaggle/working/dataset/test\")\nos.mkdir(\"/kaggle/working/dataset/test/images\")\nos.mkdir(\"/kaggle/working/dataset/test/labels\")","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:15:16.836917Z","iopub.execute_input":"2023-12-07T11:15:16.837172Z","iopub.status.idle":"2023-12-07T11:15:16.846407Z","shell.execute_reply.started":"2023-12-07T11:15:16.83715Z","shell.execute_reply":"2023-12-07T11:15:16.845623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Fully Convoluted Network (FCN-8s)","metadata":{}},{"cell_type":"code","source":"class FCN(nn.Module):\n    def __init__(self, in_channels, out_channels=1):\n        super(FCN, self).__init__()\n\n        self.encoder_1 = self.encoder_block_2(in_channels, 64) # i/p = 256, o/p = 128\n        self.encoder_2 = self.encoder_block_2(64, 128)         # i/p = 128, o/p = 64\n        self.encoder_3 = self.encoder_block_3(128, 256)        # i/p = 64, o/p = 32\n        self.encoder_4 = self.encoder_block_3(256, 512)        # i/p = 32, o/p = 16\n\n        self.mid = self.mid_block(512, 1024)                   # i/p = 16, o/p = 16\n\n        self.conv_t_32s = nn.ConvTranspose2d(1, 1, 2, 2)       # i/p = 16  , o/p = 32\n        self.conv_t_16s = nn.ConvTranspose2d(2, 1, 2, 2)       # i/p = 32,  o/p = 64\n        self.conv_t_8s  = nn.ConvTranspose2d(2, 1, 4, 4)      # i/p = 64, o/p = 256\n\n        self.x3_conv_1x1 = nn.Conv2d(256, 1, 1, 1)\n        self.x2_conv_1x1 = nn.Conv2d(128, 1, 1, 1)\n        \n        self.output = nn.Sigmoid()\n\n\n    def encoder_block_2(self, in_channels, out_channels):\n        return  nn.Sequential(\n                    nn.Conv2d(in_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.Conv2d(out_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.MaxPool2d(2, 2)\n                )\n    \n    def encoder_block_3(self, in_channels, out_channels):\n        return  nn.Sequential(\n                    nn.Conv2d(in_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.Conv2d(out_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.Conv2d(out_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.MaxPool2d(2, 2)\n                )\n    \n    def mid_block(self, in_channels, out_channels):\n        return  nn.Sequential(\n                    nn.Conv2d(in_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.Dropout2d(0.5),\n                    nn.Conv2d(out_channels, out_channels, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                    nn.Dropout2d(0.5),\n                    nn.Conv2d(out_channels, 1, 3, 1, 'same'),\n                    nn.ReLU(inplace=True),\n                )\n\n    \n    def forward(self, x):\n\n        # Encoder\n        x1 = self.encoder_1(x)\n        x2 = self.encoder_2(x1)    # conv_1x1\n        x3 = self.encoder_3(x2)    # conv_1x1\n        x4 = self.encoder_4(x3)\n\n        # Conv_1x1\n        x3_1x1 = self.x3_conv_1x1(x3)\n        x2_1x1 = self.x2_conv_1x1(x2)\n\n        # Mid-Block\n        x4 = self.mid(x4)\n\n        # FCN-32s output\n        x4 = self.conv_t_32s(x4)\n        x5 = torch.cat([x4, x3_1x1], dim=1)\n        x5 = self.conv_t_16s(x5)\n        x6 = torch.cat([x5, x2_1x1], dim=1)\n        x6 = self.conv_t_8s(x6)\n        x6 = self.output(x6)\n        \n        return x6","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:15:16.847474Z","iopub.execute_input":"2023-12-07T11:15:16.847747Z","iopub.status.idle":"2023-12-07T11:15:16.864857Z","shell.execute_reply.started":"2023-12-07T11:15:16.847726Z","shell.execute_reply":"2023-12-07T11:15:16.864142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Loading the data\n## Functions\n\n- **process_bands**: Chooses the middle frame of the bands, rounds up the DN values and converts them to 'uint8' type.\n- **create_inputs**: Combines all bands into a single file. Uses ThreadPoolExecutor for faster processing.","metadata":{}},{"cell_type":"code","source":"import concurrent.futures\n\ndef process_bands(band_path):\n    load_band = np.load(band_path)\n    load_band = np.round(load_band[:, :, 4], 0) # Rounds off DN values correctly, changing type does not do that.\n    load_band = load_band.astype(np.uint8)\n    return load_band\n\ndef create_inputs(folder_id, path, set_type):\n    npy_filepath = os.path.join(path, folder_id)\n    bands = sorted(os.listdir(npy_filepath))[:9]\n\n    combined_bands = []\n    with concurrent.futures.ThreadPoolExecutor() as executor:\n        futures = [executor.submit(process_bands, os.path.join(npy_filepath, band)) for band in bands]\n        combined_bands = [future.result() for future in concurrent.futures.as_completed(futures)]\n\n    combined_bands = np.stack(combined_bands, axis=0)\n    return combined_bands","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create train and validation inputs from bands\n\nstart = time.time()\n\ntrain_images = []\nval_images = []\n\n\ntrain_dir = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train\"\nfor file_id in tqdm(sorted(os.listdir(train_dir))):\n    train_np_band = create_inputs(file_id, train_dir, 'train')\n    train_images.append(train_np_band)\n    del train_np_band\n        \ntrain_images = torch.stack([torch.from_numpy(arr) for arr in train_images])\ntorch.save(train_images, '/kaggle/working/train_images.pt')\ndel train_images\n        \ngc.collect()\n    \nval_dir = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation\"\nfor file_id in tqdm(os.listdir(val_dir)):\n    val_np_band = create_inputs(file_id, val_dir, 'validation')\n    val_images.append(val_np_band)\n    del val_np_band\n    \nval_images = torch.stack([torch.from_numpy(arr) for arr in val_images])\ntorch.save(val_images, '/kaggle/working/val_images.pt')\ndel val_images\n    \n# test_dir = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\"\n# for file_id in tqdm(os.listdir(test_dir)):\n#     create_inputs(file_id, test_dir, 'test')\n    \nend = time.time()\n\nprint(\"Process took:\", end-start, \"seconds.\")\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create train and validation masks\n\ntrain_masks = []\nval_masks = []\n\ntrain_masks_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train\"\n    \nfor npy_file in tqdm(sorted(os.listdir(train_masks_path))):\n    load_npy_file = np.load(os.path.join(train_masks_path, npy_file, sorted(os.listdir(os.path.join(train_masks_path, npy_file)))[-1]))\n    load_npy_file = load_npy_file.astype(np.uint8)\n    train_masks.append(load_npy_file)\n    del load_npy_file\n    gc.collect()\n\ntrain_masks = torch.stack([torch.from_numpy(arr) for arr in train_masks])\ntorch.save(train_masks, '/kaggle/working/train_masks.pt')\ndel train_masks\n\n    \nval_masks_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation\"\n    \nfor npy_file in tqdm(os.listdir(val_masks_path)):\n    load_npy_file = np.load(os.path.join(val_masks_path, npy_file, sorted(os.listdir(os.path.join(val_masks_path, npy_file)))[-1]))\n    load_npy_file = load_npy_file.astype(np.uint8)\n    val_masks.append(load_npy_file)\n    del load_npy_file\n    gc.collect()\n    \nval_masks = torch.stack([torch.from_numpy(arr) for arr in val_masks])\ntorch.save(val_masks, '/kaggle/working/val_masks.pt')\ndel val_masks\n    \ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images = torch.stack([torch.from_numpy(arr) for arr in train_images])\nval_images = torch.stack([torch.from_numpy(arr) for arr in val_images])\ntrain_masks = torch.stack([torch.from_numpy(arr) for arr in train_masks])\nval_masks = torch.stack([torch.from_numpy(arr) for arr in val_masks])\n\ntorch.save(train_images, '/kaggle/working/train_images.pt')\ntorch.save(val_images, '/kaggle/working/val_images.pt')\ntorch.save(train_masks, '/kaggle/working/train_masks.pt')\ntorch.save(val_masks, '/kaggle/working/val_masks.pt')\n\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images = torch.load('/kaggle/input/private-dataset-identify-contrails/train_images.pt')\nval_images = torch.load('/kaggle/input/private-dataset-identify-contrails/val_images.pt')\ntrain_masks = torch.load('/kaggle/input/private-dataset-identify-contrails/train_masks.pt')\nval_masks = torch.load('/kaggle/input/private-dataset-identify-contrails/val_masks.pt')","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:15:16.865909Z","iopub.execute_input":"2023-12-07T11:15:16.866156Z","iopub.status.idle":"2023-12-07T11:16:51.794322Z","shell.execute_reply.started":"2023-12-07T11:15:16.866135Z","shell.execute_reply":"2023-12-07T11:16:51.793443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_masks[0].shape","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:16:51.795564Z","iopub.execute_input":"2023-12-07T11:16:51.795989Z","iopub.status.idle":"2023-12-07T11:16:51.814547Z","shell.execute_reply.started":"2023-12-07T11:16:51.795957Z","shell.execute_reply":"2023-12-07T11:16:51.813695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TensorDataset(train_images, train_masks)\nval_dataset = TensorDataset(val_images, val_masks)\n\ndel train_images\ndel train_masks\ndel val_images\ndel val_masks\n\ngc.collect()\n\ntrain_dataloader = DataLoader(train_dataset, \n                              batch_size=BATCH_SIZE, \n                              shuffle=True, \n                              drop_last=True,\n                              pin_memory=True,\n                              )\n\n\nval_dataloader = DataLoader(val_dataset, \n                            batch_size=BATCH_SIZE, \n                            shuffle=False,\n                            drop_last=False,\n                            pin_memory=True,\n                            )\n\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:16:51.815887Z","iopub.execute_input":"2023-12-07T11:16:51.816147Z","iopub.status.idle":"2023-12-07T11:16:52.085513Z","shell.execute_reply.started":"2023-12-07T11:16:51.816124Z","shell.execute_reply":"2023-12-07T11:16:52.084543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Model Training","metadata":{}},{"cell_type":"code","source":"# Defining metrics and visualisations\n\ndef dice_coefficient(predicted, target):\n\n    predicted = (predicted > 0.5).float()\n\n    intersection = torch.sum(predicted * target)\n    union = torch.sum(predicted) + torch.sum(target)\n    dice = (2.0 * intersection) / (union + 1e-8)  # Add a small epsilon to avoid division by zero\n    return dice\n\ndef jaccard_index(predicted, target):\n\n    predicted = (predicted > 0.5).float()\n\n    intersection = torch.sum(predicted * target)\n    union = torch.sum(predicted) + torch.sum(target) - intersection\n    jaccard = (intersection) / (union + 1e-8)  # Add a small epsilon to avoid division by zero\n    return jaccard\n\ndef visualize_segmentation(ground_truth, predicted):\n    fig, axes = plt.subplots(1, 2, figsize=(9, 4))\n\n    axes[0].imshow(ground_truth, cmap='gray')\n    axes[0].set_title('Ground Truth Mask')\n\n    axes[1].imshow(predicted, cmap='gray')\n    axes[1].set_title('Predicted Mask')\n\n    for ax in axes:\n        ax.axis('off')\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:16:52.088113Z","iopub.execute_input":"2023-12-07T11:16:52.088425Z","iopub.status.idle":"2023-12-07T11:16:52.097804Z","shell.execute_reply.started":"2023-12-07T11:16:52.088401Z","shell.execute_reply":"2023-12-07T11:16:52.096866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = FCN(9, 1)\n# model = nn.DataParallel(model, device_ids=[0, 1])\nmodel.to(device)\n\ncriterion = nn.BCELoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE, weight_decay=weight_decay)\n\nprint(summary(model, (9, 256, 256)))\n\n# print(model)","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:17:30.970184Z","iopub.execute_input":"2023-12-07T11:17:30.970841Z","iopub.status.idle":"2023-12-07T11:17:38.697714Z","shell.execute_reply.started":"2023-12-07T11:17:30.970808Z","shell.execute_reply":"2023-12-07T11:17:38.696744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loss = []\nval_loss = []\ndice_values = []\niou_values = []\n\nfor i in range(NUM_EPOCHS):\n    \n    # Train\n    model.train()\n    train_running_loss = 0.0\n    \n    for inputs, targets in tqdm(train_dataloader):\n        inputs, targets = inputs.to(device), targets.to(device)\n        inputs = inputs.float()\n        targets = targets.permute(0, 3, 1, 2)\n        targets = targets.float()\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        outputs = outputs.cpu()\n        targets = targets.cpu()\n        loss = criterion(outputs, targets)\n        loss.backward()\n        optimizer.step()\n        \n        train_running_loss += loss.item()\n        \n    # Validation\n    model.eval()\n    val_running_loss = 0.0\n    dice_scores = []\n    jaccard_scores = []\n    \n    for inputs, targets in val_dataloader:\n        inputs, targets = inputs.to(device), targets.to(device)\n        inputs = inputs.float()\n        targets = targets.permute(0, 3, 1, 2)\n        targets = targets.float()\n        outputs = model(inputs)\n        outputs = outputs.detach().cpu()\n        outputs = outputs.permute(0, 2, 3, 1)\n        targets = targets.cpu()\n        targets = targets.permute(0, 2, 3, 1)\n        loss = criterion(outputs, targets)\n        \n        val_running_loss += loss.item()\n        \n        for batch_idx in range(len(outputs)):\n            dice = dice_coefficient(outputs[batch_idx], targets[batch_idx])\n            jaccard = jaccard_index(outputs[batch_idx], targets[batch_idx])\n            dice_scores.append(dice.item())\n            jaccard_scores.append(jaccard.item())\n            \n#         visualize_segmentation(targets[0], outputs[0])\n        \n    train_running_loss /= len(train_dataloader)\n    val_running_loss /= len(val_dataloader)\n    avg_dice_score = sum(dice_scores) / len(dice_scores)\n    avg_jaccard_score = sum(jaccard_scores) / len(jaccard_scores)\n    \n    train_loss.append(train_running_loss)\n    val_loss.append(val_running_loss)\n    dice_values.append(avg_dice_score)\n    iou_values.append(avg_jaccard_score)\n        \n    print(f\"Epoch: {i+1}/{NUM_EPOCHS} | Training Loss: {train_running_loss:.5f} | Validation Loss: {val_running_loss:.5f} | Dice: {avg_dice_score:.5f} | IoU: {avg_jaccard_score:.5f}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-07T11:17:38.699865Z","iopub.execute_input":"2023-12-07T11:17:38.700163Z","iopub.status.idle":"2023-12-07T13:47:41.222571Z","shell.execute_reply.started":"2023-12-07T11:17:38.700137Z","shell.execute_reply":"2023-12-07T13:47:41.220676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_values = range(NUM_EPOCHS)\n\nplt.plot(x_values, train_loss, label='Train Loss')\nplt.plot(x_values, val_loss, label='Validation Loss')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T13:47:41.223635Z","iopub.status.idle":"2023-12-07T13:47:41.223961Z","shell.execute_reply.started":"2023-12-07T13:47:41.223803Z","shell.execute_reply":"2023-12-07T13:47:41.223818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_values = range(NUM_EPOCHS)\n\nplt.plot(x_values, dice_values, label='Dice Coeff.')\nplt.plot(x_values, iou_values, label='IoU')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-07T13:47:41.225322Z","iopub.status.idle":"2023-12-07T13:47:41.225698Z","shell.execute_reply.started":"2023-12-07T13:47:41.225522Z","shell.execute_reply":"2023-12-07T13:47:41.225546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}