{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":51753,"databundleVersionId":5692552,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Project: Participate in Kaggle Competition “Google Research - Identify Contrails to Reduce Global Warming”","metadata":{}},{"cell_type":"markdown","source":"## Context\nThe project consisted in entering a Kaggle competition with the objective of training a machine learning model to identify contrails in satellite images and help prevent their formation.\nFor context contrails are condensation trails, which are long, thin ice crystals clouds that form in an aircraft engine when it flies in humid areas.  Some of these contrails can linger and grow for several hours, occasionally merging into cloud formations that become visually identical to naturally occurring cirrus, making them difficult to detect and they can contribute to global warming by trapping heat in the atmosphere. The reason why it could be important to detect them is to help researchers improve the accuracy of their contrail models that predict when contrails will form and how much warming they will cause by validating these models with satellite imagery. This will help airlines avoid creating contrails and reduce their impact on climate change. ","metadata":{}},{"cell_type":"markdown","source":"## Data\nThe dataset provided by the challenge consisted of geostationary satellite images retrieved from the GOES-16 Advanced Baseline Imager (ABI). The train data consists of multiple sequences of images with one main labeled image. As contrails are easier to identify with temporal context, the challenge provides a sequence of images at 10-minute intervals, 4 before the labeled image and 3 after. For each of these sequences, there are 9 different folders that represent bands that are the infrared channel at different wavelengths and are converted to brightness temperatures. ","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import GradScaler \nimport torch.amp \nfrom tqdm.notebook import tqdm\nfrom matplotlib import animation\nimport matplotlib.pyplot as plt\nfrom IPython import display\nfrom torch.cuda.amp import autocast","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:46:04.389426Z","iopub.execute_input":"2025-12-07T17:46:04.390112Z","iopub.status.idle":"2025-12-07T17:46:04.394507Z","shell.execute_reply.started":"2025-12-07T17:46:04.39009Z","shell.execute_reply":"2025-12-07T17:46:04.393903Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Reproducibility","metadata":{}},{"cell_type":"code","source":"# Class that seeds everything to get reproducible results\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:46:12.900134Z","iopub.execute_input":"2025-12-07T17:46:12.900855Z","iopub.status.idle":"2025-12-07T17:46:12.905086Z","shell.execute_reply.started":"2025-12-07T17:46:12.900831Z","shell.execute_reply":"2025-12-07T17:46:12.904519Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset\nThe competition provides numpy arrays for each record_id, where there are 9 folders for each band (band_08 to band_16) in each record_id. The input of the model will be a version of the sequences but with only three channels (RGB) following the ash color scheme using only bands 11, 14 and 15.\nEach band folder contains an array (256, 256, 8). \nThese are a sequence of 8 images of size 256x256.\nNeed to convert data to tensors with shape (C,T,H,W)->(3,8,256,256)\nC->Channels(3 channels from the ash color scheme); T->(sequence=8 images in sequence); H->Height=256; W->Width=256\nThe label is an array (H, W, 1) that is the per pixel binary ground truth","metadata":{}},{"cell_type":"markdown","source":"### Parameters and normalization function for Ash Color Scheme\nThe ash color scheme uses the satellite bands to make images in RGB that make contrails appear in the image as dark blue. \nR = Band 15 - Band 14<br>\nG = Band 14 - Band 11<br>\nB = Band 14 <br>\n**Competition code reference:** Inversion. Ng, Joe. (2023). Visualizing Contrails [Source code]. Kaggle. https://www.kaggle.com/code/inversion/visualizing-contrails","metadata":{}},{"cell_type":"code","source":"# Bounds for Ash Color Scheme\n_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\n# Normalization function for Ash Color Scheme\ndef normalize_range(data, bounds):\n    \"\"\"Maps data to the range [0, 1].\"\"\"\n    return (data - bounds[0]) / (bounds[1] - bounds[0])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:46:30.180394Z","iopub.execute_input":"2025-12-07T17:46:30.181Z","iopub.status.idle":"2025-12-07T17:46:30.185161Z","shell.execute_reply.started":"2025-12-07T17:46:30.180975Z","shell.execute_reply":"2025-12-07T17:46:30.184351Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Quick visualization\nReplicated **Competition code reference:** Inversion. Ng, Joe. (2023). Visualizing Contrails [Source code]. Kaggle. https://www.kaggle.com/code/inversion/visualizing-contrails","metadata":{}},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/train'\nN_TIMES_BEFORE = 4\nrecord_id = '1704010292581573769'\n\nwith open(os.path.join(BASE_DIR, record_id, 'band_11.npy'), 'rb') as f:\n    band11 = np.load(f)\nwith open(os.path.join(BASE_DIR, record_id, 'band_14.npy'), 'rb') as f:\n    band14 = np.load(f)\nwith open(os.path.join(BASE_DIR, record_id, 'band_15.npy'), 'rb') as f:\n    band15 = np.load(f)\nwith open(os.path.join(BASE_DIR, record_id, 'human_pixel_masks.npy'), 'rb') as f:\n    human_pixel_mask = np.load(f)\nwith open(os.path.join(BASE_DIR, record_id, 'human_individual_masks.npy'), 'rb') as f:\n    human_individual_mask = np.load(f)\n\nr = normalize_range(band15 - band14, _TDIFF_BOUNDS)\ng = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\nb = normalize_range(band14, _T11_BOUNDS)\nfalse_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n\nimg = false_color[..., N_TIMES_BEFORE]\n\nplt.figure(figsize=(18, 6))\nax = plt.subplot(1, 3, 1)\nax.imshow(img)\nax.set_title('False color image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(human_pixel_mask, interpolation='none')\nax.set_title('Ground truth contrail mask')\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:46:38.372233Z","iopub.execute_input":"2025-12-07T17:46:38.37287Z","iopub.status.idle":"2025-12-07T17:46:39.145742Z","shell.execute_reply.started":"2025-12-07T17:46:38.372827Z","shell.execute_reply":"2025-12-07T17:46:39.144831Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Pytorch Dataset","metadata":{"execution":{"iopub.status.busy":"2025-12-07T06:05:11.928643Z","iopub.execute_input":"2025-12-07T06:05:11.928947Z","iopub.status.idle":"2025-12-07T06:05:11.93323Z","shell.execute_reply.started":"2025-12-07T06:05:11.928924Z","shell.execute_reply":"2025-12-07T06:05:11.932242Z"}}},{"cell_type":"code","source":"class ContrailDataset(Dataset):\n    \"\"\"\n    Loads the sequences, calculates Ash Color Scheme, and returns 3D tensors.\n    Input Shape: (H, W, T) from numpy files.\n    Output Shape: (C, T, H, W) for PyTorch 3D Conv.\n    \"\"\"\n    def __init__(self, data_dir, record_ids, train_mode=True):\n        self.root = Path(data_dir) \n        self.record_ids = list(record_ids)\n        self.train_mode = train_mode\n\n    def __len__(self):\n        return len(self.record_ids)\n\n    def __getitem__(self, idx):\n        rid = self.record_ids[idx]\n        rid_path = self.root / rid\n\n        # Load Bands (Shape: 256, 256, 8)\n        band11 = np.load(rid_path / \"band_11.npy\").astype(np.float32)\n        band14 = np.load(rid_path / \"band_14.npy\").astype(np.float32)\n        band15 = np.load(rid_path / \"band_15.npy\").astype(np.float32)\n\n        # Calculate Ash Color Scheme\n        # R = Band 15 - Band 14\n        r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        # G = Band 14 - Band 11\n        g = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        # B = Band 14\n        b = normalize_range(band14, _T11_BOUNDS)\n\n        # Stack to (3, 256, 256, 8)\n        rgb = np.stack([r, g, b], axis=0)\n        rgb = np.clip(rgb, 0, 1)\n\n        # Transpose from (C, H, W, T) to (C, T, H, W)\n        rgb = np.transpose(rgb, (0, 3, 1, 2)) \n\n        # Load Label (Only for Train/Validation)\n        label_path = rid_path / \"human_pixel_masks.npy\"\n        if label_path.exists():\n            label = np.load(label_path).astype(np.float32)\n            label = label[..., 0] # Remove channel (256, 256)\n            label = np.expand_dims(label, 0) # Add channel in the first position (1, 256, 256)\n            return torch.from_numpy(rgb), torch.from_numpy(label)\n        else:\n            # For Test set (no labels)\n            return torch.from_numpy(rgb), rid","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:46:55.313398Z","iopub.execute_input":"2025-12-07T17:46:55.314122Z","iopub.status.idle":"2025-12-07T17:46:55.321233Z","shell.execute_reply.started":"2025-12-07T17:46:55.314098Z","shell.execute_reply":"2025-12-07T17:46:55.32052Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model (3D U-Net with ConvLSTM in Bottle Neck)\n3D U-Net is a model that is widely used for 3D image segmentation that consists of 3D convolutional layers and a symmetrical encoder-decoder structure with skip connections to segment volumetric data. In this project the 3D data is provided by the temporal information. Instead of just doing the 3D convolutions in the bottle neck of the U-Net, to better process temporal context,  a ConvLSTM model was used to model temporal evolution of contrail formation so the model learns how features evolve over time.","metadata":{}},{"cell_type":"code","source":"# Sigle Convolution Long Short-Term Memory\nclass ConvLSTMCell(nn.Module):\n    \"\"\"\n    A single step of ConvLSTM\n    \"\"\"\n    def __init__(self, input_channels, hidden_channels, kernel_size=3):\n        \"\"\"\n        Initialize ConvLSTM cell\n\n        input_channels (int): Number of channels of input tensor.  \n        hidden_channels (int): Number of channels of hidden state.   \n        kernel_size (int): Size of the convolutional kernel.\n        \"\"\"\n        super().__init__()\n        # Initialize\n        self.input_channels = input_channels\n        self.hidden_channels = hidden_channels\n        padding = kernel_size // 2\n        # Compute input, forget, cell, and output gates\n        self.conv = nn.Conv2d(input_channels + hidden_channels, 4 * hidden_channels, kernel_size, padding=padding)\n\n    def forward(self, x, h, c):\n        \"\"\"\n        x: (B, C_in, H, W)\n        h, c: (B, C_hidden, H, W)\n        returns: h_next, c_next\n        \"\"\"\n        # Concatenate input and previous hidden state along channel axis\n        combined = torch.cat([x, h], dim=1)\n        # Convolution \n        gates = self.conv(combined)\n        # Split into input gate, forget gate, candidate, output gate\n        i, f, g, o = torch.chunk(gates, 4, dim=1)\n        # Nonlinearities\n        i = torch.sigmoid(i)\n        f = torch.sigmoid(f)\n        g = torch.tanh(g)\n        o = torch.sigmoid(o)\n        # Update\n        c_next = f * c + i * g\n        h_next = o * torch.tanh(c_next)\n        return h_next, c_next\n\nclass ConvLSTM(nn.Module):\n    \"\"\"\n    Full ConvLSTM using ConvLSTMCell\n    \"\"\"\n    def __init__(self, input_channels, hidden_channels, kernel_size=3):\n        super().__init__()\n        self.cell = ConvLSTMCell(input_channels, hidden_channels, kernel_size)\n\n    def forward(self, x):\n        \"\"\"\n        x: (Batch, Channel, Time, Height, Width)\n        \"\"\"\n        B, C, T, H, W = x.shape\n        # Initial h and c \n        h = torch.zeros(B, self.cell.hidden_channels, H, W, device=x.device)\n        c = torch.zeros(B, self.cell.hidden_channels, H, W, device=x.device)\n        \n        outputs = []\n        # Loop through each time step\n        for t in range(T):\n            # time step\n            x_t = x[:, :, t, :, :] \n            # One step of ConvLSTMCell\n            h, c = self.cell(x_t, h, c)\n            # add time dimension back (B, C, 1, H, W)\n            outputs.append(h.unsqueeze(2)) \n            \n        # Concatenate along time axis\n        return torch.cat(outputs, dim=2), (h, c)\n\n# Code from Wen, Q. (2020). ConvLSTM PyTorch implementation [Code repository]. GitHub.\n# https://github.com/ndrplz/ConvLSTM_pytorch\n\nclass Conv3DBlock(nn.Module):\n    \"\"\"\n    Standard 3D Convolution Block\n    \"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.block = nn.Sequential(\n            # First 3×3×3 conv\n            nn.Conv3d(in_channels, out_channels, 3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True),\n            # Second 3×3×3 conv\n            nn.Conv3d(out_channels, out_channels, 3, padding=1),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    def forward(self, x): return self.block(x)\n\n# Main model\nclass UNet3D_ConvLSTM(nn.Module):\n    \"\"\"\n    3D U-Net + ConvLSTM bottleneck.\n    Input:  x  (B, C=3, T=8, H=256, W=266)\n    Output: (B, 1, H, W) binary mask\n    \"\"\"\n    def __init__(self, in_channels=3, base_channels=16):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = Conv3DBlock(in_channels, base_channels)\n        self.pool1 = nn.MaxPool3d(kernel_size=(1, 2, 2)) \n        \n        self.enc2 = Conv3DBlock(base_channels, base_channels*2)\n        self.pool2 = nn.MaxPool3d(kernel_size=(1, 2, 2))\n        \n        self.enc3 = Conv3DBlock(base_channels*2, base_channels*4)\n        self.pool3 = nn.MaxPool3d(kernel_size=(1, 2, 2))\n        \n        # ConvLSTM Bottle neck\n        # Input: (B, 64, 8, 32, 32)\n        self.proj = nn.Conv3d(base_channels*4, base_channels*4, kernel_size=1)\n        self.lstm = ConvLSTM(base_channels*4, base_channels*8, kernel_size=3)\n        \n        # Decoder (Expanding Path)\n        self.up3 = nn.ConvTranspose3d(base_channels*8, base_channels*4, kernel_size=(1,2,2), stride=(1,2,2))\n        self.dec3 = Conv3DBlock(base_channels*8, base_channels*4)\n        \n        self.up2 = nn.ConvTranspose3d(base_channels*4, base_channels*2, kernel_size=(1,2,2), stride=(1,2,2))\n        self.dec2 = Conv3DBlock(base_channels*4, base_channels*2)\n        \n        self.up1 = nn.ConvTranspose3d(base_channels*2, base_channels, kernel_size=(1,2,2), stride=(1,2,2))\n        self.dec1 = Conv3DBlock(base_channels*2, base_channels)\n        \n        # Head\n        self.final = nn.Conv3d(base_channels, 1, kernel_size=1)\n\n    def forward(self, x):\n        # x: (B, 3, 8, 256, 256)\n        \n        # Encoder\n        e1 = self.enc1(x) # (16, 8, 256, 256)\n        p1 = self.pool1(e1) # (16, 8, 128, 128)\n        \n        e2 = self.enc2(p1) # (32, 8, 128, 128)\n        p2 = self.pool2(e2) # (32, 8, 64, 64)\n        \n        e3 = self.enc3(p2) # (64, 8, 64, 64)\n        p3 = self.pool3(e3) # (64, 8, 32, 32)\n        \n        # Bottleneck\n        p3 = self.proj(p3)\n        lstm_out, _ = self.lstm(p3) # (128, 8, 32, 32)\n        \n        # Decoder\n        u3 = self.up3(lstm_out) # (64, 8, 64, 64)\n        cat3 = torch.cat([u3, e3], dim=1) # Skip connection\n        d3 = self.dec3(cat3)\n        \n        u2 = self.up2(d3) # (32, 8, 128, 128)\n        cat2 = torch.cat([u2, e2], dim=1) # Skip connection\n        d2 = self.dec2(cat2)\n        \n        u1 = self.up1(d2) # (16, 8, 256, 256)\n        cat1 = torch.cat([u1, e1], dim=1) # Skip connection\n        d1 = self.dec1(cat1)\n        \n        # Final Projection\n        out_3d = self.final(d1) # (B, 1, 8, 256, 256)\n        \n        # Select the labeld image (5th image)\n        out_2d = out_3d[:, :, 4, :, :] \n        \n        return out_2d","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:47:26.527022Z","iopub.execute_input":"2025-12-07T17:47:26.527289Z","iopub.status.idle":"2025-12-07T17:47:26.544172Z","shell.execute_reply.started":"2025-12-07T17:47:26.527269Z","shell.execute_reply":"2025-12-07T17:47:26.543349Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss Function\nA combination of BCEWithLogitsLoss that combines a sigmoid layer and the Binary Cross Entropy Loss, and Dice Loss","metadata":{}},{"cell_type":"code","source":"# Dice Loss class\nclass DiceLoss(nn.Module):\n    def __init__(self, smooth=1e-6):\n        super().__init__()\n        self.smooth = smooth\n    def forward(self, logits, target):\n        probs = torch.sigmoid(logits)\n        # Intersection of prediction and label pixels\n        intersection = (probs * target).sum(dim=(1,2,3))\n        # Denominator is the union of prediction and label pixels\n        denom = probs.sum(dim=(1,2,3)) + target.sum(dim=(1,2,3)) \n        dice = (2 * intersection + self.smooth) / (denom + self.smooth)\n        return 1 - dice.mean()\n\n# Initialize loss functions\nbce_loss = nn.BCEWithLogitsLoss()\ndice_loss = DiceLoss()\n\ndef combined_loss(logits, targets, alpha=0.5):\n    \"\"\"\n    50% bce loss and 50% dice loss\n    \"\"\"\n    return alpha * bce_loss(logits, targets) + (1 - alpha) * dice_loss(logits, targets)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:47:41.829986Z","iopub.execute_input":"2025-12-07T17:47:41.830665Z","iopub.status.idle":"2025-12-07T17:47:41.837067Z","shell.execute_reply.started":"2025-12-07T17:47:41.830623Z","shell.execute_reply":"2025-12-07T17:47:41.836262Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Helper functions for one epoch","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scaler, device):\n    model.train()\n    running_loss = 0.0 # to track loss\n    # iterate through batches\n    for step, (x, y) in enumerate(loader, 1):\n        # Move to GPU\n        x = x.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n        # reset gradients\n        optimizer.zero_grad()\n        \n        # mixed precision\n        with autocast():\n            # Forward pass\n            logits = model(x)\n            # Compute loss\n            loss = combined_loss(logits, y)\n            \n        # Update\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        # Add loss to total\n        running_loss += loss.item()\n    \n    return running_loss / len(loader)\n\n\ndef validate_global_dice(model, loader, device, threshold=0.5):\n    \"\"\"\n    Calculates dice over the whole dataset.\n    \"\"\"\n    model.eval()\n    total_intersection = 0.0\n    total_union = 0.0\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(device), y.to(device)\n            logits = model(x)\n            probs = torch.sigmoid(logits)\n            preds = (probs > threshold).float()\n            \n            intersection = (preds * y).sum().item()\n            union = (preds.sum() + y.sum()).item()\n            \n            total_intersection += intersection\n            total_union += union\n    return (2 * total_intersection) / (total_union + 1e-6)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:47:49.149709Z","iopub.execute_input":"2025-12-07T17:47:49.150187Z","iopub.status.idle":"2025-12-07T17:47:49.156873Z","shell.execute_reply.started":"2025-12-07T17:47:49.150163Z","shell.execute_reply":"2025-12-07T17:47:49.156209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Prepare Data","metadata":{}},{"cell_type":"code","source":"# Reproducibility\nseed_everything(42)\n\n# Kaggle Paths\nDATA_DIR = Path('/kaggle/input/google-research-identify-contrails-reduce-global-warming')\nTRAIN_DIR = DATA_DIR / 'train'\nVALID_DIR = DATA_DIR / 'validation'\nTEST_DIR = DATA_DIR / 'test'\n\n# Record_ids\ntrain_ids = sorted(os.listdir(TRAIN_DIR))\nval_ids = sorted(os.listdir(VALID_DIR))\n\n# Convert data to tensors to input model\ntrain_ds = ContrailDataset(TRAIN_DIR, train_ids)\nval_ds = ContrailDataset(VALID_DIR, val_ids)\n\n# Data Loaders\nBATCH_SIZE = 2\nNUM_WORKERS = 2\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, \n                          num_workers=NUM_WORKERS, pin_memory=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, \n                        num_workers=NUM_WORKERS, pin_memory=True)\n\n# device used was the Kaggle GPU T4 x2\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Initialize model\nEPOCHS = 15\nBATCH_SIZE = 2 \nLR = 1e-4\nmodel = UNet3D_ConvLSTM(in_channels=3, base_channels=16).to(DEVICE)\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-2)\nscaler = GradScaler()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:48:01.884344Z","iopub.execute_input":"2025-12-07T17:48:01.885163Z","iopub.status.idle":"2025-12-07T17:48:01.939319Z","shell.execute_reply.started":"2025-12-07T17:48:01.885138Z","shell.execute_reply":"2025-12-07T17:48:01.938549Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"best_global_dice = 0.0\nsave_path = \"best_model.pth\"\n\nprint(f\"Starting Training for {EPOCHS} Epochs...\")\n\nfor epoch in range(1, EPOCHS + 1):\n    # Train\n    train_loss = train_one_epoch(model, train_loader, optimizer, scaler, DEVICE)\n    \n    # Global dice validation\n    val_global_dice = validate_global_dice(model, val_loader, DEVICE, threshold=0.5)\n\n    print(f\"Epoch {epoch} | Loss: {train_loss:.4f} | Global Dice: {val_global_dice:.4f}\")\n\n    if val_global_dice > best_global_dice:\n        best_global_dice = val_global_dice\n        torch.save(model.state_dict(), save_path)\n        print(\"  --> Saved New Best Model!\")\n\nprint(f\"Training Done. Best Global Dice: {best_global_dice:.5f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T17:48:09.559736Z","iopub.execute_input":"2025-12-07T17:48:09.560422Z","iopub.status.idle":"2025-12-07T23:35:05.957005Z","shell.execute_reply.started":"2025-12-07T17:48:09.5604Z","shell.execute_reply":"2025-12-07T23:35:05.95626Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualization of predictions","metadata":{}},{"cell_type":"code","source":"\ndef visualize_prediction(model, dataset, target=None, threshold=0.5):\n    \"\"\"\n    Visualizes a prediction against ground truth.\n    \n    model: The trained PyTorch model.\n    dataset: The validation dataset.\n    target: Can be an the index or record_id.\n    threshold: Probability threshold for binary mask.\n    \"\"\"\n    # Determine which index to load\n    idx = 0\n    if target is None:\n        idx = random.randint(0, len(dataset)-1)\n    elif isinstance(target, int):\n        idx = target\n    elif isinstance(target, str):\n        try:\n            # Find the index of the string ID in the dataset list\n            idx = dataset.record_ids.index(target)\n        except ValueError:\n            print(f\"Error: Record ID '{target}' not found in this dataset.\")\n            return\n            \n    # Load the data\n    # x: (3, 8, 256, 256), y_true: (1, 256, 256)\n    x, y_true = dataset[idx] \n    record_id = dataset.record_ids[idx]\n    \n    # Add batch dimension -> (1, 3, 8, 256, 256)\n    x_in = x.unsqueeze(0).to(DEVICE) \n    \n    # Run\n    model.eval()\n    with torch.no_grad():\n        logits = model(x_in)\n        # # Convert logits to probabilities (0 to 1)\n        probs = torch.sigmoid(logits)\n        # conver to binary mask\n        y_pred = (probs > threshold).float().cpu().numpy()[0, 0]\n        \n    # Prepare Images for Plotting\n    ash_img = x[:, 4, :, :].numpy().transpose(1, 2, 0) \n    y_true = y_true[0].numpy()\n    \n    # Plot\n    fig, ax = plt.subplots(1, 3, figsize=(15, 5))\n    \n    # Ash Color Image\n    ax[0].imshow(ash_img)\n    ax[0].set_title(f\"Record: {record_id}\\n(Index: {idx})\")\n    ax[0].axis('off')\n    \n    # Ground Truth\n    ax[1].imshow(y_true, cmap='gray')\n    ax[1].set_title(\"Ground Truth Mask\")\n    ax[1].axis('off')\n    \n    # Prediction\n    ax[2].imshow(y_pred, cmap='gray')\n    ax[2].set_title(f\"Model Prediction\\n(Threshold > {threshold})\")\n    ax[2].axis('off')\n    \n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:39:13.697102Z","iopub.execute_input":"2025-12-07T23:39:13.697429Z","iopub.status.idle":"2025-12-07T23:39:13.710048Z","shell.execute_reply.started":"2025-12-07T23:39:13.6974Z","shell.execute_reply":"2025-12-07T23:39:13.709216Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best model weights\nmodel.load_state_dict(torch.load(\"best_model.pth\"))\n\nprint(\"Visualizing Example without Contrails\")\nvisualize_prediction(model, val_ds, target='1000834164244036115')\n\n#print(\"\\n\" + \"=\"*50 + \"\\n\")\n\nprint(\"Visualizing Example with Contrails\")\nvisualize_prediction(model, val_ds, target='1087977782975734617')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:39:24.219339Z","iopub.execute_input":"2025-12-07T23:39:24.219979Z","iopub.status.idle":"2025-12-07T23:39:26.000203Z","shell.execute_reply.started":"2025-12-07T23:39:24.219949Z","shell.execute_reply":"2025-12-07T23:39:25.999518Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Generate Submission File\nBased on **Competition code reference:** Inversion. Ng, Joe. (2023). Contrails - RLE Submission [Source code]. Kaggle. https://www.kaggle.com/code/inversion/contrails-rle-submission\n","metadata":{}},{"cell_type":"code","source":"def rle_encode(x, fg_val=1):\n    \"\"\"\n    Encoding for submission.\n    x (numpy array): mask (1=contrail, 0=bg)\n    Returns: list of run lengths\n    \"\"\"\n    # 1d array with that finds the indices of pixels where there are contrails\n    dots = np.where(x.T.flatten() == fg_val)[0]\n    run_lengths = []\n    prev = -2 # Because indices start at 0\n    for b in dots:\n        # Check if the current pixel is not the neighbor of the previous pixel\n        if b > prev + 1:\n            # Add start position\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b # update\n    return run_lengths\n\ndef list_to_string(x):\n    \"\"\"\n    Converts RLE list to string for CSV\n    \"\"\"\n    if x:\n        s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n    else:\n        s = '-'\n    return s","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:39:57.171731Z","iopub.execute_input":"2025-12-07T23:39:57.172505Z","iopub.status.idle":"2025-12-07T23:39:57.17767Z","shell.execute_reply.started":"2025-12-07T23:39:57.172454Z","shell.execute_reply":"2025-12-07T23:39:57.176866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nGenerating Submission.csv...\")\n# list of test record ids\ntest_ids = sorted(os.listdir(TEST_DIR))\n# Data loader for test images\ntest_ds = ContrailDataset(TEST_DIR, test_ids)\ntest_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n\nsubmission_data = []\nmodel.eval()\n#disable gradiant\nwith torch.no_grad():\n    for x, rids in test_loader:\n        x = x.to(DEVICE)\n\n        # Get raw model outputs\n        logits = model(x)\n        # Convert logits to probabilities (0 to 1)\n        probs = torch.sigmoid(logits)\n        preds = (probs > 0.5).float().cpu().numpy()[:, 0, :, :]\n\n        # Loop through each image in the batch\n        for i, rid in enumerate(rids):\n            # Extract single image mask\n            mask = preds[i]\n            # Encode\n            rle = rle_encode(mask)\n            rle_str = list_to_string(rle)\n            \n            submission_data.append({\n                \"record_id\": rid,\n                \"encoded_pixels\": rle_str\n            })\n\ndf_sub = pd.DataFrame(submission_data)\ndf_sub.to_csv(\"submission.csv\", index=False)\nprint(\"submission.csv saved\")\nprint(df_sub.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T23:40:00.926056Z","iopub.execute_input":"2025-12-07T23:40:00.926653Z","iopub.status.idle":"2025-12-07T23:40:02.012432Z","shell.execute_reply.started":"2025-12-07T23:40:00.926631Z","shell.execute_reply":"2025-12-07T23:40:02.011612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}