{"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-12T19:54:35.669638Z","iopub.execute_input":"2026-07-12T19:54:35.66989Z","iopub.status.idle":"2026-07-12T19:54:58.295779Z","shell.execute_reply.started":"2026-07-12T19:54:35.669865Z","shell.execute_reply":"2026-07-12T19:54:58.29491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\nimport requests\nfrom io import BytesIO\nfrom torchvision import transforms\n\n# grab a sample image\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 = transform(img).unsqueeze(0).to(device)\n\nwith torch.no_grad():\n    tokens = model.patch_embed(x)\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}\")  # [batch, num_tokens, embed_dim]\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-12T19:54:58.297398Z","iopub.execute_input":"2026-07-12T19:54:58.297871Z","iopub.status.idle":"2026-07-12T19:55:00.391439Z","shell.execute_reply.started":"2026-07-12T19:54:58.297843Z","shell.execute_reply":"2026-07-12T19:55:00.390505Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Note: Halting Head — Deviation from Official A-ViT\n\n**Context:** Cell 3 implements a simplified halting head for the educational A-ViT prototype.\n\n## What we built\nA single linear layer + sigmoid, mapping each token's embedding to a scalar halting probability:\n\n```python\nclass HaltingHead(nn.Module):\n    def __init__(self, embed_dim):\n        super().__init__()\n        self.proj = nn.Linear(embed_dim, 1)\n\n    def forward(self, x):\n        return torch.sigmoid(self.proj(x)).squeeze(-1)\n```\n\n## How this differs from official A-ViT\n\n| Aspect | Our prototype | Official A-ViT |\n|---|---|---|\n| Gating mechanism | Single untrained linear layer | Dedicated gating unit trained jointly with backbone |\n| Training | None — randomly initialized | Trained end-to-end with ponder loss |\n| Loss terms | None | Ponder-loss + distributional prior regularization |\n| Halting semantics | Not meaningful yet | Learned to reflect per-token compute need |\n\n## Why this is still useful\nThe architecture illustrates the *mechanism* of adaptive halting (a per-token probability feeding into a cumulative halting rule across layers). It does not yet reproduce the *learned behavior* of official A-ViT, since the head has never seen a training signal.\n\n## Takeaway\nThis is a mechanism demo, not a functional reproduction. Halting probabilities at this stage are architecturally correct but numerically meaningless — training (with ponder loss) would be required to get real adaptive-compute behavior matching the paper.","metadata":{}},{"cell_type":"code","source":"class HaltingHead(nn.Module):\n    \"\"\"\n    Predicts a halting probability per token from its embedding.\n    Paper idea: each token accumulates halting mass across layers;\n    once cumulative mass crosses a threshold, the token stops updating.\n    \"\"\"\n    def __init__(self, embed_dim):\n        super().__init__()\n        self.proj = nn.Linear(embed_dim, 1)\n\n    def forward(self, x):\n        # x: [batch, tokens, embed_dim] -> [batch, tokens]\n        return torch.sigmoid(self.proj(x)).squeeze(-1)\n\nhalting_head = HaltingHead(model.embed_dim).to(device)\n\nwith torch.no_grad():\n    h = halting_head(tokens)\n\nprint(f\"Halting prob shape: {h.shape}\")\nprint(f\"Sample halting probs (first 10 tokens): {h[0, :10].cpu().numpy()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T19:59:26.518532Z","iopub.execute_input":"2026-07-12T19:59:26.519462Z","iopub.status.idle":"2026-07-12T19:59:26.811317Z","shell.execute_reply.started":"2026-07-12T19:59:26.519424Z","shell.execute_reply":"2026-07-12T19:59:26.810633Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_layers = len(model.blocks)\nn_tokens = tokens.shape[1]\nthreshold = 1.0  # halt once cumulative halting mass >= 1.0 (ACT-style)\n\n# fresh halting head per \"layer\" -- crude stand-in for per-layer gating\nhalting_heads = nn.ModuleList([HaltingHead(model.embed_dim) for _ in range(n_layers)]).to(device)\n\ncumulative_halt = torch.zeros(1, n_tokens).to(device)\nhalted_mask = torch.zeros(1, n_tokens, dtype=torch.bool).to(device)\nactive_counts = []\n\nx = tokens.clone()\nwith torch.no_grad():\n    for i, block in enumerate(model.blocks):\n        x = block(x)  # normal transformer block forward\n\n        h = halting_heads[i](x)  # halting prob this layer\n        h = h.masked_fill(halted_mask.squeeze(0), 0.0)  # halted tokens stop accumulating\n\n        cumulative_halt += h\n        newly_halted = (cumulative_halt >= threshold) & (~halted_mask)\n        halted_mask = halted_mask | newly_halted\n\n        active_counts.append((~halted_mask).sum().item())\n\nprint(\"Active tokens per layer:\")\nfor i, c in enumerate(active_counts):\n    print(f\"  Layer {i+1}: {c}/{n_tokens} active\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:01:36.608851Z","iopub.execute_input":"2026-07-12T20:01:36.609584Z","iopub.status.idle":"2026-07-12T20:01:37.071455Z","shell.execute_reply.started":"2026-07-12T20:01:36.609552Z","shell.execute_reply":"2026-07-12T20:01:37.070828Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scale = 0.15  # slows down accumulation so halting spreads across layers (cosmetic only)\n\ncumulative_halt = torch.zeros(1, n_tokens).to(device)\nhalted_mask = torch.zeros(1, n_tokens, dtype=torch.bool).to(device)\nactive_counts = []\n\nx = tokens.clone()\nwith torch.no_grad():\n    for i, block in enumerate(model.blocks):\n        x = block(x)\n\n        h = halting_heads[i](x) * scale\n        h = h.masked_fill(halted_mask.squeeze(0), 0.0)\n\n        cumulative_halt += h\n        newly_halted = (cumulative_halt >= threshold) & (~halted_mask)\n        halted_mask = halted_mask | newly_halted\n\n        active_counts.append((~halted_mask).sum().item())\n\nplt.figure(figsize=(7,4))\nplt.plot(range(1, n_layers+1), active_counts, marker='o')\nplt.xlabel(\"Layer\")\nplt.ylabel(\"Active tokens\")\nplt.title(\"Active tokens per layer (untrained gate, scaled)\")\nplt.grid(alpha=0.3)\nplt.show()\n\nprint(\"Active tokens per layer:\")\nfor i, c in enumerate(active_counts):\n    print(f\"  Layer {i+1}: {c}/{n_tokens} active\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:02:11.010855Z","iopub.execute_input":"2026-07-12T20:02:11.011691Z","iopub.status.idle":"2026-07-12T20:02:11.210855Z","shell.execute_reply.started":"2026-07-12T20:02:11.011658Z","shell.execute_reply":"2026-07-12T20:02:11.209958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# map patch tokens (excluding CLS) back to 14x14 grid, using final halted_mask\npatch_halted = halted_mask[0, 1:].cpu().numpy()  # drop CLS token\ngrid = patch_halted.reshape(14, 14)  # 224/16 = 14 patches per side\n\nfig, axes = plt.subplots(1, 2, figsize=(10, 5))\n\naxes[0].imshow(img.resize((224, 224)))\naxes[0].axis('off')\naxes[0].set_title(\"Input image\")\n\naxes[1].imshow(grid, cmap='Reds', vmin=0, vmax=1)\naxes[1].set_title(\"Halted patches (red = halted, final layer)\")\naxes[1].axis('off')\n\nplt.tight_layout()\nplt.show()\n\nprint(f\"Total halted patches: {patch_halted.sum()}/{len(patch_halted)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:02:32.027844Z","iopub.execute_input":"2026-07-12T20:02:32.028686Z","iopub.status.idle":"2026-07-12T20:02:32.307714Z","shell.execute_reply.started":"2026-07-12T20:02:32.028653Z","shell.execute_reply":"2026-07-12T20:02:32.306971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# recompute active_counts if not in scope (from Cell 5)\nstats_df = pd.DataFrame({\n    'Layer': range(1, n_layers+1),\n    'Active tokens': active_counts,\n    'Halted tokens': [n_tokens - c for c in active_counts],\n    '% active': [round(100*c/n_tokens, 1) for c in active_counts],\n})\n\n# crude FLOPs proxy: attention cost ~ O(N^2), MLP cost ~ O(N) per layer\n# real A-ViT FLOP savings require masking inside attention/MLP, not just bookkeeping (see Cell 4 caveat)\nstats_df['Relative attn cost (N^2)'] = [round((c/n_tokens)**2, 3) for c in active_counts]\nstats_df['Relative mlp cost (N)'] = [round(c/n_tokens, 3) for c in active_counts]\n\nprint(stats_df.to_string(index=False))\n\navg_active_frac = np.mean([c/n_tokens for c in active_counts])\nprint(f\"\\nAverage active token fraction across layers: {avg_active_frac:.2%}\")\nprint(f\"Naive compute reduction estimate (MLP-linear): {(1-avg_active_frac):.2%}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T20:03:00.053858Z","iopub.execute_input":"2026-07-12T20:03:00.054642Z","iopub.status.idle":"2026-07-12T20:03:00.361376Z","shell.execute_reply.started":"2026-07-12T20:03:00.054609Z","shell.execute_reply":"2026-07-12T20:03:00.360444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}