{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":3283978,"datasetId":1988734,"databundleVersionId":3334621}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"title","cell_type":"markdown","source":"# Gated Attention for Vision Transformers\n**Eliminating Artifact Patches via SDPA-Output Gating**\n\nNITK Surathkal · B.Tech AI · 2025-2026\n\n- **Model**: `vit_base_patch16_224` (timm, IN-21K pretrained)\n- **Classification**: ImageNet-1K ablation — G1–G5 gate positions + PNG (novel)\n- **Segmentation**: ADE20K — direct comparison with Darcet 2023 (ViT + Registers)\n- **Base paper**: Qiu et al. NeurIPS 2025 — mechanism transferred to ViT","metadata":{}},{"id":"s0","cell_type":"markdown","source":"## 0. Setup","metadata":{}},{"id":"setup","cell_type":"code","source":"import os, types, random\nfrom pathlib import Path\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nfrom scipy.stats import pearsonr\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 torchvision.datasets import ImageFolder\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\nfrom timm.layers import Attention\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nCKPT_DIR = Path('checkpoints_B')\nCKPT_DIR.mkdir(exist_ok=True)\nprint(f'torch {torch.__version__} | timm {timm.__version__} | device {DEVICE}')\nprint(f'Checkpoints → {CKPT_DIR.resolve()}')","metadata":{},"outputs":[],"execution_count":null},{"id":"s1","cell_type":"markdown","source":"## 1. Dataset Paths & Loaders","metadata":{}},{"id":"paths","cell_type":"code","source":"# ── ImageNet-1K ──────────────────────────────────────────────────────────────\nIN1K_TRAIN = Path('/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train')\nIN1K_VAL   = Path('/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val')\n\n# ── ADE20K ───────────────────────────────────────────────────────────────────\nADE_ROOT     = Path('/kaggle/input/datasets/ipythonx/ade20k-scene-parsing/ADEChallengeData2016')\nADE_TRAIN_IMG = ADE_ROOT / 'images/training'\nADE_VAL_IMG   = ADE_ROOT / 'images/validation'\nADE_TRAIN_ANN = ADE_ROOT / 'annotations/training'\nADE_VAL_ANN   = ADE_ROOT / 'annotations/validation'\n\nfor p in [IN1K_TRAIN, IN1K_VAL, ADE_VAL_IMG, ADE_VAL_ANN]:\n    print(f'  {p.name:<30} {\"Yes\" if p.exists() else \"MISSING\"}')","metadata":{},"outputs":[],"execution_count":null},{"id":"loaders","cell_type":"code","source":"MEAN, STD = (0.485, 0.456, 0.406), (0.229, 0.224, 0.225)\nBS       = 64\nSLICE    = 0.1    # flip to 1.0 for full run\nVIT_SIZE = 224    # ViT hard-requires this — do NOT change\nMASK_SIZE = 512   # mask / logit output resolution\n\ntrain_tfm = T.Compose([\n    T.RandomResizedCrop(224), T.RandomHorizontalFlip(),\n    T.ToTensor(), T.Normalize(MEAN, STD),\n])\nval_tfm = T.Compose([\n    T.Resize(256), T.CenterCrop(224),\n    T.ToTensor(), T.Normalize(MEAN, STD),\n])\n\n# ── ImageNet-1K ───────────────────────────────────────────────────────────────\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\nprint('Scanning ImageNet-1K train ...')\nall_samples = build_imagenet_samples(IN1K_TRAIN, SLICE)\nrandom.shuffle(all_samples)\nsplit       = int(len(all_samples) * 0.9)\ntrain_samp  = all_samples[:split]\nval_samp    = all_samples[split:]\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\nin1k_train = ImageSamples(train_samp, train_tfm)\nin1k_val   = ImageSamples(val_samp,   val_tfm)\nin1k_train_loader = DataLoader(in1k_train, BS, shuffle=True,  num_workers=4, pin_memory=True)\nin1k_val_loader   = DataLoader(in1k_val,   BS, shuffle=False, num_workers=4, pin_memory=True)\nprint(f'ImageNet-1K  train {len(in1k_train):,} | val {len(in1k_val):,}  ({SLICE*100:.0f}% slice, 90/10 split)')\n\n# ── ADE20K ────────────────────────────────────────────────────────────────────\n# Images resized to VIT_SIZE (224) — ViT rejects anything else.\n# Masks kept at MASK_SIZE (512) for accurate per-pixel GT.\n# Logits are upsampled to MASK_SIZE inside SegModel.forward.\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        self.img_tfm = T.Compose([\n            T.Resize((VIT_SIZE, VIT_SIZE)),\n            T.ToTensor(), T.Normalize(MEAN, STD),\n        ])\n\n    def __len__(self): return len(self.imgs)\n\n    def __getitem__(self, i):\n        img  = self.img_tfm(Image.open(self.imgs[i]).convert('RGB'))\n        mask = np.array(Image.open(self.anns[i]).resize(\n            (MASK_SIZE, MASK_SIZE), Image.NEAREST), dtype=np.int64)\n        mask = torch.from_numpy(mask).long() - 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=4, pin_memory=True)\nade_val_loader   = DataLoader(ade_val,   16, shuffle=False, num_workers=4, 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}')","metadata":{},"outputs":[],"execution_count":null},{"id":"s2","cell_type":"markdown","source":"## 2. Gate Modules — G1 through G5 + PNG","metadata":{}},{"id":"gates","cell_type":"code","source":"class GateParams(nn.Module):\n    \"\"\"Learnable parameters for one attention layer's gate.\"\"\"\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            # PNG novelty: norm-conditioned bias suppresses high-norm artifact tokens\n            self.e    = nn.Parameter(torch.ones(num_heads))\n            self.beta = nn.Parameter(torch.zeros(1))\n\n    def gate(self, x):\n        # x: (B, N, C) -> G: (B, N, H)\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    \"\"\"Single forward that covers all 5 gate positions (and PNG = G1 variant).\"\"\"\n    def forward(self, x, attn_mask=None):\n        B, N, C = x.shape\n        H, D = self.num_heads, self.head_dim\n\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\n        if pos == 'G3':  # gate Keys\n            k = k * params.gate(x).permute(0,2,1).unsqueeze(-1)\n        if pos == 'G4':  # gate Queries\n            q = q * params.gate(x).permute(0,2,1).unsqueeze(-1)\n        if pos == 'G2':  # gate Values\n            v = v * params.gate(x).permute(0,2,1).unsqueeze(-1)\n\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\n        if pos in ('G1', 'PNG'):  # gate SDPA output — best position\n            G = params.gate(x)                          # (B, N, H)\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\n        x = out.transpose(1,2).reshape(B, N, C)\n\n        if pos == 'G5':  # gate Dense output\n            x = (x.view(B,N,H,D) * params.gate(x).unsqueeze(-1)).view(B,N,C)\n\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\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\nprint('Gate code ready — covers G1/G2/G3/G4/G5/PNG')","metadata":{},"outputs":[],"execution_count":null},{"id":"s3","cell_type":"markdown","source":"## 3. Classification Helpers (ImageNet-1K)","metadata":{}},{"id":"cls_helpers","cell_type":"code","source":"def build_vit(pos='baseline', png=False, num_classes=1000, 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=num_classes)\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\n\n@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        preds = model(x).argmax(1)\n        all_preds.append(preds.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\n@torch.no_grad()\ndef top1(model, loader):\n    acc, _, _, _ = evaluate_cls(model, loader)\n    return acc\n\n\ndef train_cls(model, epochs=5, lr=1e-3, ckpt_path=None):\n    \"\"\"Train classification head (+ optional gate params).\n\n    ckpt_path: if given, the model state is saved after every epoch so that\n    training can be resumed without losing progress on an error.\n    \"\"\"\n    crit = nn.CrossEntropyLoss()\n    opt  = torch.optim.AdamW(\n        [p for p in model.parameters() if p.requires_grad], lr=lr, weight_decay=0.05)\n    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    log = {'loss': [], 'acc': [], 'f1': [], 'precision': [], 'recall': []}\n\n    for ep in range(1, epochs+1):\n        model.train()\n        for name, mod in model.named_modules():\n            if 'gate_params' not in name and 'head' not in name \\\n                    and isinstance(mod, nn.LayerNorm):\n                mod.eval()\n        total_loss = 0\n        for x, y in tqdm(in1k_train_loader, desc=f'ep{ep}', leave=False):\n            x, y = x.to(DEVICE), y.to(DEVICE)\n            opt.zero_grad()\n            loss = crit(model(x), y)\n            loss.backward()\n            opt.step()\n            total_loss += loss.item()\n        avg = total_loss / len(in1k_train_loader)\n        acc, f1, prec, rec = evaluate_cls(model, in1k_val_loader)\n        sched.step()\n        log['loss'].append(avg)\n        log['acc'].append(acc)\n        log['f1'].append(f1)\n        log['precision'].append(prec)\n        log['recall'].append(rec)\n        print(f'  ep{ep}  loss={avg:.4f}  acc={acc:.4f}  '\n              f'f1={f1:.4f}  prec={prec:.4f}  rec={rec:.4f}')\n\n        if ckpt_path is not None:\n            torch.save(model.state_dict(), ckpt_path)\n\n    return log\n\nprint('Classification helpers ready.')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s4","cell_type":"markdown","source":"## 4. Run G1–G5 + PNG Ablation (ImageNet-1K)","metadata":{}},{"id":"run_cls","cell_type":"code","source":"EPOCHS_CLS = 2   # bump to 5–10 for real results\n\nCKPT_NAMES = {\n    'Baseline':          'vit_baseline.pt',\n    'G1 — SDPA output':  'vit_g1.pt',\n    'G2 — Value proj':   'vit_g2.pt',\n    'G3 — Key proj':     'vit_g3.pt',\n    'G4 — Query proj':   'vit_g4.pt',\n    'G5 — Dense output': 'vit_g5.pt',\n    'PNG (novel)':        'vit_png.pt',\n}\n\nexperiments = [\n    ('Baseline',          'baseline', False),\n    ('G1 — SDPA output',  'G1',       False),\n    ('G2 — Value proj',   'G2',       False),\n    ('G3 — Key proj',     'G3',       False),\n    ('G4 — Query proj',   'G4',       False),\n    ('G5 — Dense output', 'G5',       False),\n    ('PNG (novel)',        'G1',       True),\n]\n\ncls_results = {}\nfor name, pos, png in experiments:\n    print(f'\\n── {name} ──')\n    ckpt_path = CKPT_DIR / CKPT_NAMES[name]\n    if ckpt_path.exists():\n        # checkpoint on disk — build skeleton without downloading pretrained weights\n        print(f'  Loading checkpoint: {ckpt_path}')\n        model = build_vit(pos, png, pretrained=False)\n        model.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\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        # no checkpoint — download pretrained weights and fine-tune\n        model = build_vit(pos, png, pretrained=True)\n        n_train = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        n_total = sum(p.numel() for p in model.parameters())\n        print(f'  trainable {n_train:,} / {n_total:,}')\n        log = train_cls(model, epochs=EPOCHS_CLS, ckpt_path=ckpt_path)\n        print(f'  Final checkpoint saved → {ckpt_path}')\n    cls_results[name] = {'model': model, 'log': log}\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s5","cell_type":"markdown","source":"## 5. Segmentation Head — Linear Probe on ADE20K","metadata":{}},{"id":"seg_head","cell_type":"code","source":"ADE_CLASSES = 150\nPATCH_GRID  = VIT_SIZE // 16   # 14  (224/16)\n\nclass LinearSegHead(nn.Module):\n    \"\"\"\n    Takes features from the last 4 ViT blocks (each 14×14 for 224px input),\n    concatenates at 14×14, applies 1×1 conv to predict per-pixel class,\n    then upsamples to MASK_SIZE.\n\n    get_intermediate_layers(reshape=False) already strips prefix tokens\n    (CLS etc.) and returns only patch tokens: (B, 196, 768).\n    No manual slicing needed.\n    \"\"\"\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) — patch tokens only\n        # Reshape each to (B, C, 14, 14), concatenate, apply head conv at 14×14,\n        # then upsample the 150-channel logits — avoids OOM from upsampling 768-ch.\n        maps = []\n        for f in feats:\n            B, N, C = f.shape   # N = PATCH_GRID² = 196\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) — logits upsampled to MASK_SIZE inside seg_head\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)  # (B, MASK_SIZE, MASK_SIZE)\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, epochs=10, lr=1e-4, freeze_vit=True):\n    seg_head = LinearSegHead(embed_dim=768, num_classes=ADE_CLASSES).to(DEVICE)\n    model = SegModel(vit_model, seg_head).to(DEVICE)\n\n    if freeze_vit:\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([p for p in model.parameters() if p.requires_grad],\n                               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()\n        if freeze_vit: model.vit.eval()\n        total = 0\n        for imgs, masks in tqdm(ade_train_loader, desc=f'seg 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()\n            opt.step()\n            total += loss.item()\n        miou = mean_iou(model, ade_val_loader)\n        sched.step()\n        log['loss'].append(total / len(ade_train_loader))\n        log['miou'].append(miou)\n        print(f'  seg ep{ep}  loss={log[\"loss\"][-1]:.4f}  mIoU={miou:.4f}')\n    return model, log\n\nprint(f'Segmentation head ready. Patch grid {PATCH_GRID}×{PATCH_GRID} → upsample to {MASK_SIZE}×{MASK_SIZE}')","metadata":{},"outputs":[],"execution_count":null},{"id":"s6","cell_type":"markdown","source":"## 6. Run Segmentation — Baseline / G1 / PNG (ADE20K)","metadata":{}},{"id":"run_seg","cell_type":"code","source":"EPOCHS_SEG = 2   # bump to 10+ for real results\n\nseg_configs = [\n    ('Baseline', 'Baseline'),\n    ('G1 Gate',  'G1 — SDPA output'),\n    ('PNG Gate', 'PNG (novel)'),\n]\n\nseg_results = {}\nfor seg_name, cls_key in seg_configs:\n    print(f'\\n── Seg: {seg_name} ──')\n    vit = cls_results[cls_key]['model']\n    seg_model, log = train_seg(vit, epochs=EPOCHS_SEG)\n    seg_results[seg_name] = {'model': seg_model, 'log': log}","metadata":{},"outputs":[],"execution_count":null},{"id":"s7","cell_type":"markdown","source":"## 7. Figures","metadata":{}},{"id":"s7a","cell_type":"markdown","source":"### Fig A — Classification Ablation Bar Chart","metadata":{}},{"id":"figA","cell_type":"code","source":"names     = list(cls_results.keys())\naccs      = [max(cls_results[n]['log']['acc'])       * 100 for n in names]\nf1s       = [max(cls_results[n]['log']['f1'])        * 100 for n in names]\nprecs     = [max(cls_results[n]['log']['precision']) * 100 for n in names]\nrecs      = [max(cls_results[n]['log']['recall'])    * 100 for n in names]\n\nx      = np.arange(len(names))\nwidth  = 0.2\nmetric_groups = [\n    (accs,  'Accuracy',  '#4C72B0'),\n    (precs, 'Precision', '#55A868'),\n    (recs,  'Recall',    '#C44E52'),\n    (f1s,   'F1',        '#DD8452'),\n]\n\nfig, ax = plt.subplots(figsize=(13, 5))\nfor i, (vals, label, color) in enumerate(metric_groups):\n    offset = (i - 1.5) * width\n    bars = ax.bar(x + offset, vals, width, label=label, color=color,\n                  edgecolor='k', linewidth=0.5)\n    for bar, v in zip(bars, vals):\n        ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.05,\n                f'{v:.1f}', ha='center', va='bottom', fontsize=6, rotation=90)\n\nax.set_xticks(x)\nax.set_xticklabels(names, rotation=18, ha='right')\nax.set_ylabel('Score (%) — ImageNet-1K val')\nax.set_title('Gate Position Ablation: Acc / Precision / Recall / F1\\n'\n             '(vit_base_patch16_224, IN-21K init)',\n             fontsize=12, fontweight='bold')\nax.legend(loc='lower right')\nax.set_ylim(0, 105)\nax.grid(axis='y', alpha=0.3)\nfig.tight_layout()\nfig.savefig('figA_cls_ablation.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figA_cls_ablation.pdf')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s7b","cell_type":"markdown","source":"### Fig B — Training Curves (Baseline / G1 / PNG)","metadata":{}},{"id":"figB","cell_type":"code","source":"show    = ['Baseline', 'G1 — SDPA output', 'PNG (novel)']\npalette = {'Baseline': '#888', 'G1 — SDPA output': '#4C72B0', 'PNG (novel)': '#DD8452'}\neps     = range(1, EPOCHS_CLS+1)\n\nfig, axes = plt.subplots(1, 4, figsize=(18, 4))\nmetric_keys   = ['acc',       'f1',  'precision', 'recall']\nmetric_labels = ['Top-1 Acc', 'F1',  'Precision', 'Recall']\n\nfor ax, key, lbl in zip(axes, metric_keys, metric_labels):\n    for n in show:\n        log = cls_results[n]['log']\n        vals = [v * 100 for v in log[key]]\n        # if only one value (loaded from ckpt), repeat it across epochs for display\n        if len(vals) < len(eps):\n            vals = vals * len(eps)\n        ax.plot(list(eps)[:len(vals)], vals, 'o-', color=palette[n], label=n)\n    ax.set_title(lbl); ax.set_xlabel('Epoch'); ax.set_ylabel('Score (%)')\n    ax.legend(fontsize=7); ax.grid(alpha=0.3)\n\nfig.suptitle('Training Dynamics — ImageNet-1K', fontsize=12, fontweight='bold')\nfig.tight_layout()\nfig.savefig('figB_cls_curves.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figB_cls_curves.pdf')\n","metadata":{},"outputs":[],"execution_count":null},{"id":"s7c","cell_type":"markdown","source":"### Fig C — ADE20K 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))\n\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\neps_seg = range(1, EPOCHS_SEG+1)\nfor n, c in zip(seg_names, colors_s):\n    ax2.plot(eps_seg, [v*100 for v in seg_results[n]['log']['miou']], '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":"s7d","cell_type":"markdown","source":"### Fig D — CLS Attention Heatmaps (Baseline vs G1 vs PNG)","metadata":{}},{"id":"figD","cell_type":"code","source":"def denorm(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\ndef get_cls_attn_map(model, batch):\n    \"\"\"Extract CLS-query attention from the last block via a register_forward_hook\n    on the QKV linear layer — works correctly with both vanilla and gated forwards.\n\n    Token layout inside the block: [0=CLS, 1-196=196 patch tokens]\n    (this model has num_prefix_tokens=1; no dist_token)\n    CLS row (index 0) to patch tokens (1:) → 196 values → 14×14 map.\n    \"\"\"\n    attn_mod = model.blocks[-1].attn\n    captured = {}\n\n    def qkv_hook(module, inp, out):\n        \"\"\"Hook on attn.qkv to capture q, k after the projection.\"\"\"\n        x = inp[0]  # (B, N, C)\n        B, N, C = x.shape\n        H, D = attn_mod.num_heads, attn_mod.head_dim\n        qkv = out.reshape(B, N, 3, H, D).permute(2, 0, 3, 1, 4)\n        q, k, _ = qkv.unbind(0)\n        q, k = attn_mod.q_norm(q), attn_mod.k_norm(k)\n        scale = D ** -0.5\n        a = (q * scale) @ k.transpose(-2, -1)\n        a = a.softmax(-1)\n        captured['a'] = a.detach()  # (B, H, N, N)\n\n    handle = attn_mod.qkv.register_forward_hook(qkv_hook)\n    model.eval()\n    with torch.no_grad():\n        model(batch)\n    handle.remove()\n\n    # CLS row (index 0), skip CLS self-attn, keep patch tokens (1:) → 196 values\n    cls = captured['a'][:, :, 0, 1:].mean(1)   # (B, 196)\n    return cls.reshape(-1, 14, 14).cpu().numpy()\n\ndef up_attn(a):\n    a = (a - a.min()) / (a.max() - a.min() + 1e-8)\n    return np.array(Image.fromarray((a*255).astype(np.uint8)).resize((224,224), Image.BILINEAR)) / 255.0\n\nvis_imgs, _ = next(iter(DataLoader(in1k_val, batch_size=4, shuffle=False)))\nvis = vis_imgs.to(DEVICE)\n\nm_b = cls_results['Baseline']['model']\nm_g = cls_results['G1 — SDPA output']['model']\nm_p = cls_results['PNG (novel)']['model']\n\na_b = get_cls_attn_map(m_b, vis)\na_g = get_cls_attn_map(m_g, vis)\na_p = get_cls_attn_map(m_p, vis)\n\nfig, axes = plt.subplots(4, 4, figsize=(12, 12))\nfor j, t in enumerate(['Original', 'Baseline', 'G1 Gate', 'PNG Gate']):\n    axes[0, j].set_title(t, fontsize=11, fontweight='bold')\nfor i in range(4):\n    orig = denorm(vis_imgs[i])\n    axes[i, 0].imshow(orig); axes[i, 0].axis('off')\n    for j, att in enumerate([a_b[i], a_g[i], a_p[i]]):\n        axes[i, j+1].imshow(orig)\n        axes[i, j+1].imshow(up_attn(att), alpha=0.55, cmap='jet')\n        axes[i, j+1].axis('off')\nfig.suptitle('CLS Attention Maps — Final Block, Head-Averaged\\n(ImageNet-1K val)',\n             fontsize=12, fontweight='bold')\nfig.tight_layout()\nfig.savefig('figD_attn_maps.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figD_attn_maps.pdf')","metadata":{},"outputs":[],"execution_count":null},{"id":"s7e","cell_type":"markdown","source":"### Fig E — Layer-wise Artifact Patch Attention","metadata":{}},{"id":"figE","cell_type":"code","source":"NORM_THRESH = 150.0\n\ndef artifact_attn_per_layer(model, imgs, thresh=NORM_THRESH):\n    \"\"\"% of CLS attention going to artifact patches (norm > thresh) per block.\n\n    Uses register_forward_hook on attn.qkv so the actual forward (including\n    gates for G1/PNG models) runs undisturbed. We compute softmax attention\n    from q, k inside the hook to capture the raw attention distribution.\n\n    Token layout inside each block: [0=CLS, 1-196=patch tokens]\n    (num_prefix_tokens=1, no dist_token for vit_base_patch16_224)\n    \"\"\"\n    # ── Register hooks on all blocks at once, run model ONCE ──\n    all_captured = {}   # block_idx → {'a': ..., 'x': ...}\n    handles = []\n\n    for block_idx, block in enumerate(model.blocks):\n        attn_mod = block.attn\n        captured = {}\n        all_captured[block_idx] = captured\n\n        def make_hook(attn, cap):\n            def qkv_hook(module, inp, out):\n                x = inp[0]  # (B, N, C)\n                B, N, C = x.shape\n                H, D = attn.num_heads, attn.head_dim\n                qkv = out.reshape(B, N, 3, H, D).permute(2, 0, 3, 1, 4)\n                q, k, _ = qkv.unbind(0)\n                q, k = attn.q_norm(q), attn.k_norm(k)\n                scale = D ** -0.5\n                a = (q * scale) @ k.transpose(-2, -1)\n                a = a.softmax(-1)\n                cap['a'] = a.detach()   # (B, H, N, N)\n                cap['x'] = x.detach()   # (B, N, C)\n            return qkv_hook\n\n        h = attn_mod.qkv.register_forward_hook(make_hook(attn_mod, captured))\n        handles.append(h)\n\n    model.eval()\n    with torch.no_grad():\n        model(imgs.to(DEVICE))\n\n    for h in handles:\n        h.remove()\n\n    # ── Compute per-layer artifact attention % ──\n    pcts = []\n    for block_idx in range(len(model.blocks)):\n        cap = all_captured[block_idx]\n        # skip CLS (0); patch tokens at 1: → (B, 196)\n        norms    = cap['x'].norm(dim=-1)[:, 1:]        # (B, 196)\n        hi       = (norms > thresh).float()\n        cls_attn = cap['a'][:, :, 0, 1:].mean(1)      # (B, 196)\n        pct      = (cls_attn * hi).sum(-1) / (cls_attn.sum(-1) + 1e-8)\n        pcts.append(pct.mean().item() * 100)\n    return pcts\n\nsample, _ = next(iter(DataLoader(in1k_val, batch_size=32, shuffle=False)))\nprint('Computing layer-wise artifact attention ...')\nart_b = artifact_attn_per_layer(m_b, sample)\nart_g = artifact_attn_per_layer(m_g, sample)\nart_p = artifact_attn_per_layer(m_p, sample)\n\nlayers = range(1, 13)\nfig, ax = plt.subplots(figsize=(9, 5))\nax.plot(layers, art_b, 'o-', color='#888',    label='Baseline')\nax.plot(layers, art_g, 's-', color='#4C72B0', label='G1 Gate')\nax.plot(layers, art_p, '^-', color='#DD8452', label='PNG Gate')\nax.set_xlabel('Transformer Block'); ax.set_ylabel('% CLS Attention to Artifact Patches')\nax.set_title(f'Layer-wise Artifact Attention (token norm > {NORM_THRESH:.0f})\\n'\n              'Answers RQ1: Does G1/PNG reduce artifact patches?',\n             fontsize=11, fontweight='bold')\nax.legend(); ax.grid(alpha=0.3)\nfig.tight_layout()\nfig.savefig('figE_artifact_attn.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figE_artifact_attn.pdf')","metadata":{},"outputs":[],"execution_count":null},{"id":"s7f","cell_type":"markdown","source":"### Fig F — PNG Gate Suppression vs Token Norm (Scatter)","metadata":{}},{"id":"figF","cell_type":"code","source":"last_attn = m_p.blocks[-1].attn\nlast_attn._capture = True\n\nall_norms, all_gates = [], []\nm_p.eval()\nwith torch.no_grad():\n    for x, _ in DataLoader(in1k_val, batch_size=64, shuffle=False):\n        m_p(x.to(DEVICE))\n        all_gates.append(last_attn._last_gate.mean(-1).cpu().flatten().numpy())\n        all_norms.append(last_attn._last_x_norm.squeeze(-1).cpu().flatten().numpy())\n        if len(all_norms) * 64 > 8000: break\n\nlast_attn._capture = False\nnorms = np.concatenate(all_norms)\ngates = np.concatenate(all_gates)\n\nidx = np.random.default_rng(42).choice(len(norms), 5000, replace=False)\nns, gs = norms[idx], gates[idx]\nr, pv  = pearsonr(ns, gs)\n\nfig, ax = plt.subplots(figsize=(6, 5))\nax.scatter(ns, gs, s=4, alpha=0.3, color='steelblue')\nxl = np.linspace(ns.min(), ns.max(), 200)\nm_, b_ = np.polyfit(ns, gs, 1)\nax.plot(xl, m_*xl+b_, 'r-', lw=2, label=f'fit  r={r:.3f}')\nax.set_xlabel('Token L2 Norm (pre-attention)')\nax.set_ylabel('Mean Gate Value (over heads)')\nax.set_title('PNG: Gate Suppresses High-Norm (Artifact) Tokens\\nFinal Block, ImageNet-1K val',\n             fontsize=11, fontweight='bold')\nax.legend()\nax.text(0.97, 0.97, f'r = {r:.3f}\\np = {pv:.1e}',\n        transform=ax.transAxes, ha='right', va='top',\n        bbox=dict(boxstyle='round', fc='wheat', alpha=0.8))\nfig.tight_layout()\nfig.savefig('figF_gate_scatter.pdf', bbox_inches='tight')\nplt.show()\nprint(f'Saved figF_gate_scatter.pdf  r={r:.3f}')","metadata":{},"outputs":[],"execution_count":null},{"id":"s7g","cell_type":"markdown","source":"### Fig G — ADE20K Segmentation Qual Examples","metadata":{}},{"id":"figG","cell_type":"code","source":"# Colormap for ADE20K (random per-class colors, fixed seed)\nrng_c = np.random.default_rng(0)\nCMAP  = np.vstack([[0,0,0], rng_c.integers(50, 255, (150, 3))]).astype(np.uint8)\n\ndef seg_to_rgb(mask_np):\n    mask_np = np.clip(mask_np + 1, 0, 150)  # shift back: -1 -> 0 (bg), 0-149 -> 1-150\n    return CMAP[mask_np]\n\ndef denorm_seg(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\n# grab 4 val samples\nvis_loader  = DataLoader(ade_val, batch_size=4, shuffle=False)\nval_imgs, val_masks = next(iter(vis_loader))\n\nm_seg_b = seg_results['Baseline']['model']\nm_seg_g = seg_results['G1 Gate']['model']\nm_seg_p = seg_results['PNG Gate']['model']\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().numpy()\n\npred_b = get_pred(m_seg_b, val_imgs)\npred_g = get_pred(m_seg_g, val_imgs)\npred_p = get_pred(m_seg_p, 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')\n\nfor i in range(4):\n    axes[i, 0].imshow(denorm_seg(val_imgs[i])); axes[i, 0].axis('off')\n    axes[i, 1].imshow(seg_to_rgb(val_masks[i].numpy())); 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('ADE20K Segmentation Predictions (512×512)', fontsize=12, fontweight='bold')\nfig.tight_layout()\nfig.savefig('figG_seg_qual.pdf', bbox_inches='tight')\nplt.show()\nprint('Saved figG_seg_qual.pdf')","metadata":{},"outputs":[],"execution_count":null},{"id":"s8","cell_type":"markdown","source":"## 8. Summary","metadata":{}},{"id":"summary","cell_type":"code","source":"print('=' * 75)\nprint('  RESULTS SUMMARY')\nprint('  Model: vit_base_patch16_224 (IN-21K pretrained)')\nprint('=' * 75)\n\nprint('\\nImageNet-1K Classification (macro, val set):')\nprint(f'  {\"Model\":<25} {\"Acc\":>8} {\"Precision\":>10} {\"Recall\":>8} {\"F1\":>8}')\nprint('  ' + '-'*62)\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    tag  = '  ← novel' if 'PNG' in name else ''\n    print(f'  {name:<25} {acc*100:>7.2f}% {prec*100:>9.2f}% {rec*100:>7.2f}% {f1*100:>7.2f}%{tag}')\n\nprint('\\nADE20K Segmentation (mIoU):')\nprint(f'  {\"Model\":<20} {\"Best mIoU\":>10}')\nprint('  ' + '-'*33)\nfor name in seg_results:\n    miou = max(seg_results[name]['log']['miou'])\n    print(f'  {name:<20} {miou*100:>9.2f}%')\n\nprint('\\nFigures:')\nfigs = [\n    ('figA_cls_ablation.pdf',  'G1–G5 + PNG Acc/Prec/Rec/F1 grouped bar chart'),\n    ('figB_cls_curves.pdf',    'Classification training curves (4 metrics)'),\n    ('figC_seg_miou.pdf',      'ADE20K mIoU bar + training curves'),\n    ('figD_attn_maps.pdf',     'CLS attention heatmaps (4 images)'),\n    ('figE_artifact_attn.pdf', 'Layer-wise artifact attention (RQ1)'),\n    ('figF_gate_scatter.pdf',  'PNG gate suppression vs token norm'),\n    ('figG_seg_qual.pdf',      'ADE20K qualitative predictions'),\n]\nfor f, d in figs:\n    ok = '✓' if Path(f).exists() else '✗'\n    print(f'  {ok} {f:<40} {d}')\nprint('=' * 75)\n","metadata":{},"outputs":[],"execution_count":null}]}