{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DLCV Project\n## [Barlow Twins: Self-Supervised Learning via Redundancy Reduction](https://proceedings.mlr.press/v139/zbontar21a.html)\nThis paper was published in the proceedings of the 38th International Conference on Machine Learning ([ICML 2021](https://icml.cc/virtual/2021/papers.html?search=Barlow+Twins:+Self-Supervised+Learning+via+Redundancy+Reduction))\n\nOriginal Implementation:\n [facebookresearch/barlowtwins](https://github.com/facebookresearch/barlowtwins)\n\n### Team details\nKeesari Vigneshwar Reddy - [22UCC052](22ucc052@lnmiit.ac.in)\n\nPalakurthy Guneeth - [22UCS144](22ucs144@lnmiit.ac.in)\n\nMeka Janaki Ram - [22UCS122](22ucs122@lnmiit.ac.in)\n\nAakash Chauhan - [22UCS001](22ucs001@lnmiit.ac.in)","metadata":{}},{"cell_type":"markdown","source":"## Introduction\n\nSelf-supervised learning (SSL) is a technique in which a model learns from unlabeled data, and is often used when the data is corrupted or if there is very little of it. A practical use for SSL is to create intermediate embeddings that are learned from the data. These embeddings are\nbased on the dataset itself, with similar images having similar embeddings, and\nvice versa. They are then attached to the rest of the model, which uses those\nembeddings as information and effectively learns and makes predictions properly.\nThese embeddings, ideally, should contain as much information and insight about\nthe data as possible, so that the model can make better predictions.\n\nA common example of self-supervised learning using Siamese networks is contrastive learning, where two augmented views of the same image are passed through identical networks to produce similar embeddings. The network is trained to minimize the distance between embeddings of similar inputs and maximize it for dissimilar ones, helping it learn meaningful representations without labels.\n\nHowever, a common problem that arises is that the model creates embeddings that are\nredundant. For example, if two images are similar, the model will create\nembeddings that are just a string of 1's, or some other value that\ncontains repeating bits of information. This is no better than a one-hot\nencoding or just having one bit as the model’s representations; it defeats the\npurpose of the embeddings, as they do not learn as much about the dataset as\npossible. For other approaches, the solution to the problem was to carefully\nconfigure the model such that it tries not to be redundant.\n\nBarlow Twins is a new approach to this problem; while other solutions mainly\ntackle the first goal of invariance (similar images have similar embeddings),\nthe Barlow Twins method also prioritizes the goal of reducing redundancy. The authors of the [paper](https://proceedings.mlr.press/v139/zbontar21a.html) propose an objective function that naturally avoids collapse by measuring the cross-correlation matrix between the outputs of two identical networks fed with distorted versions\nof a sample, and making it as close to the identity matrix as possible.\n\nIt also has the advantage of being much simpler than other methods, and its\nmodel architecture is symmetric, meaning that both twins in the model do the\nsame things. Intriguingly it benefits from very high-dimensional output vectors. BARLOW TWINS outperforms previous methods on ImageNet for semi-supervised classification in the low-data regime and near state-of-the-art on [imagenet](https://image-net.org/challenges/LSVRC/2012/).\n\n\nOne disadvantage of Barlow Twins is that it is heavily dependent on augmentation (specifically distortions), suffering major performance decreases in accuracy without them.\n\nThe paper used [ImageNet ILSVRC-2012 dataset](https://image-net.org/challenges/LSVRC/2012/) which consists of 1000 categories and 1.2 million images. According to the paper, training is distributed across 32 V100 GPUs and takes approximately 124 hours. The authors generated results using [ResNet-50](https://viso.ai/deep-learning/resnet-residual-neural-network/) on tasks like linear evaluation, semi-supervised training, image classification, object detection and image segmentation.\n\nDue to lack of compution power, dealing with such a huge dataset and multiple tasks is difficult. In this we train a Barlow Twins model and reach up to 58.4% validation accuracy on the CIFAR-10 dataset on linear evaluation task using [ResNet-18](https://viso.ai/deep-learning/resnet-residual-neural-network/).\n\n\n\n**About CIFAR-10**\n\nThe CIFAR-10 dataset consists of 60000 colour images (32x32 size) in 10 classes, with 6000 images per class. There are 50000 training images and 10000 test images.","metadata":{}},{"cell_type":"markdown","source":"![image](https://i.imgur.com/G6LnEPT.png)","metadata":{}},{"cell_type":"markdown","source":"### High-Level Theory\nThe model takes two versions of the same image (with different augmentations) as\ninput. Then it takes a prediction of each of them, creating representations.\nThey are then used to make a cross-correlation matrix.\n\nCross-correlation matrix:\n```\n(pred_1.T @ pred_2) / batch_size\n```\n\nThe cross-correlation matrix measures the correlation between the output\nneurons in the two representations made by the model predictions of the two\naugmented versions of data. Ideally, a cross-correlation matrix should look\nlike an identity matrix if the two images are the same.\n\nWhen this happens, it means that the representations:\n\n1.   Are invariant. The diagonal shows the correlation between each\nrepresentation's neurons and its corresponding augmented one. Because the two\nversions come from the same image, the diagonal of the matrix should show that\nthere is a strong correlation between them. If the images are different, there\nshouldn't be a diagonal.\n2.   Do not show signs of redundancy. If the neurons show correlation with a\nnon-diagonal neuron, it means that it is not correctly identifying similarities\nbetween the two augmented images. This means that it is redundant.\n\nHere is a good way of understanding in pseudocode(information from the original\npaper):\n\n```\nc[i][i] = 1\nc[i][j] = 0\n\nwhere:\n  c is the cross-correlation matrix\n  i is the index of one representation's neuron\n  j is the index of the second representation's neuron\n```","metadata":{}},{"cell_type":"markdown","source":"## Install, Import and Setup","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:56:02.764557Z","iopub.execute_input":"2025-04-12T05:56:02.765178Z","iopub.status.idle":"2025-04-12T05:56:03.381734Z","shell.execute_reply.started":"2025-04-12T05:56:02.765156Z","shell.execute_reply":"2025-04-12T05:56:03.381059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q  wandb \"torchmetrics>=1.0, <1.5\" \"pytorch-lightning >=2.0,<2.5\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:56:03.382865Z","iopub.execute_input":"2025-04-12T05:56:03.383084Z","iopub.status.idle":"2025-04-12T05:57:12.686053Z","shell.execute_reply.started":"2025-04-12T05:56:03.383065Z","shell.execute_reply":"2025-04-12T05:57:12.685405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from functools import partial\nfrom typing import Sequence, Tuple, Union\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pytorch_lightning as pl\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nimport torchvision.transforms.functional as VisionF\nfrom pytorch_lightning.callbacks import Callback, ModelCheckpoint\nfrom pytorch_lightning.loggers import WandbLogger\nfrom torch import Tensor\nfrom torch.utils.data import DataLoader\nfrom torchmetrics.functional import accuracy\nfrom torchvision.datasets import CIFAR10\nfrom torchvision.models.resnet import resnet34, resnet18\nfrom torchvision.utils import make_grid\nimport wandb\nwandb.login(key='44e48beef599b6c5206fe20e84b042bd22c9eaa6')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:12.686817Z","iopub.execute_input":"2025-04-12T05:57:12.687038Z","iopub.status.idle":"2025-04-12T05:57:32.917796Z","shell.execute_reply.started":"2025-04-12T05:57:12.687019Z","shell.execute_reply":"2025-04-12T05:57:32.917279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 1719\nnum_workers = 8\nmax_epochs = 300\nz_dim = 128","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:32.918979Z","iopub.execute_input":"2025-04-12T05:57:32.919338Z","iopub.status.idle":"2025-04-12T05:57:32.922218Z","shell.execute_reply.started":"2025-04-12T05:57:32.919320Z","shell.execute_reply":"2025-04-12T05:57:32.921708Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Augmentation Utilities\nThe Barlow twins algorithm is heavily reliant on\nAugmentation. One unique feature of the method is that sometimes, augmentations\nprobabilistically occur.\n\n**Augmentations**\n\n*   *RandomToGrayscale*: randomly applies grayscale to image 20% of the time\n*   *RandomColorJitter*: randomly applies color jitter 80% of the time\n*   *RandomFlip*: randomly flips image horizontally 50% of the time\n*   *RandomResizedCrop*: randomly crops an image to a random size then resizes. This\nhappens 100% of the time\n*   *RandomSolarize*: randomly applies solarization to an image 20% of the time\n*   *RandomBlur*: randomly blurs an image 20% of the time\n","metadata":{}},{"cell_type":"code","source":"class BarlowTwinsTransform:\n    def __init__(self, train=True, input_height=224, gaussian_blur=True, jitter_strength=1.0, normalize=None):\n        self.input_height = input_height\n        self.gaussian_blur = gaussian_blur\n        self.jitter_strength = jitter_strength\n        self.normalize = normalize\n        self.train = train\n\n        color_jitter = transforms.ColorJitter(\n            0.8 * self.jitter_strength,\n            0.8 * self.jitter_strength,\n            0.8 * self.jitter_strength,\n            0.2 * self.jitter_strength,\n        )\n\n        color_transform = [transforms.RandomApply([color_jitter], p=0.8), transforms.RandomGrayscale(p=0.2)]\n\n        if self.gaussian_blur:\n            kernel_size = int(0.1 * self.input_height)\n            if kernel_size % 2 == 0:\n                kernel_size += 1\n\n            color_transform.append(transforms.RandomApply([transforms.GaussianBlur(kernel_size=kernel_size)], p=0.5))\n\n        self.color_transform = transforms.Compose(color_transform)\n\n        if normalize is None:\n            self.final_transform = transforms.ToTensor()\n        else:\n            self.final_transform = transforms.Compose([transforms.ToTensor(), normalize])\n\n        self.transform = transforms.Compose(\n            [\n                transforms.RandomResizedCrop(self.input_height),\n                transforms.RandomHorizontalFlip(p=0.5),\n                self.color_transform,\n                self.final_transform,\n            ]\n        )\n\n        self.finetune_transform = None\n        if self.train:\n            self.finetune_transform = transforms.Compose(\n                [\n                    transforms.RandomCrop(32, padding=4, padding_mode=\"reflect\"),\n                    transforms.RandomHorizontalFlip(),\n                    transforms.ToTensor(),\n                ]\n            )\n        else:\n            self.finetune_transform = transforms.ToTensor()\n\n    def __call__(self, sample):\n        return self.transform(sample), self.transform(sample), self.finetune_transform(sample)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:32.922779Z","iopub.execute_input":"2025-04-12T05:57:32.922952Z","iopub.status.idle":"2025-04-12T05:57:35.029634Z","shell.execute_reply.started":"2025-04-12T05:57:32.922938Z","shell.execute_reply":"2025-04-12T05:57:35.028992Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load the CIFAR-10 dataset","metadata":{}},{"cell_type":"code","source":"def cifar10_normalization():\n    normalize = transforms.Normalize(\n        mean=[x / 255.0 for x in [125.3, 123.0, 113.9]], std=[x / 255.0 for x in [63.0, 62.1, 66.7]]\n    )\n    return normalize\n\n\ntrain_transform = BarlowTwinsTransform(\n    train=True, input_height=32, gaussian_blur=False, jitter_strength=0.5, normalize=cifar10_normalization()\n)\ntrain_dataset = CIFAR10(root=\".\", train=True, download=True, transform=train_transform)\n\nval_transform = BarlowTwinsTransform(\n    train=False, input_height=32, gaussian_blur=False, jitter_strength=0.5, normalize=cifar10_normalization()\n)\nval_dataset = CIFAR10(root=\".\", train=False, download=True, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, drop_last=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:35.030234Z","iopub.execute_input":"2025-04-12T05:57:35.030418Z","iopub.status.idle":"2025-04-12T05:57:41.699265Z","shell.execute_reply.started":"2025-04-12T05:57:35.030404Z","shell.execute_reply":"2025-04-12T05:57:41.698698Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loading","metadata":{}},{"cell_type":"code","source":"for batch in val_loader:\n    (img1, img2, _), label = batch\n    break\n\nimg_grid = make_grid(img1, normalize=True)\n\n\ndef show(imgs):\n    if not isinstance(imgs, list):\n        imgs = [imgs]\n    fix, axs = plt.subplots(ncols=len(imgs), squeeze=False)\n    for i, img in enumerate(imgs):\n        img = img.detach()\n        img = VisionF.to_pil_image(img)\n        axs[0, i].imshow(np.asarray(img))\n        axs[0, i].set(xticklabels=[], yticklabels=[], xticks=[], yticks=[])\n\n\nshow(img_grid)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:41.699857Z","iopub.execute_input":"2025-04-12T05:57:41.700071Z","iopub.status.idle":"2025-04-12T05:57:46.146941Z","shell.execute_reply.started":"2025-04-12T05:57:41.700050Z","shell.execute_reply":"2025-04-12T05:57:46.146250Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## BarlowLoss: barlow twins model's loss function\n\nAs shown in the architecture, firstly two distorted views for all images of a batch X\nsampled from a dataset. The distorted views are obtained via a distribution of data augmentations T. The two batches\nof distorted views $Y^A$ and $Y^B$\nB are then fed to a function $f_θ$, typically a deep network with trainable parameters θ, producing batches of embeddings $Z^A$ and $Z^B$ respectively.\nTo simplify notations, $Z^A$ and $Z^B$ are assumed to be mean-centered along the batch dimension, such that each unit has\nmean output 0 over the batch.\n\n\nSo the Barlow loss function is :\n\n$$\n\\mathcal{L}_{\\mathcal{BT}} \\triangleq\n\\underbrace{\\sum_i (1 - C_{ii})^2}_{\\text{invariance term}}\n+ \\lambda\n\\underbrace{\\sum_i \\sum_{j \\ne i} C_{ij}^2}_{\\text{redundancy reduction term}}\n$$\n\nwhere λ is a positive constant trading off the importance of\nthe first and second terms of the loss, and where C is the\ncross-correlation matrix computed between the outputs of\nthe two identical networks along the batch dimension:\n\n$$\nC_{ij} \\triangleq \\frac{\n\\sum_b z_{b,i}^A z_{b,j}^B\n}{\n\\sqrt{\\sum_b \\left(z_{b,i}^A\\right)^2} \\sqrt{\\sum_b \\left(z_{b,j}^B\\right)^2}\n}\n$$\n\nwhere b indexes batch samples and i, j index the vector dimension of the networks’ outputs. C is a square matrix with\nsize the dimensionality of the network’s output, and with\nvalues comprised between -1 (i.e. perfect anti-correlation)\nand 1 (i.e. perfect correlation).\n\nThe two parts to thebloss function:\n\n*   ***The invariance term***(diagonal). This part is used to make the diagonals of the\nmatrix into 1s. When this is the case, the matrix shows that the images are\ncorrelated(same).\n  * The loss function subtracts 1 from the diagonal and squares the values.\n*   ***The redundancy reduction term***(off-diagonal). Here, the barlow twins loss\nfunction aims to make these values zero. As mentioned before, it is redundant if the\nrepresentation neurons are correlated with values that are not on the diagonal.\n  * Off diagonals are squared.","metadata":{}},{"cell_type":"code","source":"class BarlowTwinsLoss(nn.Module):\n    def __init__(self, batch_size, lambda_coeff=5e-3, z_dim=128):\n        super().__init__()\n\n        self.z_dim = z_dim\n        self.batch_size = batch_size\n        self.lambda_coeff = lambda_coeff\n\n    def off_diagonal_ele(self, x):\n        # taken from: https://github.com/facebookresearch/barlowtwins/blob/main/main.py\n        # return a flattened view of the off-diagonal elements of a square matrix\n        n, m = x.shape\n        assert n == m\n        return x.flatten()[:-1].view(n - 1, n + 1)[:, 1:].flatten()\n\n    def forward(self, z1, z2):\n        # N x D, where N is the batch size and D is output dim of projection head\n        z1_norm = (z1 - torch.mean(z1, dim=0)) / torch.std(z1, dim=0)\n        z2_norm = (z2 - torch.mean(z2, dim=0)) / torch.std(z2, dim=0)\n\n        cross_corr = torch.matmul(z1_norm.T, z2_norm) / self.batch_size\n\n        on_diag = torch.diagonal(cross_corr).add_(-1).pow_(2).sum()\n        off_diag = self.off_diagonal_ele(cross_corr).pow_(2).sum()\n\n        return on_diag + self.lambda_coeff * off_diag","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.147888Z","iopub.execute_input":"2025-04-12T05:57:46.148145Z","iopub.status.idle":"2025-04-12T05:57:46.154225Z","shell.execute_reply.started":"2025-04-12T05:57:46.148120Z","shell.execute_reply":"2025-04-12T05:57:46.153678Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Barlow Twins' Model Architecture\nOrginal architecture uses [ResNet-50](https://viso.ai/deep-learning/resnet-residual-neural-network/). But due to computation power constraints we use ResNet-18\n\nTo accommodate the 32x32 CIFAR10 images, we replace the first 7x7 convolution of the Resnet backbone by a 3x3 filter. We also remove the first Maxpool layer from the network for CIFAR10 images.","metadata":{}},{"cell_type":"code","source":"encoder = resnet18()\n\n# for CIFAR10, replace the first 7x7 conv with smaller 3x3 conv and remove the first maxpool\nencoder.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)\nencoder.maxpool = nn.MaxPool2d(kernel_size=1, stride=1)\n\n# replace classification fc layer of Resnet to obtain representations from the backbone\nencoder.fc = nn.Identity()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.154927Z","iopub.execute_input":"2025-04-12T05:57:46.155207Z","iopub.status.idle":"2025-04-12T05:57:46.695927Z","shell.execute_reply.started":"2025-04-12T05:57:46.155191Z","shell.execute_reply":"2025-04-12T05:57:46.695375Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Projector network**:\n\nThe paper utilizes a 3 layer MLP with 8192 hidden dimensions and 8192 as the output dimension of the projection head. For the purposes of the tutorial, we use a smaller projection head. But, it is imperative to mention here that in practice, Barlow Twins needs to be trained using a bigger projection head as it is highly sensitive to its architecture and output dimensionality.","metadata":{}},{"cell_type":"code","source":"class ProjectionHead(nn.Module):\n    def __init__(self, input_dim=2048, hidden_dim=2048, output_dim=128):\n        super().__init__()\n\n        self.projection_head = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim, bias=True),\n            nn.BatchNorm1d(hidden_dim),\n            nn.ReLU(),\n            nn.Linear(hidden_dim, output_dim, bias=False),\n        )\n\n    def forward(self, x):\n        return self.projection_head(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.697643Z","iopub.execute_input":"2025-04-12T05:57:46.698047Z","iopub.status.idle":"2025-04-12T05:57:46.701803Z","shell.execute_reply.started":"2025-04-12T05:57:46.698029Z","shell.execute_reply":"2025-04-12T05:57:46.701313Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"For the purposes of this project, we keep things simple and use a linear warmup schedule with Adam optimizer.","metadata":{}},{"cell_type":"code","source":"def fn(warmup_steps, step):\n    if step < warmup_steps:\n        return float(step) / float(max(1, warmup_steps))\n    else:\n        return 1.0\n\n\ndef linear_warmup_decay(warmup_steps):\n    return partial(fn, warmup_steps)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.702364Z","iopub.execute_input":"2025-04-12T05:57:46.702538Z","iopub.status.idle":"2025-04-12T05:57:46.713745Z","shell.execute_reply.started":"2025-04-12T05:57:46.702524Z","shell.execute_reply":"2025-04-12T05:57:46.713118Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Pseudocode of training loop\nTaken from the paper\n![pseudocode](https://i.imgur.com/Tlrootj.png)\n\nExplaining each variable:\n\n```\ny_a: first augmented version of original image.\ny_b: second augmented version of original image.\nz_a: model representation(embeddings) of y_a.\nz_b: model representation(embeddings) of y_b.\nz_a_norm: normalized z_a.\nz_b_norm: normalized z_b.\nc: cross correlation matrix.\nc_diff: diagonal portion of loss(invariance term).\noff_diag: off-diagonal portion of loss(redundancy reduction term).\n```","metadata":{}},{"cell_type":"code","source":"class BarlowTwins(pl.LightningModule):\n    def __init__(\n        self,\n        encoder,\n        encoder_out_dim,\n        num_training_samples,\n        batch_size,\n        lambda_coeff=5e-3,\n        z_dim=128,\n        learning_rate=1e-4,\n        warmup_epochs=10,\n        max_epochs=200,\n    ):\n        super().__init__()\n\n        self.encoder = encoder\n        self.projection_head = ProjectionHead(input_dim=encoder_out_dim, hidden_dim=encoder_out_dim, output_dim=z_dim)\n        self.loss_fn = BarlowTwinsLoss(batch_size=batch_size, lambda_coeff=lambda_coeff, z_dim=z_dim)\n\n        self.learning_rate = learning_rate\n        self.warmup_epochs = warmup_epochs\n        self.max_epochs = max_epochs\n\n        self.train_iters_per_epoch = num_training_samples // batch_size\n\n    def forward(self, x):\n        return self.encoder(x)\n\n    def shared_step(self, batch):\n        (x1, x2, _), _ = batch\n\n        z1 = self.projection_head(self.encoder(x1))\n        z2 = self.projection_head(self.encoder(x2))\n\n        return self.loss_fn(z1, z2)\n\n    def training_step(self, batch, batch_idx):\n        loss = self.shared_step(batch)\n        self.log(\"train_loss\", loss, on_step=True, prog_bar=True, logger=True, on_epoch=False)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        loss = self.shared_step(batch)\n        self.log(\"val_loss\", loss, on_step=True, prog_bar=True, logger=True, on_epoch=True)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=self.learning_rate)\n\n        warmup_steps = self.train_iters_per_epoch * self.warmup_epochs\n\n        scheduler = {\n            \"scheduler\": torch.optim.lr_scheduler.LambdaLR(\n                optimizer,\n                linear_warmup_decay(warmup_steps),\n            ),\n            \"interval\": \"step\",\n            \"frequency\": 1,\n        }\n\n        return [optimizer], [scheduler]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.714474Z","iopub.execute_input":"2025-04-12T05:57:46.714639Z","iopub.status.idle":"2025-04-12T05:57:46.724724Z","shell.execute_reply.started":"2025-04-12T05:57:46.714626Z","shell.execute_reply":"2025-04-12T05:57:46.724228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Evaluation\n\n**Linear evaluation:** We define a callback which appends a linear layer on top of the encoder and trains the classification evaluation head. We make sure not to backpropagate the gradients back to the encoder while tuning the linear layer. This means that we have freezed the encoder.","metadata":{}},{"cell_type":"code","source":"class OnlineFineTuner(Callback):\n    def __init__(\n        self,\n        encoder_output_dim: int,\n        num_classes: int,\n    ) -> None:\n        super().__init__()\n\n        self.optimizer: torch.optim.Optimizer\n\n        self.encoder_output_dim = encoder_output_dim\n        self.num_classes = num_classes\n\n    def on_fit_start(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:\n        # add linear_eval layer and optimizer\n        pl_module.online_finetuner = nn.Linear(self.encoder_output_dim, self.num_classes).to(pl_module.device)\n        self.optimizer = torch.optim.Adam(pl_module.online_finetuner.parameters(), lr=1e-4)\n\n    def extract_online_finetuning_view(\n        self, batch: Sequence, device: Union[str, torch.device]\n    ) -> Tuple[Tensor, Tensor]:\n        (_, _, finetune_view), y = batch\n        finetune_view = finetune_view.to(device)\n        y = y.to(device)\n\n        return finetune_view, y\n\n    def on_train_batch_end(\n        self,\n        trainer: pl.Trainer,\n        pl_module: pl.LightningModule,\n        outputs: Sequence,\n        batch: Sequence,\n        batch_idx: int,\n    ) -> None:\n        x, y = self.extract_online_finetuning_view(batch, pl_module.device)\n\n        with torch.no_grad():\n            feats = pl_module(x)\n\n        feats = feats.detach()\n        preds = pl_module.online_finetuner(feats)\n        loss = F.cross_entropy(preds, y)\n\n        loss.backward()\n        self.optimizer.step()\n        self.optimizer.zero_grad()\n\n        acc = accuracy(F.softmax(preds, dim=1), y, task=\"multiclass\", num_classes=10)\n        pl_module.log(\"online_train_acc\", acc, on_step=True, prog_bar=True, logger=True, on_epoch=False)\n        pl_module.log(\"online_train_loss\", loss, on_step=True, prog_bar=True, logger=True, on_epoch=False)\n\n    def on_validation_batch_end(\n        self,\n        trainer: pl.Trainer,\n        pl_module: pl.LightningModule,\n        outputs: Sequence,\n        batch: Sequence,\n        batch_idx: int,\n    ) -> None:\n        x, y = self.extract_online_finetuning_view(batch, pl_module.device)\n\n        with torch.no_grad():\n            feats = pl_module(x)\n\n        feats = feats.detach()\n        preds = pl_module.online_finetuner(feats)\n        loss = F.cross_entropy(preds, y)\n\n        acc = accuracy(F.softmax(preds, dim=1), y, task=\"multiclass\", num_classes=10)\n        pl_module.log(\"online_val_acc\", acc, on_step=False, on_epoch=True, prog_bar=True, logger=True, sync_dist=True)\n        pl_module.log(\"online_val_loss\", loss, on_step=False, on_epoch=True, prog_bar=True, logger=True, sync_dist=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.725287Z","iopub.execute_input":"2025-04-12T05:57:46.725607Z","iopub.status.idle":"2025-04-12T05:57:46.739114Z","shell.execute_reply.started":"2025-04-12T05:57:46.725592Z","shell.execute_reply":"2025-04-12T05:57:46.738499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"encoder_out_dim = 512\n\nmodel = BarlowTwins(\n    encoder=encoder,\n    encoder_out_dim=encoder_out_dim,\n    num_training_samples=len(train_dataset),\n    batch_size=batch_size,\n    z_dim=z_dim,\n)\n\nonline_finetuner = OnlineFineTuner(encoder_output_dim=encoder_out_dim, num_classes=10)\ncheckpoint_callback = ModelCheckpoint(every_n_epochs=50, save_top_k=-1, save_last=True)\nwandb_logger = WandbLogger(project='DLCV Project')\n\ntrainer = pl.Trainer(\n    accelerator='gpu', \n    strategy='auto',\n    devices=1,\n    num_nodes=1,\n    logger=wandb_logger,\n    max_epochs=max_epochs,\n    log_every_n_steps=1,\n    num_sanity_val_steps=0,\n    enable_checkpointing=True,\n    enable_progress_bar=True, \n    enable_model_summary=False,\n    accumulate_grad_batches=1,\n    gradient_clip_val=None,\n    gradient_clip_algorithm=None,\n    callbacks=[online_finetuner, checkpoint_callback],\n    default_root_dir='/kaggle/working'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.739632Z","iopub.execute_input":"2025-04-12T05:57:46.739799Z","iopub.status.idle":"2025-04-12T05:57:46.842420Z","shell.execute_reply.started":"2025-04-12T05:57:46.739786Z","shell.execute_reply":"2025-04-12T05:57:46.841900Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.fit(model, train_loader, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T05:57:46.843027Z","iopub.execute_input":"2025-04-12T05:57:46.843206Z","iopub.status.idle":"2025-04-12T10:54:22.852553Z","shell.execute_reply.started":"2025-04-12T05:57:46.843192Z","shell.execute_reply":"2025-04-12T10:54:22.851914Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training curves\n![](https://i.postimg.cc/tCnsTwrJ/Barlow-loss.png)\n![](https://i.postimg.cc/RZrNnd8R/Classification-loss.png)\n![](https://i.postimg.cc/QdpVg5Zr/Accuracy-plot.png)\n","metadata":{}},{"cell_type":"markdown","source":"## Conclusion\n\n*   BARLOW TWINS learns self-supervised representations\nthrough a joint embedding of distorted images, with an objective function that maximizes similarity between the embedding vectors while reducing redundancy between their\ncomponents.\n*   With this resnet-18 model architecture, we were able to reach 58.4% validation accuracy.\n","metadata":{}},{"cell_type":"markdown","source":"## Other details\nAblation study was a crucial phase of the research paper. In this phase BARLOW TWINS was trained for 300 epochs instead of 1000 epochs. It included:\n1. Loss Function Ablations\n2. Robustness to Batch Size\n3. Effect of Removing Augmentations\n4. Projector Network Depth & Width\n5. Breaking Symmetry\n6. BYOL with a larger projector/predictor/embedding\n7. Sensitivity to λ","metadata":{}}]}