{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":6799,"databundleVersionId":4225553},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:17.670735Z","iopub.execute_input":"2026-03-03T18:29:17.671622Z","iopub.status.idle":"2026-03-03T18:29:20.833011Z","shell.execute_reply.started":"2026-03-03T18:29:17.671589Z","shell.execute_reply":"2026-03-03T18:29:20.831555Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.833546Z","iopub.status.idle":"2026-03-03T18:29:20.833798Z","shell.execute_reply.started":"2026-03-03T18:29:20.833681Z","shell.execute_reply":"2026-03-03T18:29:20.833696Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.835042Z","iopub.status.idle":"2026-03-03T18:29:20.835377Z","shell.execute_reply.started":"2026-03-03T18:29:20.835235Z","shell.execute_reply":"2026-03-03T18:29:20.83525Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.836422Z","iopub.status.idle":"2026-03-03T18:29:20.836703Z","shell.execute_reply.started":"2026-03-03T18:29:20.836582Z","shell.execute_reply":"2026-03-03T18:29:20.8366Z"}},"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):\n    m = timm.create_model('vit_base_patch16_224', pretrained=True, num_classes=num_classes)\n    if pos != 'baseline':\n        m = inject_gates(m, pos=pos, png=png)\n    # freeze everything, then unfreeze head (and gates if present)\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    # per-class TP / FP / FN using vectorised ops\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    # only classes that appear in ground truth (avoids division by zero on unseen classes)\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# keep a thin wrapper for backward compat\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        # save after each epoch so a crash doesn't lose all progress\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.837678Z","iopub.status.idle":"2026-03-03T18:29:20.837951Z","shell.execute_reply.started":"2026-03-03T18:29:20.837837Z","shell.execute_reply":"2026-03-03T18:29:20.837851Z"}},"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\n# checkpoint filename for each experiment\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    model = build_vit(pos, png)\n    if ckpt_path.exists():\n        print(f'  Loading checkpoint: {ckpt_path}')\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        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        # ckpt_path passed so the model is saved after every epoch\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.838776Z","iopub.status.idle":"2026-03-03T18:29:20.839012Z","shell.execute_reply.started":"2026-03-03T18:29:20.838903Z","shell.execute_reply":"2026-03-03T18:29:20.838917Z"}},"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    upsamples to MASK_SIZE, concatenates and predicts per-pixel class.\n    Matches the linear probe used by Darcet et al. 2023.\n\n    Token layout from get_intermediate_layers(reshape=False):\n      index 0        → CLS token\n      index 1        → dist_token\n      indices 2-197  → 196 patch tokens (14×14 grid)\n    We strip the two prefix tokens before reshaping.\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, N=198, C=768)\n        maps = []\n        for f in feats:\n            B, N, C = f.shape\n            # skip CLS token (0) and dist_token (1); keep patch tokens (2-197)\n            f = f[:, 2:, :]                                    # (B, 196, C)\n            f = f.transpose(1, 2).reshape(B, C, PATCH_GRID, PATCH_GRID)  # (B, C, 14, 14)\n            f = F.interpolate(f, size=out_size, mode='bilinear', align_corners=False)\n            maps.append(f)\n        return self.head(torch.cat(maps, dim=1))   # (B, num_classes, H, W)\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, 198, 768)  [198 = CLS + dist + 196 patches]\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}')\nprint(f'Token layout: [0=CLS, 1=dist_token, 2-{1+PATCH_GRID**2}=patches] — seg uses indices 2:{2+PATCH_GRID**2}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.839739Z","iopub.status.idle":"2026-03-03T18:29:20.839994Z","shell.execute_reply.started":"2026-03-03T18:29:20.83987Z","shell.execute_reply":"2026-03-03T18:29:20.839884Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.841214Z","iopub.status.idle":"2026-03-03T18:29:20.841436Z","shell.execute_reply.started":"2026-03-03T18:29:20.841332Z","shell.execute_reply":"2026-03-03T18:29:20.841345Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.842854Z","iopub.status.idle":"2026-03-03T18:29:20.843206Z","shell.execute_reply.started":"2026-03-03T18:29:20.843018Z","shell.execute_reply":"2026-03-03T18:29:20.843038Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.843925Z","iopub.status.idle":"2026-03-03T18:29:20.844238Z","shell.execute_reply.started":"2026-03-03T18:29:20.84407Z","shell.execute_reply":"2026-03-03T18:29:20.844085Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.848561Z","iopub.status.idle":"2026-03-03T18:29:20.848873Z","shell.execute_reply.started":"2026-03-03T18:29:20.848729Z","shell.execute_reply":"2026-03-03T18:29:20.848744Z"}},"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 temp hook.\n\n    Token layout: [0=CLS, 1=dist_token, 2-197=196 patch tokens]\n    We extract attention from CLS (row 0) to patch tokens (cols 2:),\n    giving 196 values that reshape to a 14×14 spatial map.\n    \"\"\"\n    attn_mod = model.blocks[-1].attn\n    raw = {}\n    orig_fwd = attn_mod.forward\n\n    def tmp(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        a = (q * self.scale) @ k.transpose(-2, -1)\n        a = a.softmax(-1)\n        raw['a'] = a.detach()   # (B, H, N, N)\n        x2 = (a @ v).transpose(1, 2).reshape(B, N, C)\n        return self.proj_drop(self.proj(x2))\n\n    attn_mod.forward = types.MethodType(tmp, attn_mod)\n    model.eval()\n    with torch.no_grad(): model(batch)\n    attn_mod.forward = orig_fwd\n\n    # CLS row (index 0), skip CLS self-attn (0) and dist_token (1),\n    # keep patch tokens (2:) → 196 values → reshape to 14×14\n    cls = raw['a'][:, :, 0, 2:].mean(1)   # (B, H, 196) → mean over H → (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')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.849928Z","iopub.status.idle":"2026-03-03T18:29:20.850249Z","shell.execute_reply.started":"2026-03-03T18:29:20.850118Z","shell.execute_reply":"2026-03-03T18:29:20.85014Z"}},"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    Token layout: [0=CLS, 1=dist_token, 2-197=196 patch tokens]\n    Both norms and CLS attention are computed over patch tokens only (index 2:).\n    \"\"\"\n    pcts = []\n    for block in model.blocks:\n        attn_mod = block.attn\n        raw = {}\n        orig = attn_mod.forward\n\n        def tmp(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            a = (q * self.scale) @ k.transpose(-2, -1)\n            a = a.softmax(-1)\n            raw['a'] = a.detach(); raw['x'] = x.detach()\n            x2 = (a @ v).transpose(1, 2).reshape(B, N, C)\n            return self.proj_drop(self.proj(x2))\n\n        attn_mod.forward = types.MethodType(tmp, attn_mod)\n        model.eval()\n        with torch.no_grad(): model(imgs.to(DEVICE))\n        attn_mod.forward = orig\n\n        # skip CLS (0) and dist_token (1) → 196 patch tokens\n        norms    = raw['x'].norm(dim=-1)[:, 2:]         # (B, 196)\n        hi       = (norms > thresh).float()\n        cls_attn = raw['a'][:, :, 0, 2:].mean(1)       # (B, H, 196) → (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')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.851296Z","iopub.status.idle":"2026-03-03T18:29:20.851584Z","shell.execute_reply.started":"2026-03-03T18:29:20.851425Z","shell.execute_reply":"2026-03-03T18:29:20.851446Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.852576Z","iopub.status.idle":"2026-03-03T18:29:20.85284Z","shell.execute_reply.started":"2026-03-03T18:29:20.852723Z","shell.execute_reply":"2026-03-03T18:29:20.85274Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.854375Z","iopub.status.idle":"2026-03-03T18:29:20.854639Z","shell.execute_reply.started":"2026-03-03T18:29:20.854486Z","shell.execute_reply":"2026-03-03T18:29:20.854498Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-03T18:29:20.855837Z","iopub.status.idle":"2026-03-03T18:29:20.856159Z","shell.execute_reply.started":"2026-03-03T18:29:20.85597Z","shell.execute_reply":"2026-03-03T18:29:20.855983Z"}},"outputs":[],"execution_count":null}]}