{"metadata":{"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"name":""},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":3283978,"datasetId":1988734,"databundleVersionId":3334621},{"sourceType":"kernelVersion","sourceId":301267938,"isSourceIdPinned":false}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"title","cell_type":"markdown","source":"# Option A — Fast Path: Train 3 Backbones + Segmentation\n**Self-contained. No dependency on the main notebook.**\n\nTrains only Baseline / G1 / PNG backbones (skips G2–G5), saves checkpoints after each,\nthen immediately runs ADE20K segmentation.\n\n- Images fed to ViT: **224×224** (model hard-requires this)\n- Masks / logits upsampled to: **512×512**","metadata":{}},{"id":"s0","cell_type":"markdown","source":"## 0. Install Requirements","metadata":{}},{"id":"wfdju6hajhj","cell_type":"code","source":"import subprocess, sys\n\ndef pip(*args):\n    subprocess.check_call([sys.executable, '-m', 'pip', 'install', '-q', *args])\n\n# On Kaggle, torch with CUDA is pre-installed — skip torch reinstall.\n# Only install missing packages.\npip('timm==1.0.25')\npip('kagglehub')\npip('tqdm', 'Pillow', 'numpy', 'matplotlib', 'scipy')\n\nimport torch\nprint(f'torch {torch.__version__} | CUDA available: {torch.cuda.is_available()}')\nif torch.cuda.is_available():\n    print(f'GPU: {torch.cuda.get_device_name(0)}')\nelse:\n    print('WARNING: No CUDA detected — make sure the Kaggle GPU accelerator is enabled.')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"setup","cell_type":"code","source":"import os, types, random\nfrom pathlib import Path\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\nfrom timm.layers import Attention\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nMEAN, STD = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)\n\n# ── Fetch data & checkpoints via kagglehub ────────────────────────────────────\n# On Kaggle, credentials are built-in — no need to set KAGGLE_USERNAME/KEY.\nimport kagglehub\n\nprint('Downloading ADE20K ...')\n_ade_root = Path(kagglehub.dataset_download('ipythonx/ade20k-scene-parsing'))\nprint(f'  ADE20K root: {_ade_root}')\n\nprint('Downloading backbone checkpoints ...')\n_ckpt_root = Path(kagglehub.notebook_output_download('shashwatchaturvedi35/dl-project'))\nprint(f'  Checkpoint root: {_ckpt_root}')\n\n# ── Resolve checkpoint path ───────────────────────────────────────────────────\n_ckpt_candidates = list(_ckpt_root.rglob('checkpoints_B'))\nCKPT_IN = _ckpt_candidates[0] if _ckpt_candidates else _ckpt_root\nprint(f'  Checkpoints : {CKPT_IN}')\nprint(f'  .pt files   : {[p.name for p in CKPT_IN.glob(\"*.pt\")]}')\n\n# ── Resolve ADE20K path ───────────────────────────────────────────────────────\n_ade_candidates = list(_ade_root.rglob('ADEChallengeData2016'))\nADE_BASE = _ade_candidates[0] if _ade_candidates else _ade_root\nprint(f'  ADE20K base : {ADE_BASE}')\n\nCKPT_OUT = Path('./checkpoints_seg')\nCKPT_OUT.mkdir(exist_ok=True)\n\nprint(f'\\ntorch {torch.__version__} | timm {timm.__version__} | device {DEVICE}')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s1","cell_type":"markdown","source":"## 1. Paths & constants","metadata":{}},{"id":"paths","cell_type":"code","source":"# IN1K_BASE is only defined when ImageNet was downloaded (retraining path).\n# Default to None — the train_cls cell will skip training if checkpoints exist.\nIN1K_BASE = globals().get('IN1K_BASE', None)\nIN1K_TRAIN = IN1K_BASE / 'train' if IN1K_BASE is not None else None\n\nADE_TRAIN_IMG = ADE_BASE / 'images/training'\nADE_VAL_IMG   = ADE_BASE / 'images/validation'\nADE_TRAIN_ANN = ADE_BASE / 'annotations/training'\nADE_VAL_ANN   = ADE_BASE / 'annotations/validation'\n\nSLICE       = 0.1\nEPOCHS_CLS  = 2\nEPOCHS_SEG  = 10\nVIT_SIZE    = 224\nMASK_SIZE   = 512\nADE_CLASSES = 150\n\nfor p in [ADE_TRAIN_IMG, ADE_VAL_IMG, ADE_VAL_ANN]:\n    print(f'  {p.name:<40} {\"OK\" if p.exists() else \"MISSING — check path\"}')\nif IN1K_TRAIN is not None:\n    print(f'  {\"IN1K_TRAIN\":<40} {\"OK\" if IN1K_TRAIN.exists() else \"MISSING\"}')\nelse:\n    print(f'  {\"IN1K_TRAIN\":<40} SKIPPED (checkpoints will be loaded)')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s2","cell_type":"markdown","source":"## 2. Gate code","metadata":{}},{"id":"gates","cell_type":"code","source":"class GateParams(nn.Module):\n    def __init__(self, dim, num_heads, png=False):\n        super().__init__()\n        self.W = nn.Parameter(torch.empty(dim, num_heads).normal_(std=0.02))\n        self.png = png\n        if png:\n            self.e    = nn.Parameter(torch.ones(num_heads))\n            self.beta = nn.Parameter(torch.zeros(1))\n\n    def gate(self, x):\n        g = x @ self.W\n        if self.png:\n            g = g - self.beta * (x.norm(dim=-1, keepdim=True) + 1e-6) * self.e\n        return torch.sigmoid(g)\n\n\ndef _patched_forward(mod, params, pos):\n    def forward(self, x, attn_mask=None):\n        B, N, C = x.shape\n        H, D = self.num_heads, self.head_dim\n        qkv = self.qkv(x).reshape(B, N, 3, H, D).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv.unbind(0)\n        q, k = self.q_norm(q), self.k_norm(k)\n        if pos == 'G3': k = k * params.gate(x).permute(0,2,1).unsqueeze(-1)\n        if pos == 'G4': q = q * params.gate(x).permute(0,2,1).unsqueeze(-1)\n        if pos == 'G2': v = v * params.gate(x).permute(0,2,1).unsqueeze(-1)\n        out = F.scaled_dot_product_attention(\n            q, k, v, attn_mask=attn_mask,\n            dropout_p=self.attn_drop.p if self.training else 0.)\n        if pos in ('G1', 'PNG'):\n            G = params.gate(x)\n            out = (out.transpose(1,2) * G.unsqueeze(-1)).transpose(1,2)\n            if getattr(self, '_capture', False):\n                self._last_gate   = G.detach()\n                self._last_x_norm = (x.norm(dim=-1, keepdim=True) + 1e-6).detach()\n        x = out.transpose(1,2).reshape(B, N, C)\n        if pos == 'G5':\n            x = (x.view(B,N,H,D) * params.gate(x).unsqueeze(-1)).view(B,N,C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n    return types.MethodType(forward, mod)\n\n\ndef inject_gates(model, pos='G1', png=False):\n    gate_list = []\n    p_label = 'PNG' if png else pos\n    for _, mod in model.named_modules():\n        if isinstance(mod, Attention):\n            p = GateParams(mod.qkv.in_features, mod.num_heads, png=png)\n            mod.forward = _patched_forward(mod, p, p_label)\n            gate_list.append(p)\n    model.gate_params = nn.ModuleList(gate_list)\n    return model\n\n\ndef build_vit(pos='baseline', png=False, pretrained=True):\n    \"\"\"Build ViT. Pass pretrained=False when loading from a local checkpoint\n    to avoid an unnecessary (and potentially failing) network download.\"\"\"\n    m = timm.create_model('vit_base_patch16_224', pretrained=pretrained, num_classes=1000)\n    if pos != 'baseline':\n        m = inject_gates(m, pos=pos, png=png)\n    for p in m.parameters():          p.requires_grad_(False)\n    for p in m.head.parameters():     p.requires_grad_(True)\n    if hasattr(m, 'gate_params'):\n        for p in m.gate_params.parameters(): p.requires_grad_(True)\n    return m.to(DEVICE)\n\nprint('Gate code ready.')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s3","cell_type":"markdown","source":"## 3. ImageNet-1K dataloader (for backbone fine-tuning)","metadata":{}},{"id":"in1k_loader","cell_type":"code","source":"train_tfm = T.Compose([\n    T.RandomResizedCrop(224), T.RandomHorizontalFlip(),\n    T.PILToTensor(),\n    T.ConvertImageDtype(torch.float32),\n    T.Normalize(MEAN, STD),\n])\nval_tfm = T.Compose([\n    T.Resize(256), T.CenterCrop(224),\n    T.PILToTensor(),\n    T.ConvertImageDtype(torch.float32),\n    T.Normalize(MEAN, STD),\n])\n\nclass ImageSamples(Dataset):\n    def __init__(self, samples, tfm):\n        self.samples, self.tfm = samples, tfm\n    def __len__(self): return len(self.samples)\n    def __getitem__(self, i):\n        path, label = self.samples[i]\n        return self.tfm(Image.open(path).convert('RGB')), label\n\ndef build_imagenet_samples(root, frac):\n    class_dirs   = sorted(p for p in Path(root).iterdir() if p.is_dir())\n    class_to_idx = {d.name: i for i, d in enumerate(class_dirs)}\n    samples = []\n    for d in class_dirs:\n        files = sorted(d.glob('*.JPEG')) + sorted(d.glob('*.jpg'))\n        keep  = max(1, int(len(files) * frac))\n        for f in files[:keep]:\n            samples.append((str(f), class_to_idx[d.name]))\n    return samples\n\nif IN1K_TRAIN is not None and IN1K_TRAIN.exists():\n    print('Scanning ImageNet-1K train ...')\n    all_samples = build_imagenet_samples(IN1K_TRAIN, SLICE)\n    random.shuffle(all_samples)\n    split      = int(len(all_samples) * 0.9)\n    train_samp = all_samples[:split]\n    val_samp   = all_samples[split:]\n    in1k_train_loader = DataLoader(ImageSamples(train_samp, train_tfm), 64,\n                                   shuffle=True,  num_workers=2, pin_memory=True)\n    in1k_val_loader   = DataLoader(ImageSamples(val_samp,   val_tfm),   64,\n                                   shuffle=False, num_workers=2, pin_memory=True)\n    print(f'ImageNet-1K  train {len(train_samp):,} | val {len(val_samp):,}  ({SLICE*100:.0f}% slice)')\nelse:\n    in1k_train_loader = None\n    in1k_val_loader   = None\n    print('ImageNet-1K loaders SKIPPED — checkpoints will be loaded from disk.')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s4","cell_type":"markdown","source":"## 4. ADE20K dataloader\nImages resized to **224×224** for ViT; masks kept at **512×512**.","metadata":{}},{"id":"ade_loader","cell_type":"code","source":"def pil_to_tensor(img):\n    \"\"\"Convert PIL RGB image to float tensor without numpy.\"\"\"\n    import struct\n    w, h = img.size\n    buf = img.tobytes()                                      # raw bytes\n    t = torch.frombuffer(bytearray(buf), dtype=torch.uint8) # bytearray avoids numpy\n    t = t.reshape(h, w, 3).permute(2, 0, 1).float() / 255.0\n    return t\n\ndef pil_mask_to_tensor(mask_pil):\n    \"\"\"Convert PIL palette/L mask to int64 tensor without numpy.\"\"\"\n    mask_pil = mask_pil.convert('I')                        # int32 mode\n    buf = mask_pil.tobytes()\n    t = torch.frombuffer(bytearray(buf), dtype=torch.int32)\n    return t.reshape(mask_pil.size[1], mask_pil.size[0]).long()\n\nMEAN_T = torch.tensor(MEAN).view(3, 1, 1)\nSTD_T  = torch.tensor(STD).view(3, 1, 1)\n\nclass ADE20KDataset(Dataset):\n    def __init__(self, img_dir, ann_dir, frac=1.0):\n        imgs = sorted(img_dir.glob('*.jpg'))\n        anns = sorted(ann_dir.glob('*.png'))\n        n = max(1, int(len(imgs) * frac))\n        self.imgs, self.anns = imgs[:n], anns[:n]\n\n    def __len__(self): return len(self.imgs)\n\n    def __getitem__(self, i):\n        img = Image.open(self.imgs[i]).convert('RGB').resize(\n            (VIT_SIZE, VIT_SIZE), Image.BILINEAR)\n        img = (pil_to_tensor(img) - MEAN_T) / STD_T\n\n        mask = Image.open(self.anns[i]).resize(\n            (MASK_SIZE, MASK_SIZE), Image.NEAREST)\n        mask = pil_mask_to_tensor(mask) - 1   # 1-150 → 0-149, 0 → -1 (ignore)\n        return img, mask\n\nade_train = ADE20KDataset(ADE_TRAIN_IMG, ADE_TRAIN_ANN, frac=SLICE)\nade_val   = ADE20KDataset(ADE_VAL_IMG,   ADE_VAL_ANN,   frac=SLICE)\nade_train_loader = DataLoader(ade_train, 16, shuffle=True,  num_workers=0, pin_memory=True)\nade_val_loader   = DataLoader(ade_val,   16, shuffle=False, num_workers=0, pin_memory=True)\nprint(f'ADE20K  train {len(ade_train):,} | val {len(ade_val):,}  ({SLICE*100:.0f}% slice)')\nprint(f'ViT input: {VIT_SIZE}×{VIT_SIZE}  |  Mask size: {MASK_SIZE}×{MASK_SIZE}')\n\n# Quick sanity check\n_img, _mask = ade_train[0]\nprint(f'img shape={tuple(_img.shape)} dtype={_img.dtype}  '\n      f'mask shape={tuple(_mask.shape)} dtype={_mask.dtype}  '\n      f'mask range=[{_mask.min()},{_mask.max()}]')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s5","cell_type":"markdown","source":"## 5. Train 3 backbones — with checkpoint saving after each\nIf a checkpoint already exists it is loaded instead of re-training.","metadata":{}},{"id":"train_cls","cell_type":"code","source":"@torch.no_grad()\ndef evaluate_cls(model, loader, num_classes=1000):\n    \"\"\"Returns acc, macro-F1, macro-precision, macro-recall.\"\"\"\n    model.eval()\n    all_preds, all_labels = [], []\n    for x, y in loader:\n        x, y = x.to(DEVICE), y.to(DEVICE)\n        all_preds.append(model(x).argmax(1).cpu())\n        all_labels.append(y.cpu())\n    preds  = torch.cat(all_preds)\n    labels = torch.cat(all_labels)\n\n    acc = (preds == labels).float().mean().item()\n\n    tp = torch.zeros(num_classes)\n    fp = torch.zeros(num_classes)\n    fn = torch.zeros(num_classes)\n    for c in range(num_classes):\n        pred_c = preds  == c\n        true_c = labels == c\n        tp[c]  = (pred_c & true_c).sum()\n        fp[c]  = (pred_c & ~true_c).sum()\n        fn[c]  = (~pred_c & true_c).sum()\n\n    present   = (tp + fn) > 0\n    prec_c    = tp / (tp + fp + 1e-8)\n    rec_c     = tp / (tp + fn + 1e-8)\n    f1_c      = 2 * prec_c * rec_c / (prec_c + rec_c + 1e-8)\n    precision = prec_c[present].mean().item()\n    recall    = rec_c[present].mean().item()\n    f1        = f1_c[present].mean().item()\n    return acc, f1, precision, recall\n\n\ndef train_cls(model, epochs, name, ckpt_path=None):\n    if in1k_train_loader is None:\n        raise RuntimeError('ImageNet loaders not available — cannot train. '\n                           'Make sure IN1K_BASE is set and the dataset is downloaded.')\n    crit  = nn.CrossEntropyLoss()\n    opt   = torch.optim.AdamW(\n        [p for p in model.parameters() if p.requires_grad], lr=1e-3, weight_decay=0.05)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    log   = {'loss': [], 'acc': [], 'f1': [], 'precision': [], 'recall': []}\n    for ep in range(1, epochs+1):\n        model.train()\n        for n, mod in model.named_modules():\n            if 'gate_params' not in n and 'head' not in n and isinstance(mod, nn.LayerNorm):\n                mod.eval()\n        total_loss = 0\n        for x, y in tqdm(in1k_train_loader, desc=f'{name} ep{ep}', leave=False):\n            x, y = x.to(DEVICE), y.to(DEVICE)\n            opt.zero_grad()\n            loss = nn.CrossEntropyLoss()(model(x), y)\n            loss.backward(); opt.step()\n            total_loss += loss.item()\n        acc, f1, prec, rec = evaluate_cls(model, in1k_val_loader)\n        sched.step()\n        log['loss'].append(total_loss / len(in1k_train_loader))\n        log['acc'].append(acc); log['f1'].append(f1)\n        log['precision'].append(prec); log['recall'].append(rec)\n        print(f'  {name} ep{ep}  loss={log[\"loss\"][-1]:.4f}  acc={acc:.4f}  '\n              f'f1={f1:.4f}  prec={prec:.4f}  rec={rec:.4f}')\n        if ckpt_path is not None:\n            torch.save(model.state_dict(), ckpt_path)\n    return log\n\n\n# (pos, png, checkpoint filename in CKPT_IN, display name)\nbackbone_configs = [\n    ('baseline', False, 'vit_baseline.pt', 'Baseline'),\n    ('G1',       False, 'vit_g1.pt',       'G1 Gate'),\n    ('G1',       True,  'vit_png.pt',       'PNG Gate'),\n]\n\ncls_results = {}\nfor pos, png, ckpt_name, name in backbone_configs:\n    ckpt_path = CKPT_IN / ckpt_name\n    print(f'\\n── {name} ──')\n    if ckpt_path.exists():\n        # Load from read-only input — no network needed\n        print(f'  Loading backbone: {ckpt_path}')\n        model = build_vit(pos, png, pretrained=False)\n        model.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\n        # Skip classification eval (no ImageNet val loader needed here)\n        log = {'loss': [], 'acc': [0.0], 'f1': [0.0], 'precision': [0.0], 'recall': [0.0]}\n        if in1k_val_loader is not None:\n            acc, f1, prec, rec = evaluate_cls(model, in1k_val_loader)\n            log = {'loss': [], 'acc': [acc], 'f1': [f1], 'precision': [prec], 'recall': [rec]}\n            print(f'  Loaded — acc={acc:.4f}  f1={f1:.4f}  prec={prec:.4f}  rec={rec:.4f}')\n        else:\n            print(f'  Loaded — (skipping cls eval, no ImageNet val loader)')\n    else:\n        # Fallback: train from scratch (requires ImageNet download)\n        print(f'  WARNING: {ckpt_path} not found — training from pretrained weights')\n        model = build_vit(pos, png, pretrained=True)\n        log = train_cls(model, EPOCHS_CLS, name, ckpt_path=CKPT_OUT / ckpt_name)\n    cls_results[name] = {'model': model, 'log': log}\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s6","cell_type":"markdown","source":"## 6. Segmentation head & helpers","metadata":{}},{"id":"seg_head","cell_type":"code","source":"PATCH_GRID = VIT_SIZE // 16   # 14\n\nclass LinearSegHead(nn.Module):\n    def __init__(self, embed_dim=768, num_classes=150):\n        super().__init__()\n        self.head = nn.Conv2d(embed_dim * 4, num_classes, kernel_size=1)\n\n    def forward(self, feats, out_size):\n        # feats: list of 4 tensors, each (B, 196, C=768)\n        # Reshape each to (B, C, 14, 14), concatenate, apply head conv,\n        # then upsample the 150-channel logits — avoids INT_MAX overflow.\n        maps = []\n        for f in feats:\n            B, N, C = f.shape   # N = 196 = PATCH_GRID²\n            f = f.transpose(1, 2).reshape(B, C, PATCH_GRID, PATCH_GRID)  # (B, C, 14, 14)\n            maps.append(f)\n        x = self.head(torch.cat(maps, dim=1))   # (B, num_classes, 14, 14)\n        return F.interpolate(x, size=out_size, mode='bilinear', align_corners=False)\n\n\nclass SegModel(nn.Module):\n    def __init__(self, vit, seg_head):\n        super().__init__()\n        self.vit      = vit\n        self.seg_head = seg_head\n\n    def forward(self, x):\n        # x: (B, 3, 224, 224)\n        feats = self.vit.get_intermediate_layers(x, n=[8, 9, 10, 11], reshape=False)\n        # each feat: (B, 196, 768) — prefix tokens already stripped by timm\n        return self.seg_head(feats, (MASK_SIZE, MASK_SIZE))\n\n\n@torch.no_grad()\ndef mean_iou(model, loader):\n    model.eval()\n    inter = torch.zeros(ADE_CLASSES)\n    union = torch.zeros(ADE_CLASSES)\n    for imgs, masks in loader:\n        imgs, masks = imgs.to(DEVICE), masks.to(DEVICE)\n        preds = model(imgs).argmax(1)\n        valid = masks >= 0\n        for c in range(ADE_CLASSES):\n            pred_c = (preds == c) & valid\n            true_c = (masks == c) & valid\n            inter[c] += (pred_c & true_c).sum().cpu()\n            union[c] += (pred_c | true_c).sum().cpu()\n    iou = inter / (union + 1e-6)\n    return iou[union > 0].mean().item()\n\n\ndef train_seg(vit_model, name, epochs=EPOCHS_SEG, lr=1e-4):\n    seg_head = LinearSegHead(embed_dim=768, num_classes=ADE_CLASSES).to(DEVICE)\n    model    = SegModel(vit_model, seg_head).to(DEVICE)\n\n    for p in model.vit.parameters(): p.requires_grad_(False)\n    if hasattr(model.vit, 'gate_params'):\n        for p in model.vit.gate_params.parameters(): p.requires_grad_(True)\n\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f'  trainable {trainable:,}')\n\n    opt   = torch.optim.AdamW(\n        [p for p in model.parameters() if p.requires_grad], lr=lr, weight_decay=0.01)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    crit  = nn.CrossEntropyLoss(ignore_index=-1)\n\n    log = {'loss': [], 'miou': []}\n    for ep in range(1, epochs+1):\n        model.train(); model.vit.eval()\n        total_loss = 0\n        for imgs, masks in tqdm(ade_train_loader, desc=f'seg {name} ep{ep}', leave=False):\n            imgs, masks = imgs.to(DEVICE), masks.to(DEVICE)\n            opt.zero_grad()\n            loss = crit(model(imgs), masks)\n            loss.backward(); opt.step()\n            total_loss += loss.item()\n        miou = mean_iou(model, ade_val_loader)\n        sched.step()\n        log['loss'].append(total_loss / len(ade_train_loader))\n        log['miou'].append(miou)\n        print(f'  seg {name} ep{ep}  loss={log[\"loss\"][-1]:.4f}  mIoU={miou:.4f}')\n    return model, log\n\nprint(f'Seg head ready. Patch grid {PATCH_GRID}×{PATCH_GRID} → upsample to {MASK_SIZE}×{MASK_SIZE}')\nprint(f'get_intermediate_layers output: (B, {PATCH_GRID**2}, 768) — patch tokens only')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s7","cell_type":"markdown","source":"## 7. Run segmentation","metadata":{}},{"id":"run_seg","cell_type":"code","source":"seg_results = {}\nfor name in ['Baseline', 'G1 Gate', 'PNG Gate']:\n    print(f'\\n── Seg: {name} ──')\n    vit = cls_results[name]['model']\n    seg_model, log = train_seg(vit, name, epochs=EPOCHS_SEG)\n    seg_results[name] = {'model': seg_model, 'log': log}","metadata":{},"outputs":[],"execution_count":null},{"id":"s8","cell_type":"markdown","source":"## 8. Fig C — mIoU comparison","metadata":{}},{"id":"figC","cell_type":"code","source":"seg_names = list(seg_results.keys())\nmious     = [max(seg_results[n]['log']['miou']) * 100 for n in seg_names]\ncolors_s  = ['#888', '#4C72B0', '#DD8452']\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))\nbars = ax1.bar(seg_names, mious, color=colors_s, edgecolor='k', linewidth=0.6)\nfor bar, m in zip(bars, mious):\n    ax1.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.1,\n             f'{m:.2f}', ha='center', va='bottom', fontsize=9)\nax1.set_ylabel('mIoU (%) — ADE20K val')\nax1.set_title('Segmentation: mIoU Comparison', fontweight='bold')\nax1.set_ylim(min(mious)*0.95, max(mious)*1.04)\nax1.grid(axis='y', alpha=0.3)\n\nfor n, c in zip(seg_names, colors_s):\n    ax2.plot(range(1, EPOCHS_SEG+1), [v*100 for v in seg_results[n]['log']['miou']],\n             'o-', color=c, label=n)\nax2.set_xlabel('Epoch'); ax2.set_ylabel('mIoU (%)')\nax2.set_title('Segmentation Training Curves', fontweight='bold')\nax2.legend(); ax2.grid(alpha=0.3)\n\nfig.suptitle(f'ADE20K Segmentation (Linear Probe, ViT@{VIT_SIZE}→mask@{MASK_SIZE})',\n             fontsize=12, fontweight='bold')\nfig.tight_layout()\nfig.savefig('figC_seg_miou.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figC_seg_miou.pdf')","metadata":{},"outputs":[],"execution_count":null},{"id":"s9","cell_type":"markdown","source":"## 9. Fig G — Qualitative examples","metadata":{}},{"id":"figG","cell_type":"code","source":"rng_c = torch.Generator().manual_seed(0)\nCMAP  = torch.cat([\n    torch.zeros(1, 3, dtype=torch.uint8),\n    torch.randint(50, 255, (150, 3), dtype=torch.uint8, generator=rng_c)\n], dim=0)  # (151, 3)\n\ndef seg_to_rgb(mask_t):\n    \"\"\"mask_t: (H, W) int64 tensor, values -1..149\"\"\"\n    idx = (mask_t + 1).clamp(0, 150)           # -1→0 (black), 0-149→1-150\n    return CMAP[idx].numpy()                    # (H, W, 3) uint8 — numpy only for plt\n\ndef denorm_img(t):\n    m = torch.tensor(MEAN).view(3,1,1)\n    s = torch.tensor(STD).view(3,1,1)\n    return (t.cpu()*s + m).clamp(0,1).permute(1,2,0).numpy()\n\nval_imgs, val_masks = next(iter(DataLoader(ade_val, batch_size=4, shuffle=False, num_workers=0)))\n\ndef get_pred(seg_model, imgs):\n    seg_model.eval()\n    with torch.no_grad():\n        return seg_model(imgs.to(DEVICE)).argmax(1).cpu()\n\npred_b = get_pred(seg_results['Baseline']['model'], val_imgs)\npred_g = get_pred(seg_results['G1 Gate']['model'],  val_imgs)\npred_p = get_pred(seg_results['PNG Gate']['model'],  val_imgs)\n\nfig, axes = plt.subplots(4, 5, figsize=(15, 12))\nfor j, t in enumerate(['Image', 'GT Mask', 'Baseline', 'G1 Gate', 'PNG Gate']):\n    axes[0, j].set_title(t, fontsize=10, fontweight='bold')\nfor i in range(4):\n    axes[i, 0].imshow(denorm_img(val_imgs[i]));            axes[i, 0].axis('off')\n    axes[i, 1].imshow(seg_to_rgb(val_masks[i]));           axes[i, 1].axis('off')\n    axes[i, 2].imshow(seg_to_rgb(pred_b[i]));              axes[i, 2].axis('off')\n    axes[i, 3].imshow(seg_to_rgb(pred_g[i]));              axes[i, 3].axis('off')\n    axes[i, 4].imshow(seg_to_rgb(pred_p[i]));              axes[i, 4].axis('off')\n\nfig.suptitle(f'ADE20K Segmentation Predictions (ViT@{VIT_SIZE}→mask@{MASK_SIZE})',\n             fontsize=12, fontweight='bold')\nfig.tight_layout()\nfig.savefig('figG_seg_qual.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figG_seg_qual.pdf')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s10","cell_type":"markdown","source":"## 10. Summary","metadata":{}},{"id":"summary","cell_type":"code","source":"print('=' * 70)\nprint('  OPTION A RESULTS')\nprint('=' * 70)\n\nprint(f'\\n  {\"Model\":<20} {\"Acc\":>8} {\"Precision\":>10} {\"Recall\":>8} {\"F1\":>8}  (ImageNet-1K)')\nprint('  ' + '-'*57)\nfor name in cls_results:\n    log  = cls_results[name]['log']\n    acc  = max(log['acc'])\n    prec = max(log['precision'])\n    rec  = max(log['recall'])\n    f1   = max(log['f1'])\n    print(f'  {name:<20} {acc*100:>7.2f}% {prec*100:>9.2f}% {rec*100:>7.2f}% {f1*100:>7.2f}%')\n\nprint(f'\\n  {\"Model\":<20} {\"Best mIoU\":>10}  (ADE20K)')\nprint('  ' + '-'*35)\nfor name in seg_results:\n    miou = max(seg_results[name]['log']['miou'])\n    print(f'  {name:<20} {miou*100:>9.2f}%')\n\nprint(f'\\n  Checkpoints saved in: {CKPT_DIR.resolve()}')\nprint('=' * 70)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"062e8062","cell_type":"code","source":"# Save seg checkpoints for output\nimport shutil\n\nfor name, fname in [('Baseline', 'vit_seg_baseline.pt'),\n                    ('G1 Gate',  'vit_seg_g1.pt'),\n                    ('PNG Gate', 'vit_seg_png.pt')]:\n    if name in seg_results:\n        path = CKPT_OUT / fname\n        torch.save(seg_results[name]['model'].seg_head.state_dict(), path)\n        print(f'Saved {path}')\n\nprint('Done.')\n","metadata":{},"outputs":[],"execution_count":null}]}