{"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":[{"sourceId":6799,"databundleVersionId":4225553,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":641731,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":483962,"modelId":499465}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import warnings\nfrom pydantic.warnings import UnsupportedFieldAttributeWarning\n\nwarnings.filterwarnings(\n    \"ignore\",\n    category=UnsupportedFieldAttributeWarning\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T17:46:01.736737Z","iopub.execute_input":"2026-01-27T17:46:01.737321Z","iopub.status.idle":"2026-01-27T17:46:01.740828Z","shell.execute_reply.started":"2026-01-27T17:46:01.737293Z","shell.execute_reply":"2026-01-27T17:46:01.740031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# MAGDIFF TOKEN PRUNING — FULL BENCHMARK + ANALYTICS\n# Base vs Magnitude vs MagDiff\n# ============================================================\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\nimport time, os, types, gc\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom PIL import Image\nfrom tqdm import tqdm\n\n# ------------------------------------------------------------\n# 1. CONFIG\n# ------------------------------------------------------------\n\nCONFIG = {\n    \"DEVICE\": torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"),\n    \"BATCH_SIZE\": 64,\n    \"IMG_SIZE\": 224,\n    \"DATASET_ROOT\": \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\",\n    \"PRUNE_LAYER\": 6,\n    \"RATIO\": 0.75,\n    \"IMAGES_PER_CLASS\": 10,\n    \"SEED\": 0,                      # ← ADDED\n    \"SAVE_DIR\": \"./output\"\n}\n\nos.makedirs(CONFIG[\"SAVE_DIR\"], exist_ok=True)\n\ndef clear_gpu():\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(\"Device:\", CONFIG[\"DEVICE\"])\n\n# ------------------------------------------------------------\n# 2. PRUNERS\n# ------------------------------------------------------------\n\nclass StandardMagPruner(nn.Module):\n    def __init__(self, keep_ratio):\n        super().__init__()\n        self.keep_ratio = keep_ratio\n        self.last_indices = None\n\n    def forward(self, x):\n        score = torch.norm(x, dim=-1)\n        score[:, 0] = float('inf')\n        k = int(x.shape[1] * self.keep_ratio)\n        _, idx = torch.topk(score, k, dim=1)\n        self.last_indices = idx.detach()\n        b = torch.arange(x.shape[0], device=x.device).unsqueeze(1)\n        return x[b, idx]\n\n\nclass MagDiffPruner(nn.Module):\n    def __init__(self, keep_ratio, diff_weight=1.0):\n        super().__init__()\n        self.keep_ratio = keep_ratio\n        self.diff_weight = diff_weight\n        self.last_indices = None\n\n    def forward(self, x):\n        mag = torch.norm(x, dim=-1)\n        diff = torch.norm(x - torch.roll(x, 1, 1), dim=-1)\n        score = mag + self.diff_weight * diff\n        score[:, 0] = float('inf')\n        k = int(x.shape[1] * self.keep_ratio)\n        _, idx = torch.topk(score, k, dim=1)\n        self.last_indices = idx.detach()\n        b = torch.arange(x.shape[0], device=x.device).unsqueeze(1)\n        return x[b, idx]\n\n# ------------------------------------------------------------\n# 3. SAFE FLOPs ESTIMATION\n# ------------------------------------------------------------\n\ndef estimate_vit_flops(tokens, dim=768, layers=12):\n    attn = 4 * tokens * dim * dim + 2 * tokens * tokens * dim\n    mlp = 8 * tokens * dim * dim\n    return layers * (attn + mlp) / 1e9\n\n# ------------------------------------------------------------\n# 4. PRUNER INJECTION\n# ------------------------------------------------------------\n\ndef inject_pruner(model, pruner, prune_layer):\n    blocks = model.blocks\n\n    def forward_features(self, x, *args, **kwargs):\n        x = self.patch_embed(x)\n        x = self._pos_embed(x)\n        x = self.norm_pre(x)\n\n        for i in range(prune_layer):\n            x = blocks[i](x)\n\n        if pruner is not None:\n            x = pruner(x)\n\n        for i in range(prune_layer, len(blocks)):\n            x = blocks[i](x)\n\n        return self.norm(x)\n\n    model.forward_features = types.MethodType(forward_features, model)\n    return model\n\n# ------------------------------------------------------------\n# 5. DATASET (FIXED — SEEDED, DETERMINISTIC)\n# ------------------------------------------------------------\n\nclass SimpleDataset(Dataset):\n    def __init__(self, paths, labels, transform):\n        self.paths = paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, i):\n        img = Image.open(self.paths[i]).convert(\"RGB\")\n        return self.transform(img), self.labels[i]\n\n\ndef get_loader():\n    rng = np.random.default_rng(CONFIG[\"SEED\"])\n\n    tfm = transforms.Compose([\n        transforms.Resize(256),\n        transforms.CenterCrop(224),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n    ])\n\n    paths, labels = [], []\n    classes = sorted(os.listdir(CONFIG[\"DATASET_ROOT\"]))\n\n    for i, c in enumerate(classes):\n        cls_dir = os.path.join(CONFIG[\"DATASET_ROOT\"], c)\n        files = sorted(os.listdir(cls_dir))   # ← FIXED\n        chosen = rng.choice(\n            files,\n            size=CONFIG[\"IMAGES_PER_CLASS\"],\n            replace=False\n        )\n\n        for f in chosen:\n            paths.append(os.path.join(cls_dir, f))\n            labels.append(i)\n\n    return DataLoader(\n        SimpleDataset(paths, labels, tfm),\n        batch_size=CONFIG[\"BATCH_SIZE\"],\n        shuffle=False\n    )\n\n# ------------------------------------------------------------\n# 6. BENCHMARK\n# ------------------------------------------------------------\n\ndef benchmark(model, loader):\n    correct, total = 0, 0\n    torch.cuda.synchronize()\n    start = time.time()\n\n    with torch.no_grad():\n        for x, y in tqdm(loader, leave=False):\n            x, y = x.to(CONFIG[\"DEVICE\"]), y.to(CONFIG[\"DEVICE\"])\n            out = model(x)\n            pred = out.argmax(1)\n            correct += (pred == y).sum().item()\n            total += y.size(0)\n\n    torch.cuda.synchronize()\n    elapsed = time.time() - start\n    return (\n        100 * correct / total,\n        total / elapsed,\n        elapsed / len(loader) * 1000\n    )\n# ------------------------------------------------------------\n# 7. ANALYTICS PLOTS (SHOW + SAVE)\n# ------------------------------------------------------------\n\ndef plot_token_map(indices, name):\n    mask = torch.zeros(197)\n    mask[indices[0]] = 1\n    grid = mask[1:].reshape(14,14)\n\n    plt.figure(figsize=(3,3))\n    plt.imshow(grid, cmap=\"hot\")\n    plt.title(name)\n    plt.axis(\"off\")\n    plt.savefig(f\"{CONFIG['SAVE_DIR']}/{name}.png\", dpi=300)\n    plt.show()\n    plt.close()\n\n\ndef plot_score_space(features, name):\n    mag = torch.norm(features, dim=-1).flatten().cpu()\n    diff = torch.norm(features - torch.roll(features, 1, 1), dim=-1).flatten().cpu()\n\n    plt.figure(figsize=(5,4))\n    plt.scatter(mag, diff, s=2, alpha=0.3)\n    plt.xlabel(\"Token Magnitude\")\n    plt.ylabel(\"Token Difference\")\n    plt.title(name)\n    plt.savefig(f\"{CONFIG['SAVE_DIR']}/{name}.png\", dpi=300)\n    plt.show()\n    plt.close()\n\n\ndef token_overlap(idx1, idx2):\n    overlaps = []\n    for b in range(idx1.shape[0]):\n        s1, s2 = set(idx1[b].tolist()), set(idx2[b].tolist())\n        overlaps.append(len(s1 & s2) / len(s1))\n    return np.mean(overlaps)\n\n# ------------------------------------------------------------\n# 8. MAIN EXPERIMENT\n# ------------------------------------------------------------\n\ndef run():\n    clear_gpu()\n    loader = get_loader()\n\n    base_model = timm.create_model(\n        \"vit_base_patch16_224\", pretrained=True\n    ).to(CONFIG[\"DEVICE\"]).eval()\n\n    methods = [\n        (\"ViT-Base\", None),\n        (\"Magnitude\", StandardMagPruner(CONFIG[\"RATIO\"])),\n        (\"MagDiff\", MagDiffPruner(CONFIG[\"RATIO\"]))\n    ]\n\n    results = {}\n    features_cache = {}\n\n    for name, pruner in methods:\n        print(f\"\\nRunning: {name}\")\n        model = inject_pruner(base_model, pruner, CONFIG[\"PRUNE_LAYER\"])\n        acc, fps, lat = benchmark(model, loader)\n\n        tokens = 197 if pruner is None else int(197 * CONFIG[\"RATIO\"])\n        layers = 12 if pruner is None else (12 - CONFIG[\"PRUNE_LAYER\"])\n        gflops = estimate_vit_flops(tokens, layers=layers)\n\n        results[name] = [acc, fps, lat, gflops]\n\n        if pruner is not None:\n            imgs, _ = next(iter(loader))\n            imgs = imgs.to(CONFIG[\"DEVICE\"])\n            with torch.no_grad():\n                features_cache[name] = model.forward_features(imgs)\n\n    # ------------------ RESULTS TABLE ------------------\n\n    print(\"\\n\" + \"=\"*55)\n    print(\"FINAL RESULTS\")\n    print(\"=\"*55)\n    print(f\"{'Method':<12}{'Acc(%)':>8}{'FPS':>10}{'Lat(ms)':>12}{'GFLOPs':>10}\")\n    print(\"-\"*55)\n\n    csv = []\n    for k, v in results.items():\n        print(f\"{k:<12}{v[0]:>8.2f}{v[1]:>10.2f}{v[2]:>12.2f}{v[3]:>10.2f}\")\n        csv.append([k] + v)\n\n    np.savetxt(\n        f\"{CONFIG['SAVE_DIR']}/results.csv\",\n        np.array(csv, dtype=object),\n        delimiter=\",\",\n        fmt=\"%s\"\n    )\n\n    # ------------------ DIFFERENCE ANALYSIS ------------------\n\n    mag_acc = results[\"Magnitude\"][0]\n    diff_acc = results[\"MagDiff\"][0]\n\n    mag_idx = methods[1][1].last_indices\n    diff_idx = methods[2][1].last_indices\n    overlap = token_overlap(mag_idx, diff_idx)\n\n    print(\"\\n\" + \"=\"*55)\n    print(\"DIFFERENCE ANALYSIS (Mag vs MagDiff)\")\n    print(\"=\"*55)\n    print(f\"Accuracy gain        : {diff_acc - mag_acc:+.2f}%\")\n    print(f\"Token overlap ratio  : {overlap:.3f}\")\n    print(f\"Token difference     : {(1-overlap)*100:.1f}%\")\n\n    print(\"\\nShowing spatial token maps...\")\n    plot_token_map(mag_idx, \"Magnitude_tokens\")\n    plot_token_map(diff_idx, \"MagDiff_tokens\")\n\n    print(\"Showing token score-space plots...\")\n    plot_score_space(features_cache[\"Magnitude\"], \"Magnitude_score_space\")\n    plot_score_space(features_cache[\"MagDiff\"], \"MagDiff_score_space\")\n\n    print(\"\\nAll outputs saved to:\", CONFIG[\"SAVE_DIR\"])\n\n# ------------------------------------------------------------\n# 9. RUN\n# ------------------------------------------------------------\n\nif __name__ == \"__main__\":\n    run()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-27T17:33:55.666932Z","iopub.execute_input":"2026-01-27T17:33:55.667219Z","iopub.status.idle":"2026-01-27T17:43:33.442014Z","shell.execute_reply.started":"2026-01-27T17:33:55.667198Z","shell.execute_reply":"2026-01-27T17:43:33.441433Z"}},"outputs":[],"execution_count":null}]}