{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":51753,"databundleVersionId":5692552}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q segmentation-models-pytorch albumentations","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-06T12:04:32.061595Z","iopub.execute_input":"2026-03-06T12:04:32.061834Z","iopub.status.idle":"2026-03-06T12:04:38.629288Z","shell.execute_reply.started":"2026-03-06T12:04:32.061813Z","shell.execute_reply":"2026-03-06T12:04:38.628316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install -q segmentation-models-pytorch albumentations\n\nimport os, gc, cv2, json, random, warnings\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nfrom torch.cuda.amp import GradScaler, autocast\n\nimport segmentation_models_pytorch as smp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport matplotlib.pyplot as plt\n\nwarnings.filterwarnings(\"ignore\")\n\n\n# ── Configuration ─────────────────────────────────────────────────────────────\nclass CFG:\n    BASE_DIR         = \"/kaggle/input/competitions/google-research-identify-contrails-reduce-global-warming\"\n    TRAIN_DIR        = f\"{BASE_DIR}/train\"\n    VALID_DIR        = f\"{BASE_DIR}/validation\"\n    OUTPUT_DIR       = \"/kaggle/working\"\n\n    ARCH             = \"Unet\"               # Unet | UnetPlusPlus | DeepLabV3Plus\n    BACKBONE         = \"efficientnet-b4\"    # efficientnet-b0..b7 | resnet34 | resnet50\n    IN_CHANNELS      = 3\n    NUM_CLASSES      = 1\n\n    TRAIN_SAMPLES    = 4000                 # set None to use full dataset\n    VALID_SAMPLES    = 800                  # set None to use full validation set\n\n    EPOCHS           = 15\n    BATCH_SIZE       = 16\n    VALID_BATCH      = 32\n    LR               = 2e-4\n    WEIGHT_DECAY     = 1e-5\n    IMG_SIZE         = 256\n    NUM_WORKERS      = 2\n    SEED             = 42\n    LABEL_THRESH     = 1                    # min annotators agreeing for a pixel to be contrail\n    USE_AMP          = True\n\n    DEVICE           = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(f\"Device : {CFG.DEVICE}\")\nprint(f\"GPU    : {torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU'}\")\n\n\n# ── Reproducibility ────────────────────────────────────────────────────────────\ndef seed_everything(seed):\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\nseed_everything(CFG.SEED)\n\n\n# ── Preprocessing ──────────────────────────────────────────────────────────────\ndef make_ash_rgb(record_path, frame_idx=4):\n    \"\"\"False-color ash RGB: R=B14-B15, G=B12-B14, B=B14\"\"\"\n    b14 = np.load(f\"{record_path}/band_14.npy\")[:, :, frame_idx].astype(np.float32)\n    b15 = np.load(f\"{record_path}/band_15.npy\")[:, :, frame_idx].astype(np.float32)\n    b12 = np.load(f\"{record_path}/band_12.npy\")[:, :, frame_idx].astype(np.float32)\n\n    def minmax(x):\n        return (x - x.min()) / (x.max() - x.min() + 1e-6)\n\n    return np.stack([minmax(b14 - b15), minmax(b12 - b14), minmax(b14)], axis=-1).astype(np.float32)\n\ndef load_mask(record_path, thresh=CFG.LABEL_THRESH):\n    \"\"\"Aggregate annotator masks: pixel = contrail if >= thresh annotators agreed\"\"\"\n    masks = np.load(f\"{record_path}/human_pixel_masks.npy\")   # H x W x N_annotators\n    return (masks.sum(axis=-1) >= thresh).astype(np.float32)\n\n\n# ── Dataset ────────────────────────────────────────────────────────────────────\nclass ContrailDataset(Dataset):\n    def __init__(self, records, transforms=None):\n        self.records    = records\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, idx):\n        rec   = self.records[idx].rstrip(\"/\")\n        image = make_ash_rgb(rec)   # H x W x 3\n        mask  = load_mask(rec)      # H x W\n\n        if self.transforms:\n            out   = self.transforms(image=image, mask=mask)\n            image = out[\"image\"]\n            mask  = out[\"mask\"]\n\n        if not isinstance(image, torch.Tensor):\n            image = torch.from_numpy(image.transpose(2, 0, 1))\n            mask  = torch.from_numpy(mask)\n\n        # Always ensure mask has channel dim: [H, W] → [1, H, W]\n        if mask.dim() == 2:\n            mask = mask.unsqueeze(0)\n\n        return image, mask.float()\n\n\n# ── Augmentations ──────────────────────────────────────────────────────────────\ndef train_transforms():\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=15, p=0.4),\n        A.OneOf([\n            A.GaussNoise(var_limit=(5.0, 15.0), p=1),\n            A.GaussianBlur(blur_limit=3, p=1),\n            A.RandomBrightnessContrast(p=1),\n        ], p=0.3),\n        A.CoarseDropout(max_holes=8, max_height=16, max_width=16, fill_value=0, mask_fill_value=0, p=0.2),\n        ToTensorV2(transpose_mask=True),\n    ])\n\ndef valid_transforms():\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        ToTensorV2(transpose_mask=True),\n    ])\n\n\n# ── Records & DataLoaders ──────────────────────────────────────────────────────\nall_train = sorted(glob(f\"{CFG.TRAIN_DIR}/*/\"))\nall_valid = sorted(glob(f\"{CFG.VALID_DIR}/*/\"))\n\nrandom.seed(CFG.SEED)\ntrain_records = random.sample(all_train, min(CFG.TRAIN_SAMPLES, len(all_train))) if CFG.TRAIN_SAMPLES else all_train\nvalid_records = random.sample(all_valid, min(CFG.VALID_SAMPLES, len(all_valid))) if CFG.VALID_SAMPLES else all_valid\n\nprint(f\"Train samples : {len(train_records)}\")\nprint(f\"Valid samples : {len(valid_records)}\")\n\ntrain_ds = ContrailDataset(train_records, transforms=train_transforms())\nvalid_ds = ContrailDataset(valid_records, transforms=valid_transforms())\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.BATCH_SIZE, shuffle=True,\n                          num_workers=CFG.NUM_WORKERS, pin_memory=True, drop_last=True)\nvalid_loader = DataLoader(valid_ds, batch_size=CFG.VALID_BATCH, shuffle=False,\n                          num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\n\n# ── Sanity Check: Visualize One Sample ────────────────────────────────────────\ndef visualize_sample(dataset, idx=0):\n    img, mask = dataset[idx]\n    img_np    = img.numpy().transpose(1, 2, 0)\n    mask_np   = mask.numpy().squeeze()\n\n    fig, axes = plt.subplots(1, 2, figsize=(12, 5))\n    axes[0].imshow(img_np); axes[0].set_title(\"Ash RGB\"); axes[0].axis(\"off\")\n    axes[1].imshow(img_np); axes[1].imshow(mask_np, alpha=0.5, cmap=\"Reds\")\n    axes[1].set_title(\"Image + Contrail Mask\"); axes[1].axis(\"off\")\n    plt.tight_layout()\n    plt.savefig(f\"{CFG.OUTPUT_DIR}/sample.png\", dpi=120)\n    plt.show()\n\nvisualize_sample(train_ds, idx=0)\n\n\n# ── Model ──────────────────────────────────────────────────────────────────────\ndef build_model():\n    arch_map = {\n        \"Unet\"         : smp.Unet,\n        \"UnetPlusPlus\" : smp.UnetPlusPlus,\n        \"DeepLabV3Plus\": smp.DeepLabV3Plus,\n    }\n    model = arch_map[CFG.ARCH](\n        encoder_name    = CFG.BACKBONE,\n        encoder_weights = \"imagenet\",\n        in_channels     = CFG.IN_CHANNELS,\n        classes         = CFG.NUM_CLASSES,\n        activation      = None,\n    )\n    return model.to(CFG.DEVICE)\n\nmodel = build_model()\nprint(f\"Model : {CFG.ARCH} + {CFG.BACKBONE} | Params: {sum(p.numel() for p in model.parameters())/1e6:.1f}M\")\n\n\n# ── Loss ───────────────────────────────────────────────────────────────────────\nclass DiceBCELoss(nn.Module):\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n        self.bce    = nn.BCEWithLogitsLoss()\n\n    def forward(self, logits, targets):\n        bce_loss     = self.bce(logits, targets)\n        probs        = torch.sigmoid(logits)\n        intersection = (probs * targets).sum(dim=(2, 3))\n        union        = probs.sum(dim=(2, 3)) + targets.sum(dim=(2, 3))\n        dice_loss    = 1.0 - ((2.0 * intersection + self.smooth) / (union + self.smooth)).mean()\n        return 0.5 * dice_loss + 0.5 * bce_loss\n\ncriterion = DiceBCELoss()\n\n\n# ── Metric: Global Dice (matches competition eval) ─────────────────────────────\nclass GlobalDice:\n    def __init__(self, threshold=0.5):\n        self.threshold    = threshold\n        self.intersection = 0.0\n        self.union        = 0.0\n\n    def reset(self):\n        self.intersection = 0.0\n        self.union        = 0.0\n\n    def update(self, logits, targets):\n        preds              = (torch.sigmoid(logits).detach() > self.threshold).float()\n        self.intersection += (preds * targets).sum().item()\n        self.union        += preds.sum().item() + targets.sum().item()\n\n    def compute(self):\n        return (2.0 * self.intersection + 1e-6) / (self.union + 1e-6)\n\n\n# ── Optimizer & Scheduler ──────────────────────────────────────────────────────\noptimizer = optim.AdamW(model.parameters(), lr=CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.EPOCHS, eta_min=1e-6)\nscaler    = GradScaler(enabled=CFG.USE_AMP)\n\n\n# ── Train & Validate ───────────────────────────────────────────────────────────\ndef train_one_epoch(epoch):\n    model.train()\n    metric     = GlobalDice()\n    total_loss = 0.0\n    pbar       = tqdm(train_loader, desc=f\"[Train] Epoch {epoch+1}/{CFG.EPOCHS}\")\n\n    for step, (images, masks) in enumerate(pbar):\n        images = images.to(CFG.DEVICE, non_blocking=True)\n        masks  = masks.to(CFG.DEVICE, non_blocking=True)\n\n        optimizer.zero_grad()\n        with autocast(enabled=CFG.USE_AMP):\n            logits = model(images)\n            loss   = criterion(logits, masks)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item()\n        metric.update(logits, masks)\n        pbar.set_postfix(loss=f\"{total_loss/(step+1):.4f}\", dice=f\"{metric.compute():.4f}\")\n\n    return total_loss / len(train_loader), metric.compute()\n\n\n@torch.no_grad()\ndef validate():\n    model.eval()\n    metric     = GlobalDice()\n    total_loss = 0.0\n    pbar       = tqdm(valid_loader, desc=\"[Valid]\")\n\n    for images, masks in pbar:\n        images = images.to(CFG.DEVICE, non_blocking=True)\n        masks  = masks.to(CFG.DEVICE, non_blocking=True)\n\n        with autocast(enabled=CFG.USE_AMP):\n            logits = model(images)\n            loss   = criterion(logits, masks)\n\n        total_loss += loss.item()\n        metric.update(logits, masks)\n\n    return total_loss / len(valid_loader), metric.compute()\n\n\n# ── Main Training Loop ─────────────────────────────────────────────────────────\nhistory    = {\"train_loss\": [], \"train_dice\": [], \"val_loss\": [], \"val_dice\": []}\nbest_dice  = 0.0\nmodel_path = f\"{CFG.OUTPUT_DIR}/contrail_model.pth\"\n\nfor epoch in range(CFG.EPOCHS):\n    train_loss, train_dice = train_one_epoch(epoch)\n    val_loss,   val_dice   = validate()\n    scheduler.step()\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"train_dice\"].append(train_dice)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_dice\"].append(val_dice)\n\n    print(f\"Epoch {epoch+1:02d} | Train Loss: {train_loss:.4f} | Train Dice: {train_dice:.4f} | \"\n          f\"Val Loss: {val_loss:.4f} | Val Dice: {val_dice:.4f}\")\n\n    if val_dice > best_dice:\n        best_dice = val_dice\n        torch.save({\n            \"epoch\"      : epoch + 1,\n            \"model_state\": model.state_dict(),\n            \"val_dice\"   : val_dice,\n            \"cfg\": {\n                \"arch\"    : CFG.ARCH,\n                \"backbone\": CFG.BACKBONE,\n                \"img_size\": CFG.IMG_SIZE,\n            }\n        }, model_path)\n        print(f\"  ✅ Best model saved — Val Dice: {best_dice:.4f}\")\n\n    gc.collect()\n    torch.cuda.empty_cache()\n\nprint(f\"\\nTraining complete. Best Val Dice: {best_dice:.4f} → {model_path}\")\n\n\n# ── Training Curves ────────────────────────────────────────────────────────────\ndef plot_history():\n    epochs = range(1, len(history[\"train_loss\"]) + 1)\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\n    axes[0].plot(epochs, history[\"train_loss\"], \"b-o\", label=\"Train\")\n    axes[0].plot(epochs, history[\"val_loss\"],   \"r-o\", label=\"Val\")\n    axes[0].set_title(\"Loss\"); axes[0].set_xlabel(\"Epoch\")\n    axes[0].legend(); axes[0].grid(alpha=0.3)\n\n    axes[1].plot(epochs, history[\"train_dice\"], \"b-o\", label=\"Train\")\n    axes[1].plot(epochs, history[\"val_dice\"],   \"r-o\", label=\"Val\")\n    axes[1].axhline(y=best_dice, color=\"green\", linestyle=\"--\", label=f\"Best: {best_dice:.4f}\")\n    axes[1].set_title(\"Dice Score\"); axes[1].set_xlabel(\"Epoch\")\n    axes[1].legend(); axes[1].grid(alpha=0.3)\n\n    plt.tight_layout()\n    plt.savefig(f\"{CFG.OUTPUT_DIR}/training_curves.png\", dpi=120)\n    plt.show()\n\nplot_history()\n\n\n# ── Visualize Predictions on Validation Set ────────────────────────────────────\ndef predict_and_visualize(n=5, threshold=0.5):\n    checkpoint = torch.load(model_path, map_location=CFG.DEVICE)\n    model.load_state_dict(checkpoint[\"model_state\"])\n    model.eval()\n\n    indices = random.sample(range(len(valid_ds)), n)\n    fig, axes = plt.subplots(n, 3, figsize=(15, 4 * n))\n\n    with torch.no_grad():\n        for row, idx in enumerate(indices):\n            img, mask = valid_ds[idx]\n            logit     = model(img.unsqueeze(0).to(CFG.DEVICE))\n            pred      = (torch.sigmoid(logit) > threshold).squeeze().cpu().numpy()\n            img_np    = img.numpy().transpose(1, 2, 0)\n            mask_np   = mask.squeeze().numpy()\n\n            axes[row, 0].imshow(img_np);                               axes[row, 0].set_title(\"Input\");         axes[row, 0].axis(\"off\")\n            axes[row, 1].imshow(img_np); axes[row, 1].imshow(mask_np, alpha=0.6, cmap=\"Reds\");  axes[row, 1].set_title(\"Ground Truth\"); axes[row, 1].axis(\"off\")\n            axes[row, 2].imshow(img_np); axes[row, 2].imshow(pred,    alpha=0.6, cmap=\"Blues\"); axes[row, 2].set_title(\"Prediction\");    axes[row, 2].axis(\"off\")\n\n    plt.suptitle(\"Validation — Ground Truth vs Prediction\", fontsize=14, fontweight=\"bold\")\n    plt.tight_layout()\n    plt.savefig(f\"{CFG.OUTPUT_DIR}/predictions.png\", dpi=120, bbox_inches=\"tight\")\n    plt.show()\n\npredict_and_visualize(n=5)\n\n\n# ── Save History ───────────────────────────────────────────────────────────────\nwith open(f\"{CFG.OUTPUT_DIR}/history.json\", \"w\") as f:\n    json.dump({\"history\": history, \"best_dice\": best_dice}, f, indent=2)\n\nprint(\"\\nOutput files:\")\nfor fname in os.listdir(CFG.OUTPUT_DIR):\n    size = os.path.getsize(f\"{CFG.OUTPUT_DIR}/{fname}\") / 1e6\n    print(f\"  {fname:40s} {size:.2f} MB\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-06T12:15:04.906564Z","iopub.execute_input":"2026-03-06T12:15:04.907159Z","iopub.status.idle":"2026-03-06T12:36:26.659551Z","shell.execute_reply.started":"2026-03-06T12:15:04.907126Z","shell.execute_reply":"2026-03-06T12:36:26.654451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nbase = \"/kaggle/input/competitions/google-research-identify-contrails-reduce-global-warming\"\nfor item in os.listdir(base):\n    print(item)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-06T12:11:29.707378Z","iopub.execute_input":"2026-03-06T12:11:29.707994Z","iopub.status.idle":"2026-03-06T12:11:29.713087Z","shell.execute_reply.started":"2026-03-06T12:11:29.707967Z","shell.execute_reply":"2026-03-06T12:11:29.712448Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Paste this as a NEW CELL in Kaggle and run it ────────────────────────────\n# Make sure training has finished and contrail_model.pth exists before running\n\nimport torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nfrom tqdm.notebook import tqdm\n\nTHRESHOLD = 0.5\n\n# ── Load best model ───────────────────────────────────────────────────────────\ncheckpoint = torch.load(\"/kaggle/working/contrail_model.pth\", map_location=CFG.DEVICE)\nmodel.load_state_dict(checkpoint[\"model_state\"])\nmodel.eval()\n\n# ── Collect all predictions on validation set ─────────────────────────────────\nall_preds  = []\nall_masks  = []\n\nwith torch.no_grad():\n    for images, masks in tqdm(valid_loader, desc=\"Running Inference\"):\n        images = images.to(CFG.DEVICE)\n        logits = model(images)\n        probs  = torch.sigmoid(logits).cpu().numpy()\n        preds  = (probs > THRESHOLD).astype(np.uint8)\n        masks  = masks.cpu().numpy().astype(np.uint8)\n\n        all_preds.append(preds.flatten())\n        all_masks.append(masks.flatten())\n\nall_preds = np.concatenate(all_preds)\nall_masks = np.concatenate(all_masks)\n\n# ── Compute confusion matrix values (pixel-level) ─────────────────────────────\nTP = int(((all_preds == 1) & (all_masks == 1)).sum())\nTN = int(((all_preds == 0) & (all_masks == 0)).sum())\nFP = int(((all_preds == 1) & (all_masks == 0)).sum())\nFN = int(((all_preds == 0) & (all_masks == 1)).sum())\n\ncm = np.array([[TN, FP],\n               [FN, TP]])\n\n# ── Compute metrics ────────────────────────────────────────────────────────────\nprecision = TP / (TP + FP + 1e-6)\nrecall    = TP / (TP + FN + 1e-6)\nf1        = 2 * precision * recall / (precision + recall + 1e-6)   # same as Dice\niou       = TP / (TP + FP + FN + 1e-6)\naccuracy  = (TP + TN) / (TP + TN + FP + FN + 1e-6)\nspecificity = TN / (TN + FP + 1e-6)\n\nmetrics = {\n    \"Accuracy\"   : accuracy,\n    \"Precision\"  : precision,\n    \"Recall\"     : recall,\n    \"F1 / Dice\"  : f1,\n    \"IoU\"        : iou,\n    \"Specificity\": specificity,\n}\n\n# ── Plot ───────────────────────────────────────────────────────────────────────\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\nfig.suptitle(\"Model Evaluation — Validation Set (Pixel Level)\", fontsize=16, fontweight=\"bold\", y=1.02)\n\n# ── Plot 1: Confusion Matrix ──────────────────────────────────────────────────\nlabels = np.array([[\"TN\", \"FP\"], [\"FN\", \"TP\"]])\ncolors = np.array([[0.18, 0.80, 0.44, 0.85],   # TN — green\n                   [0.90, 0.26, 0.21, 0.85],   # FP — red\n                   [0.95, 0.61, 0.07, 0.85],   # FN — orange\n                   [0.20, 0.60, 0.86, 0.85]])  # TP — blue\n\ncolor_map = np.array([[colors[0], colors[1]],\n                      [colors[2], colors[3]]])\n\nax = axes[0]\nfor i in range(2):\n    for j in range(2):\n        ax.add_patch(plt.Rectangle((j, 1-i), 1, 1,\n                     color=color_map[i][j], ec=\"white\", lw=2))\n        val = cm[i, j]\n        pct = val / cm.sum() * 100\n        ax.text(j + 0.5, 1.5 - i, f\"{labels[i,j]}\\n{val:,}\\n({pct:.1f}%)\",\n                ha=\"center\", va=\"center\", fontsize=13, fontweight=\"bold\", color=\"white\")\n\nax.set_xlim(0, 2); ax.set_ylim(0, 2)\nax.set_xticks([0.5, 1.5]); ax.set_xticklabels([\"Predicted\\nNo Contrail\", \"Predicted\\nContrail\"], fontsize=11)\nax.set_yticks([0.5, 1.5]); ax.set_yticklabels([\"Actual\\nContrail\", \"Actual\\nNo Contrail\"], fontsize=11)\nax.set_title(\"Confusion Matrix\", fontsize=14, fontweight=\"bold\", pad=12)\nax.tick_params(length=0)\n\n# ── Plot 2: Metrics Bar Chart ─────────────────────────────────────────────────\nax2 = axes[1]\nnames  = list(metrics.keys())\nvalues = list(metrics.values())\nbar_colors = [\"#3498db\", \"#2ecc71\", \"#e67e22\", \"#9b59b6\", \"#1abc9c\", \"#e74c3c\"]\n\nbars = ax2.barh(names, values, color=bar_colors, edgecolor=\"white\", height=0.55)\n\nfor bar, val in zip(bars, values):\n    ax2.text(val + 0.005, bar.get_y() + bar.get_height() / 2,\n             f\"{val:.4f}\", va=\"center\", fontsize=12, fontweight=\"bold\")\n\nax2.set_xlim(0, 1.12)\nax2.set_xlabel(\"Score\", fontsize=12)\nax2.set_title(\"Evaluation Metrics\", fontsize=14, fontweight=\"bold\", pad=12)\nax2.axvline(x=0.5, color=\"gray\", linestyle=\"--\", alpha=0.5, label=\"0.5 baseline\")\nax2.grid(axis=\"x\", alpha=0.3)\nax2.invert_yaxis()\n\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/confusion_matrix.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\n# ── Print Summary ──────────────────────────────────────────────────────────────\nprint(\"\\n\" + \"=\"*45)\nprint(\"        EVALUATION SUMMARY\")\nprint(\"=\"*45)\nfor name, val in metrics.items():\n    bar = \"█\" * int(val * 20)\n    print(f\"  {name:<12} {val:.4f}  {bar}\")\nprint(\"=\"*45)\nprint(f\"  TP: {TP:>12,}  (correct contrail pixels)\")\nprint(f\"  TN: {TN:>12,}  (correct background pixels)\")\nprint(f\"  FP: {FP:>12,}  (false alarms)\")\nprint(f\"  FN: {FN:>12,}  (missed contrails)\")\nprint(\"=\"*45)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-06T12:49:00.378814Z","iopub.execute_input":"2026-03-06T12:49:00.379161Z","iopub.status.idle":"2026-03-06T12:49:08.778662Z","shell.execute_reply.started":"2026-03-06T12:49:00.379134Z","shell.execute_reply":"2026-03-06T12:49:08.777765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}