{"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":"markdown","source":"# Phase 1 - Base Vision Transformer\n\nThis notebook implements a Vision Transformer (ViT) from scratch using PyTorch.\n\nObjectives:\n- Build a modular Vision Transformer\n- Reproduce DeiT-S architecture\n- Evaluate on ImageNet-1K\n- Use this implementation as the foundation for ATS, A-ViT and AdaViT","metadata":{}},{"cell_type":"code","source":"import os\nimport math\nimport time\nimport random\nimport numpy as np\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torchvision.datasets import ImageFolder\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ntorch.manual_seed(42)\nnp.random.seed(42)\nrandom.seed(42)\n\nprint(\"=\" * 60)\nprint(\"Base Vision Transformer\")\nprint(\"=\" * 60)\nprint(f\"Device : {device}\")\nprint(f\"PyTorch: {torch.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:44:48.465497Z","iopub.execute_input":"2026-07-12T11:44:48.465821Z","iopub.status.idle":"2026-07-12T11:44:58.106996Z","shell.execute_reply.started":"2026-07-12T11:44:48.465798Z","shell.execute_reply":"2026-07-12T11:44:58.106309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PROJECT_DIR = \"/kaggle/working/efficient-vit\"\n\nfolders = [\n    \"models\",\n    \"weights\",\n    \"results\",\n    \"utils\"\n]\n\nfor folder in folders:\n    os.makedirs(os.path.join(PROJECT_DIR, folder), exist_ok=True)\n\nprint(\"Project initialized.\")\n\nfor folder in folders:\n    print(\"✓\", folder)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:45:16.79533Z","iopub.execute_input":"2026-07-12T11:45:16.795779Z","iopub.status.idle":"2026-07-12T11:45:16.801931Z","shell.execute_reply.started":"2026-07-12T11:45:16.795746Z","shell.execute_reply":"2026-07-12T11:45:16.801092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_DIR = \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\"\n\nVAL_RAW = \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\"\n\nprint(\"Train Exists :\", os.path.exists(TRAIN_DIR))\nprint(\"Val Exists   :\", os.path.exists(VAL_RAW))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:46:05.143485Z","iopub.execute_input":"2026-07-12T11:46:05.143856Z","iopub.status.idle":"2026-07-12T11:46:05.165641Z","shell.execute_reply.started":"2026-07-12T11:46:05.14383Z","shell.execute_reply":"2026-07-12T11:46:05.164826Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Mounting Imagenet1k","metadata":{}},{"cell_type":"code","source":"import os\nimport shutil\nimport urllib.request\n\nVAL_RAW = \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val\"\nVAL_OUT = \"/kaggle/working/imagenet/val\"\n\nos.makedirs(VAL_OUT, exist_ok=True)\n\n# Download validation labels\nlabel_file = \"/kaggle/working/val_synset_labels.txt\"\n\nif not os.path.exists(label_file):\n    urllib.request.urlretrieve(\n        \"https://raw.githubusercontent.com/tensorflow/models/master/research/slim/datasets/imagenet_2012_validation_synset_labels.txt\",\n        label_file\n    )\n\nwith open(label_file) as f:\n    synsets = [x.strip() for x in f.readlines()]\n\nimages = sorted(os.listdir(VAL_RAW))\n\nassert len(images) == len(synsets) == 50000\n\nfor img, synset in zip(images, synsets):\n\n    dst = os.path.join(VAL_OUT, synset)\n    os.makedirs(dst, exist_ok=True)\n\n    shutil.copy2(\n        os.path.join(VAL_RAW, img),\n        os.path.join(dst, img)\n    )\n\nprint(\"Done!\")\nprint(\"Classes :\", len(os.listdir(VAL_OUT)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:46:24.733879Z","iopub.execute_input":"2026-07-12T11:46:24.734712Z","iopub.status.idle":"2026-07-12T11:56:20.812961Z","shell.execute_reply.started":"2026-07-12T11:46:24.734679Z","shell.execute_reply":"2026-07-12T11:56:20.812078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch Embedding\n\nThe first stage of the Vision Transformer converts an input image into a sequence of fixed-size patch embeddings.\n\nFor an input image of size **224 × 224 × 3** and a patch size of **16 × 16**, the image is divided into:\n\n- Number of patches = (224 / 16)² = 196\n- Each patch has dimensions 16 × 16 × 3 = 768\n- Each flattened patch is projected to an embedding dimension of 384 (DeiT-S)\n\nThe output of this stage is therefore:\n\n196 × 384","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Base Vision Transformer Configuration (DeiT-Small)\n# ============================================================\n\nIMAGE_SIZE = 224\nPATCH_SIZE = 16\n\nIN_CHANNELS = 3\nNUM_CLASSES = 1000\n\nEMBED_DIM = 384\nDEPTH = 12\nNUM_HEADS = 6\nMLP_RATIO = 4\n\nDROP_RATE = 0.0\nATTN_DROP_RATE = 0.0\n\nprint(\"=\" * 60)\nprint(\"Base ViT Configuration\")\nprint(\"=\" * 60)\n\nprint(f\"Image Size      : {IMAGE_SIZE}\")\nprint(f\"Patch Size      : {PATCH_SIZE}\")\nprint(f\"Embedding Dim   : {EMBED_DIM}\")\nprint(f\"Transformer Blk : {DEPTH}\")\nprint(f\"Heads           : {NUM_HEADS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:56:49.986739Z","iopub.execute_input":"2026-07-12T11:56:49.987809Z","iopub.status.idle":"2026-07-12T11:56:49.994022Z","shell.execute_reply.started":"2026-07-12T11:56:49.987772Z","shell.execute_reply":"2026-07-12T11:56:49.993206Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch Embedding\n\nUnlike CNNs, Vision Transformers do not process pixels directly.\n\nInstead, the image is divided into fixed-size patches. Each patch is projected into an embedding vector using a convolution with kernel size equal to the patch size.\n\nFor DeiT-S:\n\n- Input Image : 224 × 224 × 3\n- Patch Size  : 16 × 16\n- Number of Patches : 196\n- Embedding Dimension : 384\n\nOutput Shape:\n\n(B, 196, 384)","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Patch Embedding\n# ============================================================\n\nclass PatchEmbedding(nn.Module):\n\n    def __init__(\n        self,\n        image_size=224,\n        patch_size=16,\n        in_channels=3,\n        embed_dim=384\n    ):\n        super().__init__()\n\n        self.image_size = image_size\n        self.patch_size = patch_size\n\n        self.num_patches = (image_size // patch_size) ** 2\n\n        self.projection = nn.Conv2d(\n            in_channels=in_channels,\n            out_channels=embed_dim,\n            kernel_size=patch_size,\n            stride=patch_size\n        )\n\n    def forward(self, x):\n\n        # (B,3,224,224)\n        x = self.projection(x)\n\n        # (B,384,14,14)\n        x = x.flatten(2)\n\n        # (B,384,196)\n        x = x.transpose(1,2)\n\n        # (B,196,384)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:57:27.132918Z","iopub.execute_input":"2026-07-12T11:57:27.133221Z","iopub.status.idle":"2026-07-12T11:57:27.139362Z","shell.execute_reply.started":"2026-07-12T11:57:27.133187Z","shell.execute_reply":"2026-07-12T11:57:27.138441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Test Patch Embedding\n# ============================================================\n\npatch_embed = PatchEmbedding().to(device)\n\ndummy = torch.randn(1, 3, 224, 224).to(device)\n\noutput = patch_embed(dummy)\n\nprint(\"=\" * 60)\nprint(\"Patch Embedding Test\")\nprint(\"=\" * 60)\n\nprint(\"Input Shape       :\", dummy.shape)\nprint(\"Output Shape      :\", output.shape)\nprint(\"Number of Patches :\", patch_embed.num_patches)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:57:52.888787Z","iopub.execute_input":"2026-07-12T11:57:52.889229Z","iopub.status.idle":"2026-07-12T11:57:53.785081Z","shell.execute_reply.started":"2026-07-12T11:57:52.889199Z","shell.execute_reply":"2026-07-12T11:57:53.784295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.datasets import ImageFolder\nfrom torchvision import transforms\n\ntransform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485,0.456,0.406],\n        std=[0.229,0.224,0.225]\n    )\n])\n\nVAL_DIR = \"/kaggle/working/imagenet/val\"\n\nval_dataset = ImageFolder(\n    VAL_DIR,\n    transform=transform\n)\n\nprint(\"Validation Images :\", len(val_dataset))\nprint(\"Classes :\", len(val_dataset.classes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:58:59.934885Z","iopub.execute_input":"2026-07-12T11:58:59.935647Z","iopub.status.idle":"2026-07-12T11:59:00.072898Z","shell.execute_reply.started":"2026-07-12T11:58:59.935601Z","shell.execute_reply":"2026-07-12T11:59:00.07214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nfrom PIL import Image\nimport numpy as np\n\n# Pick one validation image\nimg_path, _ = val_dataset.samples[0]\n\nimage = Image.open(img_path).convert(\"RGB\")\nimage = image.resize((224,224))\n\nplt.figure(figsize=(6,6))\nplt.imshow(image)\n\n# Draw 16x16 grid\nfor i in range(0,225,16):\n    plt.axhline(i, color='red', linewidth=0.4)\n    plt.axvline(i, color='red', linewidth=0.4)\n\nplt.title(\"16×16 Patch Grid (196 patches)\")\nplt.axis(\"off\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:59:02.488457Z","iopub.execute_input":"2026-07-12T11:59:02.488794Z","iopub.status.idle":"2026-07-12T11:59:02.811018Z","shell.execute_reply.started":"2026-07-12T11:59:02.488765Z","shell.execute_reply":"2026-07-12T11:59:02.810151Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Class Token and Positional Embedding\n\nThe transformer operates on a sequence of tokens.\n\nAfter patch embedding:\n\n196 Patch Tokens\n\n↓\n\nA learnable classification token (CLS) is prepended.\n\n↓\n\nLearnable positional embeddings are added to preserve spatial information.\n\nFinal sequence:\n\n197 × 384","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Embedding Layer\n# ============================================================\n\nclass Embedding(nn.Module):\n\n    def __init__(\n        self,\n        num_patches=196,\n        embed_dim=384\n    ):\n        super().__init__()\n\n        # Learnable CLS token\n        self.cls_token = nn.Parameter(\n            torch.zeros(1, 1, embed_dim)\n        )\n\n        # Learnable positional embeddings\n        self.position_embedding = nn.Parameter(\n            torch.zeros(1, num_patches + 1, embed_dim)\n        )\n\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.position_embedding, std=0.02)\n\n    def forward(self, x):\n\n        B = x.shape[0]\n\n        cls = self.cls_token.expand(B, -1, -1)\n\n        x = torch.cat((cls, x), dim=1)\n\n        x = x + self.position_embedding\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T11:59:58.160673Z","iopub.execute_input":"2026-07-12T11:59:58.161097Z","iopub.status.idle":"2026-07-12T11:59:58.167164Z","shell.execute_reply.started":"2026-07-12T11:59:58.161066Z","shell.execute_reply":"2026-07-12T11:59:58.166422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"embedding = Embedding().to(device)\n\ndummy = torch.randn(1,196,384).to(device)\n\nout = embedding(dummy)\n\nprint(out.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:00:06.274934Z","iopub.execute_input":"2026-07-12T12:00:06.275525Z","iopub.status.idle":"2026-07-12T12:00:06.355756Z","shell.execute_reply.started":"2026-07-12T12:00:06.275493Z","shell.execute_reply":"2026-07-12T12:00:06.355066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Vision Transformer Configuration\n# ============================================================\n\nclass ViTConfig:\n\n    IMAGE_SIZE = 224\n    PATCH_SIZE = 16\n\n    IN_CHANNELS = 3\n    NUM_CLASSES = 1000\n\n    EMBED_DIM = 384\n\n    DEPTH = 12\n    NUM_HEADS = 6\n\n    MLP_RATIO = 4\n\n    DROP_RATE = 0.0\n    ATTN_DROP_RATE = 0.0\n\n\ncfg = ViTConfig()\n\nprint(cfg.EMBED_DIM)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:00:50.043349Z","iopub.execute_input":"2026-07-12T12:00:50.044112Z","iopub.status.idle":"2026-07-12T12:00:50.049963Z","shell.execute_reply.started":"2026-07-12T12:00:50.04408Z","shell.execute_reply":"2026-07-12T12:00:50.049124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Patch Embedding\n# ============================================================\n\nclass PatchEmbedding(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        self.num_patches = (\n            cfg.IMAGE_SIZE // cfg.PATCH_SIZE\n        ) ** 2\n\n        self.projection = nn.Conv2d(\n            in_channels=cfg.IN_CHANNELS,\n            out_channels=cfg.EMBED_DIM,\n            kernel_size=cfg.PATCH_SIZE,\n            stride=cfg.PATCH_SIZE\n        )\n\n    def forward(self, x):\n\n        x = self.projection(x)\n\n        x = x.flatten(2)\n\n        x = x.transpose(1,2)\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:01:15.829369Z","iopub.execute_input":"2026-07-12T12:01:15.829783Z","iopub.status.idle":"2026-07-12T12:01:15.835517Z","shell.execute_reply.started":"2026-07-12T12:01:15.829753Z","shell.execute_reply":"2026-07-12T12:01:15.834616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patch_embed = PatchEmbedding(cfg).to(device)\n\ndummy = torch.randn(1,3,224,224).to(device)\n\nout = patch_embed(dummy)\n\nprint(out.shape)\nprint(patch_embed.num_patches)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:01:23.196922Z","iopub.execute_input":"2026-07-12T12:01:23.197684Z","iopub.status.idle":"2026-07-12T12:01:23.211466Z","shell.execute_reply.started":"2026-07-12T12:01:23.197652Z","shell.execute_reply":"2026-07-12T12:01:23.210891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Token Embedding\n# ============================================================\n\nclass TokenEmbedding(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        num_patches = (\n            cfg.IMAGE_SIZE //\n            cfg.PATCH_SIZE\n        ) ** 2\n\n        self.cls_token = nn.Parameter(\n            torch.zeros(1,1,cfg.EMBED_DIM)\n        )\n\n        self.position_embedding = nn.Parameter(\n            torch.zeros(\n                1,\n                num_patches+1,\n                cfg.EMBED_DIM\n            )\n        )\n\n        nn.init.trunc_normal_(self.cls_token,std=0.02)\n        nn.init.trunc_normal_(self.position_embedding,std=0.02)\n\n    def forward(self,x):\n\n        B = x.shape[0]\n\n        cls = self.cls_token.expand(B,-1,-1)\n\n        x = torch.cat((cls,x),dim=1)\n\n        x = x + self.position_embedding\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:01:29.805653Z","iopub.execute_input":"2026-07-12T12:01:29.806363Z","iopub.status.idle":"2026-07-12T12:01:29.811996Z","shell.execute_reply.started":"2026-07-12T12:01:29.806337Z","shell.execute_reply":"2026-07-12T12:01:29.811116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"embedding = TokenEmbedding(cfg).to(device)\n\ndummy = torch.randn(1,196,384).to(device)\n\nout = embedding(dummy)\n\nprint(out.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:01:36.008417Z","iopub.execute_input":"2026-07-12T12:01:36.008939Z","iopub.status.idle":"2026-07-12T12:01:36.016878Z","shell.execute_reply.started":"2026-07-12T12:01:36.008909Z","shell.execute_reply":"2026-07-12T12:01:36.016239Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Query, Key and Value Projection\n\nThe self-attention mechanism first projects every token embedding into three different representations:\n\n- Query (Q)\n- Key (K)\n- Value (V)\n\nFor DeiT-S:\n\nInput:\n- Tokens = 197\n- Embedding Dimension = 384\n\nOutput:\n- Q : (B, 197, 384)\n- K : (B, 197, 384)\n- V : (B, 197, 384)\n\nThese representations are then used to compute the self-attention matrix.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Query-Key-Value Projection\n# ============================================================\n\nclass QKVProjection(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        self.embed_dim = cfg.EMBED_DIM\n\n        self.qkv = nn.Linear(\n            cfg.EMBED_DIM,\n            cfg.EMBED_DIM * 3,\n            bias=True\n        )\n\n    def forward(self, x):\n\n        B, N, C = x.shape\n\n        qkv = self.qkv(x)\n\n        q, k, v = torch.chunk(qkv, chunks=3, dim=-1)\n\n        return q, k, v","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:02:07.849479Z","iopub.execute_input":"2026-07-12T12:02:07.849855Z","iopub.status.idle":"2026-07-12T12:02:07.855681Z","shell.execute_reply.started":"2026-07-12T12:02:07.849827Z","shell.execute_reply":"2026-07-12T12:02:07.854964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"qkv = QKVProjection(cfg).to(device)\n\ndummy = torch.randn(2, 197, 384).to(device)\n\nq, k, v = qkv(dummy)\n\nprint(\"=\" * 60)\nprint(\"QKV Projection Test\")\nprint(\"=\" * 60)\n\nprint(\"Input :\", dummy.shape)\nprint(\"Q     :\", q.shape)\nprint(\"K     :\", k.shape)\nprint(\"V     :\", v.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:02:15.129729Z","iopub.execute_input":"2026-07-12T12:02:15.1307Z","iopub.status.idle":"2026-07-12T12:02:15.281785Z","shell.execute_reply.started":"2026-07-12T12:02:15.130667Z","shell.execute_reply":"2026-07-12T12:02:15.281084Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Multi-Head Self Attention\n\nMulti-Head Self Attention enables every token to interact with every other token.\n\nFor each attention head:\n\nInput Tokens\n        ↓\nQuery, Key, Value\n        ↓\nScaled Dot-Product Attention\n        ↓\nSoftmax\n        ↓\nWeighted Value Aggregation\n\nThe outputs of all attention heads are concatenated and projected back to the embedding dimension.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Multi-Head Self Attention\n# ============================================================\n\nclass MultiHeadAttention(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        self.embed_dim = cfg.EMBED_DIM\n        self.num_heads = cfg.NUM_HEADS\n        self.head_dim = cfg.EMBED_DIM // cfg.NUM_HEADS\n\n        assert self.embed_dim % self.num_heads == 0\n\n        self.qkv = nn.Linear(\n            self.embed_dim,\n            self.embed_dim * 3\n        )\n\n        self.proj = nn.Linear(\n            self.embed_dim,\n            self.embed_dim\n        )\n\n    def forward(self, x):\n\n        B, N, C = x.shape\n\n        qkv = self.qkv(x)\n\n        qkv = qkv.reshape(\n            B,\n            N,\n            3,\n            self.num_heads,\n            self.head_dim\n        )\n\n        qkv = qkv.permute(2,0,3,1,4)\n\n        q, k, v = qkv[0], qkv[1], qkv[2]\n\n        attention = (q @ k.transpose(-2,-1))\n\n        attention = attention / math.sqrt(self.head_dim)\n\n        attention = attention.softmax(dim=-1)\n\n        out = attention @ v\n\n        out = out.transpose(1,2)\n\n        out = out.reshape(B,N,C)\n\n        out = self.proj(out)\n\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:03:06.753899Z","iopub.execute_input":"2026-07-12T12:03:06.754409Z","iopub.status.idle":"2026-07-12T12:03:06.761681Z","shell.execute_reply.started":"2026-07-12T12:03:06.754379Z","shell.execute_reply":"2026-07-12T12:03:06.760795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"attention = MultiHeadAttention(cfg).to(device)\n\ndummy = torch.randn(2,197,384).to(device)\n\nout = attention(dummy)\n\nprint(out.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:03:12.57943Z","iopub.execute_input":"2026-07-12T12:03:12.58041Z","iopub.status.idle":"2026-07-12T12:03:12.70007Z","shell.execute_reply.started":"2026-07-12T12:03:12.580377Z","shell.execute_reply":"2026-07-12T12:03:12.699256Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Feed Forward Network (MLP)\n\nEach Transformer block contains a position-wise feed-forward network after the self-attention layer.\n\nThe MLP consists of:\n\n- Linear Expansion\n- GELU Activation\n- Linear Projection\n\nFor DeiT-S:\n\n384 → 1536 → 384\n\nwhere\n\nHidden Dimension = 384 × 4 = 1536","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Feed Forward Network\n# ============================================================\n\nclass MLP(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        hidden_dim = int(cfg.EMBED_DIM * cfg.MLP_RATIO)\n\n        self.fc1 = nn.Linear(\n            cfg.EMBED_DIM,\n            hidden_dim\n        )\n\n        self.act = nn.GELU()\n\n        self.fc2 = nn.Linear(\n            hidden_dim,\n            cfg.EMBED_DIM\n        )\n\n    def forward(self, x):\n\n        x = self.fc1(x)\n\n        x = self.act(x)\n\n        x = self.fc2(x)\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:03:54.916132Z","iopub.execute_input":"2026-07-12T12:03:54.916595Z","iopub.status.idle":"2026-07-12T12:03:54.922115Z","shell.execute_reply.started":"2026-07-12T12:03:54.916527Z","shell.execute_reply":"2026-07-12T12:03:54.921301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mlp = MLP(cfg).to(device)\n\ndummy = torch.randn(2,197,384).to(device)\n\nout = mlp(dummy)\n\nprint(out.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:04:02.37638Z","iopub.execute_input":"2026-07-12T12:04:02.376979Z","iopub.status.idle":"2026-07-12T12:04:02.418509Z","shell.execute_reply.started":"2026-07-12T12:04:02.376944Z","shell.execute_reply":"2026-07-12T12:04:02.417551Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Transformer Encoder Block\n\nThe Transformer Encoder Block is the fundamental building block of the Vision Transformer.\n\nEach block consists of:\n\n1. Layer Normalization\n2. Multi-Head Self Attention\n3. Residual Connection\n4. Layer Normalization\n5. Feed Forward Network\n6. Residual Connection\n\nThe DeiT-S architecture stacks 12 identical Transformer Encoder Blocks.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Transformer Encoder Block\n# ============================================================\n\nclass TransformerBlock(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        self.norm1 = nn.LayerNorm(cfg.EMBED_DIM)\n\n        self.attention = MultiHeadAttention(cfg)\n\n        self.norm2 = nn.LayerNorm(cfg.EMBED_DIM)\n\n        self.mlp = MLP(cfg)\n\n    def forward(self, x):\n\n        # Self-Attention + Residual\n        x = x + self.attention(self.norm1(x))\n\n        # MLP + Residual\n        x = x + self.mlp(self.norm2(x))\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:04:39.408907Z","iopub.execute_input":"2026-07-12T12:04:39.409348Z","iopub.status.idle":"2026-07-12T12:04:39.414905Z","shell.execute_reply.started":"2026-07-12T12:04:39.409317Z","shell.execute_reply":"2026-07-12T12:04:39.413964Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"block = TransformerBlock(cfg).to(device)\n\ndummy = torch.randn(2,197,384).to(device)\n\nout = block(dummy)\n\nprint(out.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:04:47.077446Z","iopub.execute_input":"2026-07-12T12:04:47.077867Z","iopub.status.idle":"2026-07-12T12:04:47.134657Z","shell.execute_reply.started":"2026-07-12T12:04:47.077837Z","shell.execute_reply":"2026-07-12T12:04:47.133829Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Vision Transformer\n\nThe Vision Transformer consists of:\n\n- Patch Embedding\n- Token Embedding\n- 12 Transformer Encoder Blocks\n- Final Layer Normalization\n- Classification Head\n\nThe final prediction is obtained from the CLS token.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Vision Transformer\n# ============================================================\n\nclass VisionTransformer(nn.Module):\n\n    def __init__(self, cfg):\n        super().__init__()\n\n        self.patch_embed = PatchEmbedding(cfg)\n\n        self.token_embed = TokenEmbedding(cfg)\n\n        self.blocks = nn.ModuleList([\n            TransformerBlock(cfg)\n            for _ in range(cfg.DEPTH)\n        ])\n\n        self.norm = nn.LayerNorm(cfg.EMBED_DIM)\n\n        self.head = nn.Linear(\n            cfg.EMBED_DIM,\n            cfg.NUM_CLASSES\n        )\n\n    def forward(self, x):\n\n        x = self.patch_embed(x)\n\n        x = self.token_embed(x)\n\n        for block in self.blocks:\n            x = block(x)\n\n        x = self.norm(x)\n\n        cls = x[:, 0]\n\n        logits = self.head(cls)\n\n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:05:22.157129Z","iopub.execute_input":"2026-07-12T12:05:22.157822Z","iopub.status.idle":"2026-07-12T12:05:22.163527Z","shell.execute_reply.started":"2026-07-12T12:05:22.157792Z","shell.execute_reply":"2026-07-12T12:05:22.162883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = VisionTransformer(cfg).to(device)\n\ndummy = torch.randn(2,3,224,224).to(device)\n\nlogits = model(dummy)\n\nprint(\"=\" * 60)\nprint(\"Vision Transformer Test\")\nprint(\"=\" * 60)\nprint(\"Input  :\", dummy.shape)\nprint(\"Output :\", logits.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:05:27.985468Z","iopub.execute_input":"2026-07-12T12:05:27.985799Z","iopub.status.idle":"2026-07-12T12:05:28.233986Z","shell.execute_reply.started":"2026-07-12T12:05:27.985772Z","shell.execute_reply":"2026-07-12T12:05:28.233062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Model Statistics\n# ============================================================\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(\"=\" * 60)\nprint(\"Model Statistics\")\nprint(\"=\" * 60)\nprint(f\"Total Parameters     : {total_params:,}\")\nprint(f\"Trainable Parameters : {trainable_params:,}\")\nprint(f\"Model Size           : {total_params/1e6:.2f} M\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:05:50.673087Z","iopub.execute_input":"2026-07-12T12:05:50.673518Z","iopub.status.idle":"2026-07-12T12:05:50.67997Z","shell.execute_reply.started":"2026-07-12T12:05:50.67347Z","shell.execute_reply":"2026-07-12T12:05:50.679253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!wget -q https://dl.fbaipublicfiles.com/deit/deit_small_patch16_224-cd65a155.pth \\\n    -O /kaggle/working/efficient-vit/weights/deit_small_patch16_224.pth","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:07:33.143985Z","iopub.execute_input":"2026-07-12T12:07:33.144762Z","iopub.status.idle":"2026-07-12T12:07:35.78306Z","shell.execute_reply.started":"2026-07-12T12:07:33.144708Z","shell.execute_reply":"2026-07-12T12:07:35.782232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nweight_path = \"/kaggle/working/efficient-vit/weights/deit_small_patch16_224.pth\"\n\nprint(os.path.exists(weight_path))\nprint(os.path.getsize(weight_path)/1024/1024, \"MB\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:07:41.937087Z","iopub.execute_input":"2026-07-12T12:07:41.93795Z","iopub.status.idle":"2026-07-12T12:07:41.943329Z","shell.execute_reply.started":"2026-07-12T12:07:41.937893Z","shell.execute_reply":"2026-07-12T12:07:41.94238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"state_dict = checkpoint[\"model\"]\n\nprint(type(state_dict))\nprint(len(state_dict))\n\nprint(\"\\nFirst 20 keys:\\n\")\n\nfor i, key in enumerate(state_dict.keys()):\n    print(key)\n    if i == 19:\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:08:13.445391Z","iopub.execute_input":"2026-07-12T12:08:13.446112Z","iopub.status.idle":"2026-07-12T12:08:13.45142Z","shell.execute_reply.started":"2026-07-12T12:08:13.446078Z","shell.execute_reply":"2026-07-12T12:08:13.450453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"=\"*60)\nprint(\"OUR MODEL KEYS\")\nprint(\"=\"*60)\n\nfor i, k in enumerate(model.state_dict().keys()):\n    print(k)\n    if i == 25:\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:08:53.849212Z","iopub.execute_input":"2026-07-12T12:08:53.849504Z","iopub.status.idle":"2026-07-12T12:08:53.85868Z","shell.execute_reply.started":"2026-07-12T12:08:53.849479Z","shell.execute_reply":"2026-07-12T12:08:53.85793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Convert Official DeiT Checkpoint\n# ============================================================\n\ndef convert_deit_checkpoint(state_dict):\n\n    converted = {}\n\n    for key, value in state_dict.items():\n\n        # Patch Embedding\n        key = key.replace(\n            \"patch_embed.proj\",\n            \"patch_embed.projection\"\n        )\n\n        # CLS Token\n        key = key.replace(\n            \"cls_token\",\n            \"token_embed.cls_token\"\n        )\n\n        # Positional Embedding\n        key = key.replace(\n            \"pos_embed\",\n            \"token_embed.position_embedding\"\n        )\n\n        # Attention\n        key = key.replace(\n            \".attn.\",\n            \".attention.\"\n        )\n\n        converted[key] = value\n\n    return converted","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:09:26.679058Z","iopub.execute_input":"2026-07-12T12:09:26.679788Z","iopub.status.idle":"2026-07-12T12:09:26.686072Z","shell.execute_reply.started":"2026-07-12T12:09:26.679753Z","shell.execute_reply":"2026-07-12T12:09:26.685236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"converted = convert_deit_checkpoint(state_dict)\n\nprint(list(converted.keys())[:20])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:09:34.205755Z","iopub.execute_input":"2026-07-12T12:09:34.206618Z","iopub.status.idle":"2026-07-12T12:09:34.210935Z","shell.execute_reply.started":"2026-07-12T12:09:34.206534Z","shell.execute_reply":"2026-07-12T12:09:34.21006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"missing, unexpected = model.load_state_dict(\n    converted,\n    strict=False\n)\n\nprint(\"Missing Keys\")\nprint(missing)\n\nprint()\n\nprint(\"Unexpected Keys\")\nprint(unexpected)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:09:43.489807Z","iopub.execute_input":"2026-07-12T12:09:43.490102Z","iopub.status.idle":"2026-07-12T12:09:43.526433Z","shell.execute_reply.started":"2026-07-12T12:09:43.490078Z","shell.execute_reply":"2026-07-12T12:09:43.525604Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ImageNet Validation Dataset\n# ============================================================\n\ntransform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485,0.456,0.406],\n        std =[0.229,0.224,0.225]\n    )\n])\n\nval_dataset = ImageFolder(\n    \"/kaggle/working/imagenet/val\",\n    transform=transform\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=64,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"=\"*60)\nprint(\"Validation Dataset\")\nprint(\"=\"*60)\nprint(\"Images :\", len(val_dataset))\nprint(\"Classes:\", len(val_dataset.classes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:10:34.724087Z","iopub.execute_input":"2026-07-12T12:10:34.72484Z","iopub.status.idle":"2026-07-12T12:10:34.863234Z","shell.execute_reply.started":"2026-07-12T12:10:34.724807Z","shell.execute_reply":"2026-07-12T12:10:34.86255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Put Model Into Evaluation Mode\n# ============================================================\n\nmodel.eval()\n\nprint(\"Model ready for evaluation.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:10:42.100803Z","iopub.execute_input":"2026-07-12T12:10:42.101217Z","iopub.status.idle":"2026-07-12T12:10:42.106198Z","shell.execute_reply.started":"2026-07-12T12:10:42.101187Z","shell.execute_reply":"2026-07-12T12:10:42.105584Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation Function\n\nA common evaluation pipeline is implemented to benchmark all Vision Transformer variants.\n\nMetrics:\n\n- Top-1 Accuracy\n- Top-5 Accuracy\n- Inference Time\n- Throughput (Images / Second)","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# Evaluation Function\n# ============================================================\n\n@torch.no_grad()\ndef evaluate(model, dataloader, device):\n\n    model.eval()\n\n    top1 = 0\n    top5 = 0\n    total = 0\n\n    start = time.time()\n\n    for images, labels in dataloader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        logits = model(images)\n\n        # Top-1\n        pred1 = logits.argmax(dim=1)\n        top1 += (pred1 == labels).sum().item()\n\n        # Top-5\n        pred5 = logits.topk(5, dim=1).indices\n        top5 += (\n            pred5 ==\n            labels.unsqueeze(1)\n        ).any(dim=1).sum().item()\n\n        total += labels.size(0)\n\n    elapsed = time.time() - start\n\n    return {\n        \"Top1\":100*top1/total,\n        \"Top5\":100*top5/total,\n        \"Time\":elapsed,\n        \"FPS\":total/elapsed\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:11:23.259Z","iopub.execute_input":"2026-07-12T12:11:23.259437Z","iopub.status.idle":"2026-07-12T12:11:23.266381Z","shell.execute_reply.started":"2026-07-12T12:11:23.259406Z","shell.execute_reply":"2026-07-12T12:11:23.2656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"stats = evaluate(\n    model,\n    val_loader,\n    device\n)\n\nprint(stats)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:11:29.502785Z","iopub.execute_input":"2026-07-12T12:11:29.503411Z","iopub.status.idle":"2026-07-12T12:14:56.422412Z","shell.execute_reply.started":"2026-07-12T12:11:29.50338Z","shell.execute_reply":"2026-07-12T12:14:56.421504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q fvcore\nfrom fvcore.nn import FlopCountAnalysis\n\ndummy = torch.randn(1,3,224,224).to(device)\n\nflops = FlopCountAnalysis(model, dummy)\n\nprint(\"=\"*60)\nprint(\"Model Complexity\")\nprint(\"=\"*60)\n\nprint(f\"GFLOPs     : {flops.total()/1e9:.2f}\")\nprint(f\"Parameters : {sum(p.numel() for p in model.parameters())/1e6:.2f} M\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:17:00.862744Z","iopub.execute_input":"2026-07-12T12:17:00.863632Z","iopub.status.idle":"2026-07-12T12:17:10.429464Z","shell.execute_reply.started":"2026-07-12T12:17:00.863582Z","shell.execute_reply":"2026-07-12T12:17:10.428849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Baseline Results\n# ============================================================\n\nbaseline_results = {\n    \"Model\": \"BaseViT\",\n    \"Top1\": 78.832,\n    \"Top5\": 94.482,\n    \"GFLOPs\": 4.61,\n    \"Parameters(M)\": 22.05,\n    \"FPS\": 241.65\n}\n\nprint(baseline_results)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:17:48.421477Z","iopub.execute_input":"2026-07-12T12:17:48.422069Z","iopub.status.idle":"2026-07-12T12:17:48.427746Z","shell.execute_reply.started":"2026-07-12T12:17:48.422028Z","shell.execute_reply":"2026-07-12T12:17:48.427016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.DataFrame([baseline_results])\n\ndf.to_csv(\n    \"/kaggle/working/efficient-vit/results/basevit_results.csv\",\n    index=False\n)\n\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-12T12:18:03.520527Z","iopub.execute_input":"2026-07-12T12:18:03.520901Z","iopub.status.idle":"2026-07-12T12:18:04.042796Z","shell.execute_reply.started":"2026-07-12T12:18:03.520875Z","shell.execute_reply":"2026-07-12T12:18:04.04217Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Results\n\nThe proposed Base Vision Transformer was implemented from scratch in PyTorch and validated using the official pretrained DeiT-S checkpoint. The pretrained weights were successfully adapted and loaded into the implemented architecture, confirming full architectural compatibility.\n\nResults\n\n- Parameters : **22.05 Million**\n- GFLOPs : **4.61**\n- Top-1 Accuracy : **78.83%**\n- Top-5 Accuracy : **94.48%**\n- Throughput : **241.65 Images/sec**\n\nThe reproduced BaseViT serves as the foundation for the subsequent implementation of Adaptive Token Sampling (ATS), Adaptive Vision Transformer (A-ViT), and AdaViT. All efficient Vision Transformer variants will be developed by extending this common baseline to ensure a consistent implementation and enable fair experimental comparisons.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}