{"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":"gpu","dataSources":[{"sourceId":21154,"databundleVersionId":1243559,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## VIT Paper: IMAGE IS WORTH 16X16 WORDS: TRANSFORMERS FOR IMAGE RECOGNITION AT SCALE","metadata":{}},{"cell_type":"markdown","source":"[ 3rd Jun, 2021 ]\n\n[Paper:](https://www.alphaxiv.org/abs/2010.11929) IMAGE IS WORTH 16X16 WORDS: TRANSFORMERS FOR IMAGE RECOGNITION AT SCALE\n\nArchitecture [code](https://www.notion.so/Vision-Transformer-ViT-2d7c61877d7b80a1a375f76c357e4a06?pvs=21)\n\nFinetuning [blog](https://huggingface.co/blog/fine-tune-vit)\n\nAbstract\n\nThis paper introduces the Vision Transformer (ViT), a model that adapts the standard Transformer architecture to image classification by treating images as sequences of patches. The work demonstrates that when scaled to large datasets, pure attention-based models can match or exceed the performance of convolutional neural networks (CNNs), challenging the dominant paradigm in computer vision.\n\n\n\n## Architecture and Methodology\n\nThe Vision Transformer directly applies a standard Transformer encoder to image classification with minimal modifications. The key innovation lies in how images are processed as sequences:\n\n**Image Patch Processing**: An input image is divided into fixed-size patches (typically 16×16 pixels), which are then flattened and linearly projected to create patch embeddings. This transforms the 2D image into a 1D sequence that the Transformer can process.\n\n```python\n# Image dimensions\nimage = height × width × channels\n\n# Step 2: Divide the image into patches (like tiles)\n# Example: 224 × 224 pixels\n(224 / 16) × (224 / 16) = 14 × 14 = 196 patches\n\n# Each patch dimensions\neach patch dims = 16×16×3 = 768 numbers\n```\n\nNow, each path from input is projected to the Embedding space of dims **W = D X 768** \n\nsuch that \n\n**patch_embedding** =     W      ×   patch_vector +   b\n***( D)                               (D,768)           ( 768 )           (D)***\n\n**Now, each patch embedding of D → 768 is converted to a patch embedding of dimension D** \n\nEach path will become like this\n\n```python\n[patch1_embedding,\n patch2_embedding,\n patch3_embedding,\n ...\n patch196_embedding]\n```\n\nEach output dimension is a **weighted combination of all pixels in the patch**.\n\n$$\n\\text{embedding}[0] = w_{0,1} \\cdot p_1 + w_{0,2} \\cdot p_2 + \\cdots + w_{0,768} \\cdot p_{768} + b_0\n$$\n\nAlso, remember that in the input, we give an extra learnable position embedding at the start for the **[CLS] token**. The purpose of this embedding is to learn globally about the entire patch\n\n\n\n- **Index 0** → **extra learnable [class] embedding**\n- **Index 1…N** → patch embeddings (image patches)\n\nIntuition behind this extra **[CLS] (class)** token is to learn the entire image, consider it like a summary slot for the entire image.\n\n**Sequence Construction**: Similar to BERT's approach in NLP, a learnable classification token is prepended to the sequence of patch embeddings. Position embeddings are added to retain spatial information, though the authors found that 1D position embeddings work as well as 2D-aware variants.\n\nThe mathematical formulation for the embedding process is:\n\n$$\n\\mathbf{z}_0 = [\\mathbf{x}_{\\text{class}}; \\mathbf{x}^1_p\\mathbf{E}; \\mathbf{x}^2_p\\mathbf{E}; \\cdots; \\mathbf{x}^N_p\\mathbf{E}] + \\mathbf{E}_{\\text{pos}}\n$$\n\n\n**Transformer Processing**: The embedded sequence is processed through standard Transformer encoder layers with multi-head self-attention and MLP blocks. The final representation of the classification token serves as the image representation for classification.\n\n## Inductive Bias vs. Data Scale\n\nWhat is *inductive bias*? \n\nInductive bias = built-in assumptions a model has about the world\n\nThink of it like this:\n\n- A **calculator** has the inductive bias that numbers follow math rules\n- A **CNN** has the inductive bias that images have structure (nearby pixels matter)\n- A **Vision Transformer (ViT)** has *very few* such assumptions\n\n**CNNs: Strong inductive bias (they \"understand images\" by design)**\n\n**1. Locality**\n\nNearby pixels are related.\n\n→ A CNN looks at small regions (3×3, 5×5 filters)\n\n**2. 2D structure**\n\nImages are grids (height × width)\n\n→ CNN filters slide in 2D space\n\n**3. Translation equivariance**\n\nIf an object moves slightly, it's still the same object\n\n→ A cat is a cat whether it's on the left or right\n\n**These assumptions exist in every layer of a CNN**\n\nSo CNNs don't need to *learn* these ideas—they're built in.\n\n## Vision Transformer (ViT): Weak inductive bias\n\n**ViT makes minimal assumptions about image structure.** Instead, it:\n\n1. Divides images into patches (like 16×16 squares)\n2. Treats patches as tokens (like words in NLP)\n3. Uses self-attention to relate all patches\n\n**Key difference:**\n\n- **CNN**: Starts local → builds to global\n- **ViT**: Global from the start → every patch attends to every other patch\n\n- **Small Data Regimes**: Inductive biases provide advantages when data is limited, hence CNNs are better when data is less.\n- **Large Data Regimes**: Generic architectures can learn appropriate representations directly from data, potentially surpassing specialised architectures\n- **Computational Trade-offs**: The lack of inductive bias can be compensated by scale, often with better computational efficiency\n\nVITs only perform well when pretrained on a large scale of data (300M+) images, as stated in the paper.","metadata":{}},{"cell_type":"markdown","source":"## Let's  use pre-trained VIT for image classification task here","metadata":{}},{"cell_type":"code","source":"%%capture\npip install tfrecord","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:15:03.519562Z","iopub.execute_input":"2025-12-31T14:15:03.519822Z","iopub.status.idle":"2025-12-31T14:15:06.606232Z","shell.execute_reply.started":"2025-12-31T14:15:03.519800Z","shell.execute_reply":"2025-12-31T14:15:06.605455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:57:45.713033Z","iopub.execute_input":"2025-12-31T13:57:45.713298Z","iopub.status.idle":"2025-12-31T13:57:45.988887Z","shell.execute_reply.started":"2025-12-31T13:57:45.713265Z","shell.execute_reply":"2025-12-31T13:57:45.988225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"TRANSFORMERS_NO_TF\"] = \"1\"\nkaggle_dir = \"/kaggle/input/tpu-getting-started\"\nimport io\nimport math\nimport os\nimport random\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import IterableDataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torchvision import transforms\n\nfrom sklearn.metrics import f1_score, precision_score, recall_score, confusion_matrix\n\nimport timm\nimport tfrecord\n\n\n\nfrom torch.optim import AdamW\nfrom transformers import get_cosine_schedule_with_warmup","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:57:45.989754Z","iopub.execute_input":"2025-12-31T13:57:45.990190Z","iopub.status.idle":"2025-12-31T13:57:59.060078Z","shell.execute_reply.started":"2025-12-31T13:57:45.990167Z","shell.execute_reply":"2025-12-31T13:57:59.059437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_CLASSES = [\n        'pink primrose', 'hard-leaved pocket orchid', 'canterbury bells', 'sweet pea', 'wild geranium',\n        'tiger lily', 'moon orchid', 'bird of paradise', 'monkshood', 'globe thistle',\n        'snapdragon', \"colt's foot\", 'king protea', 'spear thistle', 'yellow iris',\n        'globe-flower', 'purple coneflower', 'peruvian lily', 'balloon flower', 'giant white arum lily',\n        'fire lily', 'pincushion flower', 'fritillary', 'red ginger', 'grape hyacinth',\n        'corn poppy', 'prince of wales feathers', 'stemless gentian', 'artichoke', 'sweet william',\n        'carnation', 'garden phlox', 'love in the mist', 'cosmos', 'alpine sea holly',\n        'ruby-lipped cattleya', 'cape flower', 'great masterwort', 'siam tulip', 'lenten rose',\n        'barberton daisy', 'daffodil', 'sword lily', 'poinsettia', 'bolero deep blue',\n        'wallflower', 'marigold', 'buttercup', 'daisy', 'common dandelion',\n        'petunia', 'wild pansy', 'primula', 'sunflower', 'lilac hibiscus',\n        'bishop of llandaff', 'gaura', 'geranium', 'orange dahlia', 'pink-yellow dahlia',\n        'cautleya spicata', 'japanese anemone', 'black-eyed susan', 'silverbush', 'californian poppy',\n        'osteospermum', 'spring crocus', 'iris', 'windflower', 'tree poppy',\n        'gazania', 'azalea', 'water lily', 'rose', 'thorn apple',\n        'morning glory', 'passion flower', 'lotus', 'toad lily', 'anthurium',\n        'frangipani', 'clematis', 'hibiscus', 'columbine', 'desert-rose',\n        'tree mallow', 'magnolia', 'cyclamen ', 'watercress', 'canna lily',\n        'hippeastrum ', 'bee balm', 'pink quill', 'foxglove', 'bougainvillea',\n        'camellia', 'mallow', 'mexican petunia', 'bromelia', 'blanket flower',\n        'trumpet creeper', 'blackberry lily', 'common tulip', 'wild rose'\n    ]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T13:57:59.061593Z","iopub.execute_input":"2025-12-31T13:57:59.062005Z","iopub.status.idle":"2025-12-31T13:57:59.067534Z","shell.execute_reply.started":"2025-12-31T13:57:59.061981Z","shell.execute_reply":"2025-12-31T13:57:59.066781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Model config\n\nfrom pathlib import Path\n\n\nclass ModelConfig:\n    base_path = Path('/kaggle/input/tpu-getting-started')\n    tfrecord_dir = base_path / 'tfrecords-jpeg-512x512'\n    train_dir = tfrecord_dir / 'train'\n    val_dir = tfrecord_dir / 'val'\n    test_dir = tfrecord_dir / 'test'\n\n    num_classes = 104\n\n    batch_size = 64\n    val_batch_size = 64\n    test_batch_size = 64\n    epochs = 5\n\n    lr = 3e-4\n    min_lr = 1e-6\n    weight_decay = 1e-5\n    grad_clip = 1.0\n\n    num_workers = 0\n    seed = 2235\n\n    model_name = \"vit_base_patch16_224\"\n    drop_rate = 0.214 # dropout probability for transformer.\n    drop_path_rate = 0.118\n    use_amp = True\n\n    output_dir = Path('/kaggle/working')\n    model_path = output_dir / 'best_model.pt'\n    submission_path = output_dir / 'submission.csv'\n\n    train_limit = None\n    val_limit = None\n    test_limit = None\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    classes = MODEL_CLASSES\n\n\nModelConfig.output_dir.mkdir(parents=True, exist_ok=True)\n    \n\ndef set_seed(seed: int = ModelConfig.seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nset_seed()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:10.063494Z","iopub.execute_input":"2025-12-31T14:00:10.063833Z","iopub.status.idle":"2025-12-31T14:00:10.072695Z","shell.execute_reply.started":"2025-12-31T14:00:10.063807Z","shell.execute_reply":"2025-12-31T14:00:10.072113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import re\n\n\ndef get_tfrecord_files(directory: Path):\n    files = sorted(str(p) for p in directory.glob('*.tfrec'))\n    if not files:\n        raise FileNotFoundError(f'No TFRecord files found in {directory}')\n    return files\n\n\ndef count_data_items(filenames):\n    pattern = re.compile(r'-([0-9]*)\\.tfrec$')\n    total = 0\n    for fname in filenames:\n        match = pattern.search(fname)\n        if match:\n            total += int(match.group(1))\n    return total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:10.311429Z","iopub.execute_input":"2025-12-31T14:00:10.311894Z","iopub.status.idle":"2025-12-31T14:00:10.316390Z","shell.execute_reply.started":"2025-12-31T14:00:10.311868Z","shell.execute_reply":"2025-12-31T14:00:10.315768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### General VIT based transformers\n\nIMAGE_SIZE = 224  # ViT expects 224\n\n\ndef get_train_transforms():\n    return transforms.Compose([\n    transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.5, 0.5, 0.5],\n        std=[0.5, 0.5, 0.5]\n    )\n])\n\n\n\ndef get_eval_transforms():\n    return transforms.Compose([\n    transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.5, 0.5, 0.5],\n        std=[0.5, 0.5, 0.5]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.063422Z","iopub.execute_input":"2025-12-31T14:00:12.063908Z","iopub.status.idle":"2025-12-31T14:00:12.069122Z","shell.execute_reply.started":"2025-12-31T14:00:12.063880Z","shell.execute_reply":"2025-12-31T14:00:12.068336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## Building Pytorch dataset\nclass FlowerTFRecordDataset(IterableDataset):\n    def __init__(self, filenames, labeled=True, transform=None, seed=ModelConfig.seed, limit=None):\n        self.filenames = list(filenames)\n        self.labeled = labeled\n        self.transform = transform\n        self.seed = seed\n        self.limit = limit\n        self.epoch = 0\n        total_items = count_data_items(self.filenames)\n        self.num_samples = min(total_items, limit) if limit is not None else total_items\n\n    def set_epoch(self, epoch: int):\n        self.epoch = epoch\n\n    def __len__(self):\n        return self.num_samples\n\n    def _shuffle_files(self):\n        rng = random.Random(self.seed + self.epoch)\n        files = self.filenames.copy()\n        rng.shuffle(files)\n        return files\n\n    def _shuffle_records(self, records, *, fname):\n        rng = random.Random(self.seed + self.epoch + hash(fname) % 10_000)\n        records = list(records)\n        rng.shuffle(records)\n        return records\n\n    def __iter__(self):\n        files = self._shuffle_files()\n        worker_info = torch.utils.data.get_worker_info()\n        if worker_info is not None:\n            files = files[worker_info.id :: worker_info.num_workers]\n\n        produced = 0\n        for fname in files:\n            if self.limit is not None and produced >= self.limit:\n                break\n\n            try:\n                records = tfrecord.tfrecord_loader(fname, None, None)\n            except Exception as exc:\n                print(f'Warning: failed to read {fname}: {exc}')\n                continue\n\n            for record in self._shuffle_records(records, fname=fname):\n                if self.limit is not None and produced >= self.limit:\n                    break\n\n                image_bytes = record['image']\n                image = Image.open(io.BytesIO(image_bytes)).convert('RGB')\n                if self.transform is not None:\n                    image = self.transform(image)\n\n                if self.labeled:\n                    label_raw = record['class']\n                    if isinstance(label_raw, np.ndarray):\n                        label = int(label_raw.reshape(-1)[0])\n                    elif isinstance(label_raw, (list, tuple)):\n                        label = int(label_raw[0])\n                    else:\n                        label = int(label_raw)\n                    yield image, label\n                else:\n                    image_id = record.get('id', '')\n                    if isinstance(image_id, bytes):\n                        image_id = image_id.decode()\n                    yield image, image_id\n\n                produced += 1\n\n\ndef build_loader(directory: Path, *, labeled: bool, transform, batch_size: int, limit=None):\n    filenames = get_tfrecord_files(directory)\n    dataset = FlowerTFRecordDataset(filenames, labeled=labeled, transform=transform, limit=limit)\n\n    loader = DataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=ModelConfig.num_workers,\n        pin_memory=True,\n        drop_last=False,\n    )\n\n    return loader, dataset\n\n\ndef get_loaders():\n    train_loader, _ = build_loader(\n        ModelConfig.train_dir,\n        labeled=True,\n        transform=get_train_transforms(),\n        batch_size=ModelConfig.batch_size,\n        limit=ModelConfig.train_limit,\n    )\n\n    val_loader, _ = build_loader(\n        ModelConfig.val_dir,\n        labeled=True,\n        transform=get_eval_transforms(),\n        batch_size=ModelConfig.val_batch_size,\n        limit=ModelConfig.val_limit,\n    )\n\n    return train_loader, val_loader\n\n\ndef get_test_loader():\n    test_loader, _ = build_loader(\n        ModelConfig.test_dir,\n        labeled=False,\n        transform=get_eval_transforms(),\n        batch_size=ModelConfig.test_batch_size,\n        limit=ModelConfig.test_limit,\n    )\n    return test_loader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.223263Z","iopub.execute_input":"2025-12-31T14:00:12.223791Z","iopub.status.idle":"2025-12-31T14:00:12.236301Z","shell.execute_reply.started":"2025-12-31T14:00:12.223757Z","shell.execute_reply":"2025-12-31T14:00:12.235517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model():\n    model = timm.create_model(\n        ModelConfig.model_name,\n        pretrained=True,\n        num_classes=ModelConfig.num_classes,\n        drop_rate=ModelConfig.drop_rate,\n        drop_path_rate=ModelConfig.drop_path_rate,\n    )\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.383049Z","iopub.execute_input":"2025-12-31T14:00:12.383324Z","iopub.status.idle":"2025-12-31T14:00:12.387269Z","shell.execute_reply.started":"2025-12-31T14:00:12.383301Z","shell.execute_reply":"2025-12-31T14:00:12.386485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.optim import AdamW\nfrom transformers import get_cosine_schedule_with_warmup\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.519067Z","iopub.execute_input":"2025-12-31T14:00:12.519338Z","iopub.status.idle":"2025-12-31T14:00:12.522907Z","shell.execute_reply.started":"2025-12-31T14:00:12.519316Z","shell.execute_reply":"2025-12-31T14:00:12.522309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_optimizer(model, steps_per_epoch):\n    optimizer = AdamW(\n        model.parameters(),\n        lr=ModelConfig.lr,\n        weight_decay=ModelConfig.weight_decay,\n    )\n\n    scheduler = get_cosine_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps=int(0.1 * steps_per_epoch),\n        num_training_steps=ModelConfig.epochs * steps_per_epoch,\n    )\n    return optimizer, scheduler\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.671265Z","iopub.execute_input":"2025-12-31T14:00:12.671812Z","iopub.status.idle":"2025-12-31T14:00:12.675892Z","shell.execute_reply.started":"2025-12-31T14:00:12.671784Z","shell.execute_reply":"2025-12-31T14:00:12.675231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.auto import tqdm\nimport torch.nn.functional as F\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.827339Z","iopub.execute_input":"2025-12-31T14:00:12.827865Z","iopub.status.idle":"2025-12-31T14:00:12.831445Z","shell.execute_reply.started":"2025-12-31T14:00:12.827838Z","shell.execute_reply":"2025-12-31T14:00:12.830738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(\n    model,\n    loader,\n    optimizer,\n    scheduler,\n    epoch,\n    total_steps,\n    device,\n):\n    model.train()\n    loader.dataset.set_epoch(epoch)\n\n    total_loss = 0.0\n    total_samples = 0\n\n    # tqdm only on rank 0 (DDP-safe)\n\n    pbar = tqdm(\n        loader,\n        total=total_steps,\n        leave=False,\n    )\n    pbar.set_description(f\"Epoch {epoch + 1} [train]\")\n\n    for step, (images, labels) in enumerate(pbar):\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        outputs = model(images)\n        logits = outputs.logits if hasattr(outputs, \"logits\") else outputs\n\n        loss = F.cross_entropy(logits, labels)\n\n        optimizer.zero_grad(set_to_none=True)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n\n        batch_size = images.size(0)\n        total_loss += loss.item() * batch_size\n        total_samples += batch_size\n\n        # Live metrics\n        pbar.set_postfix({\n            \"loss\": f\"{loss.item():.4f}\",\n            \"lr\": f\"{scheduler.get_last_lr()[0]:.2e}\",\n            \"seen\": total_samples,\n        })\n\n    return total_loss / total_samples\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:12.967013Z","iopub.execute_input":"2025-12-31T14:00:12.967667Z","iopub.status.idle":"2025-12-31T14:00:12.973510Z","shell.execute_reply.started":"2025-12-31T14:00:12.967640Z","shell.execute_reply":"2025-12-31T14:00:12.972822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate(model, loader):\n    model.eval()\n\n    correct = 0\n    total = 0\n\n    for images, labels in loader:\n        images = images.to(ModelConfig.device, non_blocking=True)\n        labels = labels.to(ModelConfig.device, non_blocking=True)\n\n        outputs = model(images)\n        preds = outputs.argmax(dim=1)\n\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n    return correct / total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:13.127163Z","iopub.execute_input":"2025-12-31T14:00:13.127462Z","iopub.status.idle":"2025-12-31T14:00:13.132416Z","shell.execute_reply.started":"2025-12-31T14:00:13.127436Z","shell.execute_reply":"2025-12-31T14:00:13.131587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.manual_seed(ModelConfig.seed)\n\n# Loaders (your code)\ntrain_loader, val_loader = get_loaders()\n\n# Model\nmodel = build_model()  # or build_vit_from_scratch()\nmodel.to(ModelConfig.device)\n\nsteps_per_epoch = len(train_loader.dataset) // ModelConfig.batch_size\noptimizer, scheduler = build_optimizer(model, steps_per_epoch)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:17.491246Z","iopub.execute_input":"2025-12-31T14:00:17.491845Z","iopub.status.idle":"2025-12-31T14:00:18.922521Z","shell.execute_reply.started":"2025-12-31T14:00:17.491817Z","shell.execute_reply":"2025-12-31T14:00:18.921905Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfor epoch in range(ModelConfig.epochs):\n    steps_per_epoch = len(train_loader.dataset) // ModelConfig.batch_size\n    train_loss = train_one_epoch(\n        model,\n        train_loader,\n        optimizer,\n        scheduler,\n        epoch,\n        total_steps=steps_per_epoch,\n        device=ModelConfig.device,\n    )\n    val_acc = validate(model, val_loader)\n    # print(\"*\"*80)\n    print(\n        f\"Epoch [{epoch+1}/{ModelConfig.epochs}] \"\n        f\"| Train Loss: {train_loss:.4f} \"\n        f\"| Val Acc: {val_acc:.4f}\"\n    )\n    # print(\"*\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:00:18.923794Z","iopub.execute_input":"2025-12-31T14:00:18.924068Z","iopub.status.idle":"2025-12-31T14:12:18.198342Z","shell.execute_reply.started":"2025-12-31T14:00:18.924039Z","shell.execute_reply":"2025-12-31T14:12:18.197423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef generate_submission(model, loader):\n    model.eval()\n    all_preds = []\n    all_ids = []\n\n    pbar = tqdm(loader, leave=False)\n    pbar.set_description('Predicting test')\n\n    for images, image_ids in pbar:\n        images = images.to(ModelConfig.device, non_blocking=True)\n        outputs = model(images)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_ids.extend(image_ids)\n\n    submission = pd.DataFrame({'id': all_ids, 'label': all_preds})\n    submission.to_csv(ModelConfig.submission_path, index=False)\n    print(f'Submission saved to {ModelConfig.submission_path}')\n    return submission","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:13:23.792179Z","iopub.execute_input":"2025-12-31T14:13:23.792980Z","iopub.status.idle":"2025-12-31T14:13:23.798219Z","shell.execute_reply.started":"2025-12-31T14:13:23.792951Z","shell.execute_reply":"2025-12-31T14:13:23.797593Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_loader = get_test_loader()\nsub = generate_submission(model, test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-31T14:13:24.483856Z","iopub.execute_input":"2025-12-31T14:13:24.484373Z","iopub.status.idle":"2025-12-31T14:15:03.518353Z","shell.execute_reply.started":"2025-12-31T14:13:24.484347Z","shell.execute_reply":"2025-12-31T14:15:03.517751Z"}},"outputs":[],"execution_count":null}]}