{"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":"# [Training data-efficient image transformers & distillation through attention](https://arxiv.org/pdf/2012.12877)","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\nimport numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:42:19.623324Z","iopub.execute_input":"2026-07-20T10:42:19.623983Z","iopub.status.idle":"2026-07-20T10:42:24.817071Z","shell.execute_reply.started":"2026-07-20T10:42:19.62395Z","shell.execute_reply":"2026-07-20T10:42:24.816422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from dataclasses import dataclass\n\n@dataclass\nclass ExperimentConfig:\n    img_size: int = 224\n    patch_size: int = 16\n    in_chans: int = 3\n    num_classes: int = 1000\n    embed_dim: int = 768\n    num_layers: int = 12\n    num_heads: int = 12\n    \n    # Distillation\n    teacher_model: str = \"regnety_160\"\n    distillation_alpha: float = 1.0   # weight for distillation loss\n    \n    # Training\n    batch_size_per_gpu: int = 128\n    gradient_accumulation_steps: int = 4  # effective batch = 128 * 2 GPUs * 4 = 1024\n    epochs: int = 20\n    warmup_epochs: int = 5\n    base_lr: float = 0.001   # for reference batch size 512\n    min_lr: float = 1e-5\n    weight_decay: float = 0.05\n    \n    # Augmentation\n    mixup_alpha: float = 0.8\n    cutmix_alpha: float = 1.0\n    mixup_prob: float = 1.0\n    switch_prob: float = 0.5\n    ra_magnitude: int = 9       # RandAugment magnitude\n    ra_num_ops: int = 2\n    \n    # Model parameters\n    drop_path_rate: float = 0.1  # for deit_small/base; 0.0 for tiny\n    drop_rate: float = 0.0\n    \n    # Optimization\n    optimizer: str = \"adamw\"\n    beta1: float = 0.9\n    beta2: float = 0.999\n    eps: float = 1e-8\n    \n    # System\n    num_workers: int = 4\n    log_interval: int = 200\n\ncfg = ExperimentConfig()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:42:24.818423Z","iopub.execute_input":"2026-07-20T10:42:24.81885Z","iopub.status.idle":"2026-07-20T10:42:24.826652Z","shell.execute_reply.started":"2026-07-20T10:42:24.818824Z","shell.execute_reply":"2026-07-20T10:42:24.825904Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndata_root = '/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC'\ntrain_dir = os.path.join(data_root, 'train')\nval_dir = os.path.join(data_root, 'val')\n\nval_sorted_dir = \"/kaggle/working/imagenet_val_sorted\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:42:24.827567Z","iopub.execute_input":"2026-07-20T10:42:24.827849Z","iopub.status.idle":"2026-07-20T10:42:24.839602Z","shell.execute_reply.started":"2026-07-20T10:42:24.827818Z","shell.execute_reply":"2026-07-20T10:42:24.839021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\nimport pandas as pd\nfrom tqdm import tqdm\n\nshutil.rmtree(\"/kaggle/working/imagenet_val_sorted\", ignore_errors=True)\n\nif not os.path.exists(val_sorted_dir):\n    os.makedirs(val_sorted_dir, exist_ok=True)\n\n    annotation_file = '/kaggle/input/competitions/imagenet-object-localization-challenge/LOC_val_solution.csv'\n\n    df = pd.read_csv(annotation_file)\n    df[\"class_id\"] = df[\"PredictionString\"].apply(lambda x: x.split()[0] if isinstance(x, str) else \"unknown\")\n    \n    class_ids = df[\"class_id\"].unique()\n\n    for cid in class_ids:\n        os.makedirs(os.path.join(val_sorted_dir, cid), exist_ok=True)\n\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        img_name = row[\"ImageId\"] + \".JPEG\"\n        class_id = row[\"class_id\"]\n        src = os.path.join(val_dir, img_name)\n        dst = os.path.join(val_sorted_dir, class_id, img_name)\n        if os.path.exists(src):\n            shutil.copy(src, dst)  # copy to avoid modifying original\n    print(\"Validation data sorted.\")\nelse:\n    print(\"Validation directory already sorted.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:42:24.840471Z","iopub.execute_input":"2026-07-20T10:42:24.840804Z","iopub.status.idle":"2026-07-20T10:52:21.70371Z","shell.execute_reply.started":"2026-07-20T10:42:24.840782Z","shell.execute_reply":"2026-07-20T10:52:21.702979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\n\nmean = [0.485, 0.456, 0.406]\nstd = [0.229, 0.224, 0.225]\n\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(cfg.img_size, scale=(0.08, 1.0), ratio=(3./4., 4./3.)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandAugment(num_ops=cfg.ra_num_ops, magnitude=cfg.ra_magnitude),\n    transforms.ToTensor(),\n    transforms.Normalize(mean, std)\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize(mean, std)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:21.705734Z","iopub.execute_input":"2026-07-20T10:52:21.706066Z","iopub.status.idle":"2026-07-20T10:52:25.67233Z","shell.execute_reply.started":"2026-07-20T10:52:21.706041Z","shell.execute_reply":"2026-07-20T10:52:25.671557Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nfrom PIL import Image\n\nclass SubsetImageNet(Dataset):\n    def __init__(self, root, target_size=100_000, transform=None, shuffle_classes=True):\n        self.transform = transform\n        self.samples = []\n\n        all_class_dirs = sorted(\n            [d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]\n        )\n        # Official mapping: index = alphabetical position\n        self.class_to_idx = {cls_name: i for i, cls_name in enumerate(all_class_dirs)}\n        self.classes = all_class_dirs  # list in index order\n\n        visit_order = all_class_dirs.copy()\n        if shuffle_classes:\n            random.shuffle(visit_order)\n\n        for cls_name in visit_order:\n            class_path = os.path.join(root, cls_name)\n            idx = self.class_to_idx[cls_name]\n            all_images = [f for f in os.listdir(class_path) \n                         if f.lower().endswith(('.jpeg', '.jpg', '.png'))]\n            for img in all_images:\n                self.samples.append((os.path.join(class_path, img), idx))\n                if len(self.samples) >= target_size:\n                    break\n            if len(self.samples) >= target_size:\n                break\n\n        self.samples = self.samples[:target_size]\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        path, label = self.samples[idx]\n        image = Image.open(path).convert('RGB')\n        if self.transform:\n            image = self.transform(image)\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:25.673599Z","iopub.execute_input":"2026-07-20T10:52:25.674041Z","iopub.status.idle":"2026-07-20T10:52:25.681549Z","shell.execute_reply.started":"2026-07-20T10:52:25.674Z","shell.execute_reply":"2026-07-20T10:52:25.680829Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.datasets import ImageFolder\n\ntrain_dataset = SubsetImageNet(train_dir, target_size=128_000, transform=train_transform)\nval_dataset = ImageFolder(val_sorted_dir, transform=val_transform)\n\nprint(f\"Training samples: {len(train_dataset)}\")\nprint(f\"Validation samples: {len(val_dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:25.68231Z","iopub.execute_input":"2026-07-20T10:52:25.682651Z","iopub.status.idle":"2026-07-20T10:52:40.092853Z","shell.execute_reply.started":"2026-07-20T10:52:25.682618Z","shell.execute_reply":"2026-07-20T10:52:40.092054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from timm.data import Mixup\n\nmixup_fn = Mixup(\n    mixup_alpha=cfg.mixup_alpha,\n    cutmix_alpha=cfg.cutmix_alpha,\n    prob=cfg.mixup_prob,\n    switch_prob=cfg.switch_prob,\n    mode='batch',\n    num_classes=cfg.num_classes\n)\n\nval_loader = DataLoader(val_dataset, batch_size=cfg.batch_size_per_gpu, shuffle=False, num_workers=cfg.num_workers, pin_memory=True,\n                        prefetch_factor=2\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:40.093872Z","iopub.execute_input":"2026-07-20T10:52:40.094179Z","iopub.status.idle":"2026-07-20T10:52:44.525011Z","shell.execute_reply.started":"2026-07-20T10:52:40.094145Z","shell.execute_reply":"2026-07-20T10:52:44.524421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PatchEmbed(nn.Module):\n    def __init__(self, embed_dim=768, patch_size=16, in_chans=3):\n        super().__init__()\n        self.patch_size = patch_size\n        self.conv = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        N, C, H, W = x.shape\n        x = self.conv(x)\n        x = x.flatten(2)\n        x = x.transpose(-1, -2)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:44.525848Z","iopub.execute_input":"2026-07-20T10:52:44.526681Z","iopub.status.idle":"2026-07-20T10:52:44.531376Z","shell.execute_reply.started":"2026-07-20T10:52:44.526655Z","shell.execute_reply":"2026-07-20T10:52:44.530616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DeiTEmbedding(nn.Module):\n    def __init__(self, num_patches=196, embed_dim=768):\n        super().__init__()\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.distill_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches+2, embed_dim))\n        self.patch_embed = PatchEmbed(embed_dim=embed_dim)\n        self._init_weights()\n\n    def _init_weights(self):\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        nn.init.trunc_normal_(self.distill_token, std=0.02)\n\n    def forward(self, x):\n        B = x.shape[0]\n        x = self.patch_embed(x)\n        cls_token = self.cls_token.expand(B, -1, -1)\n        distill_token = self.distill_token.expand(B, -1, -1)\n        x = torch.cat((cls_token, distill_token, x), dim=1)\n        x = x + self.pos_embed\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:44.532311Z","iopub.execute_input":"2026-07-20T10:52:44.532654Z","iopub.status.idle":"2026-07-20T10:52:44.548704Z","shell.execute_reply.started":"2026-07-20T10:52:44.532621Z","shell.execute_reply":"2026-07-20T10:52:44.548054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DropPath(nn.Module):\n    def __init__(self, drop_path=0.):\n        super().__init__()\n        self.drop_path = drop_path\n\n    def forward(self, x):\n        if self.drop_path==0 or not self.training:\n            return x\n\n        keep_prob = 1 - self.drop_path\n        shape = (x.shape[0],) + (1,) * (x.ndim-1)\n        random_mask = keep_prob + torch.rand(shape, device=x.device, dtype=x.dtype)\n        random_mask = random_mask.floor_()\n        return x / keep_prob * random_mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:44.54942Z","iopub.execute_input":"2026-07-20T10:52:44.549687Z","iopub.status.idle":"2026-07-20T10:52:44.563151Z","shell.execute_reply.started":"2026-07-20T10:52:44.549666Z","shell.execute_reply":"2026-07-20T10:52:44.562526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EncoderBlock(nn.Module):\n    def __init__(self, num_heads=12, embed_dim=768, drop_path=0., proj_drop=0., attn_drop=0.):\n        super().__init__()\n        self.qkv = nn.Linear(embed_dim, 3*embed_dim)\n        \n        assert embed_dim % num_heads == 0\n        self.d_k = embed_dim // num_heads\n        \n        self.num_heads = num_heads\n        self.proj = nn.Linear(embed_dim, embed_dim)\n        self.ln1 = nn.LayerNorm(embed_dim)\n        \n        self.mlp = nn.Sequential(\n            nn.Linear(embed_dim, 4*embed_dim),\n            nn.GELU(),\n            nn.Dropout(proj_drop),\n            nn.Linear(4*embed_dim, embed_dim),\n            nn.Dropout(proj_drop)\n        )\n        \n        self.ln2 = nn.LayerNorm(embed_dim)\n        self.drop_path = DropPath(drop_path) if drop_path>0 else nn.Identity()\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def self_attention(self, x):\n        B, N, D = x.shape\n        qkv = self.qkv(x)\n        q, k, v = torch.chunk(qkv, 3, dim=-1) \n        \n        q = q.reshape(B, N, self.num_heads, self.d_k).permute(0, 2, 1, 3)\n        k = k.reshape(B, N, self.num_heads, self.d_k).permute(0, 2, 1, 3)\n        v = v.reshape(B, N, self.num_heads, self.d_k).permute(0, 2, 1, 3)\n        \n        attn_scores = torch.matmul(q, k.transpose(-1, -2))/(self.d_k**0.5)\n        attn_weights = torch.softmax(attn_scores, dim=-1)\n        attn_weights = self.attn_drop(attn_weights)\n        attn_outputs = torch.matmul(attn_weights, v)\n        attn_outputs = attn_outputs.permute(0, 2, 1, 3).reshape(B, N, -1)\n        \n        output = self.proj(attn_outputs)\n        output = self.proj_drop(output)\n        return output\n\n    def forward(self, x):\n        x = x + self.drop_path(self.self_attention(self.ln1(x)))\n        x = x + self.drop_path(self.mlp(self.ln2(x)))\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:44.564127Z","iopub.execute_input":"2026-07-20T10:52:44.564406Z","iopub.status.idle":"2026-07-20T10:52:44.580204Z","shell.execute_reply.started":"2026-07-20T10:52:44.564375Z","shell.execute_reply":"2026-07-20T10:52:44.579442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DeiT(nn.Module):\n    def __init__(self, num_classes=1000, num_heads=12, embed_dim=768, num_layers=12, drop_path_rate=0.1):\n        super().__init__()\n        self.embed = DeiTEmbedding(embed_dim=embed_dim)\n        self.head = nn.Linear(embed_dim, num_classes)\n        self.distill_head = nn.Linear(embed_dim, num_classes)\n\n        dpr = torch.linspace(0, drop_path_rate, num_layers)\n        \n        self.blocks = nn.ModuleList([EncoderBlock(num_heads, embed_dim, drop_path=dpr[i]) for i in range(num_layers)])\n        self.norm = nn.LayerNorm(embed_dim)\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.trunc_normal_(m.weight, std=0.02)\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.LayerNorm):\n                nn.init.constant_(m.bias, 0)\n                nn.init.constant_(m.weight, 1.0)\n\n    def forward(self, x):\n        x = self.embed(x)\n        for block in self.blocks:\n            x = block(x)\n        x = self.norm(x)\n        cls_out = self.head(x[:, 0])\n        distill_out = self.distill_head(x[:, 1])\n        if self.training:\n            return cls_out, distill_out # logits\n        else:\n            return (cls_out + distill_out) / 2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:44.580983Z","iopub.execute_input":"2026-07-20T10:52:44.581313Z","iopub.status.idle":"2026-07-20T10:52:44.595051Z","shell.execute_reply.started":"2026-07-20T10:52:44.581292Z","shell.execute_reply":"2026-07-20T10:52:44.594534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\n\ndef load_teacher(model_name, num_classes=1000):\n    model = timm.create_model(model_name, pretrained=True, num_classes=num_classes)\n    model.eval()\n    for param in model.parameters():\n        param.requires_grad = False\n    return model\n\nteacher = load_teacher(cfg.teacher_model, cfg.num_classes)\nprint(f\"Teacher model {cfg.teacher_model} loaded and frozen.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:44.597401Z","iopub.execute_input":"2026-07-20T10:52:44.597672Z","iopub.status.idle":"2026-07-20T10:52:51.895976Z","shell.execute_reply.started":"2026-07-20T10:52:44.597652Z","shell.execute_reply":"2026-07-20T10:52:51.895327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_deit(config):\n    model = DeiT(\n        num_classes=config.num_classes,\n        embed_dim=config.embed_dim,\n        num_layers=config.num_layers,\n        num_heads=config.num_heads,\n        drop_path_rate=config.drop_path_rate\n    )\n    return model\n\nmodel = create_deit(cfg)\nprint(f\"Created DeiT with {sum(p.numel() for p in model.parameters()):,} parameters\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:51.897021Z","iopub.execute_input":"2026-07-20T10:52:51.897256Z","iopub.status.idle":"2026-07-20T10:52:53.158852Z","shell.execute_reply.started":"2026-07-20T10:52:51.897233Z","shell.execute_reply":"2026-07-20T10:52:53.157766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DistillationLoss(nn.Module):\n    def __init__(self, alpha):\n        super().__init__()\n        self.alpha = alpha\n\n    def forward(self, cls_logits, dist_logits, targets, teacher_labels):\n        if targets.dim()==1: # hard labels\n            loss_ce = F.cross_entropy(cls_logits, targets)\n        else:\n            loss_ce = -torch.sum(targets * F.log_softmax(cls_logits, dim=-1), dim=-1).mean()\n            \n        loss_dist = F.cross_entropy(dist_logits, teacher_labels)\n        total_loss = loss_ce + self.alpha * loss_dist\n        return total_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:53.159943Z","iopub.execute_input":"2026-07-20T10:52:53.160243Z","iopub.status.idle":"2026-07-20T10:52:53.55012Z","shell.execute_reply.started":"2026-07-20T10:52:53.160216Z","shell.execute_reply":"2026-07-20T10:52:53.549183Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_param_groups(model, weight_decay=0.05):\n    no_decay = ['bias', 'LayerNorm.weight', 'LayerNorm.bias']\n    decay_params = []\n    no_decay_params = []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if any(nd in name for nd in no_decay):\n            no_decay_params.append(param)\n        else:\n            decay_params.append(param)\n    return [\n        {'params': decay_params, 'weight_decay': weight_decay},\n        {'params': no_decay_params, 'weight_decay': 0.0}\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:53.551194Z","iopub.execute_input":"2026-07-20T10:52:53.551471Z","iopub.status.idle":"2026-07-20T10:52:53.564621Z","shell.execute_reply.started":"2026-07-20T10:52:53.551439Z","shell.execute_reply":"2026-07-20T10:52:53.564026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WarmupCosineScheduler:\n    def __init__(self, optimizer, warmup_steps, total_steps, min_lr, base_lr):\n        self.optimizer = optimizer\n        self.warmup_steps = warmup_steps\n        self.total_steps = total_steps\n        self.min_lr = min_lr\n        self.base_lr = base_lr\n        self.current_step = 0\n\n    def step(self):\n        self.current_step += 1\n        if self.current_step <= self.warmup_steps:\n            lr = self.base_lr * self.current_step / self.warmup_steps # linear warmup\n        else: # cosine scheduler\n            progress = (self.current_step - self.warmup_steps) / (self.total_steps - self.warmup_steps)\n            lr = self.min_lr + (self.base_lr - self.min_lr) * 0.5 * (1+np.cos(np.pi * progress))\n\n        for param_group in self.optimizer.param_groups:\n            param_group['lr'] = lr\n\n    def get_lr(self):\n        return self.optimizer.param_groups[0]['lr']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:53.565447Z","iopub.execute_input":"2026-07-20T10:52:53.565801Z","iopub.status.idle":"2026-07-20T10:52:53.578192Z","shell.execute_reply.started":"2026-07-20T10:52:53.565779Z","shell.execute_reply":"2026-07-20T10:52:53.57739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = DistillationLoss(alpha=cfg.distillation_alpha)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\nteacher = teacher.to(device)\n\nif torch.cuda.device_count()>1:\n    model = nn.DataParallel(model)\n\nbase_lr = 0.001\neffective_batch = cfg.batch_size_per_gpu * torch.cuda.device_count() * cfg.gradient_accumulation_steps\nlr = base_lr * (effective_batch / 512)\n\noptimizer = torch.optim.AdamW(get_param_groups(model, cfg.weight_decay), lr=lr, betas=(cfg.beta1, cfg.beta2), eps=cfg.eps)\nscaler = torch.amp.GradScaler(enabled='cuda')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:53.57922Z","iopub.execute_input":"2026-07-20T10:52:53.579566Z","iopub.status.idle":"2026-07-20T10:52:54.377052Z","shell.execute_reply.started":"2026-07-20T10:52:53.579518Z","shell.execute_reply":"2026-07-20T10:52:54.37647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef precompute_teacher_labels(dataset, teacher, device, batch_size=512, num_workers=4):\n    if torch.cuda.device_count() > 1:\n        teacher = nn.DataParallel(teacher)\n    teacher.eval()\n    \n    loader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True)\n    \n    all_teacher_labels = []\n    for images, _ in tqdm(loader, desc=\"Precomputing teacher labels\"):\n        images = images.to(device, non_blocking=True)\n        teacher_preds = teacher(images).argmax(dim=1).cpu()\n        all_teacher_labels.append(teacher_preds)\n    return torch.cat(all_teacher_labels)\n\nif isinstance(teacher, nn.DataParallel):\n    teacher = teacher.module\n\nteacher_labels_tensor = precompute_teacher_labels(train_dataset, teacher, device, batch_size=512)\n\ntorch.save(teacher_labels_tensor, \"/kaggle/working/teacher_labels_100k.pt\")\nprint(\"Teacher labels saved. Shape:\", teacher_labels_tensor.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T10:52:54.377896Z","iopub.execute_input":"2026-07-20T10:52:54.378232Z","iopub.status.idle":"2026-07-20T12:48:14.221102Z","shell.execute_reply.started":"2026-07-20T10:52:54.378208Z","shell.execute_reply":"2026-07-20T12:48:14.220214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TeacherLabelWrapper(Dataset):\n    def __init__(self, base_dataset, teacher_labels):\n        self.base = base_dataset\n        self.teacher_labels = teacher_labels\n        assert len(self.base) == len(self.teacher_labels), \\\n            f\"Length mismatch: {len(self.base)} vs {len(self.teacher_labels)}\"\n\n    def __len__(self):\n        return len(self.base)\n\n    def __getitem__(self, idx):\n        image, target = self.base[idx]\n        teacher_label = self.teacher_labels[idx]\n        return image, target, teacher_label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T12:48:14.222256Z","iopub.execute_input":"2026-07-20T12:48:14.223436Z","iopub.status.idle":"2026-07-20T12:48:14.22976Z","shell.execute_reply.started":"2026-07-20T12:48:14.223405Z","shell.execute_reply":"2026-07-20T12:48:14.22895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"teacher_labels_tensor = torch.load(\"/kaggle/working/teacher_labels_100k.pt\")\n\ntrain_dataset_with_teacher = TeacherLabelWrapper(train_dataset, teacher_labels_tensor)\n\ntrain_loader = DataLoader(train_dataset_with_teacher, batch_size=cfg.batch_size_per_gpu, shuffle=True, num_workers=cfg.num_workers, pin_memory=True,\n                          drop_last=True, prefetch_factor=2,\n)\n\nnum_batches = len(train_loader)\nsteps_per_epoch = num_batches // cfg.gradient_accumulation_steps\ntotal_steps = cfg.epochs * steps_per_epoch\nwarmup_steps = cfg.warmup_epochs * steps_per_epoch\nscheduler = WarmupCosineScheduler(optimizer, warmup_steps, total_steps, cfg.min_lr, cfg.base_lr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T12:48:14.230711Z","iopub.execute_input":"2026-07-20T12:48:14.230965Z","iopub.status.idle":"2026-07-20T12:48:14.259407Z","shell.execute_reply.started":"2026-07-20T12:48:14.230935Z","shell.execute_reply":"2026-07-20T12:48:14.258552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(epoch, model, loader, criterion, optimizer, scheduler, scaler, cfg, device):\n    model.train()\n    losses = AverageMeter()\n    optimizer.zero_grad()\n    step_count = 0\n\n    for i, (images, targets, teacher_labels) in enumerate(loader):\n        images = images.to(device, non_blocking=True)\n        targets = targets.to(device, non_blocking=True)\n        teacher_labels = teacher_labels.to(device, non_blocking=True)\n\n        images, targets = mixup_fn(images, targets)  # targets become soft one-hot\n\n        with torch.amp.autocast(device_type='cuda'):\n            cls_logits, dist_logits = model(images)\n            loss = criterion(cls_logits, dist_logits, targets, teacher_labels)\n            loss = loss / cfg.gradient_accumulation_steps\n\n        scaler.scale(loss).backward()\n\n        if (i + 1) % cfg.gradient_accumulation_steps == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n            step_count += 1\n\n        losses.update(loss.item() * cfg.gradient_accumulation_steps, images.size(0))\n\n    return losses.avg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T12:48:14.26022Z","iopub.execute_input":"2026-07-20T12:48:14.260436Z","iopub.status.idle":"2026-07-20T12:48:14.268394Z","shell.execute_reply.started":"2026-07-20T12:48:14.260416Z","shell.execute_reply":"2026-07-20T12:48:14.26755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AverageMeter:\n    def __init__(self): \n        self.reset()\n        \n    def reset(self): \n        self.val, self.avg, self.sum, self.count = 0, 0, 0, 0\n        \n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\ndef accuracy(output, target, topk=(1,)):\n    maxk = max(topk)\n    batch_size = target.size(0)\n    _, pred = output.topk(maxk, 1, True, True)\n    pred = pred.t()\n    correct = pred.eq(target.view(1, -1).expand_as(pred))\n    res = []\n    for k in topk:\n        correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)\n        res.append(correct_k.mul_(100.0 / batch_size))\n    return res","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T12:48:14.269236Z","iopub.execute_input":"2026-07-20T12:48:14.269875Z","iopub.status.idle":"2026-07-20T12:48:14.282245Z","shell.execute_reply.started":"2026-07-20T12:48:14.269837Z","shell.execute_reply":"2026-07-20T12:48:14.281547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate(model, loader, cfg, device):\n    model.eval()\n    top1 = AverageMeter()\n    top5 = AverageMeter()\n    for images, targets in loader:\n        images = images.to(device)\n        targets = targets.to(device)\n\n        output = model(images)\n        acc1, acc5 = accuracy(output, targets, topk=(1,5))\n        top1.update(acc1.item(), images.size(0))\n        top5.update(acc5.item(), images.size(0))\n    model.train()\n    return top1.avg, top5.avg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T12:48:14.283139Z","iopub.execute_input":"2026-07-20T12:48:14.283548Z","iopub.status.idle":"2026-07-20T12:48:14.300016Z","shell.execute_reply.started":"2026-07-20T12:48:14.283498Z","shell.execute_reply":"2026-07-20T12:48:14.29907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.backends.cudnn.benchmark = True\n\nstart_epoch = 0\nbest_acc = 0.0\n\nfor epoch in range(start_epoch, cfg.epochs):\n    train_loss = train_one_epoch(epoch, model, train_loader, criterion, optimizer, scheduler, scaler, cfg, device)\n    val_acc1, val_acc5 = validate(model, val_loader, cfg, device)\n    print(f\"Epoch {epoch+1}: Train Loss {train_loss:.4f}, Val Acc1 {val_acc1:.2f}%, Acc5 {val_acc5:.2f}%\")\n\n    if val_acc1 > best_acc:\n        best_acc = val_acc1\n        print(f\"New best model with Acc1 {best_acc:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-20T13:18:12.689778Z","iopub.execute_input":"2026-07-20T13:18:12.692973Z","execution_failed":"2026-07-20T13:20:19.682Z"}},"outputs":[],"execution_count":null}]}