{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14054365,"sourceType":"datasetVersion","datasetId":8946208}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nimport os\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.notebook import tqdm\nfrom pathlib import Path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T14:54:15.206614Z","iopub.execute_input":"2025-12-08T14:54:15.206794Z","iopub.status.idle":"2025-12-08T14:54:21.625121Z","shell.execute_reply.started":"2025-12-08T14:54:15.206778Z","shell.execute_reply":"2025-12-08T14:54:21.624223Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Configuration","metadata":{}},{"cell_type":"code","source":"\n# Weights dataset \nWEIGHTS_PATH = \"/kaggle/input/cap6415-contrail-weights/best_model.pth\" \n# Competition test dataset\nDATA_DIR = Path('/kaggle/input/google-research-identify-contrails-reduce-global-warming')\nTEST_DIR = DATA_DIR / 'test'\nBATCH_SIZE = 4 \nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n# Constants\nT11_BOUNDS = (243, 303)\nCLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\nTDIFF_BOUNDS = (-4, 2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T14:54:32.060821Z","iopub.execute_input":"2025-12-08T14:54:32.061117Z","iopub.status.idle":"2025-12-08T14:54:32.148293Z","shell.execute_reply.started":"2025-12-08T14:54:32.061095Z","shell.execute_reply":"2025-12-08T14:54:32.147574Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Ash color functions","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-08T14:54:39.715152Z","iopub.execute_input":"2025-12-08T14:54:39.715723Z","iopub.status.idle":"2025-12-08T14:54:39.719916Z","shell.execute_reply.started":"2025-12-08T14:54:39.715698Z","shell.execute_reply":"2025-12-08T14:54:39.719116Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Data set","metadata":{}},{"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):\n        self.root = Path(data_dir) \n        self.record_ids = list(record_ids)\n        \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        return torch.from_numpy(rgb), rid","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T14:54:42.866964Z","iopub.execute_input":"2025-12-08T14:54:42.867251Z","iopub.status.idle":"2025-12-08T14:54:42.873436Z","shell.execute_reply.started":"2025-12-08T14:54:42.867229Z","shell.execute_reply":"2025-12-08T14:54:42.872638Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Model","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-08T14:54:46.003249Z","iopub.execute_input":"2025-12-08T14:54:46.00355Z","iopub.status.idle":"2025-12-08T14:54:46.020638Z","shell.execute_reply.started":"2025-12-08T14:54:46.003527Z","shell.execute_reply":"2025-12-08T14:54:46.019868Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### RLE functions for submission","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-08T14:54:55.054675Z","iopub.execute_input":"2025-12-08T14:54:55.055422Z","iopub.status.idle":"2025-12-08T14:54:55.060377Z","shell.execute_reply.started":"2025-12-08T14:54:55.055392Z","shell.execute_reply":"2025-12-08T14:54:55.059657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef run_inference():\n    print(\"Loading Model...\")\n    # Load model\n    model = UNet3D_ConvLSTM(in_channels=3, base_channels=16).to(DEVICE)\n    \n    # Load Weights\n    try:\n        model.load_state_dict(torch.load(WEIGHTS_PATH, map_location=DEVICE))\n        print(\"Weights loaded successfully!\")\n    except Exception as e:\n        print(f\"Error loading weights: {e}\")\n        print(\"Make sure you are pointing to the correct .pth file in your input dataset.\")\n        return\n\n    model.eval()\n    \n    print(\"Preparing Data...\")\n    if not os.path.exists(TEST_DIR):\n        print(\"Test directory not found, skipping.\")\n        return\n\n    test_ids = sorted(os.listdir(TEST_DIR))\n    test_ds = ContrailDataset(TEST_DIR, test_ids)\n    test_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\n    \n    submission_data = []\n    \n    print(\"Starting Prediction...\")\n    with torch.no_grad():\n        for x, rids in tqdm(test_loader):\n            x = x.to(DEVICE)\n            logits = model(x)\n            probs = torch.sigmoid(logits)\n            \n            # Thresholding\n            preds = (probs > 0.5).float().cpu().numpy()[:, 0, :, :]\n            \n            for i, rid in enumerate(rids):\n                mask = preds[i]\n                rle = rle_encode(mask)\n                rle_str = list_to_string(rle)\n                submission_data.append({\"record_id\": rid, \"encoded_pixels\": rle_str})\n                \n    df_sub = pd.DataFrame(submission_data)\n    df_sub.to_csv(\"submission.csv\", index=False)\n    print(\"submission.csv generated successfully!\")\n    print(df_sub.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T14:54:59.797673Z","iopub.execute_input":"2025-12-08T14:54:59.797995Z","iopub.status.idle":"2025-12-08T14:54:59.80534Z","shell.execute_reply.started":"2025-12-08T14:54:59.797971Z","shell.execute_reply":"2025-12-08T14:54:59.804571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nrun_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T14:55:04.363395Z","iopub.execute_input":"2025-12-08T14:55:04.363681Z","iopub.status.idle":"2025-12-08T14:55:06.853132Z","shell.execute_reply.started":"2025-12-08T14:55:04.363659Z","shell.execute_reply":"2025-12-08T14:55:06.852333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}