{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.12"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"5bed8bc6-887e-43ab-b3bf-52a887477cff","cell_type":"markdown","source":"# AutoDS Contrails Baseline v2\n\nUses pre-trained MaxViT U-Net with d4 TTA for contrail segmentation.\nFollows the base-unet approach from junkoda's notebook.","metadata":{}},{"id":"d37ed805-bd88-44a1-9837-0924dc1dd27c","cell_type":"code","source":"# Install from local wheels (timm 0.9.2 + SMP 0.3.3) to match pre-trained weights\n# The pre-trained MaxViT weights were saved with timm 0.9.2; latest PyPI timm has incompatible key names\n! pip install -q --no-index --find-links=/kaggle/input/contrails-base-model/pip segmentation_models_pytorch timm 2>/dev/null || pip install -q --no-index --find-links=/kaggle/input/datasets/junkoda/contrails-base-model/pip segmentation_models_pytorch timm 2>/dev/null || echo \"WARNING: local wheel install failed\"\n! cp -r /kaggle/input/contrails-base-model/src ./src 2>/dev/null || cp -r /kaggle/input/datasets/junkoda/contrails-base-model/src ./src 2>/dev/null || echo \"WARNING: src copy failed\"","metadata":{},"outputs":[],"execution_count":null},{"id":"3b7bae10-d477-4d66-afdd-b6322bc02b33","cell_type":"markdown","source":"## Imports","metadata":{}},{"id":"cca18104-3621-4cab-b04e-f8003fffbe2e","cell_type":"code","source":"import os\nimport sys\nimport glob\nimport time\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\n\nimport timm\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.base.initialization import initialize_decoder\nfrom segmentation_models_pytorch.base import modules as md\n\nprint(f\"PyTorch: {torch.__version__}\")\nprint(f\"CUDA visible: {torch.cuda.is_available()}\")\n\nSLUG = 'google-research-identify-contrails-reduce-global-warming'\n\n\ndef first_existing_path(candidates, *, required_name):\n    for path in candidates:\n        if os.path.exists(path):\n            print(f\"{required_name}: {path}\")\n            return path\n    raise FileNotFoundError(f\"{required_name} not found. Tried: {candidates}\")\n\n\ndevice = torch.device('cpu')\nprint(f\"Device: {device}\")\n\ndi = first_existing_path(\n    [\n        f'/kaggle/input/competitions/{SLUG}',\n        f'/kaggle/input/{SLUG}',\n    ],\n    required_name='competition input',\n)\nsample_path = first_existing_path(\n    [\n        os.path.join(di, 'sample_submission.csv'),\n        f'/kaggle/input/competitions/{SLUG}/sample_submission.csv',\n        f'/kaggle/input/{SLUG}/sample_submission.csv',\n    ],\n    required_name='sample_submission.csv',\n)","metadata":{},"outputs":[],"execution_count":null},{"id":"3d51e239-7fdc-4d75-adbb-86d65a7b3266","cell_type":"markdown","source":"## Configuration","metadata":{}},{"id":"196c78c7-8ff0-4972-9ad2-782db7d5e795","cell_type":"code","source":"cfg = {\n    'data': {'resize': 512},\n    'model': {\n        'encoder': 'maxvit_tiny_tf_512.in1k',\n        'pretrained': False,\n        'decoder_channels': [256, 128, 64, 32, 16],\n        'dropout': 0.0,\n    }\n}","metadata":{},"outputs":[],"execution_count":null},{"id":"873156c6-40b9-4338-81ed-2e99a6215965","cell_type":"markdown","source":"## Model Architecture","metadata":{}},{"id":"0a50e90d-8bd1-4d9c-a0db-ee3b2b6b49d7","cell_type":"code","source":"class DecoderBlock(nn.Module):\n    def __init__(self, in_channels, skip_channels, out_channels, use_batchnorm=True, dropout=0):\n        super().__init__()\n        conv_in_channels = in_channels + skip_channels\n        self.conv1 = md.Conv2dReLU(conv_in_channels, out_channels, kernel_size=3, padding=1, use_batchnorm=use_batchnorm)\n        self.conv2 = md.Conv2dReLU(out_channels, out_channels, kernel_size=3, padding=1, use_batchnorm=use_batchnorm)\n        self.dropout_skip = nn.Dropout(p=dropout)\n\n    def forward(self, x, skip=None):\n        x = F.interpolate(x, scale_factor=2, mode='nearest')\n        if skip is not None:\n            skip = self.dropout_skip(skip)\n            x = torch.cat([x, skip], dim=1)\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return x\n\n\nclass UnetDecoder(nn.Module):\n    def __init__(self, encoder_channels, decoder_channels, use_batchnorm=True, dropout=0):\n        super().__init__()\n        encoder_channels = encoder_channels[::-1]\n        head_channels = encoder_channels[0]\n        in_channels = [head_channels] + list(decoder_channels[:-1])\n        skip_channels = list(encoder_channels[1:]) + [0]\n        out_channels = decoder_channels\n        self.center = nn.Identity()\n        blocks = [DecoderBlock(in_ch, skip_ch, out_ch, use_batchnorm=use_batchnorm, dropout=dropout)\n                  for in_ch, skip_ch, out_ch in zip(in_channels, skip_channels, out_channels)]\n        self.blocks = nn.ModuleList(blocks)\n\n    def forward(self, features):\n        features = features[::-1]\n        head = features[0]\n        skips = features[1:]\n        x = self.center(head)\n        for i, decoder_block in enumerate(self.blocks):\n            skip = skips[i] if i < len(skips) else None\n            x = decoder_block(x, skip)\n        return x\n\n\ndef get_asym_conv(nc):\n    if nc == 256:\n        hidden_size = 9\n        asym_conv = nn.Sequential(\n            nn.Conv2d(1, hidden_size, kernel_size=(3, 3), padding=1, padding_mode='replicate'),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden_size, 1, kernel_size=1),\n        )\n    elif nc == 512:\n        hidden_size = 25\n        asym_conv = nn.Sequential(\n            nn.Conv2d(1, hidden_size, kernel_size=(5, 5), padding=2, stride=2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden_size, 1, kernel_size=1),\n        )\n    else:\n        raise NotImplementedError\n    return asym_conv\n\n\ndef _check_reduction(reduction_factors):\n    r_prev = 1\n    for r in reduction_factors:\n        if r / r_prev != 2:\n            raise AssertionError('Reduction assumed to increase by 2: {}'.format(reduction_factors))\n        r_prev = r\n\n\n# TTA\n_tta_config = {'d4prob': (8, True), 'd4logit': (8, False),\n               'rotprob': (4, True), 'rotlogit': (4, False), 'none': (1, False)}\n\ndef _tta_stack(x, n):\n    stack = []\n    for k in range(4):\n        xa = torch.rot90(x, k, dims=[2, 3])\n        stack.append(xa)\n        if n == 8:\n            stack.append(torch.flip(xa, dims=[3, ]))\n    return torch.cat(stack, dim=0)\n\n\ndef _tta_average(y_pred, n, prob):\n    batch_size, nch, H, W = y_pred.shape\n    y_pred = y_pred.view(n, batch_size // n, nch, H, W)\n    batch_size = batch_size // n\n    y_avg = torch.zeros((batch_size, 1, H, W), dtype=torch.float32, device=y_pred.device)\n    if prob:\n        y_pred = y_pred.sigmoid()\n    if n == 4:\n        for k in range(4):\n            y_avg += (1 / n) * torch.rot90(y_pred[k], -k, dims=[2, 3])\n    if n == 8:\n        for k in range(4):\n            y_avg += (1 / n) * torch.rot90(y_pred[2 * k], -k, dims=[2, 3])\n            y_avg += (1 / n) * torch.rot90(y_pred[2 * k + 1].flip(dims=[3, ]), -k, dims=(2, 3))\n    if prob:\n        y_avg = y_avg.clamp(1e-6, 1 - 1e-6)\n        return y_avg.logit()\n    return y_avg\n\n\nclass TTA:\n    def __init__(self, tta_str):\n        self.n, self.prob = _tta_config[tta_str]\n    def stack(self, x):\n        if self.n == 1:\n            return x\n        return _tta_stack(x, self.n)\n    def average(self, y):\n        if self.n == 1:\n            return y\n        return _tta_average(y, self.n, self.prob)\n\n\nclass Model(nn.Module):\n    def __init__(self, cfg, pretrained=False, tta=None):\n        super().__init__()\n        name = cfg['model']['encoder']\n        dropout = cfg['model']['dropout']\n        self.encoder = timm.create_model(name, features_only=True, pretrained=pretrained)\n        encoder_channels = self.encoder.feature_info.channels()\n        _check_reduction(self.encoder.feature_info.reduction())\n        decoder_channels = cfg['model']['decoder_channels']\n        self.decoder = UnetDecoder(encoder_channels=encoder_channels, decoder_channels=decoder_channels, dropout=dropout)\n        self.segmentation_head = smp.base.SegmentationHead(\n            in_channels=decoder_channels[-1], out_channels=1, activation=None, kernel_size=3)\n        initialize_decoder(self.decoder)\n        self.asym_conv = get_asym_conv(cfg['data']['resize'])\n        self.tta = TTA(tta) if tta is not None else None\n\n    def forward(self, x):\n        if self.tta is not None:\n            x = self.tta.stack(x)\n        features = self.encoder(x)\n        decoder_output = self.decoder(features)\n        y_sym = self.segmentation_head(decoder_output)\n        if self.tta is not None:\n            y_sym = self.tta.average(y_sym)\n        y_pred = self.asym_conv(y_sym)\n        return y_sym, y_pred\n\nprint(\"Model architecture defined.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"ffcc3a1d-8143-43ec-8895-de6509aee265","cell_type":"markdown","source":"## Data Loading","metadata":{}},{"id":"f305ba27-4a09-478f-aed7-d933985fd8b5","cell_type":"code","source":"def rescale_range(x, f_min, f_max):\n    return (x - f_min) / (f_max - f_min)\n\ndef ash_color(x):\n    \"\"\"False color for contrail annotation. x: (3, H, W) for bands 11, 14, 15\"\"\"\n    r = rescale_range(x[2] - x[1], -4, 2)\n    g = rescale_range(x[1] - x[0], -4, 5)\n    b = rescale_range(x[1], 243, 303)\n    x = torch.stack([r, g, b], axis=0)\n    x = 1 - x\n    return x\n\ndef load_data(path, t=4):\n    \"\"\"Load bands 11, 14, 15 from .npy files at timestep t\"\"\"\n    bands = []\n    for k in [11, 14, 15]:\n        a = np.load(os.path.join(path, f'band_{k:02d}.npy'))\n        bands.append(a)\n    x = np.stack(bands, axis=0)  # (3, H, W, T)\n    return x[:, :, :, t].copy()\n\ndef create_grid(nc, offset=0.5):\n    grid = np.zeros((nc, nc, 2), dtype=np.float32)\n    for ix in range(nc):\n        for iy in range(nc):\n            grid[ix, iy, 1] = -1 + 2 * (ix + 0.5) / nc + offset / 128\n            grid[ix, iy, 0] = -1 + 2 * (iy + 0.5) / nc + offset / 128\n    grid = torch.from_numpy(grid).unsqueeze(0)\n    return grid\n\nprint(\"Data functions defined.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"4ceaad14-0536-4a14-8d01-23bb9a9b9c91","cell_type":"markdown","source":"## RLE Encoding","metadata":{}},{"id":"622098da-d276-4ec7-b380-a113d32eab45","cell_type":"code","source":"def to_str(x):\n    if x:\n        s = ' '.join(str(int(v)) for v in x)\n    else:\n        s = '-'\n    return s\n\ndef encode(x, th=0.5):\n    \"\"\"Encode segmentation mask to RLE\"\"\"\n    dots = np.where(x.T.flatten() > th)[0]\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return to_str(run_lengths)\n\nprint(\"RLE encoding defined.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"3e9f8a1e-ada7-4e95-b638-d96ab5c38dec","cell_type":"markdown","source":"## Load Pre-trained Model","metadata":{}},{"id":"13f2c6e4-6e92-4de8-8496-35642620eed6","cell_type":"code","source":"# Debug: list /kaggle/input to find actual mount paths\nprint(\"Listing /kaggle/input/:\")\nfor entry in sorted(os.listdir('/kaggle/input')):\n    full = os.path.join('/kaggle/input', entry)\n    if os.path.isdir(full):\n        print(f\"  [dir] {entry}\")\n        try:\n            for sub in sorted(os.listdir(full))[:10]:\n                print(f\"    {sub}\")\n        except:\n            pass\n    else:\n        print(f\"  [file] {entry}\")\n\n# Load pre-trained weights - use exact path matching original notebook\nweights_path = None\nweight_candidates = [\n    '/kaggle/input/contrails-base-model/weights/maxvit_tiny_wd/model0.pytorch',\n    '/kaggle/input/datasets/junkoda/contrails-base-model/weights/maxvit_tiny_wd/model0.pytorch',\n]\nfor candidate in weight_candidates:\n    if os.path.exists(candidate):\n        weights_path = candidate\n        break\n\nmodel = None\nif weights_path is not None:\n    # Create model with TTA inside (matching original approach)\n    model = Model(cfg, pretrained=False, tta='d4prob')\n    model.load_state_dict(torch.load(weights_path, map_location='cpu'))\n    print(f\"Loaded weights from {weights_path}\")\n    model.to(device)\n    model.eval()\n    print(\"Model loaded and ready.\")\nelse:\n    print(\"WARNING: pre-trained weights not found; using empty-mask fallback submission.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"0bc4a8fa-63c1-4f6f-9e93-cf5fca0c4987","cell_type":"markdown","source":"## Run Inference","metadata":{}},{"id":"d1545489-6e3a-4d07-9031-2ee8ace558a4","cell_type":"code","source":"test_dir = os.path.join(di, 'test')\ntest_dirs = sorted(glob.glob(os.path.join(test_dir, '*')))\nprint(f\"Found {len(test_dirs)} test samples\")\nsample = pd.read_csv(sample_path, dtype={'record_id': str})\nexpected_ids = sample['record_id'].astype(str).tolist()\nif not test_dirs:\n    raise FileNotFoundError(f\"No test samples found under {test_dir}\")\n\nresize = cfg['data']['resize']\nth = 0.45\npreds = []\n\nif model is None:\n    for file_id in expected_ids:\n        preds.append({'file_id': file_id, 'y_pred': np.zeros((1, 256, 256), dtype=np.float32)})\nelse:\n    test_by_id = {os.path.basename(path): path for path in test_dirs}\n    for file_id in tqdm(expected_ids, desc='Inference'):\n        test_path = test_by_id.get(file_id)\n        if test_path is None:\n            print(f'WARNING: missing test folder for {file_id}; writing empty mask')\n            preds.append({'file_id': file_id, 'y_pred': np.zeros((1, 256, 256), dtype=np.float32)})\n            continue\n\n        try:\n            # Load bands and apply ash color\n            x = load_data(test_path, t=4)\n            x = torch.from_numpy(x)\n            x = ash_color(x)\n\n            # Resize to model input size\n            x = F.interpolate(x.unsqueeze(0), size=(resize, resize), mode='bilinear', align_corners=False)\n            x = x.to(device)\n\n            # Inference with TTA inside model\n            with torch.no_grad():\n                _, y_pred = model(x)\n                y_pred = y_pred.sigmoid().cpu().numpy().squeeze()\n\n            # Resize back to 256x256\n            if y_pred.shape[0] != 256:\n                y_pred = F.interpolate(\n                    torch.from_numpy(y_pred).unsqueeze(0).unsqueeze(0),\n                    size=(256, 256), mode='bilinear', align_corners=False\n                ).numpy().squeeze()\n\n            preds.append({'file_id': file_id, 'y_pred': y_pred.reshape(1, 256, 256)})\n\n        except Exception as e:\n            print(f'Error processing {file_id}: {e}')\n            preds.append({'file_id': file_id, 'y_pred': np.zeros((1, 256, 256), dtype=np.float32)})\n\nprint(f\"Processed {len(preds)} test samples\")","metadata":{},"outputs":[],"execution_count":null},{"id":"5142275b-ae2b-430d-85b1-0ba8717cddb9","cell_type":"markdown","source":"## Generate Submission","metadata":{}},{"id":"22d7fbc3-c87e-438d-8498-2108f0470f82","cell_type":"code","source":"file_ids = []\nencoded_pixels = []\n\nfor pred in preds:\n    file_ids.append(pred['file_id'])\n    encoded = encode(pred['y_pred'], th)\n    encoded_pixels.append(encoded)\n\nsubmission = pd.DataFrame({\n    'record_id': file_ids,\n    'encoded_pixels': encoded_pixels\n})\n\n# Verify format\nsubmission = sample[['record_id']].merge(submission, on='record_id', how='left')\nsubmission['encoded_pixels'] = submission['encoded_pixels'].fillna('-')\nassert list(submission.columns) == list(sample.columns), (submission.columns, sample.columns)\nassert len(submission) == len(sample), (len(submission), len(sample))\nassert submission['record_id'].astype(str).tolist() == sample['record_id'].astype(str).tolist()\nassert submission['encoded_pixels'].notna().all()\n\nsubmission.to_csv('submission.csv', index=False)\nprint(f\"Submission saved: {submission.shape}\")\nprint(submission.head())\n\nprint(f\"\\nSample submission shape: {sample.shape}\")\nprint(f\"Submission shape: {submission.shape}\")\nprint(f\"Columns match: {list(submission.columns) == list(sample.columns)}\")\nprint(f\"Row count match: {len(submission) == len(sample)}\")\n\nprint(\"\\nDone!\")","metadata":{},"outputs":[],"execution_count":null}]}