{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ── CELL 1: PACKAGES ──────────────────────────────────────────────────────────\n\nimport subprocess, sys\n\ndef _pip(pkg):\n    r = subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", pkg, \"-q\"],\n                       capture_output=True, text=True)\n    print(f\"  {'[OK]' if r.returncode == 0 else '[FAIL]'}  pip install {pkg}\")\n    if r.returncode != 0:\n        print(f\"         {r.stderr.strip()[:300]}\")\n\nprint(\"=\" * 60)\n_pip(\"timm\")\nprint(\"=\" * 60)\n\n\n# ── CELL 2: IMPORTS ───────────────────────────────────────────────────────────\n\nimport os, csv, gc, time, types, math, json, uuid, warnings\nfrom pathlib  import Path\nfrom datetime import datetime\n\nimport numpy as np\nimport torch\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset\nfrom PIL  import Image\nfrom tqdm import tqdm\nimport timm\n\nwarnings.filterwarnings(\"ignore\", category=UserWarning, module=\"timm\")\n\nSEED = 42\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\ndef gpu_is_supported():\n    if not torch.cuda.is_available():\n        return False\n    major, minor = torch.cuda.get_device_capability(0)\n    return major >= 7\n\nUSE_SAFE_CUDA = gpu_is_supported()\n\nif USE_SAFE_CUDA:\n    torch.cuda.manual_seed_all(SEED)\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = False\n\nprint(\"Imports OK\")\nprint(f\"  PyTorch : {torch.__version__}\")\nprint(f\"  timm    : {timm.__version__}\")\n\nif torch.cuda.is_available():\n    major, minor = torch.cuda.get_device_capability(0)\n    print(f\"  CUDA visible : True  — {torch.cuda.get_device_name(0)} (sm_{major}{minor})\")\n    print(f\"  CUDA usable  : {USE_SAFE_CUDA}\")\nelse:\n    print(\"  CUDA visible : False\")\n    print(\"  CUDA usable  : False\")\n\n\n# ── CELL 3: MODEL PROFILES AND LAYER PRESETS ──────────────────────────────────\n\nMODEL_PROFILES = {\n    \"vit_base\": {\n        \"timm_name\" : \"vit_base_patch16_224\",\n        \"n_layers\"  : 12,\n        \"embed_dim\" : 768,\n        \"patch_size\": 16,\n        \"batch_size\": 32,\n    },\n    \"vit_large\": {\n        \"timm_name\" : \"vit_large_patch16_224\",\n        \"n_layers\"  : 24,\n        \"embed_dim\" : 1024,\n        \"patch_size\": 16,\n        \"batch_size\": 16,\n    },\n    \"deit_base\": {\n        \"timm_name\" : \"deit_base_patch16_224\",\n        \"n_layers\"  : 12,\n        \"embed_dim\" : 768,\n        \"patch_size\": 16,\n        \"batch_size\": 32,\n    },\n}\n\nLAYER_PRESET_REGISTRY = {\n    \"vit_base\": {\n        \"core_4\": {\"layers\": [4, 6, 8, 10], \"ratios\": [0.85, 0.73, 0.68, 0.62]},\n        \"core_3\": {\"layers\": [4, 7, 10],    \"ratios\": [0.85, 0.70, 0.60]},\n        \"late_3\": {\"layers\": [6, 8, 10],    \"ratios\": [0.73, 0.62, 0.52]},\n    },\n    \"vit_large\": {\n        \"core_4\": {\"layers\": [8, 12, 16, 20], \"ratios\": [0.85, 0.73, 0.68, 0.62]},\n        \"core_3\": {\"layers\": [8, 14, 20],     \"ratios\": [0.85, 0.70, 0.60]},\n    },\n    \"deit_base\": {\n        \"core_4\": {\"layers\": [4, 6, 8, 10], \"ratios\": [0.85, 0.73, 0.68, 0.62]},\n        \"core_3\": {\"layers\": [4, 7, 10],    \"ratios\": [0.85, 0.70, 0.60]},\n    },\n}\n\nprint(\"Model profiles:\")\nfor k, v in MODEL_PROFILES.items():\n    print(f\"  {k:<12}  {v['timm_name']:<35}  L={v['n_layers']}  d={v['embed_dim']}\")\n\nprint(\"\\nLayer presets:\")\nfor mk in LAYER_PRESET_REGISTRY:\n    for pk, pv in LAYER_PRESET_REGISTRY[mk].items():\n        print(f\"  {mk}/{pk:<10}  layers={pv['layers']}  ratios={pv['ratios']}\")\n\n\n# ── CELL 4: USER CONFIGURATION ────────────────────────────────────────────────\n\nMODEL_KEY        = \"vit_large\"\nLAYER_PRESET_KEY = \"core_4\"\n\nRHO_SWEEP     = [0.1, 0.2, 0.3, 0.4, 0.5]\nABLATION_RHOS = [0.3]\nRUN_ABLATION  = True\n\n# Asymmetric routing: hard and easy budgets are independent (no zero-sum constraint).\n# rho=0.0 is equivalent to Fixed TopK (no adaptive adjustment).\nHARD_BONUS_RATIO = 0.15   # hard images receive +15% of base keep\nEASY_CUT_RATIO   = 0.10   # easy images receive -10% of base keep\n\n# Layer selectivity: L6 is disabled after calibration confirmed negligible signal.\n# L8 and L10 carry reliable routing signal (D1→D10 spread ~10%, AUC~0.57).\nHARD_THRESHOLDS = {6: 0.0,  8: 0.15, 10: 0.15}\nEASY_THRESHOLDS = {6: 1.0,  8: 0.85, 10: 0.85}\n\nUSE_FUSED    = True\n\n_BASE            = \"/kaggle/input/competitions/imagenet-object-localization-challenge\"\nVAL_IMG_DIR      = f\"{_BASE}/ILSVRC/Data/CLS-LOC/val\"\nVAL_SOLUTION_CSV = f\"{_BASE}/LOC_val_solution.csv\"\nSYNSET_MAP_TXT   = f\"{_BASE}/LOC_synset_mapping.txt\"\n\nNUM_WORKERS    = 4\nMIN_KEEP       = 20\nWARMUP_BATCHES = 5\nCAL_BATCHES    = 300\n\nTRACK_ROUTING    = True\nTRACK_SCORE_DIST = True\n\nprint(f\"Model    : {MODEL_KEY}  ({MODEL_PROFILES[MODEL_KEY]['timm_name']})\")\nprint(f\"Preset   : {LAYER_PRESET_KEY}  ->  {LAYER_PRESET_REGISTRY[MODEL_KEY][LAYER_PRESET_KEY]['layers']}\")\nprint(f\"Rho sweep: {RHO_SWEEP}\")\nprint(f\"Ablation : {ABLATION_RHOS}  (USE_FUSED=False)\")\nprint(f\"Hard thresholds: {HARD_THRESHOLDS}\")\nprint(f\"Easy thresholds: {EASY_THRESHOLDS}\")\nprint(f\"Hard bonus ratio: {HARD_BONUS_RATIO}  Easy cut ratio: {EASY_CUT_RATIO}\")\n\n\n# ── CELL 5: AUTO-DERIVED CONFIG ───────────────────────────────────────────────\n\nassert MODEL_KEY in MODEL_PROFILES\nassert MODEL_KEY in LAYER_PRESET_REGISTRY\nassert LAYER_PRESET_KEY in LAYER_PRESET_REGISTRY[MODEL_KEY]\n\n_mp = MODEL_PROFILES[MODEL_KEY]\n_lp = LAYER_PRESET_REGISTRY[MODEL_KEY][LAYER_PRESET_KEY]\n\nassert len(_lp[\"layers\"]) == len(_lp[\"ratios\"])\nfor _li in _lp[\"layers\"]:\n    assert 0 < _li < _mp[\"n_layers\"], f\"Layer {_li} out of bounds\"\n\ndef gpu_is_supported():\n    if not torch.cuda.is_available():\n        return False\n    major, minor = torch.cuda.get_device_capability(0)\n    return major >= 7\n\nUSE_SAFE_CUDA = gpu_is_supported()\n\nCONFIG = {\n    \"DEVICE\"           : torch.device(\"cuda\" if USE_SAFE_CUDA else \"cpu\"),\n    \"MODEL_KEY\"        : MODEL_KEY,\n    \"MODEL_NAME\"       : _mp[\"timm_name\"],\n    \"N_LAYERS\"         : _mp[\"n_layers\"],\n    \"EMBED_DIM\"        : _mp[\"embed_dim\"],\n    \"PATCH_SIZE\"       : _mp[\"patch_size\"],\n    \"BATCH_SIZE\"       : _mp[\"batch_size\"],\n    \"N_PATCHES\"        : (224 // _mp[\"patch_size\"]) ** 2,\n    \"N_TOKENS\"         : (224 // _mp[\"patch_size\"]) ** 2 + 1,\n    \"LAYER_PRESET_KEY\" : LAYER_PRESET_KEY,\n    \"PRUNE_LAYERS\"     : _lp[\"layers\"],\n    \"BASE_SCHEDULE\"    : dict(zip(_lp[\"layers\"], _lp[\"ratios\"])),\n    \"HARD_THRESHOLDS\"  : HARD_THRESHOLDS,\n    \"EASY_THRESHOLDS\"  : EASY_THRESHOLDS,\n    \"HARD_BONUS_RATIO\" : HARD_BONUS_RATIO,\n    \"EASY_CUT_RATIO\"   : EASY_CUT_RATIO,\n    \"VAL_IMG_DIR\"      : VAL_IMG_DIR,\n    \"VAL_SOLUTION_CSV\" : VAL_SOLUTION_CSV,\n    \"SYNSET_MAP_TXT\"   : SYNSET_MAP_TXT,\n    \"NUM_WORKERS\"      : NUM_WORKERS,\n    \"MIN_KEEP\"         : MIN_KEEP,\n    \"USE_FUSED\"        : USE_FUSED,\n    \"WARMUP_BATCHES\"   : WARMUP_BATCHES,\n    \"CAL_BATCHES\"      : CAL_BATCHES,\n    \"TRACK_ROUTING\"    : TRACK_ROUTING,\n    \"TRACK_SCORE_DIST\" : TRACK_SCORE_DIST,\n    \"RHO_SWEEP\"        : RHO_SWEEP,\n    \"ABLATION_RHOS\"    : ABLATION_RHOS,\n    \"RUN_ABLATION\"     : RUN_ABLATION,\n}\n\nLAYER_CAL_SORTED = {}\n\ndef reset_layer_cal():\n    global LAYER_CAL_SORTED\n    LAYER_CAL_SORTED = {li: None for li in CONFIG[\"PRUNE_LAYERS\"]}\n\nreset_layer_cal()\n\nprint(\"=\" * 68)\nprint(\"ASPP-ViT — Active Configuration\")\nprint(\"=\" * 68)\nprint(f\"  Model        : {CONFIG['MODEL_NAME']}\")\nprint(f\"  Architecture : {CONFIG['N_LAYERS']} blocks  embed_dim={CONFIG['EMBED_DIM']}\")\nprint(f\"  Spatial      : {CONFIG['N_PATCHES']} patches  ({CONFIG['N_TOKENS']} tokens incl. CLS)\")\nprint(f\"  Batch size   : {CONFIG['BATCH_SIZE']}\")\nprint(f\"  Prune layers : {CONFIG['PRUNE_LAYERS']}\")\nprint(f\"  Base schedule: {CONFIG['BASE_SCHEDULE']}\")\nprint(f\"  Hard thresholds : {CONFIG['HARD_THRESHOLDS']}\")\nprint(f\"  Easy thresholds : {CONFIG['EASY_THRESHOLDS']}\")\nprint(f\"  Hard bonus ratio: {CONFIG['HARD_BONUS_RATIO']}  \"\n      f\"Easy cut ratio: {CONFIG['EASY_CUT_RATIO']}\")\nprint(f\"  Rho sweep    : {CONFIG['RHO_SWEEP']}\")\nprint(f\"  Ablation rhos: {CONFIG['ABLATION_RHOS']}  USE_FUSED=False\")\nprint(f\"  Cal images   : {CONFIG['CAL_BATCHES']} x {CONFIG['BATCH_SIZE']} = \"\n      f\"{CONFIG['CAL_BATCHES'] * CONFIG['BATCH_SIZE']:,}\")\nprint(f\"  Device       : {CONFIG['DEVICE']}\")\nprint(\"=\" * 68)\n\n\n# ── CELL 6: RESULTS INFRASTRUCTURE ───────────────────────────────────────────\n\nOUTPUT_DIR   = Path(\"/kaggle/working\")\nMASTER_CSV   = OUTPUT_DIR / \"ASPP_ViT_master_results.csv\"\nRUN_MANIFEST = OUTPUT_DIR / \"ASPP_ViT_runs.json\"\n\n_CSV_COLS = [\n    \"run_id\", \"timestamp\", \"model_key\", \"layer_preset_key\", \"prune_layers\",\n    \"pass_type\", \"rho\", \"use_fused\",\n    \"top1\", \"top5\", \"gflops\", \"fps\", \"latency_ms\",\n    \"delta_acc_vs_baseline\", \"delta_acc_vs_fixed\", \"delta_gf_vs_fixed\",\n    \"avg_keeps\", \"cal_batches\", \"min_keep\", \"batch_size\", \"is_ablation\",\n]\n\ndef _make_run_id():\n    return datetime.now().strftime(\"%Y%m%d_%H%M%S\") + \"_\" + uuid.uuid4().hex[:6]\n\ndef _npz_path(run_id):\n    return OUTPUT_DIR / f\"ASPP_ViT_{CONFIG['MODEL_KEY']}_{CONFIG['LAYER_PRESET_KEY']}_{run_id}.npz\"\n\ndef _fmt_avg_keeps(avg_keeps):\n    if not avg_keeps:\n        return \"\"\n    return \"|\".join(f\"L{li}:{v:.1f}\" for li, v in sorted(avg_keeps.items()))\n\ndef append_csv(row):\n    file_exists = MASTER_CSV.exists()\n    with open(MASTER_CSV, \"a\", newline=\"\") as f:\n        w = csv.DictWriter(f, fieldnames=_CSV_COLS, extrasaction=\"ignore\")\n        if not file_exists:\n            w.writeheader()\n        w.writerow({k: row.get(k, \"\") for k in _CSV_COLS})\n\ndef save_manifest(run_id, config_snapshot):\n    records = []\n    if RUN_MANIFEST.exists():\n        try:\n            with open(RUN_MANIFEST) as f:\n                records = json.load(f)\n        except (json.JSONDecodeError, OSError):\n            records = []\n    records.append({\"run_id\": run_id, \"timestamp\": datetime.now().isoformat(),\n                    \"config\": config_snapshot})\n    with open(RUN_MANIFEST, \"w\") as f:\n        json.dump(records, f, indent=2, default=str)\n\nprint(\"Results infrastructure ready.\")\nprint(f\"  CSV      : {MASTER_CSV}\")\nprint(f\"  Manifest : {RUN_MANIFEST}\")\n\n\n# ── CELL 7: GPU UTILITIES ─────────────────────────────────────────────────────\n\ndef clear_gpu():\n    if torch.cuda.is_available():\n        torch.cuda.synchronize()\n        torch.cuda.empty_cache()\n    gc.collect()\n\ndef load_model():\n    model = timm.create_model(CONFIG[\"MODEL_NAME\"], pretrained=True) \\\n               .to(CONFIG[\"DEVICE\"]).eval()\n    n = sum(p.numel() for p in model.parameters()) / 1e6\n    print(f\"  Loaded {CONFIG['MODEL_NAME']}  ({n:.1f}M params)\")\n    return model\n\nprint(\"GPU utilities ready.\")\n\n\n# ── CELL 8: IMAGENET VALIDATION DATASET ──────────────────────────────────────\n\ndef _build_synset_to_idx(path):\n    s2i = {}\n    with open(path) as f:\n        for idx, line in enumerate(f):\n            syn = line.strip().split()[0]\n            if syn:\n                s2i[syn] = idx\n    assert len(s2i) == 1000\n    return s2i\n\ndef _build_val_pairs(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\n                continue\n            paths.append(path)\n            labels.append(s2i[syn])\n    print(f\"  Val set: {len(paths):,} images  ({missing} skipped)\")\n    assert len(paths) >= 49_000\n    return paths, labels\n\nclass ImageNetValDataset(Dataset):\n    _MEAN = [0.485, 0.456, 0.406]\n    _STD  = [0.229, 0.224, 0.225]\n\n    def __init__(self, paths, labels):\n        self.paths, self.labels = paths, labels\n        self.transform = transforms.Compose([\n            transforms.Resize(256, interpolation=transforms.InterpolationMode.BICUBIC),\n            transforms.CenterCrop(224),\n            transforms.ToTensor(),\n            transforms.Normalize(self._MEAN, self._STD),\n        ])\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, i):\n        return self.transform(Image.open(self.paths[i]).convert(\"RGB\")), self.labels[i]\n\ndef get_val_loader():\n    print(\"\\nBuilding validation DataLoader...\")\n    s2i = _build_synset_to_idx(CONFIG[\"SYNSET_MAP_TXT\"])\n    paths, labels = _build_val_pairs(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):,}  |  BS: {CONFIG['BATCH_SIZE']}\")\n    return loader\n\nprint(\"Dataset utilities ready.\")\n\n\n# ── CELL 9: ATTENTION PATCHING ────────────────────────────────────────────────\n\ndef patch_attention(model):\n    # Monkeypatch all attention blocks to expose attn_weights and accept alive_mask.\n    for blk in model.blocks:\n        attn = blk.attn\n        if getattr(attn, \"_aspp_patched\", False):\n            continue\n\n        def _make_fwd():\n            def forward(self, x, alive_mask=None, attn_mask=None, **kwargs):\n                B, N, C  = x.shape\n                head_dim = C // self.num_heads\n                qkv = (self.qkv(x)\n                       .reshape(B, N, 3, self.num_heads, head_dim)\n                       .permute(2, 0, 3, 1, 4))\n                q, k, v = qkv.unbind(0)\n                aw = (q @ k.transpose(-2, -1)) * (head_dim ** -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 = torch.nan_to_num(aw.softmax(dim=-1), nan=0.0)\n                aw = self.attn_drop(aw)\n                self.attn_weights = aw.mean(dim=1)   # [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\n        attn.forward       = types.MethodType(_make_fwd(), attn)\n        attn._aspp_patched = True\n\nprint(\"Attention patching ready.\")\n\n\ndef compute_cls_focus_score(attn_weights, Np, has_fused=False):\n    \"\"\"\n    CLS focus score: complement of normalised Shannon entropy over CLS-to-patch\n    attention weights. Excludes the fused residual token when present to avoid\n    score inflation from accumulated dropped-patch attention mass.\n    \"\"\"\n    eps     = 1e-9\n    cls_raw = attn_weights[:, 0, 1:Np + 1].clamp(min=eps)\n    if has_fused:\n        cls_raw = cls_raw[:, :-1]\n    cls_p = cls_raw / (cls_raw.sum(-1, keepdim=True) + eps)\n    H     = -(cls_p * cls_p.clamp(min=eps).log()).sum(-1)\n    H_max = math.log(max(CONFIG[\"N_PATCHES\"], 2))\n    return (1.0 - H / H_max).clamp(0.0, 1.0)\n\n\ndef normalize_score(score, layer_idx):\n    \"\"\"\n    Percentile normalisation against the empirical CDF built during calibration.\n    Maps raw focus scores to a uniform [0,1] distribution, making routing\n    thresholds portable across layers and backbone architectures.\n    \"\"\"\n    ref = LAYER_CAL_SORTED.get(layer_idx)\n    if ref is None:\n        raise ValueError(\n            f\"No calibration data for layer {layer_idx}. Run Pass 0 first.\")\n    score_np   = score.detach().cpu().numpy()\n    rank       = np.searchsorted(ref, score_np, side=\"left\")\n    percentile = (rank + 0.5) / len(ref)\n    return torch.tensor(percentile, dtype=torch.float32, device=score.device)\n\n\ndef compute_adaptive_keep_ratios(score_norm, layer_idx, rho, entering_rc):\n    \"\"\"\n    Coarse three-way routing with asymmetric token adjustment.\n\n    Routing relies on bin membership rather than score magnitude, consistent\n    with the empirical observation that AUC~0.57 makes fine-grained ordering\n    unreliable. Only the tail bins (hard/easy, ~15% each) are adjusted;\n    mid-bin images follow the fixed schedule unchanged.\n\n    Hard  (score < HARD_THRESHOLD): keep += round(rho * HARD_BONUS_RATIO * base_keep)\n    Easy  (score > EASY_THRESHOLD): keep -= round(rho * EASY_CUT_RATIO  * base_keep)\n    Mid   (otherwise)             : keep  = base_keep\n\n    Layers with HARD_THRESHOLD=0.0 and EASY_THRESHOLD=1.0 are effectively\n    disabled and behave identically to Fixed TopK.\n    \"\"\"\n    base      = CONFIG[\"BASE_SCHEDULE\"][layer_idx]\n    rc        = entering_rc.float()\n    base_keep = torch.round(base * rc).int()\n\n    hard_thresh = CONFIG[\"HARD_THRESHOLDS\"].get(layer_idx, 0.15)\n    easy_thresh = CONFIG[\"EASY_THRESHOLDS\"].get(layer_idx, 0.85)\n\n    if rho == 0.0 or (hard_thresh == 0.0 and easy_thresh == 1.0):\n        keep_counts = base_keep\n        keep_counts = keep_counts.clamp(min=CONFIG[\"MIN_KEEP\"])\n        keep_counts = torch.min(keep_counts, entering_rc.to(dtype=torch.int32))\n        return keep_counts.float() / rc.clamp(min=1.0)\n\n    hard_bonus_ratio = CONFIG[\"HARD_BONUS_RATIO\"]\n    easy_cut_ratio   = CONFIG[\"EASY_CUT_RATIO\"]\n\n    hard_mask = (score_norm < hard_thresh).int()\n    easy_mask = (score_norm > easy_thresh).int()\n\n    hard_delta = torch.round(\n        rho * hard_bonus_ratio * base_keep.float()\n    ).int() * hard_mask\n\n    easy_delta = torch.round(\n        rho * easy_cut_ratio * base_keep.float()\n    ).int() * easy_mask\n\n    keep_counts = base_keep + hard_delta - easy_delta\n    keep_counts = keep_counts.clamp(min=CONFIG[\"MIN_KEEP\"])\n    keep_counts = torch.min(keep_counts, entering_rc.to(dtype=torch.int32))\n    return keep_counts.float() / rc.clamp(min=1.0)\n\n\nprint(\"Score functions ready.\")\n\n\n# ── CELL 11: TOKEN REDUCTION ──────────────────────────────────────────────────\n\ndef adaptive_topk_reduce(\n        x, attn_weights, score_norm, layer_idx, rho,\n        min_keep, use_fused, entering_rc, has_old_fused):\n    B, N, C = x.shape\n    device  = x.device\n\n    rc     = entering_rc.to(device=device, dtype=torch.int32)\n    max_rc = rc.max().item()\n\n    scale = CONFIG.get(\"RHO_SCALE\", {li: 1.0 for li in CONFIG[\"PRUNE_LAYERS\"]})\n    effective_rho = rho * scale.get(layer_idx, 1.0)\n\n    if effective_rho == 0:\n        base_keep_ratio = CONFIG[\"BASE_SCHEDULE\"][layer_idx]\n\n        cls = x[:, 0:1, :]\n\n        if has_old_fused:\n            real_patches   = x[:, 1 : 1 + max_rc, :]\n            old_fused_tok  = x[:, 1 + max_rc :, :]\n            cls_attn_real  = attn_weights[:, 0, 1 : 1 + max_rc]\n            old_fused_attn = attn_weights[:, 0, 1 + max_rc : N]\n        else:\n            real_patches   = x[:, 1:, :]\n            old_fused_tok  = None\n            cls_attn_real  = attn_weights[:, 0, 1:]\n            old_fused_attn = None\n\n        keep_counts = torch.round(base_keep_ratio * rc.float()).int().clamp(min=min_keep)\n        keep_counts = torch.min(keep_counts, rc).clamp(min=min_keep)\n        max_keep    = keep_counts.max().item()\n\n        imp        = cls_attn_real / (cls_attn_real.sum(1, keepdim=True) + 1e-6)\n        sorted_idx = imp.argsort(dim=1, descending=True)\n        rank_mask  = (torch.arange(max_rc, device=device).unsqueeze(0)\n                      < keep_counts.unsqueeze(1))\n        alive_real = torch.zeros(B, max_rc, dtype=torch.bool, device=device)\n        alive_real.scatter_(1, sorted_idx, rank_mask)\n\n        if use_fused:\n            drop_attn = cls_attn_real * (~alive_real).float()\n            if has_old_fused:\n                all_attn   = torch.cat([drop_attn, old_fused_attn], dim=1)\n                all_tokens = torch.cat([real_patches, old_fused_tok], dim=1)\n            else:\n                all_attn   = drop_attn\n                all_tokens = real_patches\n\n            dw        = all_attn / (all_attn.sum(1, keepdim=True) + 1e-9)\n            new_fused = torch.nan_to_num(\n                (dw.unsqueeze(-1) * all_tokens).sum(1, keepdim=True), nan=0.0)\n\n        pos    = torch.where(\n            alive_real,\n            torch.arange(max_rc, device=device).unsqueeze(0).expand(B, -1),\n            torch.full((B, max_rc), max_rc, device=device, dtype=torch.long))\n        top_sp = pos.sort(dim=1).values[:, :max_keep]\n        valid  = (top_sp < max_rc)\n        g_idx  = top_sp.clamp(0, max_rc - 1)\n        out_p  = real_patches.gather(\n            1, g_idx.unsqueeze(-1).expand(B, max_keep, C)\n        ) * valid.unsqueeze(-1).to(real_patches.dtype)\n\n        if use_fused:\n            x_out = torch.cat([cls, out_p, new_fused], dim=1)\n        else:\n            x_out = torch.cat([cls, out_p], dim=1)\n\n        max_len    = x_out.shape[1]\n        alive_mask = (torch.arange(max_len, device=device).unsqueeze(0)\n                      < (keep_counts + 1).unsqueeze(1))\n        alive_mask[:, 0] = True\n        if use_fused:\n            alive_mask[:, -1] = True\n\n        avg_keep = keep_counts.float().mean().item()\n        return x_out, alive_mask, keep_counts, avg_keep, \\\n               keep_counts.min().item(), keep_counts.max().item()\n\n    # Adaptive path (effective_rho > 0)\n    keep_ratios = compute_adaptive_keep_ratios(score_norm, layer_idx, rho, rc)\n    keep_counts = torch.round(keep_ratios * rc.float()).int().clamp(min=min_keep)\n    keep_counts = torch.min(keep_counts, rc).clamp(min=min_keep, max=max_rc)\n    max_keep    = keep_counts.max().item()\n\n    cls = x[:, 0:1, :]\n\n    if has_old_fused:\n        real_patches   = x[:, 1 : 1 + max_rc, :]\n        old_fused_tok  = x[:, 1 + max_rc :, :]\n        cls_attn_real  = attn_weights[:, 0, 1 : 1 + max_rc]\n        old_fused_attn = attn_weights[:, 0, 1 + max_rc : N]\n    else:\n        real_patches   = x[:, 1:, :]\n        old_fused_tok  = None\n        cls_attn_real  = attn_weights[:, 0, 1:]\n        old_fused_attn = None\n\n    imp        = cls_attn_real / (cls_attn_real.sum(1, keepdim=True) + 1e-6)\n    sorted_idx = imp.argsort(dim=1, descending=True)\n    rank_mask  = (torch.arange(max_rc, device=device).unsqueeze(0)\n                  < keep_counts.unsqueeze(1))\n    alive_real = torch.zeros(B, max_rc, dtype=torch.bool, device=device)\n    alive_real.scatter_(1, sorted_idx, rank_mask)\n\n    if use_fused:\n        drop_attn = cls_attn_real * (~alive_real).float()\n        if has_old_fused:\n            all_attn   = torch.cat([drop_attn, old_fused_attn], dim=1)\n            all_tokens = torch.cat([real_patches, old_fused_tok], dim=1)\n        else:\n            all_attn   = drop_attn\n            all_tokens = real_patches\n\n        dw        = all_attn / (all_attn.sum(1, keepdim=True) + 1e-9)\n        new_fused = torch.nan_to_num(\n            (dw.unsqueeze(-1) * all_tokens).sum(1, keepdim=True), nan=0.0)\n\n    pos    = torch.where(\n        alive_real,\n        torch.arange(max_rc, device=device).unsqueeze(0).expand(B, -1),\n        torch.full((B, max_rc), max_rc, device=device, dtype=torch.long))\n    top_sp = pos.sort(dim=1).values[:, :max_keep]\n    valid  = (top_sp < max_rc)\n    g_idx  = top_sp.clamp(0, max_rc - 1)\n    out_p  = real_patches.gather(\n        1, g_idx.unsqueeze(-1).expand(B, max_keep, C)\n    ) * valid.unsqueeze(-1).to(real_patches.dtype)\n\n    if use_fused:\n        x_out = torch.cat([cls, out_p, new_fused], dim=1)\n    else:\n        x_out = torch.cat([cls, out_p], dim=1)\n\n    max_len    = x_out.shape[1]\n    alive_mask = (torch.arange(max_len, device=device).unsqueeze(0)\n                  < (keep_counts + 1).unsqueeze(1))\n    alive_mask[:, 0] = True\n    if use_fused:\n        alive_mask[:, -1] = True\n\n    avg_keep = keep_counts.float().mean().item()\n    return x_out, alive_mask, keep_counts, avg_keep, \\\n           keep_counts.min().item(), keep_counts.max().item()\n\n\ndef fixed_topk_reduce(x, attn_weights, keep_ratio, min_keep, entering_rc,\n                      add_fused=False):\n    B, N, C = x.shape\n    device  = x.device\n\n    rc          = entering_rc.to(device=device, dtype=torch.int32)\n    keep_counts = torch.round(keep_ratio * rc.float()).int().clamp(min=min_keep)\n    keep_counts = torch.min(keep_counts, rc).clamp(min=min_keep)\n    max_keep    = keep_counts.max().item()\n\n    cls      = x[:, 0:1, :]\n    patches  = x[:, 1:1 + rc.max().item(), :]\n    Np       = patches.shape[1]\n    cls_attn = attn_weights[:, 0, 1 : 1 + Np]\n    imp      = cls_attn / (cls_attn.sum(1, keepdim=True) + 1e-6)\n\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_p = torch.zeros(B, Np, dtype=torch.bool, device=device)\n    alive_p.scatter_(1, sorted_idx, rank_mask)\n\n    if add_fused:\n        drop_attn = cls_attn * (~alive_p).float()\n        dw        = drop_attn / (drop_attn.sum(1, keepdim=True) + 1e-9)\n        new_fused = torch.nan_to_num(\n            (dw.unsqueeze(-1) * patches).sum(1, keepdim=True), nan=0.0)\n\n    pos    = torch.where(\n        alive_p,\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    out_p  = patches.gather(\n        1, g_idx.unsqueeze(-1).expand(B, max_keep, C)\n    ) * valid.unsqueeze(-1).to(patches.dtype)\n\n    if add_fused:\n        x_out = torch.cat([cls, out_p, new_fused], dim=1)\n    else:\n        x_out = torch.cat([cls, out_p], dim=1)\n\n    max_len    = x_out.shape[1]\n    alive_mask = (torch.arange(max_len, device=device).unsqueeze(0)\n                  < (keep_counts + 1).unsqueeze(1))\n    alive_mask[:, 0] = True\n    if add_fused:\n        alive_mask[:, -1] = True\n\n    avg_keep = keep_counts.float().mean().item()\n    return x_out, alive_mask, keep_counts, avg_keep, \\\n           keep_counts.min().item(), keep_counts.max().item()\n\n\nprint(\"Token reduction functions ready.\")\n\n\n# ── CELL 12: GFLOPS ESTIMATION ────────────────────────────────────────────────\n\ndef block_gflops(n):\n    d = CONFIG[\"EMBED_DIM\"]\n    return (4 * n * d**2 + 2 * n**2 * d + 8 * n * d**2) / 1e9\n\ndef block_gflops_split(nm, nf):\n    \"\"\"\n    Split-block cost: MHSA operates on pre-prune token count (nm),\n    MLP on post-prune count (nf). Consistent with EViT/ToMe reporting.\n    \"\"\"\n    d = CONFIG[\"EMBED_DIM\"]\n    return ((4 * nm * d**2 + 2 * nm**2 * d) + (8 * nf * d**2)) / 1e9\n\ndef baseline_gflops():\n    return sum(block_gflops(CONFIG[\"N_TOKENS\"]) for _ in range(CONFIG[\"N_LAYERS\"]))\n\ndef compute_multilayer_gflops(avg_keeps, use_fused=True):\n    fused_bonus = 2 if use_fused else 1\n    prune_set   = set(CONFIG[\"PRUNE_LAYERS\"])\n    cur         = float(CONFIG[\"N_TOKENS\"])\n    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 + fused_bonus\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_base_gf = baseline_gflops()\nprint(f\"GFLOPs ready.  Baseline ({CONFIG['MODEL_NAME']}): {_base_gf:.3f}\")\n\n\n# ── CELL 13: ROUTING TRACKER ──────────────────────────────────────────────────\n\nclass RoutingTracker:\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      = CONFIG[\"BASE_SCHEDULE\"][li]\n        # Use np.round to match torch.round rounding convention in the model.\n        fixed_ref = np.maximum(MIN_KEEP,\n                               np.round(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        self._decisions[li].append(\n            np.where(delta < 0, 0, np.where(delta > 0, 2, 1)))\n\n    def summary(self, rho):\n        if not CONFIG[\"TRACK_ROUTING\"]:\n            return\n        print(f\"\\n  -- Routing (rho={rho}) --\")\n        print(f\"  {'Layer':>6}  {'Below':>12}  {'Equal':>12}  {'Above':>12}  {'AvgΔ':>7}\")\n        print(f\"  {'─' * 62}\")\n        for li in CONFIG[\"PRUNE_LAYERS\"]:\n            if not self._decisions[li]:\n                continue\n            dec = np.concatenate(self._decisions[li])\n            dlt = np.concatenate(self._deltas[li])\n            nb, ne, na = (dec == 0).sum(), (dec == 1).sum(), (dec == 2).sum()\n            n = len(dec)\n            print(f\"  {li:>6}  {nb:>7}({nb/n*100:>4.1f}%)  \"\n                  f\"{ne:>7}({ne/n*100:>4.1f}%)  \"\n                  f\"{na:>7}({na/n*100:>4.1f}%)  \"\n                  f\"{dlt.mean():>+6.1f}\")\n\nprint(\"RoutingTracker ready.\")\n\n\n# ── CELL 14: SCORE-NORM DISTRIBUTION TRACKER ─────────────────────────────────\n\nclass ScoreNormTracker:\n    def __init__(self):\n        self._above = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n\n    def record(self, li, score_norm):\n        if not CONFIG[\"TRACK_SCORE_DIST\"]:\n            return\n        self._above[li].append((score_norm > 0.5).float().mean().item())\n\n    def summary(self):\n        if not CONFIG[\"TRACK_SCORE_DIST\"]:\n            return\n        print(\"\\n  -- Score-norm distribution (ideal ≈ 50% above 0.5) --\")\n        print(f\"  {'Layer':>6}  {'%Above':>8}  {'%Below':>8}  Status\")\n        print(f\"  {'─' * 42}\")\n        for li in CONFIG[\"PRUNE_LAYERS\"]:\n            if not self._above[li]:\n                continue\n            pct_a = np.mean(self._above[li]) * 100\n            pct_b = 100 - pct_a\n            dev   = abs(pct_a - 50)\n            status = \"✓ good\" if dev < 5 else (\"~ ok\" if dev < 10 else \"⚠ skewed\")\n            print(f\"  {li:>6}  {pct_a:>7.1f}%  {pct_b:>7.1f}%  {status}\")\n\nprint(\"ScoreNormTracker ready.\")\n\n\n# ── CELL 15: SHARED BLOCK FORWARD BUILDERS ────────────────────────────────────\n\ndef make_normal_block(block):\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    def forward(x, **kwargs):\n        B, N, _ = x.shape\n        alive   = state[\"alive_mask\"]\n\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\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\nprint(\"Shared block builders ready.\")\n\n\n# ── CELL 16: PASS 0 — CALIBRATION ────────────────────────────────────────────\n#\n# The focus score at each pruning layer is computed from that layer's own\n# attention weights, after MHSA runs on the full pre-prune token sequence.\n# This provides the richest available attention context for the scoring signal.\n\ndef run_pass0_calibration(loader):\n    reset_layer_cal()\n\n    device      = CONFIG[\"DEVICE\"]\n    cal_batches = CONFIG[\"CAL_BATCHES\"]\n    prune_set   = set(CONFIG[\"PRUNE_LAYERS\"])\n\n    cal_scores = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n    logged     = {li: False for li in CONFIG[\"PRUNE_LAYERS\"]}\n\n    cal_state = {\n        \"entering_rc\": None,\n        \"has_fused\"  : False,\n        \"alive_mask\" : None,\n    }\n\n    def _reconcile(alive, B, N, dev):\n        if alive is None or alive.shape[1] == N:\n            return alive\n        if alive.shape[1] > N:\n            return alive[:, :N]\n        return torch.cat([alive,\n                          torch.zeros(B, N - alive.shape[1],\n                                      dtype=torch.bool, device=dev)], dim=1)\n\n    def make_prune_cal(block, li):\n        keep_r = CONFIG[\"BASE_SCHEDULE\"][li]\n\n        def forward(x, **kwargs):\n            B, N, _ = x.shape\n            alive   = _reconcile(cal_state[\"alive_mask\"], B, N, x.device)\n            cal_state[\"alive_mask\"] = alive\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            Np = N - 1\n            s  = compute_cls_focus_score(\n                block.attn.attn_weights, Np, has_fused=cal_state[\"has_fused\"])\n            cal_scores[li].append(s.detach().cpu().numpy())\n\n            if cal_state[\"entering_rc\"] is None:\n                rc_t = torch.full((B,), CONFIG[\"N_PATCHES\"],\n                                   dtype=torch.int32, device=x.device)\n            else:\n                rc_t = cal_state[\"entering_rc\"]\n\n            x_out, alive_mask, keep_counts, avg_k, kmin, kmax = fixed_topk_reduce(\n                x, block.attn.attn_weights, keep_r, CONFIG[\"MIN_KEEP\"],\n                rc_t, add_fused=CONFIG[\"USE_FUSED\"])\n\n            cal_state[\"entering_rc\"] = keep_counts\n            cal_state[\"has_fused\"]   = CONFIG[\"USE_FUSED\"]\n            cal_state[\"alive_mask\"]  = alive_mask\n\n            if not logged[li]:\n                print(f\"  Cal L{li}: seq {N} -> {x_out.shape[1]}  \"\n                      f\"avg={avg_k:.0f}  [{kmin},{kmax}]\")\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\n        return forward\n\n    model = load_model()\n    patch_attention(model)\n\n    first_prune = min(CONFIG[\"PRUNE_LAYERS\"])\n    for idx, blk in enumerate(model.blocks):\n        if idx in prune_set:\n            blk.forward = make_prune_cal(blk, idx)\n        elif idx > first_prune:\n            blk.forward = make_post_prune_block(blk, cal_state)\n        else:\n            blk.forward = make_normal_block(blk)\n\n    print(f\"\\nPass 0 — Calibration  \"\n          f\"({cal_batches} x {CONFIG['BATCH_SIZE']} = \"\n          f\"{cal_batches * CONFIG['BATCH_SIZE']:,} images)\")\n\n    with torch.no_grad():\n        for bi, (x, _) in enumerate(\n                tqdm(loader, desc=\"  Calibrating\", total=cal_batches)):\n            if bi >= cal_batches:\n                break\n            cal_state[\"entering_rc\"] = None\n            cal_state[\"has_fused\"]   = False\n            cal_state[\"alive_mask\"]  = None\n            model(x.to(device))\n\n    del model\n    clear_gpu()\n\n    print(f\"\\n  {'L(prune)':<18}  {'n':>7}  {'median':>8}  \"\n          f\"{'p25':>7}  {'p75':>7}  {'%>median':>9}\")\n    print(f\"  {'─' * 64}\")\n\n    for li in CONFIG[\"PRUNE_LAYERS\"]:\n        arr = np.concatenate(cal_scores[li]) if cal_scores[li] else np.array([0.5])\n        if cal_scores[li]:\n            LAYER_CAL_SORTED[li] = np.sort(arr)\n        else:\n            LAYER_CAL_SORTED[li] = None\n            print(f\"  WARNING: no scores collected for L{li}.\")\n            continue\n\n        med = float(np.median(arr))\n        p25 = float(np.percentile(arr, 25))\n        p75 = float(np.percentile(arr, 75))\n        pct = float((arr > med).mean() * 100)\n        print(f\"  L{li:<16}  {len(arr):>7,}  {med:>8.4f}  \"\n              f\"{p25:>7.4f}  {p75:>7.4f}  {pct:>8.1f}%\")\n\n    print(\"  LAYER_CAL_SORTED updated.\")\n\n\nprint(\"Pass 0 function ready.\")\n\n\n# ── CELL 17: PASS 1 — FULL-MODEL BASELINE ────────────────────────────────────\n\ndef run_pass1_baseline(loader):\n    device = CONFIG[\"DEVICE\"]\n    warmup = CONFIG[\"WARMUP_BATCHES\"]\n    top1 = top5 = total = 0\n    tt = 0.0; ti = 0\n\n    model = load_model()\n    patch_attention(model)\n    for blk in model.blocks:\n        blk.forward = make_normal_block(blk)\n\n    print(\"\\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            is_warmup = bi < warmup\n            if not is_warmup:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                ts = time.perf_counter()\n            out = model(x)\n            if not is_warmup:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                tt += time.perf_counter() - ts\n                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    del model; clear_gpu()\n\n    acc1 = top1 / total * 100\n    acc5 = top5 / total * 100\n    fps  = ti / tt        if tt > 0 else 0.0\n    lat  = tt / ti * 1000 if ti > 0 else 0.0\n    print(f\"  Baseline  Top-1={acc1:.2f}%  Top-5={acc5:.2f}%  FPS={fps:.1f}  Lat={lat:.2f}ms\")\n    return acc1, acc5, fps, lat\n\n\nprint(\"Pass 1 function ready.\")\n\n\n# ── CELL 18: PASS 2a — FIXED TOP-K BASELINE ──────────────────────────────────\n\ndef run_pass2a_fixed(loader, base_gf, use_fused=False):\n    device      = CONFIG[\"DEVICE\"]\n    warmup      = CONFIG[\"WARMUP_BATCHES\"]\n    prune_set   = set(CONFIG[\"PRUNE_LAYERS\"])\n    first_prune = min(CONFIG[\"PRUNE_LAYERS\"])\n\n    state = {\n        \"alive_mask\" : None,\n        \"entering_rc\": None,\n        \"kc\": {li: {\"sum\": 0, \"min\": 9999, \"max\": 0, \"n\": 0}\n               for li in CONFIG[\"PRUNE_LAYERS\"]},\n    }\n    logged = {li: False for li in CONFIG[\"PRUNE_LAYERS\"]}\n    fixed_correct_log = []\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_fixed(block, li):\n        keep_r = CONFIG[\"BASE_SCHEDULE\"][li]\n\n        def forward(x, **kwargs):\n            B, N, _ = x.shape\n            alive   = state[\"alive_mask\"]\n            if alive is not None and alive.shape[1] != N:\n                alive = alive[:, :N] if alive.shape[1] > N else \\\n                        torch.cat([alive, torch.zeros(B, N - alive.shape[1],\n                                   dtype=torch.bool, device=x.device)], dim=1)\n                state[\"alive_mask\"] = alive\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            rc_t = (state[\"entering_rc\"]\n                    if state[\"entering_rc\"] is not None\n                    else torch.full((B,), CONFIG[\"N_PATCHES\"],\n                                    dtype=torch.int32, device=x.device))\n\n            x_out, alive_mask, keep_counts, avg_k, kmin, kmax = fixed_topk_reduce(\n                x, block.attn.attn_weights, keep_r, CONFIG[\"MIN_KEEP\"],\n                rc_t, add_fused=use_fused)\n\n            state[\"alive_mask\"]  = alive_mask\n            state[\"entering_rc\"] = keep_counts\n            _upd(li, keep_counts.cpu().numpy())\n\n            if not logged[li]:\n                print(f\"  Fixed L{li}: seq {N} -> {x_out.shape[1]}  \"\n                      f\"avg={avg_k:.0f}  [{kmin},{kmax}]  ratio={keep_r}\")\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\n        return forward\n\n    top1 = top5 = total = 0\n    tt = 0.0; ti = 0\n\n    model = load_model()\n    patch_attention(model)\n    for idx, blk in enumerate(model.blocks):\n        if idx in prune_set:\n            blk.forward = make_fixed(blk, idx)\n        elif idx > first_prune:\n            blk.forward = make_post_prune_block(blk, state)\n        else:\n            blk.forward = make_normal_block(blk)\n\n    mode = \"with fused token\" if use_fused else \"NO fused token\"\n    print(f\"\\nPass 2a — Fixed Top-K  {CONFIG['BASE_SCHEDULE']}  ({mode})\")\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            state[\"alive_mask\"]  = None\n            state[\"entering_rc\"] = None\n\n            is_warmup = bi < warmup\n            if not is_warmup:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                ts = time.perf_counter()\n            out = model(x)\n            if not is_warmup:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                tt += time.perf_counter() - ts\n                ti += x.shape[0]\n\n            correct1 = (out.argmax(1) == y).cpu().numpy()\n            fixed_correct_log.extend(correct1.tolist())\n\n            top1  += correct1.sum()\n            top5  += (out.topk(5, 1).indices ==\n                      y.unsqueeze(1)).any(1).sum().item()\n            total += y.shape[0]\n\n    del model; clear_gpu()\n\n    acc1 = top1 / total * 100\n    acc5 = top5 / total * 100\n    fps  = ti / tt        if tt > 0 else 0.0\n    lat  = tt / ti * 1000 if ti > 0 else 0.0\n\n    avg_keeps = {li: state[\"kc\"][li][\"sum\"] / max(state[\"kc\"][li][\"n\"], 1)\n                 for li in CONFIG[\"PRUNE_LAYERS\"]}\n    fixed_gf  = compute_multilayer_gflops(avg_keeps, use_fused=True)\n\n    fixed_correct_arr = np.array(fixed_correct_log, dtype=np.bool_)\n    np.save(str(OUTPUT_DIR / \"fixed_correct.npy\"), fixed_correct_arr)\n    print(f\"  [Saved] fixed_correct.npy  \"\n          f\"({fixed_correct_arr.sum():,} / {len(fixed_correct_arr):,} correct)\")\n\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}  (baseline delta={fixed_gf - base_gf:+.3f})\")\n    print(\"  Avg keep: \" +\n          \" \".join(f\"L{li}={avg_keeps[li]:.1f}\" for li in CONFIG[\"PRUNE_LAYERS\"]))\n    return acc1, acc5, fps, lat, avg_keeps, fixed_gf\n\n\nprint(\"Pass 2a function ready.\")\n\n\n# ── CELL 19: PASS 3 — ADAPTIVE ASPP-ViT ──────────────────────────────────────\n\ndef run_pass3_adaptive(loader, rho, fixed_gf, use_fused=None):\n    use_fused_val = CONFIG[\"USE_FUSED\"] if use_fused is None else use_fused\n\n    device      = CONFIG[\"DEVICE\"]\n    warmup      = CONFIG[\"WARMUP_BATCHES\"]\n    prune_set   = set(CONFIG[\"PRUNE_LAYERS\"])\n    first_prune = min(CONFIG[\"PRUNE_LAYERS\"])\n    logged      = {li: False for li in CONFIG[\"PRUNE_LAYERS\"]}\n\n    state = {\n        \"alive_mask\"        : None,\n        \"entering_rc\"       : None,\n        \"has_fused\"         : False,\n        \"kc\"                : {li: {\"sum\": 0, \"min\": 9999, \"max\": 0, \"n\": 0}\n                               for li in CONFIG[\"PRUNE_LAYERS\"]},\n        \"batch_score_norm\"  : {li: None for li in CONFIG[\"PRUNE_LAYERS\"]},\n        \"batch_keep\"        : {li: None for li in CONFIG[\"PRUNE_LAYERS\"]},\n        \"batch_entering_rc\" : {li: None for li in CONFIG[\"PRUNE_LAYERS\"]},\n    }\n\n    score_log       = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n    keep_log        = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n    entering_rc_log = {li: [] for li in CONFIG[\"PRUNE_LAYERS\"]}\n    correct_log     = []\n    topk_log        = []\n\n    rt = RoutingTracker()   if CONFIG[\"TRACK_ROUTING\"]    else None\n    sn = ScoreNormTracker() if CONFIG[\"TRACK_SCORE_DIST\"] 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 _reconcile(alive, B, N, dev):\n        if alive is None or alive.shape[1] == N:\n            return alive\n        if alive.shape[1] > N:\n            return alive[:, :N]\n        return torch.cat([alive,\n                          torch.zeros(B, N - alive.shape[1],\n                                      dtype=torch.bool, device=dev)], dim=1)\n\n    def make_adaptive(block, li):\n        def forward(x, **kwargs):\n            B, N, _ = x.shape\n            alive   = _reconcile(state[\"alive_mask\"], B, N, x.device)\n            state[\"alive_mask\"] = alive\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            Np         = N - 1\n            score_raw  = compute_cls_focus_score(\n                block.attn.attn_weights, Np, has_fused=state[\"has_fused\"])\n            score_norm = normalize_score(score_raw, li)\n            state[\"batch_score_norm\"][li] = score_norm.detach().cpu()\n\n            if sn is not None:\n                sn.record(li, score_norm)\n\n            rc_t = (state[\"entering_rc\"]\n                    if state[\"entering_rc\"] is not None\n                    else torch.full((B,), CONFIG[\"N_PATCHES\"],\n                                    dtype=torch.int32, device=x.device))\n\n            # Store entering_rc before it is overwritten by the prune step.\n            state[\"batch_entering_rc\"][li] = rc_t.detach().cpu()\n\n            x_out, alive_mask, keep_counts, avg_k, kmin, kmax = adaptive_topk_reduce(\n                x, block.attn.attn_weights, score_norm,\n                li, rho, CONFIG[\"MIN_KEEP\"], use_fused_val,\n                rc_t, has_old_fused=state[\"has_fused\"])\n\n            state[\"batch_keep\"][li]  = keep_counts.detach().cpu()\n            state[\"entering_rc\"]     = keep_counts\n            state[\"has_fused\"]       = use_fused_val\n            state[\"alive_mask\"]      = alive_mask\n\n            kc_np = keep_counts.detach().cpu().numpy()\n            rc_np = rc_t.detach().cpu().numpy()\n            _upd(li, kc_np)\n            if rt is not None:\n                rt.record(li, kc_np, rc_np)\n\n            if not logged[li]:\n                scale_d  = CONFIG.get(\"RHO_SCALE\", {})\n                eff_rho  = rho * scale_d.get(li, 1.0)\n                fixed_ref = max(CONFIG[\"MIN_KEEP\"],\n                                int(rc_np.mean() * CONFIG[\"BASE_SCHEDULE\"][li]))\n                print(f\"  Adaptive L{li} (rho={rho}, eff_rho={eff_rho:.3f}, \"\n                      f\"fused={use_fused_val}): seq {N} -> out={x_out.shape[1]}  \"\n                      f\"avg={avg_k:.0f}  [{kmin},{kmax}]  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\n        return forward\n\n    top1 = top5 = total = 0\n    tt = 0.0; ti = 0\n\n    model = load_model()\n    patch_attention(model)\n    for idx, blk in enumerate(model.blocks):\n        if idx in prune_set:\n            blk.forward = make_adaptive(blk, idx)\n        elif idx > first_prune:\n            blk.forward = make_post_prune_block(blk, state)\n        else:\n            blk.forward = make_normal_block(blk)\n\n    label = f\"rho={rho}\" + (\"\" if use_fused_val else \" [no-fused]\")\n    print(f\"\\nPass 3 — Adaptive {label}\")\n    with torch.no_grad():\n        for bi, (x, y) in enumerate(tqdm(loader, desc=f\"  {label}\")):\n            x, y = x.to(device), y.to(device)\n            state[\"alive_mask\"]  = None\n            state[\"entering_rc\"] = None\n            state[\"has_fused\"]   = False\n            for li in CONFIG[\"PRUNE_LAYERS\"]:\n                state[\"batch_score_norm\"][li]  = None\n                state[\"batch_keep\"][li]        = None\n                state[\"batch_entering_rc\"][li] = None\n\n            is_warmup = bi < warmup\n            if not is_warmup:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                ts = time.perf_counter()\n\n            out = model(x)\n\n            if not is_warmup:\n                if device.type == \"cuda\": torch.cuda.synchronize()\n                tt += time.perf_counter() - ts\n                ti += x.shape[0]\n\n            correct1 = (out.argmax(1) == y).cpu().numpy()\n            correct5 = (out.topk(5, 1).indices ==\n                        y.unsqueeze(1)).any(1).cpu().numpy()\n            correct_log.extend(correct1.tolist())\n            topk_log.extend(correct5.tolist())\n            top1  += correct1.sum()\n            top5  += correct5.sum()\n            total += y.shape[0]\n\n            for li in CONFIG[\"PRUNE_LAYERS\"]:\n                if state[\"batch_score_norm\"][li]  is not None:\n                    score_log[li].extend(state[\"batch_score_norm\"][li].numpy().tolist())\n                if state[\"batch_keep\"][li]         is not None:\n                    keep_log[li].extend(state[\"batch_keep\"][li].numpy().tolist())\n                if state[\"batch_entering_rc\"][li]  is not None:\n                    entering_rc_log[li].extend(\n                        state[\"batch_entering_rc\"][li].numpy().tolist())\n\n    del model; clear_gpu()\n\n    acc1 = top1 / total * 100\n    acc5 = top5 / total * 100\n    fps  = ti / tt        if tt > 0 else 0.0\n    lat  = tt / ti * 1000 if ti > 0 else 0.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, use_fused=use_fused_val)\n\n    print(f\"\\n  {label}  Top-1={acc1:.2f}%  Top-5={acc5:.2f}%  \"\n          f\"FPS={fps:.1f}  Lat={lat:.2f}ms\")\n    print(f\"  GFLOPs: {ada_gf:.4f}  (Δ vs fixed={ada_gf - fixed_gf:+.4f})\")\n    print(f\"  Note: GFLOPs tracks average compute. Hard images receive real \"\n          f\"extra tokens (see token-allocation analysis below). \"\n          f\"Overall average delta is near-neutral by design.\")\n\n    print(\"  Avg keep: \" +\n          \" \".join(f\"L{li}={avg_keeps[li]:.1f}\" for li in CONFIG[\"PRUNE_LAYERS\"]))\n\n    if rt is not None: rt.summary(rho)\n    if sn is not None: sn.summary()\n\n    _print_per_image_diagnostic(\n        score_log, keep_log, entering_rc_log, correct_log, topk_log, label, rho)\n\n    rho_str   = str(rho).replace(\".\", \"p\")\n    fused_str = \"fused\" if use_fused_val else \"nofused\"\n    tag       = f\"{rho_str}_{fused_str}\"\n    last_li   = CONFIG[\"PRUNE_LAYERS\"][-1]\n\n    np.save(str(OUTPUT_DIR / f\"ada_correct_{tag}.npy\"),\n            np.array(correct_log, dtype=np.bool_))\n    np.save(str(OUTPUT_DIR / f\"ada_scores_L{last_li}_{tag}.npy\"),\n            np.array(score_log[last_li], dtype=np.float32))\n    np.save(str(OUTPUT_DIR / f\"ada_keeps_L{last_li}_{tag}.npy\"),\n            np.array(keep_log[last_li], dtype=np.float32))\n    np.save(str(OUTPUT_DIR / f\"ada_entering_rc_L{last_li}_{tag}.npy\"),\n            np.array(entering_rc_log[last_li], dtype=np.float32))\n    print(f\"  [Saved] per-image arrays  tag={tag}\")\n\n    return acc1, acc5, fps, lat, avg_keeps, ada_gf\n\n\ndef _print_per_image_diagnostic(score_log, keep_log, entering_rc_log,\n                                 correct_log, topk_log, label, rho):\n    correct = np.array(correct_log, dtype=np.float32)\n    topk    = np.array(topk_log,    dtype=np.float32)\n    n_total = len(correct)\n\n    if n_total == 0:\n        return\n\n    print(f\"\\n  {'─'*76}\")\n    print(f\"  PER-IMAGE DIAGNOSTIC  [{label}]  (n={n_total:,})\")\n    print(f\"  {'─'*76}\")\n\n    for li in CONFIG[\"PRUNE_LAYERS\"]:\n        scores      = np.array(score_log[li])      if score_log[li]      else None\n        keeps       = np.array(keep_log[li])        if keep_log[li]       else None\n        entering_rc = np.array(entering_rc_log[li]) if entering_rc_log[li] else None\n\n        if scores is None or len(scores) != n_total:\n            print(f\"  L{li}: score log mismatch, skipping\")\n            continue\n\n        # np.round matches torch.round rounding convention used in model.\n        if entering_rc is not None:\n            base_keep_per_img = np.round(\n                CONFIG[\"BASE_SCHEDULE\"][li] * entering_rc.astype(float)\n            ).astype(int)\n            base_keep_mean = base_keep_per_img.mean()\n        else:\n            base_keep_per_img = None\n            base_keep_mean    = CONFIG[\"BASE_SCHEDULE\"][li] * CONFIG[\"N_PATCHES\"]\n\n        q25, q50, q75 = np.percentile(scores, [25, 50, 75])\n\n        print(f\"\\n  Layer {li}  \"\n              f\"(avg_base_keep={base_keep_mean:.1f}  \"\n              f\"score p25={q25:.3f} p50={q50:.3f} p75={q75:.3f})\")\n        print(f\"  {'Quartile':<26}  {'n':>6}  {'Top-1':>7}  {'Top-5':>7}  \"\n              f\"{'AvgKeep':>8}  {'ΔKeep':>7}  {'ΔKeep%':>7}\")\n        print(f\"  {'─'*72}\")\n\n        quartiles = [\n            (\"Q1 hard  (score 0–25%)\",   scores <  q25),\n            (\"Q2       (score 25–50%)\",  (scores >= q25) & (scores < q50)),\n            (\"Q3       (score 50–75%)\",  (scores >= q50) & (scores < q75)),\n            (\"Q4 easy  (score 75–100%)\", scores >= q75),\n        ]\n\n        for qname, mask in quartiles:\n            if mask.sum() == 0:\n                continue\n            q_acc1 = correct[mask].mean() * 100\n            q_acc5 = topk[mask].mean()    * 100\n            q_keep = keeps[mask].mean()   if keeps is not None else float(\"nan\")\n            if base_keep_per_img is not None:\n                q_base  = base_keep_per_img[mask].mean()\n                q_delta = q_keep - q_base\n                q_dpct  = q_delta / q_base * 100 if q_base > 0 else 0.0\n            else:\n                q_delta = float(\"nan\")\n                q_dpct  = float(\"nan\")\n            print(f\"  {qname:<26}  {mask.sum():>6,}  {q_acc1:>6.2f}%  \"\n                  f\"{q_acc5:>6.2f}%  {q_keep:>8.1f}  {q_delta:>+7.1f}  \"\n                  f\"{q_dpct:>+6.1f}%\")\n\n        hard_thresh = CONFIG[\"HARD_THRESHOLDS\"].get(li, 0.15)\n        easy_thresh = CONFIG[\"EASY_THRESHOLDS\"].get(li, 0.85)\n        easy_bin = scores > easy_thresh\n        hard_bin = scores < hard_thresh\n        mid_bin  = ~easy_bin & ~hard_bin\n\n        print(f\"\\n  L{li} coarse bin split \"\n              f\"(hard<{hard_thresh}  mid  easy>{easy_thresh}):\")\n        for bname, bmask in [\n                (f\"Hard bin (score<{hard_thresh})\", hard_bin),\n                (f\"Mid  bin ({hard_thresh}–{easy_thresh})\", mid_bin),\n                (f\"Easy bin (score>{easy_thresh})\", easy_bin)]:\n            if bmask.sum() == 0:\n                continue\n            b_acc  = correct[bmask].mean() * 100\n            b_keep = keeps[bmask].mean()   if keeps is not None else float(\"nan\")\n            if base_keep_per_img is not None:\n                b_ref   = base_keep_per_img[bmask].mean()\n                b_delta = b_keep - b_ref\n            else:\n                b_delta = float(\"nan\")\n            print(f\"    {bname:<30}  n={bmask.sum():>6,}  \"\n                  f\"Top-1={b_acc:.2f}%  avg_keep={b_keep:.1f}  \"\n                  f\"Δ={b_delta:+.1f}\")\n\n        # np.round used for base_keep_per_img; mid_bin used as fallback reference\n        # when got_same is too small for reliable accuracy estimation.\n        if keeps is not None and base_keep_per_img is not None and rho > 0:\n            got_more = keeps > base_keep_per_img\n            got_less = keeps < base_keep_per_img\n            got_same = keeps == base_keep_per_img\n\n            if got_more.sum() > 100 and got_less.sum() > 100:\n                acc_more = correct[got_more].mean() * 100\n                acc_less = correct[got_less].mean() * 100\n\n                if got_same.sum() > 50:\n                    acc_same    = correct[got_same].mean() * 100\n                    mid_ref_src = \"got_same\"\n                elif mid_bin.sum() > 50:\n                    acc_same    = correct[mid_bin].mean() * 100\n                    mid_ref_src = \"mid_bin(fallback)\"\n                else:\n                    acc_same    = float(\"nan\")\n                    mid_ref_src = \"unavailable\"\n\n                delta_vs_mid = (acc_less - acc_same\n                                if not np.isnan(acc_same) else float(\"nan\"))\n\n                print(f\"\\n  L{li} token-allocation analysis (rho={rho}):\")\n                print(f\"    Got BONUS (hard, n={got_more.sum():,}): \"\n                      f\"Top-1={acc_more:.2f}%  \"\n                      f\"avg_keep={keeps[got_more].mean():.1f}  \"\n                      f\"avg_base={base_keep_per_img[got_more].mean():.1f}\")\n                print(f\"    Got SAME  (mid,  n={got_same.sum():,}  \"\n                      f\"ref={mid_ref_src}): Top-1={acc_same:.2f}%\")\n                print(f\"    Got CUT   (easy, n={got_less.sum():,}): \"\n                      f\"Top-1={acc_less:.2f}%  \"\n                      f\"avg_keep={keeps[got_less].mean():.1f}  \"\n                      f\"avg_base={base_keep_per_img[got_less].mean():.1f}\")\n                print(f\"    Easy images tolerated cut: \"\n                      f\"acc_cut={acc_less:.2f}% vs mid_ref={acc_same:.2f}%  \"\n                      f\"Δ={delta_vs_mid:+.2f}%\")\n                print(f\"    [Target: Δ > -0.30%  (easy images lose <0.3% acc)]\")\n\n    print(f\"\\n  {'─'*76}\")\n    print(f\"  SCORE–ACCURACY CORRELATION  [{label}]\")\n    print(f\"  {'─'*76}\")\n\n    last_li = CONFIG[\"PRUNE_LAYERS\"][-1]\n    if score_log[last_li] and len(score_log[last_li]) == n_total:\n        scores_last = np.array(score_log[last_li])\n        corr = np.corrcoef(scores_last, correct)[0, 1]\n        print(f\"  Pearson corr(score_norm_L{last_li}, correct): {corr:+.4f}\")\n        print(f\"  Expected: +0.08 to +0.12  (weak but real, AUC~0.57)\")\n\n        print(f\"\\n  Accuracy by score_norm decile (L{last_li}):\")\n        print(f\"  {'Decile':<20}  {'score range':>14}  {'n':>6}  \"\n              f\"{'Top-1':>7}  {'ΔvsD5':>7}\")\n        d5_acc = None\n        decile_accs = []\n        for d in range(10):\n            lo   = np.percentile(scores_last, d * 10)\n            hi   = np.percentile(scores_last, (d + 1) * 10)\n            mask = (scores_last >= lo) & (scores_last < hi) if d < 9 \\\n                   else (scores_last >= lo)\n            d_acc = correct[mask].mean() * 100 if mask.sum() > 0 else None\n            decile_accs.append((lo, hi, mask, d_acc))\n            if d == 4:\n                d5_acc = d_acc\n\n        for d, (lo, hi, mask, d_acc) in enumerate(decile_accs):\n            if d_acc is None:\n                continue\n            tag_d = \"hard\" if d < 2 else (\"easy\" if d > 7 else \"    \")\n            dvs5  = f\"{d_acc - d5_acc:+.2f}%\" if d5_acc is not None else \"  —\"\n            print(f\"  D{d+1:02d} ({tag_d}) [{lo:.3f},{hi:.3f}]  \"\n                  f\"{mask.sum():>6,}  {d_acc:>6.2f}%  {dvs5:>7}\")\n\n        if decile_accs[0][3] is not None and decile_accs[-1][3] is not None:\n            spread = decile_accs[-1][3] - decile_accs[0][3]\n            print(f\"\\n  D1→D10 spread: {spread:+.2f}%  \"\n                  f\"({'strong enough for coarse bins' if spread > 8 else 'marginal signal'})\")\n\n    print(f\"  {'─'*76}\")\n\n\nprint(\"Pass 3 function ready.\")\n\n\ndef run_option_a_comparison(rho, use_fused=True):\n    rho_str   = str(rho).replace(\".\", \"p\")\n    fused_str = \"fused\" if use_fused else \"nofused\"\n    tag       = f\"{rho_str}_{fused_str}\"\n    last_li   = CONFIG[\"PRUNE_LAYERS\"][-1]\n\n    fixed_path = OUTPUT_DIR / \"fixed_correct.npy\"\n    ada_path   = OUTPUT_DIR / f\"ada_correct_{tag}.npy\"\n    score_path = OUTPUT_DIR / f\"ada_scores_L{last_li}_{tag}.npy\"\n    keep_path  = OUTPUT_DIR / f\"ada_keeps_L{last_li}_{tag}.npy\"\n    rc_path    = OUTPUT_DIR / f\"ada_entering_rc_L{last_li}_{tag}.npy\"\n\n    for p in [fixed_path, ada_path, score_path]:\n        if not p.exists():\n            print(f\"  Missing: {p.name} — run Pass 2a and Pass 3 first.\")\n            return\n\n    fixed_correct = np.load(str(fixed_path)).astype(np.float32)\n    ada_correct   = np.load(str(ada_path)).astype(np.float32)\n    scores        = np.load(str(score_path)).astype(np.float32)\n    keeps         = np.load(str(keep_path)).astype(np.float32) \\\n                    if keep_path.exists() else None\n    entering_rc   = np.load(str(rc_path)).astype(np.float32) \\\n                    if rc_path.exists() else None\n\n    n = len(fixed_correct)\n    assert len(ada_correct) == n and len(scores) == n\n\n    stayed_correct = (fixed_correct == 1) & (ada_correct == 1)\n    stayed_wrong   = (fixed_correct == 0) & (ada_correct == 0)\n    flipped_wrong  = (fixed_correct == 1) & (ada_correct == 0)\n    flipped_right  = (fixed_correct == 0) & (ada_correct == 1)\n\n    hard_thresh = CONFIG[\"HARD_THRESHOLDS\"].get(last_li, 0.15)\n    easy_thresh = CONFIG[\"EASY_THRESHOLDS\"].get(last_li, 0.85)\n    easy_bin    = scores > easy_thresh\n    hard_bin    = scores < hard_thresh\n    mid_bin     = ~easy_bin & ~hard_bin\n\n    SEP = \"=\" * 80\n\n    print(f\"\\n{SEP}\")\n    print(f\"  OPTION A — PER-IMAGE FLIP ANALYSIS  \"\n          f\"[rho={rho} fused={use_fused}]\")\n    print(f\"  Fixed baseline vs Adaptive  (n={n:,} images)\")\n    print(SEP)\n\n    print(f\"\\n  Overall flip counts:\")\n    print(f\"  {'Category':<34}  {'n':>7}  {'%':>6}\")\n    print(f\"  {'─'*48}\")\n    print(f\"  {'Both correct (TT)':<34}  \"\n          f\"{stayed_correct.sum():>7,}  {stayed_correct.mean()*100:>5.1f}%\")\n    print(f\"  {'Both wrong   (FF)':<34}  \"\n          f\"{stayed_wrong.sum():>7,}  {stayed_wrong.mean()*100:>5.1f}%\")\n    print(f\"  {'Fixed correct→Ada wrong (HURT)':<34}  \"\n          f\"{flipped_wrong.sum():>7,}  {flipped_wrong.mean()*100:>5.2f}%\")\n    print(f\"  {'Fixed wrong→Ada correct (HELPED)':<34}  \"\n          f\"{flipped_right.sum():>7,}  {flipped_right.mean()*100:>5.2f}%\")\n    net = int(flipped_right.sum()) - int(flipped_wrong.sum())\n    print(f\"  {'Net flip (helped - hurt)':<34}  {net:>+7,}\")\n    print(f\"  {'Net accuracy change':<34}  \"\n          f\"{(ada_correct.mean() - fixed_correct.mean())*100:>+6.2f}%\")\n\n    print(f\"\\n  Per-difficulty-bin breakdown:\")\n    print(f\"  {'Bin':<26}  {'n':>6}  {'Hurt(TF)':>10}  \"\n          f\"{'Helped(FT)':>11}  {'NetFlip':>8}  \"\n          f\"{'HurtRate':>9}  {'HelpRate':>9}  Note\")\n    print(f\"  {'─'*98}\")\n\n    bins = [\n        (f\"Hard (score<{hard_thresh})\",         hard_bin),\n        (f\"Mid  ({hard_thresh}–{easy_thresh})\",  mid_bin),\n        (f\"Easy (score>{easy_thresh})\",          easy_bin),\n    ]\n    for bname, bmask in bins:\n        if bmask.sum() == 0:\n            continue\n        n_bin     = bmask.sum()\n        n_hurt    = (flipped_wrong & bmask).sum()\n        n_helped  = (flipped_right & bmask).sum()\n        net_bin   = int(n_helped) - int(n_hurt)\n        hurt_rate = n_hurt   / n_bin * 100\n        help_rate = n_helped / n_bin * 100\n        if \"Easy\" in bname:\n            note = (\"✓ good\" if hurt_rate < 1.5 else\n                    \"⚠ caution\" if hurt_rate < 3.0 else \"✗ too many hurt\")\n        elif \"Hard\" in bname:\n            note = \"✓ good\" if help_rate > 0.2 else \"~ marginal\"\n        else:\n            note = \"~ expected\"\n        print(f\"  {bname:<26}  {n_bin:>6,}  {n_hurt:>8,}  \"\n              f\"{n_helped:>9,}  {net_bin:>+8,}  \"\n              f\"{hurt_rate:>8.2f}%  {help_rate:>8.2f}%  {note}\")\n\n    # np.round used to match torch.round convention for base_keep computation.\n    easy_cut = easy_bin.copy()\n    if keeps is not None and entering_rc is not None:\n        base_keep_per_img = np.round(\n            CONFIG[\"BASE_SCHEDULE\"][last_li] * entering_rc.astype(float)\n        ).astype(int)\n        easy_cut = easy_bin & (keeps < base_keep_per_img)\n\n    n_easy_cut  = easy_cut.sum()\n    n_easy_hurt = (flipped_wrong & easy_cut).sum()\n    if n_easy_cut > 0:\n        print(f\"\\n  Key paper metric — easy images that received a token cut:\")\n        print(f\"    n cut          : {n_easy_cut:,}\")\n        print(f\"    n hurt (TF)    : {n_easy_hurt:,}  \"\n              f\"({n_easy_hurt/n_easy_cut*100:.2f}%)\")\n        print(f\"    n helped (FT)  : {(flipped_right & easy_cut).sum():,}  \"\n              f\"({(flipped_right & easy_cut).sum()/n_easy_cut*100:.2f}%)\")\n        print(f\"    Net accuracy   : \"\n              f\"{(ada_correct[easy_cut].mean() - fixed_correct[easy_cut].mean())*100:+.3f}%\")\n        print(f\"    Interpretation : \"\n              f\"{'✓ cut is safe' if n_easy_hurt/n_easy_cut < 0.015 else '⚠ cut is hurting'}\")\n\n    if flipped_wrong.sum() > 10:\n        hurt_scores = scores[flipped_wrong]\n        help_scores = scores[flipped_right]\n        print(f\"\\n  Score distribution of flipped images:\")\n        print(f\"    Hurt  (TF): mean={hurt_scores.mean():.3f}  \"\n              f\"median={np.median(hurt_scores):.3f}  \"\n              f\"%easy={(hurt_scores > easy_thresh).mean()*100:.1f}%  \"\n              f\"%hard={(hurt_scores < hard_thresh).mean()*100:.1f}%\")\n        print(f\"    Helped(FT): mean={help_scores.mean():.3f}  \"\n              f\"median={np.median(help_scores):.3f}  \"\n              f\"%easy={(help_scores > easy_thresh).mean()*100:.1f}%  \"\n              f\"%hard={(help_scores < hard_thresh).mean()*100:.1f}%\")\n        print(f\"    [Hurt images should cluster at low scores — \"\n              f\"high mean here would indicate easy-image routing errors]\")\n\n    easy_hurt_rate = (flipped_wrong & easy_bin).sum() / easy_bin.sum() * 100\n    hard_help_rate = (flipped_right & hard_bin).sum() / hard_bin.sum() * 100\n    net_easy = (ada_correct[easy_bin].mean() - fixed_correct[easy_bin].mean()) * 100\n    net_hard = (ada_correct[hard_bin].mean() - fixed_correct[hard_bin].mean()) * 100\n    print(f\"\\n  Paper claim check:\")\n    print(f\"    Easy images: hurt_rate={easy_hurt_rate:.2f}%  \"\n          f\"net_acc={net_easy:+.3f}%  \"\n          f\"{'✓ CLAIM HOLDS' if easy_hurt_rate < 1.5 else '✗ TOO MANY HURT'}\")\n    print(f\"    Hard images: help_rate={hard_help_rate:.2f}%  \"\n          f\"net_acc={net_hard:+.3f}%  \"\n          f\"{'✓ BONUS HELPS' if hard_help_rate > 0.3 else '~ marginal benefit'}\")\n    print(SEP)\n\n\nprint(\"Option A comparison function ready.\")\n\n\n# ── CELL 20: ANALYSIS AND PAPER TABLES ───────────────────────────────────────\n\ndef analyze(b_acc1, b_acc5, b_fps, b_lat,\n            f_acc1, f_acc5, f_fps, f_lat, fixed_gf,\n            adaptive_results, ablation_results, base_gf):\n\n    SEP  = \"=\" * 100\n    DASH = \"─\" * 96\n\n    last_li = CONFIG[\"PRUNE_LAYERS\"][-1]\n    hard_t  = CONFIG[\"HARD_THRESHOLDS\"].get(last_li, 0.15)\n    easy_t  = CONFIG[\"EASY_THRESHOLDS\"].get(last_li, 0.85)\n\n    print(f\"\\n{SEP}\")\n    print(f\"  TABLE 1 — Main Results  (ImageNet-1k Val, training-free token pruning)\")\n    print(f\"  Primary claim: hard images recover accuracy via protective token budgets;\")\n    print(f\"  easy images tolerate proportional reductions with minimal accuracy loss.\")\n    print(f\"{SEP}\")\n    print(f\"  {'Method':<32}  {'Top-1':>7}  {'Hard Top-1':>11}  \"\n          f\"{'Easy Top-1':>11}  {'GFLOPs':>9}  {'FPS':>6}\")\n    print(f\"  {DASH}\")\n\n    print(f\"  {'Fixed Top-K (no routing)':<32}  {f_acc1:>7.2f}%  \"\n          f\"{'(uniform)':>11}  {'(uniform)':>11}  \"\n          f\"{fixed_gf:>9.4f}  {f_fps:>6.1f}\")\n\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in adaptive_results:\n        rho_str = str(rho).replace(\".\", \"p\")\n        tag     = f\"{rho_str}_fused\"\n        score_p = OUTPUT_DIR / f\"ada_scores_L{last_li}_{tag}.npy\"\n        ada_p   = OUTPUT_DIR / f\"ada_correct_{tag}.npy\"\n\n        if score_p.exists() and ada_p.exists():\n            scores  = np.load(str(score_p)).astype(np.float32)\n            correct = np.load(str(ada_p)).astype(np.float32)\n            hard_m  = scores < hard_t\n            easy_m  = scores > easy_t\n            h_acc   = correct[hard_m].mean() * 100 if hard_m.sum() > 0 else float(\"nan\")\n            e_acc   = correct[easy_m].mean() * 100 if easy_m.sum() > 0 else float(\"nan\")\n            h_str   = f\"{h_acc:.2f}%\"\n            e_str   = f\"{e_acc:.2f}%\"\n        else:\n            h_str = e_str = \"N/A\"\n\n        star = \" ★\" if acc1 >= f_acc1 - 0.10 else \"  \"\n        print(f\"  {'ASPP-ViT rho='+str(rho):<32}  {acc1:>7.2f}%{star}  \"\n              f\"{h_str:>11}  {e_str:>11}  \"\n              f\"{gf:>9.4f}  {fps:>6.1f}\")\n\n    print(f\"  ★ = within 0.10% of Fixed TopK overall accuracy\")\n    print(f\"  Hard = score_norm < {hard_t} at L{last_li}  (~15% of images)\")\n    print(f\"  Easy = score_norm > {easy_t} at L{last_li}  (~15% of images)\")\n    print(f\"  GFLOPs may exceed Fixed TopK — routing intentionally allocates \"\n          f\"extra compute to hard images.\")\n\n    print(f\"\\n{SEP}\")\n    print(f\"  TABLE 2 — Per-Difficulty Accuracy  \"\n          f\"(Hard=score<{hard_t}, Easy=score>{easy_t} at L{last_li})\")\n    print(f\"{SEP}\")\n    print(f\"  {'Method':<32}  {'Overall':>8}  {'Hard Top-1':>11}  \"\n          f\"{'Mid Top-1':>10}  {'Easy Top-1':>11}  \"\n          f\"{'Hard AvgKeep':>13}  {'Easy AvgKeep':>13}\")\n    print(f\"  {DASH}\")\n\n    print(f\"  {'Fixed Top-K (no routing)':<32}  {f_acc1:>7.2f}%  \"\n          f\"{'— (same)':>11}  {'— (same)':>10}  {'— (same)':>11}  \"\n          f\"{'base':>13}  {'base':>13}\")\n\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in adaptive_results:\n        rho_str  = str(rho).replace(\".\", \"p\")\n        tag      = f\"{rho_str}_fused\"\n        score_p  = OUTPUT_DIR / f\"ada_scores_L{last_li}_{tag}.npy\"\n        keep_p   = OUTPUT_DIR / f\"ada_keeps_L{last_li}_{tag}.npy\"\n        rc_p     = OUTPUT_DIR / f\"ada_entering_rc_L{last_li}_{tag}.npy\"\n        ada_p    = OUTPUT_DIR / f\"ada_correct_{tag}.npy\"\n\n        if not all(p.exists() for p in [score_p, keep_p, rc_p, ada_p]):\n            print(f\"  ASPP-ViT rho={rho:<28}  {acc1:>7.2f}%  \"\n                  f\"(per-image arrays missing)\")\n            continue\n\n        scores     = np.load(str(score_p)).astype(np.float32)\n        keeps      = np.load(str(keep_p)).astype(np.float32)\n        entering_r = np.load(str(rc_p)).astype(np.float32)\n        correct    = np.load(str(ada_p)).astype(np.float32)\n\n        hard_m = scores < hard_t\n        easy_m = scores > easy_t\n        mid_m  = ~hard_m & ~easy_m\n\n        # np.round matches torch.round convention for base keep computation.\n        base_keep_img = np.round(\n            CONFIG[\"BASE_SCHEDULE\"][last_li] * entering_r.astype(float)\n        ).astype(int)\n\n        h_acc  = correct[hard_m].mean() * 100 if hard_m.sum() > 0 else float(\"nan\")\n        m_acc  = correct[mid_m].mean()  * 100 if mid_m.sum()  > 0 else float(\"nan\")\n        e_acc  = correct[easy_m].mean() * 100 if easy_m.sum() > 0 else float(\"nan\")\n        h_keep = keeps[hard_m].mean() if hard_m.sum() > 0 else float(\"nan\")\n        e_keep = keeps[easy_m].mean() if easy_m.sum() > 0 else float(\"nan\")\n        h_base = base_keep_img[hard_m].mean() if hard_m.sum() > 0 else float(\"nan\")\n        e_base = base_keep_img[easy_m].mean() if easy_m.sum() > 0 else float(\"nan\")\n\n        h_delta    = h_keep - h_base if not np.isnan(h_base) else float(\"nan\")\n        e_delta    = e_keep - e_base if not np.isnan(e_base) else float(\"nan\")\n        h_keep_str = f\"{h_keep:.1f}({h_delta:+.1f})\" if not np.isnan(h_delta) else \"nan\"\n        e_keep_str = f\"{e_keep:.1f}({e_delta:+.1f})\" if not np.isnan(e_delta) else \"nan\"\n\n        print(f\"  {'ASPP-ViT rho='+str(rho):<32}  {acc1:>7.2f}%  \"\n              f\"{h_acc:>10.2f}%  {m_acc:>9.2f}%  {e_acc:>10.2f}%  \"\n              f\"{h_keep_str:>13}  {e_keep_str:>13}\")\n\n    print(f\"\\n{SEP}\")\n    print(f\"  TABLE 3 — Ablation: Component Decomposition\")\n    print(f\"{SEP}\")\n    print(f\"  {'Config':<44}  {'Top-1':>7}  {'GFLOPs':>9}  \"\n          f\"{'ΔTop-1':>8}  {'ΔGFLOPs':>9}\")\n    print(f\"  {DASH}\")\n\n    def _arow(name, acc1, gf):\n        print(f\"  {name:<44}  {acc1:>7.2f}%  {gf:>9.4f}  \"\n              f\"{acc1 - f_acc1:>+7.2f}%  {gf - fixed_gf:>+8.4f}\")\n\n    _arow(\"Fixed Top-K (no routing, with fused)\", f_acc1, fixed_gf)\n    for rho, acc1, acc5, fps, lat, ak, gf in ablation_results:\n        _arow(f\"+ Routing only      (rho={rho}, no fused)\", acc1, gf)\n        for rho2, acc1f, acc5f, fps2, lat2, ak2, gf2 in adaptive_results:\n            if rho2 == rho:\n                _arow(f\"  Full ASPP-ViT     (rho={rho}, with fused)\", acc1f, gf2)\n                break\n\n    print(f\"\\n{SEP}\")\n    print(f\"  TABLE 4 — rho Sweep  \"\n          f\"(Hard bonus = rho × {CONFIG['HARD_BONUS_RATIO']} × base_keep  \"\n          f\"Easy cut = rho × {CONFIG['EASY_CUT_RATIO']} × base_keep)\")\n    print(f\"{SEP}\")\n    hdr = (f\"  {'rho':>5}  {'fused':>5}  {'Top-1':>7}  {'Top-5':>7}  \"\n           f\"{'GFLOPs':>9}  {'FPS':>7}  {'ΔAcc':>8}  {'ΔGF':>9}  \"\n           + \"  \".join(f\"L{li}\" for li in CONFIG[\"PRUNE_LAYERS\"]))\n    print(hdr)\n    print(f\"  {'─' * 94}\")\n\n    all_rows = []\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in adaptive_results:\n        all_rows.append((rho, True, acc1, acc5, fps, lat, avg_keeps, gf))\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in ablation_results:\n        all_rows.append((rho, False, acc1, acc5, fps, lat, avg_keeps, gf))\n    all_rows.sort(key=lambda r: (r[0], not r[1]))\n\n    for rho, fused, acc1, acc5, fps, lat, avg_keeps, gf in all_rows:\n        da   = acc1 - f_acc1\n        dg   = gf   - fixed_gf\n        star = \" ★\" if da > -0.10 else \"  \"\n        row  = (f\"  {rho:>5}  {'T' if fused else 'F':>5}  \"\n                f\"{acc1:>7.2f}%  {acc5:>6.2f}%  {gf:>9.4f}  \"\n                f\"{fps:>7.1f}  {da:>+7.2f}%  {dg:>+8.4f}{star}  \"\n                + \"  \".join(f\"{avg_keeps.get(li, 0):>4.0f}\"\n                             for li in CONFIG[\"PRUNE_LAYERS\"]))\n        print(row)\n    print(f\"  ★ = Overall accuracy within 0.10% of Fixed TopK\")\n    print(f\"  Note: GFLOPs above Fixed TopK reflects extra compute allocated \"\n          f\"to hard images — this is the intended routing behaviour.\")\n\n    print(f\"\\n{SEP}\")\n    print(f\"  TABLE 5 — Throughput\")\n    print(f\"{SEP}\")\n    print(f\"  {'Method':<36}  {'Top-1':>7}  {'FPS':>7}  {'Lat ms':>8}  {'GFLOPs':>9}\")\n    print(f\"  {'─' * 78}\")\n    print(f\"  {'Baseline (no pruning)':<36}  {b_acc1:>7.2f}%  {b_fps:>7.1f}  \"\n          f\"{b_lat:>7.2f}  {base_gf:>9.4f}\")\n    print(f\"  {'Fixed Top-K':<36}  {f_acc1:>7.2f}%  {f_fps:>7.1f}  \"\n          f\"{f_lat:>7.2f}  {fixed_gf:>9.4f}\")\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in adaptive_results:\n        print(f\"  {'ASPP-ViT rho='+str(rho):<36}  {acc1:>7.2f}%  {fps:>7.1f}  \"\n              f\"{lat:>7.2f}  {gf:>9.4f}\")\n\n    best = max(adaptive_results, key=lambda r: r[1])\n    best_rho, best_acc1 = best[0], best[1]\n\n    print(f\"\\n{SEP}\")\n    print(f\"  VERDICT\")\n    print(f\"{SEP}\")\n    print(f\"  Fixed TopK GFLOPs reduction vs baseline : \"\n          f\"{(1 - fixed_gf / base_gf)*100:.1f}%  \"\n          f\"({base_gf:.4f} -> {fixed_gf:.4f})\")\n    print(f\"  Best overall accuracy : rho={best_rho}  Top-1={best_acc1:.2f}%  \"\n          f\"(Δ={best_acc1 - f_acc1:+.2f}% vs Fixed)\")\n    print(f\"\\n  Paper claim evaluation:\")\n    print(f\"  → CLS attention entropy: training-free difficulty proxy \"\n          f\"(AUC~0.57, D1→D10 spread ~10%)\")\n    print(f\"  → Percentile normalization: cross-layer, cross-dataset \"\n          f\"score comparability\")\n    print(f\"  → Asymmetric allocation: hard images receive protective \"\n          f\"budgets above fixed schedule\")\n    print(f\"  → Easy images tolerate proportional reductions \"\n          f\"(see TABLE 2 and Option A flip analysis)\")\n    print(f\"  → GFLOPs exceeding Fixed TopK confirms hard images \"\n          f\"receive real extra compute — this is the mechanism, not a flaw\")\n    print(SEP)\n\n\nprint(\"Analysis function ready.\")\n\n\n# ── CELL 21: SAVE RESULTS (NPZ + CSV + MANIFEST) ─────────────────────────────\n\ndef save_all_results(run_id, b_acc1, b_acc5, b_fps, b_lat,\n                     f_acc1, f_acc5, f_fps, f_lat, fixed_gf,\n                     adaptive_results, ablation_results, base_gf):\n\n    npz = _npz_path(run_id)\n    sd  = {\n        \"run_id\"       : np.array([run_id]),\n        \"model_key\"    : np.array([CONFIG[\"MODEL_KEY\"]]),\n        \"preset_key\"   : np.array([CONFIG[\"LAYER_PRESET_KEY\"]]),\n        \"prune_layers\" : np.array(CONFIG[\"PRUNE_LAYERS\"]),\n        \"base_ratios\"  : np.array([CONFIG[\"BASE_SCHEDULE\"][l]\n                                   for l in CONFIG[\"PRUNE_LAYERS\"]]),\n        \"base_gf\"      : np.array([base_gf]),\n        \"baseline_acc1\": np.array([b_acc1]),\n        \"baseline_acc5\": np.array([b_acc5]),\n        \"baseline_fps\" : np.array([b_fps]),\n        \"baseline_lat\" : np.array([b_lat]),\n        \"fixed_acc1\"   : np.array([f_acc1]),\n        \"fixed_acc5\"   : np.array([f_acc5]),\n        \"fixed_fps\"    : np.array([f_fps]),\n        \"fixed_lat\"    : np.array([f_lat]),\n        \"fixed_gf\"     : np.array([fixed_gf]),\n        \"cal_medians\"  : np.array([float(np.median(LAYER_CAL_SORTED[l]))\n                                   for l in CONFIG[\"PRUNE_LAYERS\"]]),\n    }\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in adaptive_results:\n        k = \"rho\" + str(rho).replace(\".\", \"p\")\n        sd[f\"{k}_acc1\"]      = np.array([acc1])\n        sd[f\"{k}_acc5\"]      = np.array([acc5])\n        sd[f\"{k}_fps\"]       = np.array([fps])\n        sd[f\"{k}_lat\"]       = np.array([lat])\n        sd[f\"{k}_gf\"]        = np.array([gf])\n        sd[f\"{k}_avg_keeps\"] = np.array([avg_keeps[l] for l in CONFIG[\"PRUNE_LAYERS\"]])\n    for rho, acc1, acc5, fps, lat, avg_keeps, gf in ablation_results:\n        k = \"abl_rho\" + str(rho).replace(\".\", \"p\")\n        sd[f\"{k}_acc1\"]      = np.array([acc1])\n        sd[f\"{k}_acc5\"]      = np.array([acc5])\n        sd[f\"{k}_fps\"]       = np.array([fps])\n        sd[f\"{k}_gf\"]        = np.array([gf])\n        sd[f\"{k}_avg_keeps\"] = np.array([avg_keeps[l] for l in CONFIG[\"PRUNE_LAYERS\"]])\n\n    np.savez(str(npz), **sd)\n    print(f\"  [NPZ]      {npz.name}\")\n\n    def _row(pass_type, rho, acc1, acc5, fps, lat, gf,\n             avg_keeps=None, uf=True, is_abl=False):\n        return {\n            \"run_id\"               : run_id,\n            \"timestamp\"            : datetime.now().isoformat(),\n            \"model_key\"            : CONFIG[\"MODEL_KEY\"],\n            \"layer_preset_key\"     : CONFIG[\"LAYER_PRESET_KEY\"],\n            \"prune_layers\"         : str(CONFIG[\"PRUNE_LAYERS\"]),\n            \"pass_type\"            : pass_type,\n            \"rho\"                  : rho,\n            \"use_fused\"            : uf,\n            \"top1\"                 : round(acc1, 4),\n            \"top5\"                 : round(acc5, 4),\n            \"gflops\"               : round(gf, 4),\n            \"fps\"                  : round(fps, 2),\n            \"latency_ms\"           : round(lat, 3),\n            \"delta_acc_vs_baseline\": round(acc1 - b_acc1, 4),\n            \"delta_acc_vs_fixed\"   : round(acc1 - f_acc1, 4)\n                                     if pass_type not in (\"baseline\", \"fixed\") else \"\",\n            \"delta_gf_vs_fixed\"    : round(gf - fixed_gf, 4)\n                                     if pass_type not in (\"baseline\", \"fixed\") else \"\",\n            \"avg_keeps\"            : _fmt_avg_keeps(avg_keeps or {}),\n            \"cal_batches\"          : CONFIG[\"CAL_BATCHES\"],\n            \"min_keep\"             : CONFIG[\"MIN_KEEP\"],\n            \"batch_size\"           : CONFIG[\"BATCH_SIZE\"],\n            \"is_ablation\"          : is_abl,\n        }\n\n    rows = [\n        _row(\"baseline\", \"\", b_acc1, b_acc5, b_fps, b_lat, base_gf, uf=False),\n        _row(\"fixed\",    \"\", f_acc1, f_acc5, f_fps, f_lat, fixed_gf, uf=False),\n    ]\n    for rho, acc1, acc5, fps, lat, ak, gf in adaptive_results:\n        rows.append(_row(\"adaptive\", rho, acc1, acc5, fps, lat, gf,\n                         avg_keeps=ak, uf=CONFIG[\"USE_FUSED\"]))\n    for rho, acc1, acc5, fps, lat, ak, gf in ablation_results:\n        rows.append(_row(\"adaptive_no_fused\", rho, acc1, acc5, fps, lat, gf,\n                         avg_keeps=ak, uf=False, is_abl=True))\n\n    for r in rows:\n        append_csv(r)\n    print(f\"  [CSV]      {len(rows)} rows -> {MASTER_CSV.name}\")\n\n    snap = {k: v for k, v in CONFIG.items()\n            if isinstance(v, (str, int, float, bool, list))}\n    save_manifest(run_id, snap)\n    print(f\"  [MANIFEST] {RUN_MANIFEST.name}\")\n    print(f\"\\n  Kaggle tip: save /kaggle/working/ as a Dataset to persist across sessions.\")\n\n\nprint(\"Saving function ready.\")\n\n\n# ── CELL 22: MAIN RUN ─────────────────────────────────────────────────────────\n\ndef run():\n    run_id = _make_run_id()\n    print(f\"\\n{'#'*68}\")\n    print(f\"  ASPP-ViT  —  Run ID: {run_id}\")\n    print(f\"  Model  : {CONFIG['MODEL_NAME']}\")\n    print(f\"  Preset : {CONFIG['LAYER_PRESET_KEY']}  ->  \"\n          f\"layers={CONFIG['PRUNE_LAYERS']}\")\n    print(f\"{'#'*68}\")\n\n    clear_gpu()\n    base_gf = baseline_gflops()\n    print(f\"\\n  Baseline GFLOPs ({CONFIG['MODEL_NAME']}): {base_gf:.3f}\")\n\n    loader = get_val_loader()\n\n    print(f\"\\n{'#'*60}\\n  PASS 0 — CALIBRATION\\n{'#'*60}\")\n    run_pass0_calibration(loader)\n\n    print(f\"\\n{'#'*60}\\n  PASS 1 — FULL-MODEL BASELINE\\n{'#'*60}\")\n    b_acc1, b_acc5, b_fps, b_lat = run_pass1_baseline(loader)\n\n    print(f\"\\n{'#'*60}\\n  PASS 2a — FIXED TOP-K BASELINE\\n{'#'*60}\")\n    f_acc1, f_acc5, f_fps, f_lat, f_avg_keeps, fixed_gf =  \\\n        run_pass2a_fixed(loader, base_gf)\n\n    adaptive_results = []\n    for rho in CONFIG[\"RHO_SWEEP\"]:\n        print(f\"\\n{'#'*60}\\n  PASS 3 — ADAPTIVE  rho={rho}\\n{'#'*60}\")\n        res = run_pass3_adaptive(loader, rho, fixed_gf, use_fused=None)\n        adaptive_results.append((rho, *res))\n\n    ablation_results = []\n    if CONFIG[\"RUN_ABLATION\"]:\n        for rho in CONFIG[\"ABLATION_RHOS\"]:\n            print(f\"\\n{'#'*60}\\n  PASS 3b — ABLATION  rho={rho}  \"\n                  f\"[no fused]\\n{'#'*60}\")\n            res = run_pass3_adaptive(loader, rho, fixed_gf, use_fused=False)\n            ablation_results.append((rho, *res))\n    else:\n        print(\"\\n  Pass 3b skipped (RUN_ABLATION=False)\")\n\n    print(f\"\\n{'#'*60}\\n  OPTION A — PER-IMAGE FLIP ANALYSIS\\n{'#'*60}\")\n    for rho, *_ in adaptive_results:\n        run_option_a_comparison(rho, use_fused=True)\n    for rho, *_ in ablation_results:\n        run_option_a_comparison(rho, use_fused=False)\n\n    analyze(b_acc1, b_acc5, b_fps, b_lat,\n            f_acc1, f_acc5, f_fps, f_lat, fixed_gf,\n            adaptive_results, ablation_results, base_gf)\n\n    save_all_results(run_id,\n                     b_acc1, b_acc5, b_fps, b_lat,\n                     f_acc1, f_acc5, f_fps, f_lat, fixed_gf,\n                     adaptive_results, ablation_results, base_gf)\n\n    print(f\"\\n{'#'*68}\")\n    print(f\"  ASPP-ViT COMPLETE  —  Run ID: {run_id}\")\n    print(f\"{'#'*68}\")\n\n\n# ── CELL 23: EXECUTE ──────────────────────────────────────────────────────────\n\nif __name__ == \"__main__\":\n    run()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}