{"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":"# [Emerging Properties in Self-Supervised Vision Transformers](https://arxiv.org/pdf/2104.14294)","metadata":{}},{"cell_type":"code","source":"!pip uninstall torch -y\n!pip uninstall torchvision torchaudio -y\n!pip install torch==2.9.0 torch-xla==2.9.0\n!pip install torchvision==0.24.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T09:10:17.707703Z","iopub.execute_input":"2026-07-25T09:10:17.70873Z","iopub.status.idle":"2026-07-25T09:14:06.582377Z","shell.execute_reply.started":"2026-07-25T09:10:17.708668Z","shell.execute_reply":"2026-07-25T09:14:06.580818Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile train_dino.py\n\nimport os\nos.environ[\"TF_CPP_MIN_LOG_LEVEL\"] = \"3\"\nos.environ[\"PJRT_DEVICE\"] = \"TPU\"\nos.environ[\"XLA_USE_BF16\"] = \"1\"\n\nimport 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\n\nimport torch_xla\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.runtime as xr\nimport torch_xla.core.xla_model as xm\nfrom torch_xla.amp import autocast\nimport torch_xla.distributed.xla_multiprocessing as xmp\n\nconfig = {\n    # Dataset\n    \"data_dir\": \"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC\", \n    \"subset_size\": 128_000,                 # total images for training\n    \"num_workers\": 4,                     \n    # Model\n    \"arch\": \"vit_small\",\n    \"patch_size\": 16,\n    \"img_size\": 224,\n    \"out_dim\": 65536,                     # K prototypes\n    \"hidden_dim\": 384,                    \n    \"depth\": 12,\n    \"heads\": 6,\n    \"mlp_dim\": 1536,\n    \"bottleneck_dim\": 256,\n    \"head_hidden_dim\": 2048,\n    # DINO\n    \"momentum_teacher\": 0.996,\n    \"tau_teacher\": 0.04,\n    \"tau_student\": 0.1,\n    \"center_momentum\": 0.9,\n    # Training\n    \"batch_size_per_core\": 16,            # total batch = 128 (8 cores)\n    \"epochs\": 20,    \n    \"warmup_epochs\": 5,\n    \"base_lr\": 0.0005,                    \n    \"weight_decay\": 0.04,\n    \"clip_grad\": 3.0,\n    # Multi-crop\n    \"n_global_crops\": 2,\n    \"n_local_crops\": 4,\n    \"global_crop_scale\": (0.4, 1.0),\n    \"local_crop_scale\": (0.05, 0.4),\n}\n\nimport glob\nimport random\n\ndef create_subset(data_dir, subset_size):\n    class_dirs = os.listdir(data_dir)\n    images = []\n    for cls in class_dirs:\n        cls_path = os.path.join(data_dir, cls)\n        image_paths = sorted(glob.glob(os.path.join(cls_path, '*.JPEG')))\n        sampled = random.sample(image_paths, subset_size//len(class_dirs))\n        images.extend(sampled)\n\n    return images\n\nsubset_size = config['subset_size']\nsubset_images = create_subset(os.path.join(config['data_dir'],'train'), subset_size)\nprint(f\"Total images in subset: {len(subset_images)}\")\n\nimport torchvision.transforms as T\n\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD  = (0.229, 0.224, 0.225)\n\ndef base_transform(crop_size, scale, flip_prob=0.5):\n    return T.Compose([\n        T.RandomResizedCrop(crop_size, scale=scale, interpolation=T.InterpolationMode.BICUBIC),\n        T.RandomHorizontalFlip(p=flip_prob),\n    ])\n\ndef color_transform(s=1.0, blur_radius=1.0):\n    transforms = [\n        T.RandomApply([T.ColorJitter(0.8*s, 0.8*s, 0.4*s, 0.2*s)], p=0.8),\n        T.RandomGrayscale(p=0.2),\n    ]\n    if blur_radius>0:\n        transforms.append(T.RandomApply([T.GaussianBlur(kernel_size=23, sigma=(0.1, 2.0))], p=0.5))\n    transforms.append(T.toTensor())\n    transforms.append(T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD))\n    return T.Compose(transforms)\n\nimport torchvision\nfrom torchvision import datasets, transforms\n\nclass MultiCropTransform:\n    def __init__(self, global_crop_scale, local_crop_scale, n_global, n_local, crop_size=224):\n        self.n_global = n_global\n        self.n_local = n_local\n        self.crop_size = crop_size\n\n        self.global_transform1 = T.Compose([\n            T.RandomResizedCrop(crop_size, scale=global_crop_scale, interpolation=T.InterpolationMode.BICUBIC),\n            T.RandomHorizontalFlip(p=0.5),\n            T.RandomApply([T.ColorJitter(0.8, 0.8, 0.4, 0.2)], p=0.8),\n            T.RandomGrayscale(p=0.2),\n            T.RandomApply([T.GaussianBlur(kernel_size=23, sigma=(0.1, 2.0))], p=0.5),\n            T.RandomSolarize(threshold=128, p=0.2),   # solarization for one view\n            T.ToTensor(),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ])\n\n        self.global_transform2 = T.Compose([\n            T.RandomResizedCrop(crop_size, scale=global_crop_scale, interpolation=T.InterpolationMode.BICUBIC),\n            T.RandomHorizontalFlip(p=0.5),\n            T.RandomApply([T.ColorJitter(0.8, 0.8, 0.4, 0.2)], p=0.8),\n            T.RandomGrayscale(p=0.2),\n            T.RandomApply([T.GaussianBlur(kernel_size=23, sigma=(0.1, 2.0))], p=0.5),\n            # no solarization here\n            T.ToTensor(),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ])\n\n        self.local_transform = T.Compose([\n            T.RandomResizedCrop(crop_size, scale=local_crop_scale, interpolation=T.InterpolationMode.BICUBIC),\n            T.RandomHorizontalFlip(p=0.5),\n            T.RandomApply([T.ColorJitter(0.8, 0.8, 0.4, 0.2)], p=0.8),\n            T.RandomGrayscale(p=0.2),\n            # No GaussianBlur on locals per DINO paper\n            T.ToTensor(),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ])\n\n    def __call__(self, img):\n        views = []\n        views.append(self.global_transform1(img))\n        views.append(self.global_transform2(img))\n        \n        for _ in range(self.n_local):\n            views.append(self.local_transform(img))\n        return views\n\nfrom PIL import Image\n\nclass DINOImageNet(Dataset):\n    def __init__(self, img_path, multicrop_transform):\n        self.img_path = img_path\n        self.multicrop_transform = multicrop_transform\n\n    def __len__(self):\n        return len(self.img_path)\n\n    def __getitem__(self, idx):\n        path = self.img_path[idx]\n        \n        with Image.open(path) as img:\n            img = img.convert('RGB')\n        views = self.multicrop_transform(img)\n        return views\n\ndef multicrop_collate(batch):\n    n_views = len(batch[0])\n    collated = []\n    for idx in range(n_views):\n        tensors=[item[idx] for item in batch]\n        collated.append(torch.stack(tensors))\n    return tuple(collated)\n\ndef create_dino_dataloader(image_paths, global_scale, local_scale, n_global, n_local, batch_size_per_core, num_workers=4, crop_size=224):\n    transform = MultiCropTransform(global_scale, local_scale, n_global, n_local, crop_size)\n    dataset = DINOImageNet(image_paths, transform)\n\n    sampler = torch.utils.data.distributed.DistributedSampler(dataset, num_replicas=xr.world_size(), rank=xr.global_ordinal(), shuffle=True)\n\n    dataloader = DataLoader(dataset, batch_size=batch_size_per_core, shuffle=False, num_workers=num_workers, \n                            collate_fn=multicrop_collate, drop_last=True, sampler=sampler)\n    return dataloader\n\nclass PatchEmbed(nn.Module):\n    def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=384):\n        super().__init__()\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.num_patches = (img_size // patch_size) ** 2  # 196\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)\n\n    def forward(self, x):\n        x = self.proj(x)\n        x = x.flatten(2)\n        x = x.transpose(1, 2)\n        return x\n\nclass Attention(nn.Module):\n    def __init__(self, dim, num_heads=6, qkv_bias=False, attn_drop=0., proj_drop=0.):\n        super().__init__()\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = head_dim ** -0.5\n\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n\n    def forward(self, x, return_attention=False):\n        B, N, C = x.shape\n        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads)\n        qkv = qkv.permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(dim=-1)\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n\n        if return_attention:\n            cls_attn = attn[:, :, 0, 1:].mean(dim=1) # attention of CLS (index 0) to all patches (1..N-1) and attn shape (B, num_heads, N, N)\n            return x, cls_attn\n        return x\n\nclass Mlp(nn.Module):\n    def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n\nclass Block(nn.Module):\n    def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, drop=0., attn_drop=0.):\n        super().__init__()\n        self.norm1 = nn.LayerNorm(dim)\n        self.attn = Attention(dim, num_heads=num_heads, qkv_bias=qkv_bias, \n                              attn_drop=attn_drop, proj_drop=drop)\n        self.norm2 = nn.LayerNorm(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, drop=drop)\n\n    def forward(self, x, return_attention=False):\n        if return_attention:\n            attn_out, attn_map = self.attn(self.norm1(x), return_attention=True)\n            x = x + attn_out\n            x = x + self.mlp(self.norm2(x))\n            return x, attn_map\n        else:\n            x = x + self.attn(self.norm1(x))\n            x = x + self.mlp(self.norm2(x))\n            return x\n\nclass VisionTransformer(nn.Module):\n    def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=0,\n                 embed_dim=384, depth=12, num_heads=6, mlp_ratio=4., qkv_bias=False,\n                 drop_rate=0., attn_drop_rate=0.):\n        super().__init__()\n        self.embed_dim = embed_dim\n        self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim)\n        num_patches = self.patch_embed.num_patches  # 196\n\n        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))\n        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        self.blocks = nn.ModuleList([\n            Block(dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, \n                  qkv_bias=qkv_bias, drop=drop_rate, attn_drop=attn_drop_rate)\n            for _ in range(depth)\n        ])\n        self.norm = nn.LayerNorm(embed_dim)\n\n        self.head = nn.Identity()  # default\n\n        nn.init.trunc_normal_(self.pos_embed, std=0.02)\n        nn.init.trunc_normal_(self.cls_token, std=0.02)\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\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, return_attention=False):\n        B = x.shape[0]\n        x = self.patch_embed(x)\n\n        cls_tokens = self.cls_token.expand(B, -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n\n        x = x + self.pos_embed\n        x = self.pos_drop(x)\n\n        attn_map = None\n        for i, blk in enumerate(self.blocks):\n            if i == len(self.blocks) - 1 and return_attention:\n                x, attn_map = blk(x, return_attention=True)\n            else:\n                x = blk(x)\n\n        x = self.norm(x)\n        cls_token_out = x[:, 0]\n\n        if return_attention:\n            return cls_token_out, attn_map\n        return cls_token_out\n\nclass DINOHead(nn.Module):\n    def __init__(self, in_dim=384, out_dim=65536, hidden_dim=2048, bottleneck_dim=256):\n        super().__init__()\n        self.mlp = nn.Sequential(\n            nn.Linear(in_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, bottleneck_dim),\n        )\n        self.last_layer = nn.utils.weight_norm(nn.Linear(bottleneck_dim, out_dim, bias=False), dim=1)\n\n    def forward(self, x):\n        x = self.mlp(x)\n        x = F.normalize(x, p=2, dim=-1)\n        x = self.last_layer(x)\n        return x\n\n    def normalize_weights(self):\n        with torch.no_grad():\n            self.last_layer.weight_g.data.fill_(1.0)\n            self.last_layer.weight_v.data = F.normalize(self.last_layer.weight_v.data, dim=1)\n\nimport copy\n\nclass DINO(nn.Module):\n    def __init__(self, student_vit, teacher_vit, head, momentum_teacher=0.996, tau_teacher=0.04, tau_student=0.1,\n                 center_momentum=0.9, out_dim=65536):\n        super().__init__()\n        self.student_vit = student_vit\n        self.student_head = head\n        self.teacher_vit = teacher_vit\n\n        self.teacher_head = DINOHead(in_dim=384, out_dim=65536, hidden_dim=2048, bottleneck_dim=256)\n        self.teacher_head.load_state_dict(self.student_head.state_dict())\n        \n        for p in self.teacher_head.parameters():\n            p.requires_grad = False\n\n        self.teacher_vit.load_state_dict(self.student_vit.state_dict())\n        \n        for p in self.teacher_vit.parameters():\n            p.requires_grad = False\n\n        self.momentum_teacher = momentum_teacher\n        self.tau_teacher = tau_teacher\n        self.tau_student = tau_student\n        self.center_momentum = center_momentum\n        self.register_buffer('center', torch.zeros(1, out_dim))\n\n    @torch.no_grad()\n    def update_teacher(self):\n        for ps, pt in zip(self.student_vit.parameters(), self.teacher_vit.parameters()):\n            pt.mul_(self.momentum_teacher)\n            pt.add_(ps, alpha=1-self.momentum_teacher)\n    \n        for ps, pt in zip(self.student_head.parameters(), self.teacher_head.parameters()):\n            pt.mul_(self.momentum_teacher)\n            pt.add_(ps, alpha=1-self.momentum_teacher)\n\n    @torch.no_grad()\n    def update_center(self, teacher_logits):\n        batch_center = teacher_logits.mean(dim=0, keepdim=True)\n        batch_center = xm.all_reduce(xm.REDUCE_SUM, batch_center)\n        \n        batch_center /= xr.world_size()\n        self.center = (self.center * self.center_momentum + batch_center * (1-self.center_momentum))\n\n    def forward(self, views):\n        B = views[0].shape[0]\n        n_global = 2\n\n        student_input = torch.cat(views, dim=0)         \n        cls_emb_student = self.student_vit(student_input) \n        student_logits_all = self.student_head(cls_emb_student)\n        student_logits_list = list(torch.split(student_logits_all, B, dim=0))\n\n        with torch.no_grad():\n            teacher_input = torch.cat(views[:n_global], dim=0) \n            cls_emb_teacher = self.teacher_vit(teacher_input).float()\n            teacher_logits = self.teacher_head(cls_emb_teacher)  \n            teacher_logits_global = list(torch.split(teacher_logits, B, dim=0))\n\n            self.update_center(teacher_logits) \n            teacher_logits_centered = [tl - self.center for tl in teacher_logits_global]\n            teacher_probs = [F.softmax(tlc / self.tau_teacher, dim=-1) for tlc in teacher_logits_centered]\n\n        student_log_probs = [F.log_softmax(sl / self.tau_student, dim=-1) for sl in student_logits_list]\n\n        loss = 0.\n        n_loss_terms = 0\n        for ig, teacher_prob in enumerate(teacher_probs):\n            for iv, student_log_prob in enumerate(student_log_probs):\n                if ig == iv:          # skip self-pair\n                    continue\n                loss += - (teacher_prob * student_log_prob).sum(dim=-1).mean()\n                n_loss_terms += 1\n        loss /= n_loss_terms\n        return loss\n\n    def step_teacher(self):\n        self.update_teacher()\n        if hasattr(self.student_head, 'normalize_weights'):\n            self.student_head.normalize_weights()\n\ndef get_params_groups(model):\n    regularized, not_regularized = [], []\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if name.endswith(\".bias\") or \"layernorm\" in name.lower() or \"ln\" in name or \"weight_g\" in name:\n            not_regularized.append(param)\n        else:\n            regularized.append(param)\n    return [\n        {'params': regularized, 'weight_decay': config['weight_decay']},\n        {'params': not_regularized, 'weight_decay': 0.0}\n    ]\n\ndef create_scheduler(optimizer, config, steps_per_epoch):\n    warmup_steps = config['warmup_epochs'] * steps_per_epoch\n    total_steps = config['epochs'] * steps_per_epoch\n\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step + 1) / float(max(1, warmup_steps))\n        else:\n            progress = float(step - warmup_steps) / float(max(1, total_steps - warmup_steps))\n            return 0.5 * (1.0 + np.cos(np.pi * progress))\n\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)\n\n@torch.no_grad()\ndef extract_features(model, dataloader, device):\n    model.eval()\n    features_list = []\n    labels_list = []\n    for image, label in dataloader:\n        image = image.to(device, non_blocking=True)\n        cls_embed = model(image)\n        features_list.append(cls_embed.cpu())\n        labels_list.append(label)\n    features = torch.cat(features_list, dim=0)\n    labels = torch.cat(labels_list, dim=0)\n    return features, labels\n\n@torch.no_grad()\ndef knn_classifier(train_features, train_labels, val_features, val_labels, k=10, T=0.07, device='cpu'):\n    train_features = train_features.to(device)\n    train_labels = train_labels.to(device)\n    train_features = F.normalize(train_features, dim=1)\n    val_features = F.normalize(val_features, dim=1)\n\n    batch_size = 256\n    total_correct = 0\n    total = 0\n\n    for i in range(0, len(val_features), batch_size):\n        val_batch = val_features[i:i+batch_size].to(device)\n        B = val_batch.size(0)\n        sim = torch.mm(val_batch, train_features.t()) / T\n        topk = torch.topk(sim, k, dim=1)\n        topk_indices = topk.indices\n        topk_similarities = topk.values\n\n        topk_labels = train_labels[topk_indices]      # (B, k)\n        weights = torch.exp(topk_similarities)\n        num_classes = 1000\n        one_hot = torch.zeros(B, k, num_classes, device=weights.device)\n        one_hot.scatter_(2, topk_labels.unsqueeze(-1), 1)\n        scores = (one_hot * weights.unsqueeze(-1)).sum(dim=1)\n        pred = scores.argmax(dim=1)\n\n        gt_labels = val_labels[i:i+batch_size].to(device)\n        total_correct += (pred == gt_labels).sum().item()\n        total += B\n    acc = total_correct / total * 100\n    return acc\n\nimport pandas as pd\n\nclass ImageNetValDataset(Dataset):\n    def __init__(self, val_dir, labels_csv, synset_file, transform=None):\n        self.val_dir = val_dir\n        self.transform = transform\n\n        synset_to_idx = {}\n        with open(synset_file) as f:\n            for idx, line in enumerate(f):\n                synset = line.split()[0]\n                synset_to_idx[synset] = idx\n\n        self.df = pd.read_csv(labels_csv)\n\n        self.df[\"synset\"] = self.df[\"PredictionString\"].str.split().str[0]\n\n        self.df[\"label\"] = self.df[\"synset\"].map(synset_to_idx)\n        assert self.df[\"label\"].notna().all(), \"Some synsets not found in synset_file!\"\n\n        self.df[\"image\"] = self.df[\"ImageId\"].apply(\n            lambda x: x if x.endswith(\".JPEG\") else x + \".JPEG\"\n        )\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.val_dir, row[\"image\"])\n        img = Image.open(img_path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        return img, int(row[\"label\"])\n\ndef main(index):\n    device = torch_xla.device()\n    rank = xr.global_ordinal()\n    world_size = xr.world_size()\n\n    subset_images_set=set(subset_images)\n\n    config['effective_batch_size'] = config['batch_size_per_core'] * world_size\n    config['learning_rate'] = config['base_lr'] * config['effective_batch_size'] / 256\n\n    student_vit = VisionTransformer(img_size=224, patch_size=16, embed_dim=384, depth=12, num_heads=6)\n    teacher_vit = VisionTransformer(img_size=224, patch_size=16, embed_dim=384, depth=12, num_heads=6)\n    head = DINOHead(in_dim=384, out_dim=65536, hidden_dim=2048, bottleneck_dim=256)\n\n    dino = DINO(student_vit, teacher_vit, head, momentum_teacher=config['momentum_teacher'], tau_teacher=config['tau_teacher'],\n        tau_student=config['tau_student'], center_momentum=config['center_momentum'], out_dim=config['out_dim']).to(device)\n\n    params_groups = get_params_groups(dino.student_vit) + get_params_groups(dino.student_head)\n    optimizer = torch.optim.AdamW(params_groups, lr=config['learning_rate'], betas=(0.9, 0.999))\n\n    train_loader = create_dino_dataloader(subset_images, global_scale=config['global_crop_scale'], local_scale=config['local_crop_scale'],\n        n_global=config['n_global_crops'], n_local=config['n_local_crops'], batch_size_per_core=config['batch_size_per_core'],\n        num_workers=config['num_workers'], crop_size=config['img_size'])\n    mp_train_loader = pl.MpDeviceLoader(train_loader, device)\n\n    steps_per_epoch = len(mp_train_loader)\n    scheduler = create_scheduler(optimizer, config, steps_per_epoch)\n\n    for epoch in range(config['epochs']):\n        dino.train()\n        mp_train_loader._loader.sampler.set_epoch(epoch)\n        total_loss = 0.0\n        for step, views in enumerate(mp_train_loader):\n            optimizer.zero_grad()\n\n            with autocast(device, dtype=torch.bfloat16):\n                loss = dino(views)\n                \n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(dino.parameters(), config[\"clip_grad\"])\n\n            xm.optimizer_step(optimizer)\n            scheduler.step()\n            dino.step_teacher()\n            total_loss += loss.item()\n            if step % 10 == 0 and rank == 0:\n                print(f\"Epoch {epoch+1}, step {step}, loss {loss.item():.4f}, lr {scheduler.get_last_lr()[0]:.6f}\")\n        avg_loss = total_loss / (step+1)\n        if rank == 0:\n            print(f\"Epoch {epoch+1} average loss: {avg_loss:.4f}\")\n\n    if rank == 0:\n        print(\"Training complete. Starting evaluation...\")\n\n        transform_eval = T.Compose([\n            T.Resize(256),\n            T.CenterCrop(224),\n            T.ToTensor(),\n            T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)\n        ])\n\n        train_dir = os.path.join(config['data_dir'], 'train')\n        train_eval_ds = torchvision.datasets.ImageFolder(root=train_dir, transform=transform_eval)\n        subset_indices = [i for i, (path, _) in enumerate(train_eval_ds.imgs) if path in subset_images_set]\n        train_eval_ds = torch.utils.data.Subset(train_eval_ds, subset_indices)\n        train_loader_eval = DataLoader(train_eval_ds, batch_size=256, shuffle=False, num_workers=4)\n\n        val_img_dir = os.path.join(config['data_dir'], 'val')\n        val_csv_path = '/kaggle/input/competitions/imagenet-object-localization-challenge/LOC_val_solution.csv'\n        synset_file = '/kaggle/input/competitions/imagenet-object-localization-challenge/LOC_synset_mapping.txt'\n        val_ds = ImageNetValDataset(val_img_dir, val_csv_path, synset_file, transform=transform_eval)\n        val_loader_eval = DataLoader(val_ds, batch_size=256, shuffle=False, num_workers=4)\n        \n        dino.student_vit.eval()\n        train_feats, train_labels = extract_features(dino.student_vit, train_loader_eval, device)\n        val_feats, val_labels = extract_features(dino.student_vit, val_loader_eval, device)\n\n        acc_k10 = knn_classifier(train_feats, train_labels, val_feats, val_labels, k=10, T=0.07, device='cpu')\n        acc_k20 = knn_classifier(train_feats, train_labels, val_feats, val_labels, k=20, T=0.07, device='cpu')\n        print(f\"k-NN top-1 (k=10): {acc_k10:.2f}%\")\n        print(f\"k-NN top-1 (k=20): {acc_k20:.2f}%\")\n\nif __name__ == \"__main__\":\n    xmp.spawn(main, args=())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T09:14:06.586647Z","iopub.execute_input":"2026-07-25T09:14:06.587145Z","iopub.status.idle":"2026-07-25T09:14:06.609399Z","shell.execute_reply.started":"2026-07-25T09:14:06.587088Z","shell.execute_reply":"2026-07-25T09:14:06.608462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python train_dino.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-25T09:14:06.61068Z","iopub.execute_input":"2026-07-25T09:14:06.611134Z","iopub.status.idle":"2026-07-25T09:15:49.744322Z","shell.execute_reply.started":"2026-07-25T09:14:06.611094Z","shell.execute_reply":"2026-07-25T09:15:49.743074Z"}},"outputs":[],"execution_count":null}]}