{"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\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-09T06:20:38.907818Z","iopub.execute_input":"2026-03-09T06:20:38.908158Z","iopub.status.idle":"2026-03-09T06:20:44.8681Z","shell.execute_reply.started":"2026-03-09T06:20:38.908131Z","shell.execute_reply":"2026-03-09T06:20:44.867231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, json, glob\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.losses import DiceLoss, FocalLoss\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\n\n# ── 1.  CONFIG ────────────────────────────────────────────────────────────────\nclass CFG:\n    # paths (Kaggle default)\n    DATA_DIR   = \"/kaggle/input/competitions/google-research-identify-contrails-reduce-global-warming\"\n    TRAIN_DIR  = os.path.join(DATA_DIR, \"train\")\n    VALID_DIR  = os.path.join(DATA_DIR, \"validation\")\n    SAVE_PATH  = \"/kaggle/working/contrail_model.pth\"\n    HIST_PATH  = \"/kaggle/working/history.json\"\n\n    # model\n    ENCODER    = \"efficientnet-b3\"   # pretrained ImageNet backbone\n    IN_CHANS   = 3                   # single Ash-RGB frame (faster, simpler)\n    NUM_CLASSES= 1\n\n    # training\n    IMAGE_SIZE   = 256\n    BATCH_SIZE   = 16\n    EPOCHS       = 15\n    LR           = 1e-4\n    WEIGHT_DECAY = 1e-4\n    TRAIN_SIZE   = 4000   # number of train records to use (matches previous run)\n    VALID_SIZE   = 500    # number of validation records to use\n\n    # inference\n    THRESHOLD  = 0.35                # tuned for sparse contrail pixels\n\n    DEVICE     = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nprint(f\"[CFG] Device: {CFG.DEVICE}  |  Encoder: {CFG.ENCODER}  |  Epochs: {CFG.EPOCHS}\")\n\n\n# ── 2.  ASH-COLOR HELPER ─────────────────────────────────────────────────────\n# Standard meteorological enhancement that makes contrails highly visible.\n# Bands are loaded as raw numpy arrays (float32) from each record folder.\n#   R = Band15 - Band14  →  clamp to [-4,  2]\n#   G = Band14 - Band11  →  clamp to [-4,  5]\n#   B = Band14           →  clamp to [243, 303]\n# Each channel is then normalised to [0, 1].\n\nBAND_RANGES = {\n    \"r\": (-4.0,  2.0),\n    \"g\": (-4.0,  5.0),\n    \"b\": (243.0, 303.0),\n}\n\ndef load_band(record_dir: str, band: str, t: int) -> np.ndarray:\n    \"\"\"Load a single timestamp slice from a band file using memory-mapped IO.\n    Returns float32 array of shape (H, W) for the given timestamp t.\"\"\"\n    path = os.path.join(record_dir, f\"{band}.npy\")\n    arr = np.load(path, mmap_mode='r')               # lazy, no full read\n    return arr[..., t].astype(np.float32)            # (H, W)\n\ndef ash_color(record_dir: str, t: int = 4) -> np.ndarray:\n    \"\"\"\n    Compute the 3-channel Ash-Color image for timestamp t.\n    Returns float32 of shape (H, W, 3) in [0, 1].\n    Only loads 3 bands (fast: 3 file reads per sample).\n    \"\"\"\n    b11 = load_band(record_dir, \"band_11\", t)\n    b14 = load_band(record_dir, \"band_14\", t)\n    b15 = load_band(record_dir, \"band_15\", t)\n\n    R = b15 - b14\n    G = b14 - b11\n    B = b14\n\n    def _norm(x, lo, hi):\n        return np.clip((x - lo) / (hi - lo), 0.0, 1.0)\n\n    R = _norm(R, *BAND_RANGES[\"r\"])\n    G = _norm(G, *BAND_RANGES[\"g\"])\n    B = _norm(B, *BAND_RANGES[\"b\"])\n\n    return np.stack([R, G, B], axis=-1)               # (H, W, 3)\n\n\n# ── 3.  DATASET ───────────────────────────────────────────────────────────────\nclass ContrailDataset(Dataset):\n    \"\"\"\n    Loads Ash-Color images for timestamps (t-1, t, t+1) and stacks them\n    into a 9-channel tensor.  Ground-truth mask = human_pixel_masks.npy.\n    \"\"\"\n\n    def __init__(self, data_dir: str, transform=None):\n        # Only keep subdirectories (each record_id is a folder)\n        all_paths = glob.glob(os.path.join(data_dir, \"*\"))\n        self.records = sorted([p for p in all_paths if os.path.isdir(p)])\n        if len(self.records) == 0:\n            raise RuntimeError(\n                f\"No record folders found in: {data_dir}\\n\"\n                \"Make sure the Kaggle dataset is attached and the path is correct.\"\n            )\n        # Apply size limit so training matches previous run speed\n        if \"train\" in data_dir and CFG.TRAIN_SIZE:\n            self.records = self.records[:CFG.TRAIN_SIZE]\n        elif \"valid\" in data_dir and CFG.VALID_SIZE:\n            self.records = self.records[:CFG.VALID_SIZE]\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.records)\n\n    def __getitem__(self, idx):\n        rec = self.records[idx]\n\n        # --- single Ash image at t=4 (current, labelled timestamp)\n        image = ash_color(rec, t=4)                    # (256,256,3) float32\n\n        # --- ground truth mask\n        mask_path = os.path.join(rec, \"human_pixel_masks.npy\")\n        mask = np.load(mask_path).astype(np.float32)   # (256,256)\n        if mask.ndim == 3:\n            mask = (mask.mean(axis=-1) > 0).astype(np.float32)\n\n        # --- augment\n        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image = augmented[\"image\"]                 # tensor (3,H,W)\n            mask  = augmented[\"mask\"]\n        else:\n            image = torch.from_numpy(image.transpose(2, 0, 1))\n            mask  = torch.from_numpy(mask)\n\n        mask = mask.unsqueeze(0)                       # (1, H, W)\n        return image, mask\n\n\ndef get_transforms(train: bool):\n    if train:\n        return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.Affine(translate_percent=0.05, scale=(0.9, 1.1),\n                     rotate=(-15, 15), p=0.4),\n            A.RandomBrightnessContrast(brightness_limit=0.2,\n                                       contrast_limit=0.2, p=0.4),\n            ToTensorV2(),\n        ])\n    else:\n        return A.Compose([ToTensorV2()])\n\n\n# ── 4.  MODEL ─────────────────────────────────────────────────────────────────\ndef get_model():\n    model = smp.Unet(\n        encoder_name   = CFG.ENCODER,\n        encoder_weights= \"imagenet\",\n        in_channels    = CFG.IN_CHANS,\n        classes        = CFG.NUM_CLASSES,\n        activation     = None,             # raw logits; we apply sigmoid manually\n    )\n    return model.to(CFG.DEVICE)\n\n\n# ── 5.  LOSS ──────────────────────────────────────────────────────────────────\nclass CombinedLoss(nn.Module):\n    \"\"\"0.5 × Dice + 0.5 × Focal.  Handles extreme class imbalance well.\"\"\"\n    def __init__(self):\n        super().__init__()\n        self.dice  = DiceLoss(mode=\"binary\",  from_logits=True, smooth=1.0)\n        self.focal = FocalLoss(mode=\"binary\", gamma=2.0)\n\n    def forward(self, logits, targets):\n        return 0.5 * self.dice(logits, targets) + 0.5 * self.focal(logits, targets)\n\n\n# ── 6.  METRICS ───────────────────────────────────────────────────────────────\ndef dice_score(preds: torch.Tensor, targets: torch.Tensor,\n               threshold: float = CFG.THRESHOLD, eps: float = 1e-6) -> float:\n    preds   = (torch.sigmoid(preds) > threshold).float()\n    inter   = (preds * targets).sum()\n    return (2.0 * inter / (preds.sum() + targets.sum() + eps)).item()\n\n\n# ── 7.  TRAINING LOOP ─────────────────────────────────────────────────────────\ndef train_one_epoch(model, loader, optimizer, criterion):\n    model.train()\n    total_loss, total_dice = 0.0, 0.0\n    for images, masks in tqdm(loader, desc=\"Train\", leave=False):\n        images, masks = images.to(CFG.DEVICE), masks.to(CFG.DEVICE)\n        optimizer.zero_grad()\n        logits = model(images)\n        loss   = criterion(logits, masks)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n        total_dice += dice_score(logits.detach(), masks.detach())\n    n = len(loader)\n    return total_loss / n, total_dice / n\n\n\n@torch.no_grad()\ndef validate(model, loader, criterion):\n    model.eval()\n    total_loss, total_dice = 0.0, 0.0\n    for images, masks in tqdm(loader, desc=\"Valid\", leave=False):\n        images, masks = images.to(CFG.DEVICE), masks.to(CFG.DEVICE)\n        logits = model(images)\n        total_loss += criterion(logits, masks).item()\n        total_dice += dice_score(logits, masks)\n    n = len(loader)\n    return total_loss / n, total_dice / n\n\n\ndef train():\n    # ── datasets & loaders\n    train_ds = ContrailDataset(CFG.TRAIN_DIR, transform=get_transforms(True))\n    valid_ds = ContrailDataset(CFG.VALID_DIR, transform=get_transforms(False))\n    print(f\"[Data] Train: {len(train_ds)} records | Valid: {len(valid_ds)} records\")\n\n    train_loader = DataLoader(train_ds, batch_size=CFG.BATCH_SIZE,\n                              shuffle=True,  num_workers=2, pin_memory=True)\n    valid_loader = DataLoader(valid_ds, batch_size=CFG.BATCH_SIZE,\n                              shuffle=False, num_workers=2, pin_memory=True)\n\n    # ── model, optimiser, scheduler, loss\n    model     = get_model()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.LR,\n                                  weight_decay=CFG.WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n                    optimizer, T_max=CFG.EPOCHS, eta_min=1e-6)\n    criterion = CombinedLoss()\n\n    history   = {\"train_loss\": [], \"train_dice\": [],\n                 \"val_loss\":   [], \"val_dice\":   []}\n    best_dice = 0.0\n\n    print(f\"\\n{'Epoch':>6} {'TrainLoss':>10} {'TrainDice':>10} \"\n          f\"{'ValLoss':>10} {'ValDice':>10} {'LR':>10}\")\n    print(\"─\" * 60)\n\n    for epoch in range(1, CFG.EPOCHS + 1):\n        tr_loss, tr_dice = train_one_epoch(model, train_loader, optimizer, criterion)\n        vl_loss, vl_dice = validate(model, valid_loader, criterion)\n        scheduler.step()\n\n        history[\"train_loss\"].append(tr_loss)\n        history[\"train_dice\"].append(tr_dice)\n        history[\"val_loss\"].append(vl_loss)\n        history[\"val_dice\"].append(vl_dice)\n\n        lr = scheduler.get_last_lr()[0]\n        print(f\"{epoch:>6} {tr_loss:>10.4f} {tr_dice:>10.4f} \"\n              f\"{vl_loss:>10.4f} {vl_dice:>10.4f} {lr:>10.2e}\")\n\n        if vl_dice > best_dice:\n            best_dice = vl_dice\n            torch.save({\"model_state\": model.state_dict(),\n                        \"epoch\":       epoch,\n                        \"best_dice\":   best_dice,\n                        \"cfg\":         {k: str(v) for k, v in vars(CFG).items()\n                                        if not k.startswith(\"_\")}},\n                       CFG.SAVE_PATH)\n            print(f\"  ✓ Saved best model  (val Dice = {best_dice:.4f})\")\n\n    # save history\n    with open(CFG.HIST_PATH, \"w\") as f:\n        json.dump({\"history\": history, \"best_dice\": best_dice}, f, indent=2)\n\n    print(f\"\\n[Done] Best Val Dice = {best_dice:.4f}\")\n    plot_training_curves(history, best_dice)\n    return model, valid_loader\n\n\n# ── 8.  TRAINING CURVES ───────────────────────────────────────────────────────\ndef plot_training_curves(history: dict, best_dice: float):\n    epochs = range(1, len(history[\"train_loss\"]) + 1)\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))\n    fig.suptitle(\"Training History\", fontsize=14, fontweight=\"bold\")\n\n    ax1.plot(epochs, history[\"train_loss\"], \"b-o\", ms=4, label=\"Train\")\n    ax1.plot(epochs, history[\"val_loss\"],   \"r-o\", ms=4, label=\"Val\")\n    ax1.set_title(\"Loss (Dice + Focal)\")\n    ax1.set_xlabel(\"Epoch\"); ax1.set_ylabel(\"Loss\")\n    ax1.legend(); ax1.grid(alpha=0.3)\n\n    ax2.plot(epochs, history[\"train_dice\"], \"b-o\", ms=4, label=\"Train\")\n    ax2.plot(epochs, history[\"val_dice\"],   \"r-o\", ms=4, label=\"Val\")\n    ax2.axhline(best_dice, color=\"green\", ls=\"--\",\n                label=f\"Best: {best_dice:.4f}\")\n    ax2.set_title(\"Dice Score\")\n    ax2.set_xlabel(\"Epoch\"); ax2.set_ylabel(\"Dice\")\n    ax2.legend(); ax2.grid(alpha=0.3)\n\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/training_curves.png\", dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(\"Saved → training_curves.png\")\n\n\n# ── 9.  EVALUATION / INFERENCE ───────────────────────────────────────────────\n@torch.no_grad()\ndef evaluate(model=None, valid_loader=None):\n    \"\"\"\n    Load best saved weights, run inference on validation set, compute\n    pixel-level metrics, and visualise results.\n    Can be run standalone after training by calling evaluate() with no args.\n    \"\"\"\n    # ── rebuild loader if not provided\n    if valid_loader is None:\n        valid_ds = ContrailDataset(CFG.VALID_DIR, transform=get_transforms(False))\n        valid_loader = DataLoader(valid_ds, batch_size=CFG.BATCH_SIZE,\n                                  shuffle=False, num_workers=2, pin_memory=True)\n\n    # ── load best model weights\n    if model is None:\n        model = get_model()\n\n    checkpoint = torch.load(CFG.SAVE_PATH, map_location=CFG.DEVICE)\n    model.load_state_dict(checkpoint[\"model_state\"])\n    print(f\"[Eval] Loaded best model from epoch {checkpoint['epoch']}  \"\n          f\"(saved Dice = {checkpoint['best_dice']:.4f})\")\n    model.eval()\n\n    # ── collect predictions\n    all_preds, all_masks, sample_batches = [], [], []\n    for i, (images, masks) in enumerate(tqdm(valid_loader, desc=\"Inference\")):\n        images = images.to(CFG.DEVICE)\n        logits = model(images)\n        probs  = torch.sigmoid(logits).cpu()\n        preds  = (probs > CFG.THRESHOLD).float()\n\n        all_preds.append(preds.flatten().numpy())\n        all_masks.append(masks.flatten().numpy())\n\n        if i < 5:   # store up to 5 batches for visualisation\n            sample_batches.append((images.cpu(), masks.cpu(),\n                                   probs.cpu(), preds.cpu()))\n\n    all_preds = np.concatenate(all_preds).astype(np.uint8)\n    all_masks = np.concatenate(all_masks).astype(np.uint8)\n\n    # ── pixel-level confusion matrix values\n    TP = int(((all_preds == 1) & (all_masks == 1)).sum())\n    TN = int(((all_preds == 0) & (all_masks == 0)).sum())\n    FP = int(((all_preds == 1) & (all_masks == 0)).sum())\n    FN = int(((all_preds == 0) & (all_masks == 1)).sum())\n    EPS = 1e-6\n\n    metrics = {\n        \"Accuracy\"   : (TP + TN) / (TP + TN + FP + FN + EPS),\n        \"Precision\"  : TP / (TP + FP + EPS),\n        \"Recall\"     : TP / (TP + FN + EPS),\n        \"F1 / Dice\"  : 2 * TP / (2 * TP + FP + FN + EPS),\n        \"IoU\"        : TP / (TP + FP + FN + EPS),\n        \"Specificity\": TN / (TN + FP + EPS),\n    }\n\n    # ── print summary\n    print(\"\\n\" + \"=\" * 50)\n    print(\"         EVALUATION SUMMARY\")\n    print(\"=\" * 50)\n    for name, val in metrics.items():\n        bar = \"█\" * int(val * 25)\n        print(f\"  {name:<12} {val:.4f}  {bar}\")\n    print(\"=\" * 50)\n    print(f\"  TP: {TP:>12,}   (correct contrail pixels)\")\n    print(f\"  TN: {TN:>12,}   (correct background pixels)\")\n    print(f\"  FP: {FP:>12,}   (false alarms)\")\n    print(f\"  FN: {FN:>12,}   (missed contrails)\")\n    print(\"=\" * 50)\n\n    # ── plot confusion matrix + metrics\n    _plot_confusion_and_metrics(TP, TN, FP, FN, metrics)\n\n    # ── plot sample predictions\n    _plot_sample_predictions(sample_batches)\n\n    return metrics\n\n\ndef _plot_confusion_and_metrics(TP, TN, FP, FN, metrics):\n    cm = np.array([[TN, FP], [FN, TP]])\n    labels  = np.array([[\"TN\", \"FP\"], [\"FN\", \"TP\"]])\n    colours = [\"#2ecc71\", \"#e74c3c\", \"#e67e22\", \"#3498db\"]   # TN FP FN TP\n\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n    fig.suptitle(\"Model Evaluation — Validation Set (Pixel Level)\",\n                 fontsize=15, fontweight=\"bold\", y=1.02)\n\n    # confusion matrix\n    flat_colours = [colours[0], colours[1], colours[2], colours[3]]\n    for idx, (i, j) in enumerate([(0,0),(0,1),(1,0),(1,1)]):\n        ax1.add_patch(plt.Rectangle((j, 1-i), 1, 1,\n                      color=flat_colours[idx], ec=\"white\", lw=2))\n        val = cm[i, j]\n        pct = val / cm.sum() * 100\n        ax1.text(j+0.5, 1.5-i,\n                 f\"{labels[i,j]}\\n{val:,}\\n({pct:.1f}%)\",\n                 ha=\"center\", va=\"center\",\n                 fontsize=12, fontweight=\"bold\", color=\"white\")\n\n    ax1.set_xlim(0, 2); ax1.set_ylim(0, 2)\n    ax1.set_xticks([0.5, 1.5])\n    ax1.set_xticklabels([\"Predicted\\nNo Contrail\", \"Predicted\\nContrail\"], fontsize=11)\n    ax1.set_yticks([0.5, 1.5])\n    ax1.set_yticklabels([\"Actual\\nContrail\", \"Actual\\nNo Contrail\"], fontsize=11)\n    ax1.set_title(\"Confusion Matrix\", fontsize=13, fontweight=\"bold\", pad=12)\n    ax1.tick_params(length=0)\n\n    # metrics bar chart\n    names  = list(metrics.keys())\n    values = list(metrics.values())\n    bar_colors = [\"#3498db\",\"#2ecc71\",\"#e67e22\",\"#9b59b6\",\"#1abc9c\",\"#e74c3c\"]\n    bars = ax2.barh(names, values, color=bar_colors, edgecolor=\"white\", height=0.55)\n    for 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=11, fontweight=\"bold\")\n    ax2.set_xlim(0, 1.15)\n    ax2.set_xlabel(\"Score\", fontsize=12)\n    ax2.set_title(\"Evaluation Metrics\", fontsize=13, fontweight=\"bold\", pad=12)\n    ax2.axvline(0.5, color=\"gray\", ls=\"--\", alpha=0.5, label=\"0.5 baseline\")\n    ax2.grid(axis=\"x\", alpha=0.3)\n    ax2.invert_yaxis()\n\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/confusion_matrix.png\", dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(\"Saved → confusion_matrix.png\")\n\n\ndef _plot_sample_predictions(sample_batches):\n    \"\"\"Show up to 5 sample prediction comparisons.\"\"\"\n    n = min(5, len(sample_batches))\n    fig, axes = plt.subplots(n, 3, figsize=(12, 4 * n))\n    if n == 1:\n        axes = [axes]\n    fig.suptitle(\"Validation — Ground Truth vs Prediction\",\n                 fontsize=14, fontweight=\"bold\")\n\n    for row, (images, masks, probs, preds) in enumerate(sample_batches[:n]):\n        img = images[0]          # (9, H, W) – use middle Ash frame (channels 3-5)\n        ash = img[0:3].permute(1, 2, 0).numpy()   # (H, W, 3) – ash color\n        gt  = masks[0, 0].numpy()                  # (H, W)\n        pr  = probs[0, 0].numpy()                  # (H, W) probability map\n\n        ax_img, ax_gt, ax_pr = axes[row]\n\n        ax_img.imshow(ash)\n        ax_img.set_title(\"Input (Ash RGB, t=0)\")\n        ax_img.axis(\"off\")\n\n        ax_gt.imshow(ash, alpha=0.7)\n        ax_gt.imshow(gt,  cmap=\"Reds\", alpha=0.5, vmin=0, vmax=1)\n        ax_gt.set_title(\"Ground Truth\")\n        ax_gt.axis(\"off\")\n\n        ax_pr.imshow(ash, alpha=0.7)\n        ax_pr.imshow(pr,  cmap=\"Blues\", alpha=0.5, vmin=0, vmax=1)\n        ax_pr.set_title(f\"Prediction (thr={CFG.THRESHOLD})\")\n        ax_pr.axis(\"off\")\n\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/predictions.png\", dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(\"Saved → predictions.png\")\n\n\n# ── 10.  ENTRY POINT ─────────────────────────────────────────────────────────\nif __name__ == \"__main__\":\n    # Run full training then evaluation\n    print(\"=\" * 60)\n    print(\"  CONTRAIL DETECTION — Training Pipeline\")\n    print(f\"  Encoder : {CFG.ENCODER}\")\n    print(f\"  Device  : {CFG.DEVICE}\")\n    print(f\"  Epochs  : {CFG.EPOCHS}  |  Batch: {CFG.BATCH_SIZE}\")\n    print(f\"  Loss    : 0.5×Dice + 0.5×Focal\")\n    print(\"=\" * 60 + \"\\n\")\n\n    trained_model, val_loader = train()\n    evaluate(trained_model, val_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-09T06:30:02.119489Z","iopub.execute_input":"2026-03-09T06:30:02.11982Z","iopub.status.idle":"2026-03-09T06:58:05.410533Z","shell.execute_reply.started":"2026-03-09T06:30:02.119794Z","shell.execute_reply":"2026-03-09T06:58:05.409838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}