{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\nimport matplotlib.pyplot as plt\nimport numpy as np\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\nmodel = timm.create_model('deit_small_patch16_224', pretrained=True)\nmodel.eval().to(device)\n\nprint(f\"Device: {device}\")\nprint(f\"Embed dim: {model.embed_dim}\")\nprint(f\"Depth: {len(model.blocks)}\")\nprint(f\"Num heads: {model.blocks[0].attn.num_heads}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:12:51.282635Z","iopub.execute_input":"2026-07-12T20:12:51.282896Z","iopub.status.idle":"2026-07-12T20:13:07.88232Z","shell.execute_reply.started":"2026-07-12T20:12:51.282864Z","shell.execute_reply":"2026-07-12T20:13:07.881582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\nimport requests\nfrom io import BytesIO\nfrom torchvision import transforms\n\nurl = \"https://github.com/pytorch/hub/raw/master/images/dog.jpg\"\nimg = Image.open(BytesIO(requests.get(url).content)).convert('RGB')\n\ntransform = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\nx_img = transform(img).unsqueeze(0).to(device)\n\nwith torch.no_grad():\n    tokens = model.patch_embed(x_img)\n    cls_token = model.cls_token.expand(tokens.shape[0], -1, -1)\n    tokens = torch.cat((cls_token, tokens), dim=1)\n    tokens = tokens + model.pos_embed\n    tokens = model.pos_drop(tokens)\n\nprint(f\"Token tensor shape: {tokens.shape}\")\nprint(f\"Num patch tokens: {tokens.shape[1] - 1} (+1 CLS token)\")\n\nplt.imshow(img)\nplt.axis('off')\nplt.title(\"Input image\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:15:09.372602Z","iopub.execute_input":"2026-07-12T20:15:09.373462Z","iopub.status.idle":"2026-07-12T20:15:10.625397Z","shell.execute_reply.started":"2026-07-12T20:15:09.373431Z","shell.execute_reply":"2026-07-12T20:15:10.624639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PolicyNetwork(nn.Module):\n    \"\"\"\n    AdaViT idea: a lightweight policy predicts, per token, per layer,\n    whether to use: (1) the block at all, (2) each attention head, (3) MLP dims.\n    Real AdaViT uses Gumbel-Softmax for differentiable discrete decisions.\n    Here: one policy net per layer, outputs 3 decision types.\n    \"\"\"\n    def __init__(self, embed_dim, num_heads):\n        super().__init__()\n        self.num_heads = num_heads\n        # 3 decisions: use-block, use-head (per head), use-mlp-dim (coarse, single gate)\n        self.block_gate = nn.Linear(embed_dim, 2)   # [skip, keep] logits\n        self.head_gate = nn.Linear(embed_dim, num_heads * 2)  # [skip, keep] per head\n        self.mlp_gate = nn.Linear(embed_dim, 2)\n\n    def forward(self, x, tau=1.0, hard=True):\n        # x: [batch, tokens, embed_dim] -> use CLS token as global decision driver\n        cls = x[:, 0]  # [batch, embed_dim]\n\n        block_logits = self.block_gate(cls)                      # [batch, 2]\n        head_logits = self.head_gate(cls).view(-1, self.num_heads, 2)  # [batch, heads, 2]\n        mlp_logits = self.mlp_gate(cls)                           # [batch, 2]\n\n        block_decision = torch.nn.functional.gumbel_softmax(block_logits, tau=tau, hard=hard)\n        head_decision = torch.nn.functional.gumbel_softmax(head_logits, tau=tau, hard=hard, dim=-1)\n        mlp_decision = torch.nn.functional.gumbel_softmax(mlp_logits, tau=tau, hard=hard)\n\n        return {\n            'use_block': block_decision[:, 1],       # 1 = keep, 0 = skip\n            'use_head': head_decision[:, :, 1],       # [batch, heads]\n            'use_mlp': mlp_decision[:, 1],\n        }\n\npolicies = nn.ModuleList([\n    PolicyNetwork(model.embed_dim, model.blocks[0].attn.num_heads) for _ in range(len(model.blocks))\n]).to(device)\n\nwith torch.no_grad():\n    decisions = policies[0](tokens)\n\nprint(\"Layer 0 policy decisions:\")\nprint(f\"  use_block: {decisions['use_block'].item():.0f}\")\nprint(f\"  use_head (per head): {decisions['use_head'][0].cpu().numpy()}\")\nprint(f\"  use_mlp: {decisions['use_mlp'].item():.0f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:17:21.733777Z","iopub.execute_input":"2026-07-12T20:17:21.734313Z","iopub.status.idle":"2026-07-12T20:17:22.052362Z","shell.execute_reply.started":"2026-07-12T20:17:21.734287Z","shell.execute_reply":"2026-07-12T20:17:22.051623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_layers = len(model.blocks)\nlayer_decisions = []\n\nx = tokens.clone()\nwith torch.no_grad():\n    for i, block in enumerate(model.blocks):\n        decisions = policies[i](x)\n\n        if decisions['use_block'].item() > 0.5:\n            x = block(x)  # full block compute (heads/mlp gating not actually applied to compute, see caveat)\n        # else: skip block entirely, x passes through unchanged\n\n        layer_decisions.append({\n            'layer': i + 1,\n            'block_used': decisions['use_block'].item(),\n            'heads_used': decisions['use_head'][0].sum().item(),\n            'total_heads': model.blocks[0].attn.num_heads,\n            'mlp_used': decisions['use_mlp'].item(),\n        })\n\nfor d in layer_decisions:\n    print(f\"Layer {d['layer']:2d}: block={'keep' if d['block_used'] else 'SKIP':4s} | \"\n          f\"heads={int(d['heads_used'])}/{d['total_heads']} | \"\n          f\"mlp={'keep' if d['mlp_used'] else 'SKIP'}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:17:45.610213Z","iopub.execute_input":"2026-07-12T20:17:45.610648Z","iopub.status.idle":"2026-07-12T20:17:45.714388Z","shell.execute_reply.started":"2026-07-12T20:17:45.610619Z","shell.execute_reply":"2026-07-12T20:17:45.71376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"layers = [d['layer'] for d in layer_decisions]\nblock_used = [d['block_used'] for d in layer_decisions]\nheads_used = [d['heads_used'] for d in layer_decisions]\nmlp_used = [d['mlp_used'] for d in layer_decisions]\n\nfig, axes = plt.subplots(3, 1, figsize=(8, 8), sharex=True)\n\naxes[0].bar(layers, block_used, color=['green' if b else 'red' for b in block_used])\naxes[0].set_ylabel(\"Block used\")\naxes[0].set_yticks([0, 1])\naxes[0].set_title(\"Block-level skip decisions per layer\")\n\naxes[1].bar(layers, heads_used, color='steelblue')\naxes[1].set_ylabel(\"Heads used\")\naxes[1].axhline(model.blocks[0].attn.num_heads, color='gray', linestyle='--', linewidth=1)\n\naxes[2].bar(layers, mlp_used, color=['green' if m else 'red' for m in mlp_used])\naxes[2].set_ylabel(\"MLP used\")\naxes[2].set_yticks([0, 1])\naxes[2].set_xlabel(\"Layer\")\n\nplt.tight_layout()\nplt.show()\n\nn_blocks_skipped = sum(1 for b in block_used if b == 0)\nprint(f\"Blocks skipped: {n_blocks_skipped}/{n_layers}\")\nprint(f\"Avg heads used per layer: {np.mean(heads_used):.2f}/{model.blocks[0].attn.num_heads}\")\nprint(f\"MLPs skipped: {sum(1 for m in mlp_used if m == 0)}/{n_layers}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:18:58.503877Z","iopub.execute_input":"2026-07-12T20:18:58.504148Z","iopub.status.idle":"2026-07-12T20:18:58.966036Z","shell.execute_reply.started":"2026-07-12T20:18:58.504126Z","shell.execute_reply":"2026-07-12T20:18:58.965345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.DataFrame(layer_decisions)\ndf['relative_compute'] = df.apply(\n    lambda r: 0.0 if r['block_used'] == 0 else (r['heads_used']/r['total_heads'] + int(r['mlp_used']==1)) / 2,\n    axis=1\n)\n\nprint(df[['layer', 'block_used', 'heads_used', 'mlp_used', 'relative_compute']].to_string(index=False))\n\navg_compute = df['relative_compute'].mean()\nprint(f\"\\nAverage relative compute per layer: {avg_compute:.2%}\")\nprint(f\"Naive compute reduction estimate: {(1 - avg_compute):.2%}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:19:55.368481Z","iopub.execute_input":"2026-07-12T20:19:55.369385Z","iopub.status.idle":"2026-07-12T20:19:55.65341Z","shell.execute_reply.started":"2026-07-12T20:19:55.369344Z","shell.execute_reply":"2026-07-12T20:19:55.652783Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# AdaViT Prototype — What We Built\n\nAn educational reproduction of the *core idea* behind AdaViT (Meng et al.), built on top of a pretrained DeiT-Small backbone. **Not the official AdaViT implementation** — no official checkpoint, no training.\n\n## Steps\n\n1. **Loaded pretrained DeiT-Small** (`timm`, ImageNet weights) as the backbone.\n2. **Extracted token embeddings** for a sample image (patch embed + CLS token + positional embedding).\n3. **Built a per-layer policy network** — a small linear head reading the CLS token, producing three Gumbel-Softmax decisions per layer:\n   - Skip or keep the entire transformer block\n   - Skip or keep each attention head individually\n   - Skip or keep the MLP\n4. **Ran a forward pass applying these decisions**:\n   - Block-skip is **real** — skipped blocks are bypassed in the forward pass.\n   - Head/MLP skip decisions are **logged only**, not actually masked out of compute.\n5. **Visualized decisions across all 12 layers** (chart below).\n6. **Computed a naive relative-compute estimate** from the logged decisions.\n\n## Results\n\n| Layer | Block | Heads used (/6) | MLP |\n|---|---|---|---|\n| 1  | keep | 3 | skip |\n| 2  | keep | 4 | keep |\n| 3  | keep | 4 | skip |\n| 4  | **skip** | 2 | skip |\n| 5  | **skip** | 0 | keep |\n| 6  | **skip** | 4 | skip |\n| 7  | keep | 3 | skip |\n| 8  | **skip** | 2 | skip |\n| 9  | keep | 2 | skip |\n| 10 | keep | 4 | keep |\n| 11 | **skip** | 3 | skip |\n| 12 | **skip** | 3 | keep |\n\n- Blocks skipped: 6/12\n- Avg heads used per layer: 2.83/6\n- MLPs skipped: 8/12\n- Naive relative-compute estimate: 22.2% (i.e. \"77.8% reduction\")\n\n## Important caveats (read before presenting)\n\n- **Policy is untrained.** Decisions come from random Gumbel-Softmax noise, not learned behavior — there's no layer-wise trend (e.g. more skipping in later layers), which real AdaViT does show after training.\n- **No budget loss.** Real AdaViT trains with a target-FLOP-budget regularizer that trades off accuracy vs. speed. Without it, the policy has no signal for what's actually redundant.\n- **The 77.8% \"reduction\" number is not a real efficiency measurement.** It's a proxy formula applied to random decisions, and is driven mostly by 6/12 blocks randomly skipping. Treat it as illustrative of the *bookkeeping mechanism*, not a performance claim.\n- **Head/MLP skipping isn't actually applied to compute** in this notebook — only block-level skipping produces real FLOP savings here.\n- **Decisions are CLS-token-driven only** (one global decision per image per layer), simpler than some AdaViT variants that decide per-token.\n\n## Bottom line\n\nThis demonstrates the *architecture* of AdaViT's learned structural pruning (per-layer block/head/MLP gating via differentiable discrete decisions) — not a working reproduction of its trained behavior or measured compute savings.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}