{"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":"<div style=\"background-color:#5D73F2; color:#19180F; font-size:40px; font-family:Arial; padding:10px; border: 5px solid #19180F; border-radius:10px\"> MAnet for segmentation of remote sensing images </div>\n<div style=\"background-color:#D5D9F2; color:#19180F; font-size:15px; font-family:Arial; padding:10px; border: 5px solid #19180F; border-radius:10px\"> \n📌\nResearch paper - \n    <a href=\"https://arxiv.org/abs/2009.02130\"> Link </a><br>\n📌 Refer this <a href=\"https://www.kaggle.com/competitions/google-research-identify-contrails-reduce-global-warming/discussion/425901\"> Byte sized primer on choosing the right segmentation model(Competition's Discussion forum) </a><br>\n</div>\n<div style=\"background-color:#A8B4F6; color:#19180F; font-size:30px; font-family:Arial; padding:10px; border: 5px solid #19180F; border-radius:10px\">Contents of the notebook </div>\n<div style=\"background-color:#D5D9F2; color:#19180F; font-size:15px; font-family:Arial; padding:10px; border: 5px solid #19180F; border-radius:10px\"> \n1. Visualizing few band images along with their segmentations.<br>\n2. Preprocessing and creating dataloaders.<br>\n3. Defining MANet architecture.<br>\n4. Training and validation loop along with checkpointing the model.<br>\n5. Generate predictions <br>\n<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Arial; padding:10px; border: 5px solid #19180F; border-radius:10px\"> \n    📌 <b>What's New: 💡</b> <br>\n    1. Data Augmentation.<br>\n    2. Visual validation of input-output pairs generated by the dataloaders.<br>\n    3. Weighted Loss(nn.BCEwithLogitsLoss with pos_weight arg)<br>\n    4. Learning rate scheduler.<br>\n    5. MAnet model conceptualized as Ideal for the task.<br>\n    6. Training on complete train/val split of data with false color generated image by blending multiple channels & their mask.<br></div>\n    \n ","metadata":{}},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\nImporting modules    </div>","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.transforms import ToTensor\nfrom matplotlib import animation\nimport matplotlib.pyplot as plt\nfrom IPython import display\nimport json\nfrom torchvision import transforms\nimport pandas as pd\nfrom torch.optim.lr_scheduler import StepLR\nfrom tqdm import tqdm\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T14:47:43.295524Z","iopub.execute_input":"2023-08-08T14:47:43.296098Z","iopub.status.idle":"2023-08-08T14:47:47.799715Z","shell.execute_reply.started":"2023-08-08T14:47:43.296069Z","shell.execute_reply":"2023-08-08T14:47:47.798391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\n1. Visualizing data    </div>","metadata":{}},{"cell_type":"code","source":"# Load metadata information\nwith open('/kaggle/input/google-research-identify-contrails-reduce-global-warming/train_metadata.json', 'r') as f:\n    metadata = json.load(f)\n\n# Helper function to load data\ndef load_data(record_id, mode = \"train\"):\n    if mode == \"train\":\n        folder_path = f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/{record_id}/'\n    elif mode ==\"val\":\n        folder_path = f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation/{record_id}/'\n    else:\n        folder_path = f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/{record_id}/'\n\n            \n        \n    bands = [np.load(f'{folder_path}band_{i:02d}.npy') for i in range(8, 17)]\n    human_pixel_masks = np.load(f'{folder_path}human_pixel_masks.npy')\n    return bands, human_pixel_masks\n\n# Visualize sample images and their segmentation masks\nrecord_ids = ['1000216489776414077', '10016536018877742', '10038729395249389']\n\nfig, axes = plt.subplots(len(record_ids), 2, figsize=(12, 6*len(record_ids)))\n\nfor i, record_id in enumerate(record_ids):\n    bands, human_pixel_masks = load_data(record_id)\n\n    # Choose a frame to visualize\n    frame_index = 5\n    image = bands[0][:, :, frame_index]  # Choose band 08 for visualization\n    mask = human_pixel_masks[:, :, 0]\n\n    axes[i, 0].imshow(image)\n    axes[i, 0].set_title(f'Infrared Image - {record_id}')\n    axes[i, 1].imshow(mask, cmap='jet')\n    axes[i, 1].set_title(f'Contrail Segmentation Mask - {record_id}')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T14:47:47.801637Z","iopub.execute_input":"2023-08-08T14:47:47.802229Z","iopub.status.idle":"2023-08-08T14:47:51.006702Z","shell.execute_reply.started":"2023-08-08T14:47:47.802195Z","shell.execute_reply":"2023-08-08T14:47:51.005492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\n2. Preprocessing the data and creating dataloaders    </div>","metadata":{}},{"cell_type":"code","source":"# Define image transformations\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.RandomRotation(degrees=15),  # Example: Random rotation up to 15 degrees\n    transforms.RandomHorizontalFlip(p=0.5),  # Example: Random horizontal flip with 50% probability\n])\n\n# Custom Dataset class with false color image generation\nclass ContrailsDataset(Dataset):\n    def __init__(self, data_dir):\n        self.data_dir = data_dir\n        self.record_ids = os.listdir(data_dir)\n\n    def __len__(self):\n        return len(self.record_ids)\n\n    def __getitem__(self, idx):\n        record_id = self.record_ids[idx]\n        record_dir = os.path.join(self.data_dir, record_id)\n\n        bands = []\n        for i in range(8, 17):\n            band_path = os.path.join(record_dir, f\"band_{i:02d}.npy\")\n            with open(band_path, 'rb') as f:\n                band = np.load(f)\n                bands.append(band)\n\n        with open(os.path.join(record_dir, 'human_pixel_masks.npy'), 'rb') as f:\n            human_pixel_mask = np.load(f)\n\n        # Combine bands into a false color image\n        _T11_BOUNDS = (243, 303)\n        _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n        _TDIFF_BOUNDS = (-4, 2)\n\n        def normalize_range(data, bounds):\n            \"\"\"Maps data to the range [0, 1].\"\"\"\n            return (data - bounds[0]) / (bounds[1] - bounds[0])\n\n        r = normalize_range(bands[7] - bands[6], _TDIFF_BOUNDS)\n        g = normalize_range(bands[6] - bands[3], _CLOUD_TOP_TDIFF_BOUNDS)\n        b = normalize_range(bands[6], _T11_BOUNDS)\n        false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n\n        return {'bands': false_color, 'mask': human_pixel_mask}\n\n\ntrain_data_dir = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/train'\nvalid_data_dir = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation'\nbatch_size = 1\nnum_epochs = 10\nlearning_rate = 0.001\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ntrain_dataset = ContrailsDataset(train_data_dir)\nvalid_dataset = ContrailsDataset(valid_data_dir)\n\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=batch_size, shuffle=True)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T14:47:51.007724Z","iopub.execute_input":"2023-08-08T14:47:51.008042Z","iopub.status.idle":"2023-08-08T14:47:51.453586Z","shell.execute_reply.started":"2023-08-08T14:47:51.008015Z","shell.execute_reply":"2023-08-08T14:47:51.45259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\nPerforming sanity check of the dataloaders by visually inspecting the I/O pairs   </div>","metadata":{}},{"cell_type":"code","source":"count=0\nfor batch in train_loader:\n    bands, mask =  batch['bands'].squeeze(0), batch['mask'].squeeze(0)\n    if len(np.unique(mask)) > 1:\n        print(bands[:,:,:,0].shape,mask.shape)\n\n        plt.imshow(bands[:,:,:,0])\n        plt.show()\n        plt.imshow(mask)\n        plt.show()\n        count+=1\n    if count>5:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-08-08T14:47:51.456091Z","iopub.execute_input":"2023-08-08T14:47:51.457035Z","iopub.status.idle":"2023-08-08T14:47:58.655487Z","shell.execute_reply.started":"2023-08-08T14:47:51.457Z","shell.execute_reply":"2023-08-08T14:47:58.654474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"count=0\nfor batch in valid_loader:\n    bands, mask =  batch['bands'].squeeze(0), batch['mask'].squeeze(0)\n    if len(np.unique(mask)) > 1:\n        print(bands[:,:,:,0].shape,mask.shape)\n\n        plt.imshow(bands[:,:,:,0])\n        plt.show()\n        plt.imshow(mask)\n        plt.show()\n        count+=1\n    if count>5:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-08-08T14:47:58.657056Z","iopub.execute_input":"2023-08-08T14:47:58.657412Z","iopub.status.idle":"2023-08-08T14:48:07.835783Z","shell.execute_reply.started":"2023-08-08T14:47:58.657379Z","shell.execute_reply":"2023-08-08T14:48:07.834631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\n3. Defining the MANet architecture with Spatial Attention    </div>","metadata":{}},{"cell_type":"code","source":"class MAnet(nn.Module):\n    def __init__(self):\n        super(MAnet, self).__init__()\n        \n        # First convolutional layer\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1)\n        self.relu = nn.ReLU(inplace=True)\n        \n        # First attention block\n        self.attention1 = self._make_attention_block(64)\n        # Second attention block\n        self.attention2 = self._make_attention_block(64)\n        \n        # Second convolutional layer\n        self.conv2 = nn.Conv2d(64, 1, kernel_size=1, stride=1)\n        self.sigmoid = nn.Sigmoid()\n        \n    def _make_attention_block(self, in_channels):\n        # Attention block consists of three convolutional layers followed by ReLU and Sigmoid activations\n        return nn.Sequential(\n            nn.Conv2d(in_channels, in_channels // 8, kernel_size=1, stride=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(in_channels // 8, in_channels // 8, kernel_size=3, stride=1, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(in_channels // 8, in_channels, kernel_size=1, stride=1),\n            nn.Sigmoid()\n        )\n        \n    def forward(self, x):\n        # Apply the first convolutional layer\n        out = self.conv1(x)\n        out = self.relu(out)\n        \n        # Apply the first attention block\n        attention1 = self.attention1(out)\n        # Apply the second attention block\n        attention2 = self.attention2(out)\n        \n        # Combine the output using element-wise multiplication with attention masks\n        out = out * attention1 + out * attention2\n        \n        # Apply the second convolutional layer followed by sigmoid activation\n        out = self.conv2(out)\n        out = self.sigmoid(out)\n        \n        return out\n\n# Create an instance of the MAnet model\nmodel = MAnet()\n\n# Test the model with a random input\ninput = torch.randn(2, 3, 256, 256)  # Batch size of 2\noutput = model(input)\nprint(output.shape)  # Should print: torch.Size([2, 1, 256, 256])\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T14:49:24.822504Z","iopub.execute_input":"2023-08-08T14:49:24.822901Z","iopub.status.idle":"2023-08-08T14:49:25.171476Z","shell.execute_reply.started":"2023-08-08T14:49:24.822873Z","shell.execute_reply":"2023-08-08T14:49:25.170557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\n4. Training the model on MANet data along with validating and saving the model   </div>","metadata":{}},{"cell_type":"code","source":"\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n\"\"\" Weighted Loss computation takes time.. Uncomment to process\n#Calculate the class weights based on the training data\ncontrail_pixels = sum([np.sum(mask) for _, mask in tqdm(train_dataset)])\ntotal_pixels = sum([np.prod(mask.shape) for _, mask in tqdm(train_dataset)])\nweight_non_contrail = total_pixels / (total_pixels - contrail_pixels)\nweight_contrail = total_pixels / contrail_pixels\n\"\"\"\n# # Create the weighted loss function\n# criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(weight_contrail))\ncriterion = nn.BCEWithLogitsLoss()\n\n# Initialize the model and optimizer\nmodel = MAnet().to(device)\noptimizer = optim.Adam(model.parameters(), lr=0.001)\nscheduler = StepLR(optimizer, step_size=5, gamma=0.1)  # Reduce LR by 0.1 every 5 epochs\n\n# Training loop\nnum_epochs = 10\nfor epoch in tqdm(range(num_epochs)):\n    model.train()\n    total_loss = 0.0\n    for step, batch in enumerate(train_loader):\n        images, masks = batch['bands'][:,:,:,:,0].to(device), batch['mask'].to(device)\n        images, masks = images.permute(0,3,1,2), masks.permute(0,3,1,2)\n        #images = images[:,:,:,]\n        #print(images,masks)\n\n        # Forward pass\n        outputs = model(images)  # Add channel dimension\n        #print(outputs)\n\n        # Calculate loss\n        loss = criterion(outputs, masks.float())\n        total_loss += loss.item()\n\n        # Backward and optimize\n        optimizer.zero_grad()\n        loss.backward()\n        scheduler.step()\n        if step%100==0:\n            print(\"Step-{}, Loss-{}\".format(step, loss.item()))\n\n    avg_loss = total_loss / len(train_loader)\n    print(f'Epoch [{epoch+1}/{num_epochs}], Loss: {avg_loss:.4f}')\n\n# Validation loop\nmodel.eval()\nwith torch.no_grad():\n    total_val_loss = 0.0\n    for images, masks in (valid_loader):\n        images, masks = batch['bands'][:,:,:,:,0].to(device), batch['mask'].to(device)\n\n        # Forward pass\n        outputs = model(images)\n\n        # Calculate loss\n        val_loss = criterion(outputs, masks.unsqueeze(1).float())\n        total_val_loss += val_loss.item()\n\n    avg_val_loss = total_val_loss / len(valid_loader)\n    print(f'Validation Loss: {avg_val_loss:.4f}')\n\n# Save the model checkpoint\ntorch.save(model.state_dict(), 'MANet_contrail_segmentation.pth')\n","metadata":{"_kg_hide-output":false,"execution":{"iopub.status.busy":"2023-08-08T14:50:05.355172Z","iopub.execute_input":"2023-08-08T14:50:05.355549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<div style=\"background-color:#F0E3D2; color:#19180F; font-size:15px; font-family:Verdana; padding:10px; border: 2px solid #19180F; border-radius:10px\"> \n📌\n5. Generate predictions   </div>","metadata":{}},{"cell_type":"code","source":"def predict_on_test_set(model, test_record_ids):\n    model.eval()\n    test_predictions = {}\n    with torch.no_grad():\n        for record_id in test_record_ids:\n            folder_path = f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/{record_id}/'\n            bands = [np.load(f'{folder_path}band_{i:02d}.npy') for i in range(8, 17)]\n            test_images = torch.tensor(bands[0][:, :, 5], dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(device)\n            test_outputs = model(test_images)\n            test_predictions[record_id] = test_outputs.squeeze(0).squeeze(0).cpu().numpy()\n\n    return test_predictions\n\n\ntest_record_ids = ['1000834164244036115', '1002653297254493116']  # Replace with rd ids\nmodel = MANetWithSpatialAttention().to(device)\nmodel.load_state_dict(torch.load('MANet_contrail_segmentation.pth'))\n\n# Generate test set predictions\ntest_predictions = predict_on_test_set(model, test_record_ids)\n\ndef rle_encode(mask):\n    pixels = mask.flatten()\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\nsubmission = pd.DataFrame(columns=['record_id', 'prediction'])\nfor record_id, prediction in test_predictions.items():\n    encoded_prediction = rle_encode((prediction > 0.5).astype(np.int))\n    submission = submission.append({'record_id': record_id, 'prediction': encoded_prediction}, ignore_index=True)\n\nsubmission.to_csv('submission.csv', index=False)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}