{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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"},"papermill":{"default_parameters":{},"duration":1087.035813,"end_time":"2023-01-22T10:17:16.597092","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-01-22T09:59:09.561279","version":"2.3.4"},"colab":{"provenance":[]},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":39272,"databundleVersionId":4629629},{"sourceType":"datasetVersion","sourceId":4619805,"datasetId":2688675,"databundleVersionId":4681402}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"## importing and installing stuff","metadata":{"id":"Y3iW0S4HfaSb"}},{"cell_type":"code","source":"!pip install pydicom -q","metadata":{"id":"5nMAz_AQt3N0","outputId":"5a367117-0401-41e1-9838-07b1a238ccec","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:26:56.815988Z","iopub.execute_input":"2026-07-10T11:26:56.816254Z","iopub.status.idle":"2026-07-10T11:27:01.362588Z","shell.execute_reply.started":"2026-07-10T11:26:56.816223Z","shell.execute_reply":"2026-07-10T11:27:01.361692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2 as cv\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pydicom\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\nfrom ignite.engine import Events, create_supervised_trainer, create_supervised_evaluator\nfrom ignite.metrics import Accuracy, Loss, RunningAverage\nfrom ignite.contrib.handlers import ProgressBar\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import models, transforms","metadata":{"papermill":{"duration":3.715291,"end_time":"2023-01-22T09:59:22.251424","exception":false,"start_time":"2023-01-22T09:59:18.536133","status":"completed"},"tags":[],"id":"d51b515d","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:27:17.659935Z","iopub.execute_input":"2026-07-10T11:27:17.660541Z","iopub.status.idle":"2026-07-10T11:27:27.446516Z","shell.execute_reply.started":"2026-07-10T11:27:17.660503Z","shell.execute_reply":"2026-07-10T11:27:27.445897Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## theoretical background\n\nUntil now, you have played around with Convolutional Neural Networks. For years, they were the dominant force in Computer Vision. Until a groundbreaking paper introduced transformers. Initially introduced for Natural Language Processing tasks, they were then also employed in Computer Vision as well. This is that original paper:\n\n[Attention is All you Need](https://proceedings.neurips.cc/paper_files/paper/2017/file/3f5ee243547dee91fbd053c1c4a845aa-Paper.pdf)\n\nIf you didn't understand too much from that, worry not. Go through below resources as well:\n\n[Basic overview (from an NLP viewpoint)](https://medium.com/inside-machine-learning/what-is-a-transformer-d07dd1fbec04)\n\n[A bit more in depth (again, from an NLP viewpoint)](https://towardsdatascience.com/transformers-141e32e69591)\n\n\nTo get more background on how exactly vision transformers work:\n\n[For visual learners](https://youtu.be/qU7wO02urYU?si=bj8Xj-DG2qDbwdnH)\n\n[Read through this as well](https://www.v7labs.com/blog/vision-transformer-guide)\n\nTransformers found their initial applications in natural language processing (NLP) tasks. To use this NLP model for computer vision tasks, we have to divide our input image into patches. After flattening the patches, we can treat each flattened patches as single word. We add positional embeddings to the linear projection of flattened patches. An extra token is added at the beginning for classification tasks. In BERT model, this token is called [CLS] token.\n\nSo if our input image size is (512, 512), after dividing the image into patches of size (16, 16), we get 1024 (32 times 32) patches. After flattening the patches and projecting the flattened patches, we have 1024 tokens. After adding positional embeddings and concatenating classification token at the beginning, we have 1025 tokens.\n\nWe then feed our tokens into the transformer encoder. Transformer encoder is made up of self attention and feedforward network. This [video](https://www.youtube.com/watch?v=_UVfwBqcnbM) by AssemblyAI explains the transformer architecture beautifully.\n\nThe number of tokens in the output of the transformer encoder is equal to number of input tokens. We take the first token from the output (corresponds to the classification token) and feed the token in a multilayer perceptron head for classification.","metadata":{"papermill":{"duration":0.005588,"end_time":"2023-01-22T09:59:18.497792","exception":false,"start_time":"2023-01-22T09:59:18.492204","status":"completed"},"tags":[],"id":"7cdc0644"}},{"cell_type":"markdown","source":"## dataset info\nThe dataset was contributed by mammography screening programs in Australia and the U.S. It includes detailed labels, with radiologists’ evaluations and follow-up pathology results for suspected malignancies.\n\nThe dataset is stored in dicom formats. Converting dicom data to png/jpg just by rescaling it will harm the quality of the data. [This notebook](https://www.kaggle.com/code/raddar/convert-dicom-to-np-array-the-correct-way/notebook) is an awesome resource for anyone working with dicom files for X-Ray.\n\nTo know more about how dicom images work, go through [this](https://towardsdatascience.com/understanding-dicom-bce665e62b72). They are specially used for X - Ray images.","metadata":{"papermill":{"duration":0.005403,"end_time":"2023-01-22T09:59:18.50875","exception":false,"start_time":"2023-01-22T09:59:18.503347","status":"completed"},"tags":[],"id":"0aed2f4a"}},{"cell_type":"markdown","source":"## utility functions\nWe write some utility functions beforehand. To know how to work with PyDicom funtions, read [this](https://towardsdatascience.com/introducing-pydicom-its-classes-methods-and-attributes-518c1d71162).","metadata":{"papermill":{"duration":0.005416,"end_time":"2023-01-22T09:59:18.530519","exception":false,"start_time":"2023-01-22T09:59:18.525103","status":"completed"},"tags":[],"id":"d95bbc52"}},{"cell_type":"code","source":"def read_xray(file_path, img_size=None):\n    dcm_data = pydicom.dcmread(file_path)\n    img = dcm_data.pixel_array\n\n    if dcm_data.PhotometricInterpretation == \"MONOCHROME1\": \n        img = np.max(img) - img\n    \n    if img_size:\n        img = cv.resize(img,img_size)\n        \n    img = np.expand_dims(img,axis=0)\n    img = img.astype(np.float32)\n    img = img/(np.max(img))\n    return img","metadata":{"papermill":{"duration":0.015247,"end_time":"2023-01-22T09:59:22.272633","exception":false,"start_time":"2023-01-22T09:59:22.257386","status":"completed"},"tags":[],"id":"221fb680","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:27:31.825943Z","iopub.execute_input":"2026-07-10T11:27:31.826508Z","iopub.status.idle":"2026-07-10T11:27:31.83191Z","shell.execute_reply.started":"2026-07-10T11:27:31.826479Z","shell.execute_reply":"2026-07-10T11:27:31.831038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def patchify(batch, patch_size):\n    b, c, h, w = batch.shape\n    ph, pw = patch_size\n    nh, nw = h // ph, w // pw\n    x = batch.reshape(b, c, nh, ph, nw, pw)\n    batch_patches = x.permute(0, 2, 4, 3, 5, 1)\n    return batch_patches","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:27:34.325322Z","iopub.execute_input":"2026-07-10T11:27:34.32583Z","iopub.status.idle":"2026-07-10T11:27:34.330499Z","shell.execute_reply.started":"2026-07-10T11:27:34.325798Z","shell.execute_reply":"2026-07-10T11:27:34.329656Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We test our `patchify` function on a single image.","metadata":{"papermill":{"duration":0.00537,"end_time":"2023-01-22T09:59:22.30469","exception":false,"start_time":"2023-01-22T09:59:22.29932","status":"completed"},"tags":[],"id":"645afc03"}},{"cell_type":"code","source":"FILE_PATH = ('/kaggle/input/rsna-breast-cancer-detection/'\n             'train_images/10006/1459541791.dcm')\n\nimg = read_xray(FILE_PATH, img_size=(512, 512))\nprint(\"Image shape:\",img.shape)\nbatch = torch.tensor(img[None])\nprint(\"Batch shape:\", *batch.shape)\n\npatch_size = (16, 16)\nbatch_patches = patchify(batch, patch_size)\nprint(\"Batch patches shape:\", *batch_patches.shape)\n\npatches = batch_patches[0]\nprint(\"Patches shape:\", *patches.shape)\n\nnh, nw, ph, pw, c = patches.shape\nprint(f\"nh, nw, ph, pw, c = {nh, nw, ph, pw, c}\")\n\nplt.figure(figsize=(5, 5))\nplt.imshow(img[0], cmap=\"gray\")\nplt.axis(\"off\")\n\nprint(f\"nh, nw = {nh, nw}\")\nplt.figure(figsize=(5,5))\n\nfor i in range(nh):\n    for j in range(nw):\n        plt.subplot(nh, nw, i * nw + j + 1)\n        #plt.imshow(patches[0, i, j], cmap=\"gray\")\n        plt.imshow(patches[i, j], cmap=\"gray\")\n        plt.axis(\"off\")","metadata":{"papermill":{"duration":48.930892,"end_time":"2023-01-22T10:00:11.241361","exception":false,"start_time":"2023-01-22T09:59:22.310469","status":"completed"},"tags":[],"id":"522d21b9","outputId":"92211244-22b1-4e50-fe3d-c0f167d522e8","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:27:37.126781Z","iopub.execute_input":"2026-07-10T11:27:37.127212Z","iopub.status.idle":"2026-07-10T11:27:53.478157Z","shell.execute_reply.started":"2026-07-10T11:27:37.12718Z","shell.execute_reply":"2026-07-10T11:27:53.477221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_mlp(in_features, hidden_units, out_features):\n    dims = [in_features] + hidden_units + [out_features]\n    layers = []\n    for dim1, dim2 in zip(dims[:-2], dims[1:-1]):\n        layers.append(nn.Linear(dim1, dim2))\n        layers.append(nn.ReLU())\n\n    layers.append(nn.Linear(dims[-2], dims[-1]))\n    return nn.Sequential(*layers)","metadata":{"papermill":{"duration":0.017061,"end_time":"2023-01-22T10:00:11.2657","exception":false,"start_time":"2023-01-22T10:00:11.248639","status":"completed"},"tags":[],"id":"5b559866","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:01.111099Z","iopub.execute_input":"2026-07-10T11:28:01.111542Z","iopub.status.idle":"2026-07-10T11:28:01.116019Z","shell.execute_reply.started":"2026-07-10T11:28:01.111513Z","shell.execute_reply":"2026-07-10T11:28:01.115392Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## image to sequence block\nThis Block takes a batch of image as input and returns a batch of sequences. Later on we feed this sequences into the transformer encoder.","metadata":{"papermill":{"duration":0.006314,"end_time":"2023-01-22T10:00:11.278755","exception":false,"start_time":"2023-01-22T10:00:11.272441","status":"completed"},"tags":[],"id":"6cb491e3"}},{"cell_type":"code","source":"class Img2Seq(nn.Module):\n    def __init__(self, img_size, patch_size, n_channels, d_model):\n        super().__init__()\n        self.patch_size = patch_size\n        self.img_size = img_size\n\n        nh, nw = img_size[0] // patch_size[0], img_size[1] // patch_size[1]\n        n_tokens = nh * nw\n        token_dim = patch_size[0] * patch_size[1] * n_channels\n        \n        self.linear = nn.Linear(token_dim, d_model)\n        self.cls_token = nn.Parameter(torch.randn(1, 1, d_model))\n        self.pos_emb = nn.Parameter(torch.randn(n_tokens + 1, d_model))\n        \n    def forward(self, batch):\n        batch = patchify(batch, self.patch_size)\n        b, nh, nw, ph, pw, c = batch.shape\n        batch = batch.reshape(b, nh * nw, -1)\n        emb = self.linear(batch)\n        Cls = self.cls_token.expand(b, -1, -1)\n        x = torch.cat([Cls, emb], dim=1) \n        return x + self.pos_emb","metadata":{"papermill":{"duration":0.018674,"end_time":"2023-01-22T10:00:11.304195","exception":false,"start_time":"2023-01-22T10:00:11.285521","status":"completed"},"tags":[],"id":"58ae05ba","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:04.42346Z","iopub.execute_input":"2026-07-10T11:28:04.423888Z","iopub.status.idle":"2026-07-10T11:28:04.430681Z","shell.execute_reply.started":"2026-07-10T11:28:04.423857Z","shell.execute_reply":"2026-07-10T11:28:04.429831Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## visual transformer module\nThis modules wraps up everything. We can divide this module into 3 parts:\n* An image to sequence encoder\n* Transformer encoder\n* Multilayer perceptron head classification\n\nWe use `torch.nn.TransformerEncoder` and `torch.nn.TransformerEncoderLayer` to implement our transformer encoder. We highly recommend to read the official documentation to [learn more about the layers](https://pytorch.org/docs/stable/generated/torch.nn.TransformerEncoder.html).\nFor the loss function, we will be using [gelu](https://arxiv.org/abs/1606.08415). You can also use this video to [learn about it](https://www.youtube.com/watch?v=FWhMkpo9yuM).","metadata":{"papermill":{"duration":0.006736,"end_time":"2023-01-22T10:00:11.317817","exception":false,"start_time":"2023-01-22T10:00:11.311081","status":"completed"},"tags":[],"id":"26d7647d"}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ViT(nn.Module):\n    def __init__(self,img_size,patch_size,n_channels,d_model,nhead,dim_feedforward,blocks,mlp_head_units,n_classes,):\n        super().__init__()\n        self.img2seq = Img2Seq(img_size, patch_size, n_channels, d_model)\n        encoder_layer = nn.TransformerEncoderLayer(d_model=d_model,nhead=nhead,dim_feedforward=dim_feedforward,activation='gelu',batch_first=True)\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=blocks)\n        self.mlp = get_mlp(d_model, mlp_head_units, n_classes)\n        self.output = nn.Sigmoid() if n_classes == 1 else nn.Softmax(dim=-1)\n        \n    def forward(self, batch):\n        batch = self.img2seq(batch)\n        batch = self.transformer_encoder(batch)\n        batch = batch[:, 0, :]\n        batch = self.mlp(batch)\n        output = self.output(batch)\n        return output","metadata":{"papermill":{"duration":0.018775,"end_time":"2023-01-22T10:00:11.34348","exception":false,"start_time":"2023-01-22T10:00:11.324705","status":"completed"},"tags":[],"id":"00b544d2","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:08.287613Z","iopub.execute_input":"2026-07-10T11:28:08.288049Z","iopub.status.idle":"2026-07-10T11:28:08.294687Z","shell.execute_reply.started":"2026-07-10T11:28:08.288006Z","shell.execute_reply":"2026-07-10T11:28:08.293847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"papermill":{"duration":0.137093,"end_time":"2023-01-22T10:00:11.514575","exception":false,"start_time":"2023-01-22T10:00:11.377482","status":"completed"},"tags":[],"id":"6a5ac8c3","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:10.832681Z","iopub.execute_input":"2026-07-10T11:28:10.833269Z","iopub.status.idle":"2026-07-10T11:28:11.110082Z","shell.execute_reply.started":"2026-07-10T11:28:10.83324Z","shell.execute_reply":"2026-07-10T11:28:11.109124Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## dataset loading\nWe will be using the [CLAHE](https://www.geeksforgeeks.org/clahe-histogram-eqalization-opencv/) algorithm to improve contrast between the tiled images. Apart from this, we are using the pytorch data handling modules like DataLoaded which you can read about [here](https://pytorch.org/docs/stable/data.html).","metadata":{"papermill":{"duration":0.007039,"end_time":"2023-01-22T10:00:11.528747","exception":false,"start_time":"2023-01-22T10:00:11.521708","status":"completed"},"tags":[],"id":"4587c40c"}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, df, img_path):\n        self.df = df\n        self.img_path = img_path\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        patient_id, image_id, cancer = self.df.iloc[idx][['patient_id', 'image_id', 'cancer']]\n        file = os.path.join(self.img_path, f'{patient_id}_{image_id}.png')\n        \n        img = cv.imread(file, cv.IMREAD_GRAYSCALE)\n        clahe = cv.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        img = clahe.apply(img)\n        img = img.astype('float32') / 255.0\n        X = torch.tensor(img[np.newaxis], dtype=torch.float32)\n        y = torch.tensor([cancer], dtype=torch.float32)\n        return X, y","metadata":{"papermill":{"duration":0.018483,"end_time":"2023-01-22T10:00:11.554179","exception":false,"start_time":"2023-01-22T10:00:11.535696","status":"completed"},"tags":[],"id":"8349a5ae","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:11.284193Z","iopub.execute_input":"2026-07-10T11:28:11.285093Z","iopub.status.idle":"2026-07-10T11:28:11.291056Z","shell.execute_reply.started":"2026-07-10T11:28:11.28506Z","shell.execute_reply":"2026-07-10T11:28:11.290175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:13.92267Z","iopub.execute_input":"2026-07-10T11:28:13.923098Z","iopub.status.idle":"2026-07-10T11:28:14.043268Z","shell.execute_reply.started":"2026-07-10T11:28:13.923069Z","shell.execute_reply":"2026-07-10T11:28:14.04264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"counts = df.cancer.value_counts()\ndf['weights'] = df['cancer'].apply(lambda x: 1/counts[x])\ntrain_df, val_df = train_test_split(df, test_size=0.25, stratify=df['cancer'])","metadata":{"papermill":{"duration":0.432908,"end_time":"2023-01-22T10:00:11.993993","exception":false,"start_time":"2023-01-22T10:00:11.561085","status":"completed"},"tags":[],"id":"c8cdd3ff","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:14.222827Z","iopub.execute_input":"2026-07-10T11:28:14.223121Z","iopub.status.idle":"2026-07-10T11:28:14.374138Z","shell.execute_reply.started":"2026-07-10T11:28:14.223095Z","shell.execute_reply":"2026-07-10T11:28:14.373492Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:15.183299Z","iopub.execute_input":"2026-07-10T11:28:15.183743Z","iopub.status.idle":"2026-07-10T11:28:15.195431Z","shell.execute_reply.started":"2026-07-10T11:28:15.18371Z","shell.execute_reply":"2026-07-10T11:28:15.194611Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_path = '/kaggle/input/rsna-breast-cancer-512-pngs'\ntrain_samples = len(train_df)\nval_samples = len(val_df)\nBATCH_SIZE = 16\n\ntrain_ds = RSNADataset(train_df, img_path)\nval_ds = RSNADataset(val_df, img_path)\n\ntrain_sampler = WeightedRandomSampler(train_df['weights'].values, train_samples)\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=train_sampler)\n\nval_sampler = WeightedRandomSampler(val_df['weights'].values, val_samples)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, sampler=val_sampler)","metadata":{"papermill":{"duration":0.018775,"end_time":"2023-01-22T10:00:12.019912","exception":false,"start_time":"2023-01-22T10:00:12.001137","status":"completed"},"tags":[],"id":"b2c20004","trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:15.841071Z","iopub.execute_input":"2026-07-10T11:28:15.841512Z","iopub.status.idle":"2026-07-10T11:28:15.848133Z","shell.execute_reply.started":"2026-07-10T11:28:15.841483Z","shell.execute_reply":"2026-07-10T11:28:15.847415Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## model creation and training start","metadata":{"papermill":{"duration":0.007287,"end_time":"2023-01-22T10:00:12.034465","exception":false,"start_time":"2023-01-22T10:00:12.027178","status":"completed"},"tags":[],"id":"87621c6e"}},{"cell_type":"code","source":"model = ViT(img_size = (512, 512),patch_size = (32, 32),n_channels = 1,d_model = 512,nhead = 8,dim_feedforward = 1024,blocks = 4,mlp_head_units = [256],n_classes = 1,).to(device)\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\ncriterion = nn.BCELoss()\n\ntrainer = create_supervised_trainer(model, optimizer, criterion, device=device)\n\nval_metrics = {\n    \"bce\": Loss(criterion),\n    \"accuracy\": Accuracy(output_transform=lambda out: (torch.round(out[0]), out[1]))\n}\nevaluator = create_supervised_evaluator(model, metrics=val_metrics, device=device)\n\nlog_interval = 10\nmax_epochs = 10\nglobal best_loss\nbest_loss = float('inf')\n\nRunningAverage(output_transform=lambda x: x).attach(trainer, 'loss')\n\n# FIX: Local iteration math so it never goes past the max batch count\n@trainer.on(Events.ITERATION_COMPLETED(every=log_interval))\ndef log_training_loss(trainer):\n    epoch_len = len(train_loader)\n    local_iteration = (trainer.state.iteration - 1) % epoch_len + 1\n    print(f\"Epoch[{trainer.state.epoch}] Iteration[{local_iteration}/{epoch_len}] | Loss: {trainer.state.output:.4f}\")\n\n# KEEP: Your exact validation loop block\n@trainer.on(Events.EPOCH_COMPLETED)\ndef log_validation_results(trainer):\n    global best_loss\n    print(\"\\n--> Running Validation...\")\n    evaluator.run(val_loader)\n    loss = evaluator.state.metrics['bce']\n    acc = evaluator.state.metrics['accuracy']\n    \n    if loss < best_loss:\n        best_loss = loss\n        print(f\"New best loss: {best_loss:.4f}! Saving checkpoint...\")\n        torch.save(model.state_dict(), 'best_model_vit.pt')\n        \n    print(f\"Validation Results - Epoch: {trainer.state.epoch} Accuracy: {acc:.4f} Avg loss: {loss:.4f}\\n\")\n\nprint(\"Starting fresh engine run with fast ViT...\")\noutput_state = trainer.run(train_loader, max_epochs=max_epochs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-10T11:28:19.67357Z","iopub.execute_input":"2026-07-10T11:28:19.674228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<h1>Research Task:</h1>\n\n(AB TAK KA TASK ITSELF WAS KINDA RESEARCH TASK LOL, but still)\n\n\nRead any research paper which involves usage of Vision Transformer. Ensure it satisfies the following conditions\n(Paper ko access karna from college ka website)\n\n1. Recent paper, published within last 5 years\n2. Not used for basic object detection and image segmentation\n\nAnd you have to explain\n\n1. Purpose of the papers\n2. Novelty of the papers\n3. Comparison with some other papers\n4. Methodology\n\nAlso, what improvement you feel can be done in the paper - either pre-processing POV, dataset handling, methodology optimization, more visually simple diagrams etc (Ek simple addition which u feel can make that paper more better, don't worry just ek simple exercise to make u think aur kn)","metadata":{}},{"cell_type":"markdown","source":"# **End of Task**\n\n> ©Synapse 2025 - 2026\n","metadata":{}}]}