{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ================================================================\n# ASPP-ViT: Adaptive Score-Proportional Pruning for Vision Transformers\n# ================================================================\n\nimport os, csv, gc, time, types, math\nimport torch\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\nimport numpy as np\nfrom PIL import Image\nfrom tqdm import tqdm\n\n\n# ================================================================\n# CONFIG\n# ================================================================\n\nBASE = \"/kaggle/input/competitions/imagenet-object-localization-challenge\"\n\n# Per-layer base keep ratios (cascaded schedule)\nBASE_SCHEDULE = {4: 0.85, 6: 0.73, 8: 0.68, 10: 0.62}\n\n# Populated during calibration pass\nLAYER_NORM = {li: {\"mean\": 0.0, \"std\": 1.0} for li in [4, 6, 8, 10]}\n\nCONFIG = {\n    \"DEVICE\":           torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n    \"BATCH_SIZE\":       32,\n    \"NUM_WORKERS\":      4,\n    \"VAL_IMG_DIR\":      f\"{BASE}/ILSVRC/Data/CLS-LOC/val\",\n    \"VAL_SOLUTION_CSV\": f\"{BASE}/LOC_val_solution.csv\",\n    \"SYNSET_MAP_TXT\":   f\"{BASE}/LOC_synset_mapping.txt\",\n    \"MODEL_NAME\":       \"vit_base_patch16_224\",\n    \"N_LAYERS\":         12,\n    \"EMBED_DIM\":        768,\n    \"MIN_KEEP\":         20,\n    \"PRUNE_LAYERS\":     [4, 6, 8, 10],\n    \"RHO_SWEEP\":        [0.0, 0.1, 0.2, 0.3, 0.4, 0.5],\n    \"USE_FUSED\":        True,\n    \"WARMUP_BATCHES\":   5,\n    \"CAL_BATCHES\":      150,\n    \"NORM_K\":           1.5,       # clipping half-width in units of std\n    \"TRACK_ROUTING\":    True,\n}\n\nprint(\"=\" * 65)\nprint(\"ASPP-ViT: Adaptive Score-Proportional Pruning\")\nprint(f\"Device        : {CONFIG['DEVICE']}\")\nprint(f\"Prune layers  : {CONFIG['PRUNE_LAYERS']}\")\nprint(f\"Base schedule : {BASE_SCHEDULE}\")\nprint(f\"Rho sweep     : {CONFIG['RHO_SWEEP']}\")\nprint(f\"Norm          : mean ± {CONFIG['NORM_K']}×std\")\nprint(f\"Track routing : {CONFIG['TRACK_ROUTING']}\")\nprint(\"=\" * 65)\n\n\ndef clear_gpu():\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    gc.collect()\n\n\n# ================================================================\n# DATASET\n# ================================================================\n\ndef build_synset_to_idx(path):\n    s2i = {}\n    with open(path) as f:\n        for idx, line in enumerate(f):\n            s = line.strip().split(\" \")[0]\n            if s: s2i[s] = idx\n    assert len(s2i) == 1000\n    return s2i\n\n\ndef build_val_dataset(img_dir, csv_path, s2i):\n    paths, labels, missing = [], [], 0\n    with open(csv_path) as f:\n        for row in csv.DictReader(f):\n            syn  = row[\"PredictionString\"].strip().split(\" \")[0]\n            path = os.path.join(img_dir, row[\"ImageId\"].strip() + \".JPEG\")\n            if syn not in s2i or not os.path.exists(path):\n                missing += 1; continue\n            paths.append(path); labels.append(s2i[syn])\n    print(f\"  Val: {len(paths)} images, {missing} skipped\")\n    assert len(paths) >= 49000\n    return paths, labels\n\n\nclass ImageNetValDataset(Dataset):\n    def __init__(self, paths, labels):\n        self.paths, self.labels = paths, labels\n        self.T = transforms.Compose([\n            transforms.Resize(256,\n                interpolation=transforms.InterpolationMode.BICUBIC),\n            transforms.CenterCrop(224),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485, 0.456, 0.406],\n                                  [0.229, 0.224, 0.225]),\n        ])\n    def __len__(self): return len(self.paths)\n    def __getitem__(self, i):\n        return self.T(Image.open(self.paths[i]).convert(\"RGB\")), self.labels[i]\n\n\ndef get_val_loader():\n    print(\"\\nBuilding validation loader...\")\n    s2i = build_synset_to_idx(CONFIG[\"SYNSET_MAP_TXT\"])\n    paths, labels = build_val_dataset(\n        CONFIG[\"VAL_IMG_DIR\"], CONFIG[\"VAL_SOLUTION_CSV\"], s2i)\n    loader = DataLoader(\n        ImageNetValDataset(paths, labels),\n        batch_size=CONFIG[\"BATCH_SIZE\"], shuffle=False,\n        num_workers=CONFIG[\"NUM_WORKERS\"], pin_memory=True,\n        persistent_workers=True, drop_last=False,\n    )\n    print(f\"  Batches: {len(loader)} | Images: {len(loader.dataset)}\")\n    return loader\n\n\n# ================================================================\n# ATTENTION PATCHING\n# ================================================================\n\ndef patch_attention(model):\n    \"\"\"Replace block attention forward to expose per-head attention weights.\"\"\"\n    for blk in model.blocks:\n        attn = blk.attn\n        if hasattr(attn, \"_aspp_patched\"): continue\n        def _make():\n            def forward(self, x, alive_mask=None, attn_mask=None, **kwargs):\n                B, N, C = x.shape\n                qkv = self.qkv(x).reshape(\n                    B, N, 3, self.num_heads, C // self.num_heads\n                ).permute(2, 0, 3, 1, 4)\n                q, k, v = qkv.unbind(0)\n                aw = (q @ k.transpose(-2, -1)) * (q.shape[-1] ** -0.5)\n                if attn_mask is not None:\n                    aw = aw + attn_mask\n                if alive_mask is not None:\n                    aw = aw.masked_fill(\n                        (~alive_mask).unsqueeze(1).unsqueeze(2), float('-inf'))\n                aw = aw.softmax(-1)\n                aw = torch.nan_to_num(aw, nan=0.0)\n                aw = self.attn_drop(aw)\n                self.attn_weights = aw.mean(dim=1)   # head-averaged, (B, N, N)\n                x = (aw @ v).transpose(1, 2).reshape(B, N, C)\n                return self.proj_drop(self.proj(x))\n            return forward\n        attn.forward = types.MethodType(_make(), attn)\n        attn._aspp_patched = True\n\n\n# ================================================================\n# CLS FOCUS SCORE\n# ================================================================\n\ndef compute_cls_focus_score(attn_weights, Np):\n    \"\"\"\n    Compute per-image CLS focus score s ∈ [0, 1].\n\n    s = 1 - H(a_cls) / log(Np), where H is entropy of the CLS→patch\n    attention distribution. High s indicates peaked (focused) attention;\n    low s indicates diffuse (uniform) attention.\n    \"\"\"\n    eps     = 1e-9\n    cls_raw = attn_weights[:, 0, 1:Np+1].clamp(min=eps)\n    cls_p   = cls_raw / cls_raw.sum(dim=-1, keepdim=True)\n    H       = -(cls_p * cls_p.clamp(min=eps).log()).sum(dim=-1)\n    H_max   = math.log(max(Np, 2))\n    return (1.0 - H / H_max).clamp(0.0, 1.0)\n\n\n# ================================================================\n# SCORE NORMALIZATION\n# ================================================================\n\ndef normalize_score(score, layer_idx):\n    \"\"\"\n    Affine normalization to [0, 1] with calibrated mean and std.\n\n    Maps [mu - k*std, mu + k*std] linearly onto [0, 1], so\n    E[score_norm] ≈ 0.5, guaranteeing E[keep_ratio] = base_keep.\n    \"\"\"\n    mu  = LAYER_NORM[layer_idx][\"mean\"]\n    std = LAYER_NORM[layer_idx][\"std\"]\n    k   = CONFIG[\"NORM_K\"]\n    return ((score - mu) / (2.0 * k * std + 1e-9) + 0.5).clamp(0.0, 1.0)\n\n\n# ================================================================\n# ADAPTIVE KEEP RATIO\n# ================================================================\n\ndef compute_adaptive_keep_ratios(score_norm, layer_idx, rho):\n    \"\"\"\n    Per-image adaptive keep ratio:\n        r(s) = base + rho * base * (0.5 - s_norm)\n\n    Easy images (s_norm > 0.5) receive r < base (fewer tokens kept).\n    Hard images (s_norm < 0.5) receive r > base (more tokens kept).\n    E[r] = base when E[s_norm] = 0.5 (guaranteed by normalization).\n    \"\"\"\n    base = BASE_SCHEDULE[layer_idx]\n    return (base + rho * base * (0.5 - score_norm)).clamp(0.30, 0.95)\n\n\n# ================================================================\n# VECTORIZED ADAPTIVE TOP-K REDUCTION\n# ================================================================\n\ndef adaptive_topk_reduce(x, attn_weights, score_norm, layer_idx,\n                          rho, min_keep, use_fused, entering_rc):\n    \"\"\"\n    Per-image adaptive top-k token reduction with optional fused token.\n\n    Keep counts are derived from per-image entering_rc (real patch count\n    entering this layer), ensuring the adaptive and fixed reference share\n    the same denominator for iso-compute comparison.\n\n    Args:\n        x           : (B, N, C) token sequence (CLS + patches)\n        attn_weights: (B, N, N) head-averaged attention\n        score_norm  : (B,) normalized CLS focus score in [0, 1]\n        layer_idx   : pruning layer index\n        rho         : adaptivity strength\n        min_keep    : minimum patches to retain per image\n        use_fused   : if True, append one aggregated fused token\n        entering_rc : (B,) int tensor of real patch counts entering this layer\n\n    Returns:\n        x_out, alive_mask, keep_counts, avg_keep, keep_min, keep_max\n    \"\"\"\n    B, N, C = x.shape\n    Np      = N - 1          # padded patch count (batch max)\n    device  = x.device\n\n    keep_ratios = compute_adaptive_keep_ratios(score_norm, layer_idx, rho)\n\n    # Keep counts use per-image entering_rc so that keep_counts[b] ≤ real count\n    if entering_rc is not None:\n        rc          = entering_rc.to(device=device, dtype=torch.int32)\n        keep_counts = (keep_ratios * rc.float()).int()\n        keep_counts = keep_counts.clamp(min=min_keep)\n        keep_counts = torch.min(keep_counts, rc)\n    else:\n        keep_counts = (keep_ratios * Np).int().clamp(min=min_keep, max=Np)\n\n    keep_counts = keep_counts.clamp(min=min_keep, max=Np)\n    max_keep    = keep_counts.max().item()\n\n    cls      = x[:, 0:1, :]\n    patches  = x[:, 1:1+Np, :]\n    cls_attn = attn_weights[:, 0, 1:Np+1]\n    imp      = cls_attn / (cls_attn.sum(1, keepdim=True) + 1e-6)\n\n    # Select top-k patches by CLS attention importance\n    sorted_idx    = imp.argsort(dim=1, descending=True)\n    rank_mask     = (torch.arange(Np, device=device).unsqueeze(0)\n                     < keep_counts.unsqueeze(1))\n    alive_patches = torch.zeros(B, Np, dtype=torch.bool, device=device)\n    alive_patches.scatter_(1, sorted_idx, rank_mask)\n\n    # Fused token: attention-weighted mean of pruned patches\n    if use_fused:\n        dw    = cls_attn * (~alive_patches).float()\n        dw    = dw / (dw.sum(1, keepdim=True) + 1e-9)\n        fused = (dw.unsqueeze(-1) * patches).sum(1, keepdim=True)\n        fused = torch.nan_to_num(fused, nan=0.0)\n\n    # Gather surviving patches in original spatial order\n    pos = torch.where(\n        alive_patches,\n        torch.arange(Np, device=device).unsqueeze(0).expand(B, -1),\n        torch.full((B, Np), Np, device=device, dtype=torch.long))\n    top_sp   = pos.sort(dim=1).values[:, :max_keep]\n    valid    = (top_sp < Np)\n    g_idx    = top_sp.clamp(0, Np - 1)\n    gathered = patches.gather(1, g_idx.unsqueeze(-1).expand(B, max_keep, C))\n    out_p    = gathered * valid.unsqueeze(-1).to(gathered.dtype)\n\n    if use_fused:\n        x_out     = torch.cat([cls, out_p, fused], dim=1)\n        real_lens = keep_counts + 2   # CLS + kept + fused\n    else:\n        x_out     = torch.cat([cls, out_p], dim=1)\n        real_lens = keep_counts + 1   # CLS + kept\n\n    max_len    = x_out.shape[1]\n    alive_mask = (torch.arange(max_len, device=device).unsqueeze(0)\n                  < real_lens.unsqueeze(1))\n    alive_mask[:, 0] = True          # CLS always alive\n\n    avg_keep = keep_counts.float().mean().item()\n    return (x_out, alive_mask, keep_counts,\n            avg_keep, keep_counts.min().item(), keep_counts.max().item())\n\n\n# ================================================================\n# FIXED TOP-K REDUCTION (baseline)\n# ================================================================\n\ndef fixed_topk_reduce(x, attn_weights, keep_ratio, min_keep, Np):\n    \"\"\"Top-k token selection with a fixed per-layer keep ratio.\"\"\"\n    n_keep  = max(min_keep, int(Np * keep_ratio))\n    cls     = x[:, 0:1, :]\n    patches = x[:, 1:1+Np, :]\n    imp     = attn_weights[:, 0, 1:Np+1]\n    imp     = imp / (imp.sum(1, keepdim=True) + 1e-6)\n    _, idx  = imp.topk(n_keep, dim=1)\n    idx, _  = idx.sort(dim=1)\n    kept    = patches.gather(1, idx.unsqueeze(-1).expand(\n        x.shape[0], n_keep, x.shape[2]))\n    return torch.cat([cls, kept], dim=1)\n\n\n# ================================================================\n# GFLOPS ESTIMATION\n# ================================================================\n\ndef block_gflops(n, d=768):\n    \"\"\"GFLOPs for a full ViT block with n tokens.\"\"\"\n    return (4*n*d**2 + 2*n**2*d + 8*n*d**2) / 1e9\n\ndef block_gflops_split(nm, nf, d=768):\n    \"\"\"GFLOPs for a pruning block: attention over nm, MLP over nf tokens.\"\"\"\n    return ((4*nm*d**2 + 2*nm**2*d) + (8*nf*d**2)) / 1e9\n\ndef baseline_gflops():\n    return sum(block_gflops(197) for _ in range(CONFIG[\"N_LAYERS\"]))\n\ndef compute_multilayer_gflops(avg_keeps):\n    \"\"\"Propagate average keep counts through the layer schedule.\"\"\"\n    prune_set = set(CONFIG[\"PRUNE_LAYERS\"])\n    cur = 197.0; total = 0.0\n    for i in range(CONFIG[\"N_LAYERS\"]):\n        if i in prune_set:\n            avg_k = avg_keeps.get(i, cur - 1)\n            n_out = avg_k + 2\n            total += block_gflops_split(int(round(cur)), int(round(n_out)))\n            cur = n_out\n        else:\n            total += block_gflops(int(round(cur)))\n    return total\n\n\n# ================================================================\n# ROUTING TRACKER\n# ================================================================\n\nclass RoutingTracker:\n    \"\"\"\n    Records per-layer routing decisions relative to the fixed-ratio baseline.\n\n    delta[b] = keep_counts[b] - fixed_ref[b]\n             ≈ (ratio[b] - base) * entering_rc[b]\n\n    Negative delta: easy image, fewer tokens than fixed baseline.\n    Positive delta: hard image, more tokens than fixed baseline.\n    \"\"\"\n    def __init__(self):\n        self.decisions = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n        self.deltas    = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n\n    def record(self, li, keep_counts_np, entering_rc_np):\n        base      = BASE_SCHEDULE[li]\n        min_keep  = CONFIG[\"MIN_KEEP\"]\n        fixed_ref = np.maximum(\n            min_keep,\n            (entering_rc_np.astype(float) * base).astype(int))\n        delta     = keep_counts_np.astype(int) - fixed_ref\n        self.deltas[li].append(delta.copy())\n        decisions = np.where(delta < 0, 0, np.where(delta > 0, 2, 1))\n        self.decisions[li].append(decisions)\n\n    def summary(self, rho):\n        if not CONFIG[\"TRACK_ROUTING\"]: return\n        print(f\"\\n  Routing analysis (ρ={rho})\")\n        print(f\"  Below=easy (fewer tokens than fixed)  \"\n              f\"Above=hard (more tokens than fixed)\")\n        print(f\"  {'Layer':>7}  {'Below':>12}  {'Equal':>8}  \"\n              f\"{'Above':>12}  {'Avg Δ':>8}  Regime\")\n        print(f\"  {'─'*72}\")\n\n        for li in CONFIG[\"PRUNE_LAYERS\"]:\n            if not self.decisions[li]: continue\n            dec = np.concatenate(self.decisions[li])\n            dlt = np.concatenate(self.deltas[li])\n            nb  = (dec==0).sum(); ne = (dec==1).sum(); na = (dec==2).sum()\n            n   = len(dec)\n            pb=nb/n*100; pe=ne/n*100; pa=na/n*100; ad=dlt.mean()\n\n            if   pb > 35 and pa > 35: regime = \"balanced adaptive\"\n            elif pa > 70:             regime = \"predominantly hard\"\n            elif pb > 70:             regime = \"predominantly easy\"\n            elif abs(ad) < 1:         regime = \"near-fixed\"\n            else:                     regime = \"mixed\"\n\n            print(f\"  {li:>7}  {nb:>7}({pb:>4.1f}%)  \"\n                  f\"{ne:>5}({pe:>4.1f}%)  \"\n                  f\"{na:>7}({pa:>4.1f}%)  \"\n                  f\"{ad:>+7.1f}  {regime}\")\n\n        # Cross-layer delta correlation\n        li0 = CONFIG[\"PRUNE_LAYERS\"][1]\n        li1 = CONFIG[\"PRUNE_LAYERS\"][2]\n        if self.deltas[li0] and self.deltas[li1]:\n            d0 = np.concatenate(self.deltas[li0])\n            d1 = np.concatenate(self.deltas[li1])\n            mn = min(len(d0), len(d1))\n            r  = np.corrcoef(d0[:mn], d1[:mn])[0, 1]\n            print(f\"\\n  Cross-layer (L{li0}↔L{li1}) Δ correlation: r={r:+.3f}\")\n            if   r > 0.3: print(f\"  → Consistent routing across layers\")\n            elif r > 0.1: print(f\"  → Weak cross-layer consistency\")\n            else:         print(f\"  → Independent per-layer routing\")\n\n\n# ================================================================\n# BLOCK FORWARD BUILDERS\n# ================================================================\n\ndef make_normal_block(block):\n    \"\"\"Standard ViT block forward (no pruning).\"\"\"\n    def forward(x, **kwargs):\n        ao = block.attn(block.norm1(x), **kwargs)\n        if hasattr(block, \"ls1\"): ao = block.ls1(ao)\n        x = x + block.drop_path1(ao)\n        mo = block.mlp(block.norm2(x))\n        if hasattr(block, \"ls2\"): mo = block.ls2(mo)\n        return x + block.drop_path2(mo)\n    return forward\n\n\ndef make_post_prune_block(block, state):\n    \"\"\"ViT block forward with alive_mask applied to attention.\"\"\"\n    def forward(x, **kwargs):\n        B, N, C = x.shape\n        alive = state[\"alive_mask\"]\n        if alive is not None and alive.shape[1] != N:\n            if alive.shape[1] > N:\n                alive = alive[:, :N]\n            else:\n                pad   = torch.zeros(B, N - alive.shape[1],\n                                    dtype=torch.bool, device=x.device)\n                alive = torch.cat([alive, pad], dim=1)\n            state[\"alive_mask\"] = alive\n        ao = block.attn(block.norm1(x), alive_mask=alive, **kwargs)\n        if hasattr(block, \"ls1\"): ao = block.ls1(ao)\n        x = x + block.drop_path1(ao)\n        mo = block.mlp(block.norm2(x))\n        if hasattr(block, \"ls2\"): mo = block.ls2(mo)\n        return x + block.drop_path2(mo)\n    return forward\n\n\n# ================================================================\n# PASS 0 — CALIBRATION\n# ================================================================\n\ndef run_pass0_calibration(model, loader):\n    \"\"\"\n    Calibration pass: run fixed-ratio pruning to collect per-layer\n    CLS focus score statistics (mean, std) used for normalization.\n    \"\"\"\n    device       = CONFIG[\"DEVICE\"]\n    prune_layers = set(CONFIG[\"PRUNE_LAYERS\"])\n    min_keep     = CONFIG[\"MIN_KEEP\"]\n    cal_batches  = CONFIG[\"CAL_BATCHES\"]\n    k            = CONFIG[\"NORM_K\"]\n    cal_scores   = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n    logged       = {li: False for li in prune_layers}\n\n    def make_cal(block, li):\n        keep_r = BASE_SCHEDULE[li]\n        def forward(x, **kwargs):\n            B, N, C = x.shape; Np = N - 1\n            ao = block.attn(block.norm1(x), **kwargs)\n            if hasattr(block, \"ls1\"): ao = block.ls1(ao)\n            x = x + block.drop_path1(ao)\n            s = compute_cls_focus_score(block.attn.attn_weights, Np)\n            cal_scores[li].append(s.cpu().numpy())\n            x = fixed_topk_reduce(x, block.attn.attn_weights,\n                                   keep_r, min_keep, Np)\n            if not logged[li]:\n                print(f\"  Cal L{li}: in={N} → out={x.shape[1]}  \"\n                      f\"keep_ratio={keep_r}\")\n                logged[li] = True\n            mo = block.mlp(block.norm2(x))\n            if hasattr(block, \"ls2\"): mo = block.ls2(mo)\n            return x + block.drop_path2(mo)\n        return forward\n\n    for idx, blk in enumerate(model.blocks):\n        blk.forward = (make_cal(blk, idx) if idx in prune_layers\n                       else make_normal_block(blk))\n\n    print(f\"\\nPass 0 — Calibration \"\n          f\"({cal_batches}×{CONFIG['BATCH_SIZE']}=\"\n          f\"{cal_batches * CONFIG['BATCH_SIZE']} images)\")\n    with torch.no_grad():\n        for bi, (x, _) in enumerate(\n                tqdm(loader, desc=\"  Cal\", total=cal_batches)):\n            if bi >= cal_batches: break\n            _ = model(x.to(device))\n\n    print(f\"\\n  {'L':>4}  {'mean':>8}  {'std':>8}  \"\n          f\"{'2k·std':>8}  {'E[s_norm]':>10}  OK?\")\n    print(f\"  {'─'*48}\")\n    for li in sorted(CONFIG[\"PRUNE_LAYERS\"]):\n        s   = np.concatenate(cal_scores[li])\n        mu  = float(s.mean())\n        std = float(s.std())\n        LAYER_NORM[li][\"mean\"] = mu\n        LAYER_NORM[li][\"std\"]  = max(std, 1e-6)\n        sn  = np.clip((s - mu) / (2 * k * std + 1e-9) + 0.5, 0, 1)\n        en  = float(sn.mean())\n        ok  = \"✓\" if abs(en - 0.5) < 0.03 else \"⚠\"\n        print(f\"  {li:>4}  {mu:>8.4f}  {std:>8.4f}  \"\n              f\"{2*k*std:>8.4f}  {en:>10.4f}  {ok}\")\n    print(f\"  Normalization statistics updated.\")\n\n\n# ================================================================\n# PASS 1 — BASELINE\n# ================================================================\n\ndef run_pass1_baseline(model, loader):\n    \"\"\"Full ViT-B/16 evaluation without any token reduction.\"\"\"\n    device = CONFIG[\"DEVICE\"]; warmup = CONFIG[\"WARMUP_BATCHES\"]\n    top1 = top5 = total = 0; tt = 0.0; ti = 0\n    for blk in model.blocks: blk.forward = make_normal_block(blk)\n    print(f\"\\nPass 1 — Baseline (no pruning)\")\n    with torch.no_grad():\n        for bi, (x, y) in enumerate(tqdm(loader, desc=\"  Baseline\")):\n            x, y = x.to(device), y.to(device)\n            iw   = bi < warmup\n            if not iw:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                ts = time.perf_counter()\n            out = model(x)\n            if not iw:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                tt += time.perf_counter() - ts; ti += x.shape[0]\n            top1  += (out.argmax(1) == y).sum().item()\n            top5  += (out.topk(5, 1).indices == y.unsqueeze(1)).any(1).sum().item()\n            total += y.shape[0]\n    acc1 = top1/total*100; acc5 = top5/total*100\n    fps  = ti/tt if tt > 0 else 0; lat = tt/ti*1000 if ti > 0 else 0\n    print(f\"  Baseline  Top-1={acc1:.2f}%  Top-5={acc5:.2f}%  \"\n          f\"FPS={fps:.1f}  Lat={lat:.2f}ms\")\n    return acc1, acc5, fps, lat\n\n\n# ================================================================\n# PASS 2 — FIXED KEEP-RATIO BASELINE\n# ================================================================\n\ndef run_pass2_fixed(model, loader, base_gf):\n    \"\"\"Fixed per-layer keep-ratio pruning baseline (HTK schedule).\"\"\"\n    device       = CONFIG[\"DEVICE\"]\n    prune_layers = set(CONFIG[\"PRUNE_LAYERS\"])\n    min_keep     = CONFIG[\"MIN_KEEP\"]\n    warmup       = CONFIG[\"WARMUP_BATCHES\"]\n    logged       = {li: False for li in prune_layers}\n    keep_info    = {}\n    top1 = top5 = total = 0; tt = 0.0; ti = 0\n\n    def make_fp(block, li):\n        keep_r = BASE_SCHEDULE[li]\n        def forward(x, **kwargs):\n            B, N, C = x.shape; Np = N - 1\n            ao = block.attn(block.norm1(x), **kwargs)\n            if hasattr(block, \"ls1\"): ao = block.ls1(ao)\n            x = x + block.drop_path1(ao)\n            x = fixed_topk_reduce(x, block.attn.attn_weights,\n                                   keep_r, min_keep, Np)\n            if not logged[li]:\n                nk = x.shape[1] - 1\n                print(f\"\\n  Fixed L{li}: in={N} → out={x.shape[1]}  \"\n                      f\"keep_ratio={keep_r}\")\n                logged[li] = True; keep_info[li] = nk\n            mo = block.mlp(block.norm2(x))\n            if hasattr(block, \"ls2\"): mo = block.ls2(mo)\n            return x + block.drop_path2(mo)\n        return forward\n\n    for idx, blk in enumerate(model.blocks):\n        blk.forward = (make_fp(blk, idx) if idx in prune_layers\n                       else make_normal_block(blk))\n\n    print(f\"\\nPass 2 — Fixed keep-ratio {BASE_SCHEDULE}\")\n    with torch.no_grad():\n        for bi, (x, y) in enumerate(tqdm(loader, desc=\"  Fixed\")):\n            x, y = x.to(device), y.to(device)\n            iw   = bi < warmup\n            if not iw:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                ts = time.perf_counter()\n            out = model(x)\n            if not iw:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                tt += time.perf_counter() - ts; ti += x.shape[0]\n            top1  += (out.argmax(1) == y).sum().item()\n            top5  += (out.topk(5, 1).indices == y.unsqueeze(1)).any(1).sum().item()\n            total += y.shape[0]\n\n    acc1 = top1/total*100; acc5 = top5/total*100\n    fps  = ti/tt if tt > 0 else 0; lat = tt/ti*1000 if ti > 0 else 0\n    fixed_avg = {li: keep_info.get(li, int(BASE_SCHEDULE[li] * 196))\n                 for li in CONFIG[\"PRUNE_LAYERS\"]}\n    fixed_gf  = compute_multilayer_gflops(fixed_avg)\n    print(f\"  Fixed  Top-1={acc1:.2f}%  Top-5={acc5:.2f}%  \"\n          f\"FPS={fps:.1f}  Lat={lat:.2f}ms\")\n    print(f\"  GFLOPs: {fixed_gf:.3f}  \"\n          f\"(saved={(1 - fixed_gf/base_gf)*100:.1f}% vs baseline)\")\n    return acc1, acc5, fps, lat, keep_info, fixed_gf\n\n\n# ================================================================\n# PASS 3 — ADAPTIVE SWEEP\n# ================================================================\n\ndef run_pass3_adaptive(model, loader, rho, fixed_gf):\n    \"\"\"\n    Adaptive ASPP-ViT evaluation for a given adaptivity strength ρ.\n\n    Per-image keep ratios are computed from the normalized CLS focus score.\n    State (alive_mask, entering_rc) is propagated across pruning layers\n    and reset at each batch boundary.\n    \"\"\"\n    device       = CONFIG[\"DEVICE\"]\n    prune_layers = set(CONFIG[\"PRUNE_LAYERS\"])\n    min_keep     = CONFIG[\"MIN_KEEP\"]\n    use_fused    = CONFIG[\"USE_FUSED\"]\n    warmup       = CONFIG[\"WARMUP_BATCHES\"]\n    first_prune  = min(prune_layers)\n    track        = CONFIG[\"TRACK_ROUTING\"]\n    logged       = {li: False for li in prune_layers}\n    top1 = top5 = total = 0; tt = 0.0; ti = 0\n\n    state = {\n        \"alive_mask\":  None,\n        \"entering_rc\": None,   # real patch count entering the next prune layer\n        \"kc\": {li: {\"sum\": 0, \"min\": 9999, \"max\": 0, \"n\": 0}\n               for li in CONFIG[\"PRUNE_LAYERS\"]},\n    }\n\n    tracker = RoutingTracker() if track else None\n\n    def upd(li, kc_np):\n        s = state[\"kc\"][li]\n        s[\"sum\"] += int(kc_np.sum())\n        s[\"min\"]  = min(int(s[\"min\"]), int(kc_np.min()))\n        s[\"max\"]  = max(int(s[\"max\"]), int(kc_np.max()))\n        s[\"n\"]   += len(kc_np)\n\n    def make_ap(block, li):\n        def forward(x, **kwargs):\n            B, N, C = x.shape; Np = N - 1\n            alive = state[\"alive_mask\"]\n\n            ao = block.attn(block.norm1(x), alive_mask=alive, **kwargs)\n            if hasattr(block, \"ls1\"): ao = block.ls1(ao)\n            x = x + block.drop_path1(ao)\n\n            score_raw  = compute_cls_focus_score(block.attn.attn_weights, Np)\n            score_norm = normalize_score(score_raw, li)\n\n            # Capture entering_rc before updating state\n            entering_rc_now = state[\"entering_rc\"]\n            if entering_rc_now is None:\n                # First pruning layer: all 196 patches are real\n                entering_rc_tensor = torch.full(\n                    (B,), 196, dtype=torch.int32, device=device)\n                entering_rc_np = np.full(B, 196, dtype=int)\n            else:\n                entering_rc_tensor = entering_rc_now\n                entering_rc_np     = entering_rc_now.cpu().numpy().astype(int)\n\n            x_out, alive_mask, keep_counts, avg_k, kmin, kmax = \\\n                adaptive_topk_reduce(\n                    x, block.attn.attn_weights, score_norm,\n                    li, rho, min_keep, use_fused,\n                    entering_rc_tensor)\n\n            # Propagate real patch count to the next pruning layer\n            new_rc = keep_counts + (1 if use_fused else 0)\n            state[\"alive_mask\"]  = alive_mask\n            state[\"entering_rc\"] = new_rc\n\n            kc_np = keep_counts.cpu().numpy()\n            upd(li, kc_np)\n\n            if track and tracker is not None:\n                tracker.record(li, kc_np, entering_rc_np)\n\n            if not logged[li]:\n                fixed_ref = max(min_keep, int(\n                    entering_rc_np.mean() * BASE_SCHEDULE[li]))\n                print(f\"\\n  Adaptive L{li} (ρ={rho}): \"\n                      f\"in={N} → out_max={x_out.shape[1]}  \"\n                      f\"avg={avg_k:.0f}  [{kmin},{kmax}]  \"\n                      f\"fixed_ref≈{fixed_ref}\")\n                logged[li] = True\n\n            mo = block.mlp(block.norm2(x_out))\n            if hasattr(block, \"ls2\"): mo = block.ls2(mo)\n            return x_out + block.drop_path2(mo)\n        return forward\n\n    for idx, blk in enumerate(model.blocks):\n        if idx in prune_layers:   blk.forward = make_ap(blk, idx)\n        elif idx > first_prune:   blk.forward = make_post_prune_block(blk, state)\n        else:                     blk.forward = make_normal_block(blk)\n\n    print(f\"\\nPass 3 — Adaptive ρ={rho}\")\n    with torch.no_grad():\n        for bi, (x, y) in enumerate(tqdm(loader,\n                desc=f\"  ρ={rho}\", leave=True)):\n            x, y = x.to(device), y.to(device)\n            state[\"alive_mask\"]  = None\n            state[\"entering_rc\"] = None\n\n            iw = bi < warmup\n            if not iw:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                ts = time.perf_counter()\n            out = model(x)\n            if not iw:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                tt += time.perf_counter() - ts; ti += x.shape[0]\n\n            top1  += (out.argmax(1) == y).sum().item()\n            top5  += (out.topk(5, 1).indices == y.unsqueeze(1)).any(1).sum().item()\n            total += y.shape[0]\n\n            if (bi + 1) % 200 == 0:\n                acc_now = top1/total*100\n                k_str   = \" | \".join(\n                    f\"L{li}[{state['kc'][li]['min']},\"\n                    f\"{state['kc'][li]['sum']//max(state['kc'][li]['n'],1)},\"\n                    f\"{state['kc'][li]['max']}]\"\n                    for li in CONFIG[\"PRUNE_LAYERS\"])\n                print(f\"  [{bi+1}/{len(loader)}] acc={acc_now:.2f}%  \"\n                      f\"keep[min,avg,max]: {k_str}\")\n\n    acc1 = top1/total*100; acc5 = top5/total*100\n    fps  = ti/tt if tt > 0 else 0; lat = tt/ti*1000 if ti > 0 else 0\n\n    avg_keeps = {li: state[\"kc\"][li][\"sum\"] / max(state[\"kc\"][li][\"n\"], 1)\n                 for li in CONFIG[\"PRUNE_LAYERS\"]}\n    ada_gf = compute_multilayer_gflops(avg_keeps)\n    diff   = ada_gf - fixed_gf\n    flag   = \"✓\" if abs(diff) < 0.3 else \"✗\"\n\n    print(f\"\\n  ρ={rho}  Top-1={acc1:.2f}%  Top-5={acc5:.2f}%  \"\n          f\"FPS={fps:.1f}  Lat={lat:.2f}ms\")\n    print(f\"  GFLOPs: {ada_gf:.3f}  (Δ={diff:+.3f} vs fixed  {flag})\")\n    print(f\"  Avg keep: \" +\n          \" \".join(f\"L{li}={avg_keeps[li]:.1f}\"\n                   for li in CONFIG[\"PRUNE_LAYERS\"]))\n\n    if track and tracker is not None:\n        tracker.summary(rho)\n\n    return acc1, acc5, fps, lat, avg_keeps, ada_gf\n\n\n# ================================================================\n# RESULTS ANALYSIS\n# ================================================================\n\ndef analyze(b_acc1, b_acc5, b_fps, b_lat,\n            f_acc1, f_acc5, f_fps, f_lat,\n            fixed_keep_info, fixed_gf,\n            adaptive_results, base_gf):\n\n    print(f\"\\n{'='*95}\")\n    print(f\"  ASPP-ViT — Results Summary\")\n    print(f\"{'='*95}\")\n    print(f\"  GFLOPs: theoretical estimate. FPS: wall-clock with CUDA sync.\")\n    print(f\"  Pareto ✓ = ΔAcc > +0.05% AND |ΔGFLOPs| < 0.3 vs Fixed baseline.\")\n    print(f\"{'='*95}\")\n\n    print(f\"\\n  {'Method':<22} {'Top-1':>7} {'Top-5':>7} \"\n          f\"{'FPS':>7} {'Lat(ms)':>8} {'GFLOPs':>8} \"\n          f\"{'Δ Acc':>8} {'Δ GF':>8}  Note\")\n    print(f\"  {'─'*92}\")\n    print(f\"  {'Baseline':<22} {b_acc1:>7.2f}% {b_acc5:>6.2f}% \"\n          f\"{b_fps:>7.1f} {1000/b_fps:>7.2f} {base_gf:>8.3f}\")\n    print(f\"  {'Fixed keep-ratio':<22} {f_acc1:>7.2f}% {f_acc5:>6.2f}% \"\n          f\"{f_fps:>7.1f} {1000/f_fps:>7.2f} {fixed_gf:>8.3f} \"\n          f\"{f_acc1-b_acc1:>+7.2f}%\")\n\n    best_acc = f_acc1; best_rho = None; pareto = []\n\n    for rho, acc1, acc5, fps, lat, avg_keeps, ada_gf in adaptive_results:\n        d_acc = acc1 - f_acc1; d_gf = ada_gf - fixed_gf\n        matched = abs(d_gf) < 0.3; is_par = d_acc > 0.05 and matched\n        if acc1 > best_acc: best_acc = acc1; best_rho = rho\n        if is_par: pareto.append(rho)\n        note = \"✓ Pareto\" if is_par else (\"↑ GFLOPs\" if d_gf > 0.3 else \"\")\n        print(f\"  {'ASPP ρ='+str(rho):<22} {acc1:>7.2f}% {acc5:>6.2f}% \"\n              f\"{fps:>7.1f} {lat:>7.2f} {ada_gf:>8.3f} \"\n              f\"{d_acc:>+7.2f}% {d_gf:>+7.3f}  {note}\")\n\n    print(f\"\\n  Per-layer average patch counts:\")\n    hdr = f\"  {'L':>3}  {'base r':>6}  {'fixed':>6}\"\n    for r, *_ in adaptive_results: hdr += f\"  {'ρ'+str(r):>7}\"\n    print(hdr); print(f\"  {'─'*72}\")\n    for li in CONFIG[\"PRUNE_LAYERS\"]:\n        fk  = fixed_keep_info.get(li, int(BASE_SCHEDULE[li] * 196))\n        row = f\"  {li:>3}  {BASE_SCHEDULE[li]:>6.2f}  {fk:>6}\"\n        for rho, a, a5, fps, lat, avg_k, gf in adaptive_results:\n            row += f\"  {avg_k[li]:>7.1f}\"\n        print(row)\n\n    print(f\"\\n  GFLOPs iso-compute check (fixed={fixed_gf:.3f}, tol ±0.3):\")\n    for rho, a, a5, fps, lat, avg_k, gf in adaptive_results:\n        diff = gf - fixed_gf; ok = \"✓\" if abs(diff) < 0.3 else \"✗\"\n        print(f\"  ρ={rho}: {gf:.3f}  Δ={diff:+.3f}  {ok}\")\n\n    print(f\"\\n  Throughput (FPS):  Baseline={b_fps:.0f}  Fixed={f_fps:.0f}\", end=\"\")\n    for rho, a, a5, fps, lat, avg_k, gf in adaptive_results:\n        print(f\"  ρ={rho}:{fps:.0f}\", end=\"\")\n    print()\n\n    print(f\"\\n  Verdict:\")\n    if pareto:\n        print(f\"  Pareto improvement at ρ={pareto}\")\n        print(f\"  Higher accuracy at matched GFLOPs vs fixed-ratio baseline\")\n        print(f\"  Best: ρ={best_rho}  Top-1={best_acc:.2f}%  \"\n              f\"(+{best_acc-f_acc1:.2f}% over fixed baseline)\")\n    else:\n        print(f\"  No single ρ dominates; report full Pareto curve.\")\n    print(f\"{'='*95}\")\n\n\n# ================================================================\n# MAIN\n# ================================================================\n\ndef run():\n    clear_gpu()\n    loader  = get_val_loader()\n    base_gf = baseline_gflops()\n    print(f\"\\n  Baseline GFLOPs: {base_gf:.3f}\")\n\n    print(f\"\\n{'#'*65}\\n  PASS 0 — CALIBRATION\\n{'#'*65}\")\n    model = timm.create_model(CONFIG[\"MODEL_NAME\"], pretrained=True\n                              ).to(CONFIG[\"DEVICE\"]).eval()\n    patch_attention(model)\n    run_pass0_calibration(model, loader)\n    del model; clear_gpu()\n\n    print(f\"\\n{'#'*65}\\n  PASS 1 — BASELINE\\n{'#'*65}\")\n    model = timm.create_model(CONFIG[\"MODEL_NAME\"], pretrained=True\n                              ).to(CONFIG[\"DEVICE\"]).eval()\n    patch_attention(model)\n    b_acc1, b_acc5, b_fps, b_lat = run_pass1_baseline(model, loader)\n    del model; clear_gpu()\n\n    print(f\"\\n{'#'*65}\\n  PASS 2 — FIXED KEEP-RATIO BASELINE\\n{'#'*65}\")\n    model = timm.create_model(CONFIG[\"MODEL_NAME\"], pretrained=True\n                              ).to(CONFIG[\"DEVICE\"]).eval()\n    patch_attention(model)\n    f_acc1, f_acc5, f_fps, f_lat, fixed_keep_info, fixed_gf = \\\n        run_pass2_fixed(model, loader, base_gf)\n    del model; clear_gpu()\n\n    adaptive_results = []\n    for rho in CONFIG[\"RHO_SWEEP\"]:\n        print(f\"\\n{'#'*65}\\n  PASS 3 — ADAPTIVE ρ={rho}\\n{'#'*65}\")\n        model = timm.create_model(CONFIG[\"MODEL_NAME\"], pretrained=True\n                                  ).to(CONFIG[\"DEVICE\"]).eval()\n        patch_attention(model)\n        acc1, acc5, fps, lat, avg_keeps, ada_gf = \\\n            run_pass3_adaptive(model, loader, rho, fixed_gf)\n        adaptive_results.append((rho, acc1, acc5, fps, lat, avg_keeps, ada_gf))\n        del model; clear_gpu()\n\n    analyze(b_acc1, b_acc5, b_fps, b_lat,\n            f_acc1, f_acc5, f_fps, f_lat,\n            fixed_keep_info, fixed_gf,\n            adaptive_results, base_gf)\n\n    out  = \"/kaggle/working/ASPP_ViT_results.npz\"\n    save = {\n        \"baseline_acc\": np.array([b_acc1]),\n        \"baseline_fps\": np.array([b_fps]),\n        \"fixed_acc\":    np.array([f_acc1]),\n        \"fixed_fps\":    np.array([f_fps]),\n        \"fixed_gf\":     np.array([fixed_gf]),\n        \"base_gf\":      np.array([base_gf]),\n        \"cal_mean\": np.array([LAYER_NORM[l][\"mean\"]\n                               for l in CONFIG[\"PRUNE_LAYERS\"]]),\n        \"cal_std\":  np.array([LAYER_NORM[l][\"std\"]\n                               for l in CONFIG[\"PRUNE_LAYERS\"]]),\n    }\n    for rho, acc1, acc5, fps, lat, avg_keeps, ada_gf in adaptive_results:\n        k = str(rho).replace(\".\", \"p\")\n        save[f\"acc_rho{k}\"]       = np.array([acc1])\n        save[f\"fps_rho{k}\"]       = np.array([fps])\n        save[f\"gf_rho{k}\"]        = np.array([ada_gf])\n        save[f\"avg_keeps_rho{k}\"] = np.array(\n            [avg_keeps[l] for l in CONFIG[\"PRUNE_LAYERS\"]])\n    np.savez(out, **save)\n    print(f\"\\n  Results saved → {out}\")\n    print(f\"\\n{'#'*65}\\n  ASPP-ViT COMPLETE\\n{'#'*65}\")\n\nif __name__ == \"__main__\":\n    run()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}