{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"},{"sourceId":1062313,"sourceType":"datasetVersion","datasetId":589173},{"sourceId":1677248,"sourceType":"datasetVersion","datasetId":993628},{"sourceId":1695720,"sourceType":"datasetVersion","datasetId":1003815}],"dockerImageVersionId":30299,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!cp -r ../input/vittutorialillustrations/* ./ \n\n!pip install nb_black\n%load_ext nb_black","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-21T16:33:16.793817Z","iopub.execute_input":"2024-04-21T16:33:16.794448Z","iopub.status.idle":"2024-04-21T16:33:37.800802Z","shell.execute_reply.started":"2024-04-21T16:33:16.794372Z","shell.execute_reply":"2024-04-21T16:33:37.799569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Introduction\n\nThis notebook essentially has two parts as seen from the title. \n\n<a href=\"#Vision-Transformers:-A-gentle-introduction\">1. Vision Transformer: A gentle introduction</a> <br>\n<a href=\"#Vision-Transformer-Implementation-in-PyTorch\">2. Implementation in PyTorch</a>\n\n**I'll briefly try to explain the fundamental ideas behind Vision Transformers and how it works before getting into the implementation of ViT in PyTorch for this competition.**\n\nSo if you are only interested in code, you can feel free to skip straight ahead to the second section of this notebook. The implementation isn't a whole lot different, thanks to [rwightman/pytorch-image-models](https://github.com/rwightman/pytorch-image-models) library which contains all the model implementations including the pretrained weights for us to use. ","metadata":{}},{"cell_type":"markdown","source":"# <font size=4 color='blue'>If you find this notebook useful, leave an upvote, that motivates me to write more such notebooks.</font>","metadata":{}},{"cell_type":"markdown","source":"# Vision Transformers: A gentle introduction\n\nVision Transformers were first introduced in the paper [AN IMAGE IS WORTH 16X16 WORDS:\nTRANSFORMERS FOR IMAGE RECOGNITION AT SCALE](https://arxiv.org/pdf/2010.11929.pdf) by the Google Brain team in late October 2020. \n\nTo understand how ViT works, obviously you must have prior knowledge about how Transformers work and what problems it solved. I'll briefly introduce you to how transformers work before getting into the details of the topic at hand - ViT.\n\n![ViT-Illustration](vision-transformer.png)\n\nIf you are new to NLP and interested in learning more about the transformer models and get a fair intuition of how they actually work, I recomment checking out the fantastic blog posts of [Jay Allamar](https://jalammar.github.io/). (The image above was also inspired from one of his blog posts)","metadata":{}},{"cell_type":"markdown","source":"## Transformers: A brief overview\n\n> **If you already understand Transformers, feel free to skip ahead to the next section.**\n\nTransformer models well and truly revolutionized Natural Language Processing as we know. When they were first introduced they broke multiple NLP records and were pushing the then State of the Art. Now, they have become a de-facto standard for modern NLP tasks and they bring spectacular performance gains when compared to the previous generation of models like LSTMs and GRUs.  \n\nBy far the most important paper that transformed the NLP landscape is the [\"Attention is all you need\"](https://arxiv.org/pdf/1706.03762.pdf) paper. The transformer architecture was introduced in this paper. \n\n### **Motivations:**\n\nThe existing models at that time for sequence and NLP tasks mostly involved RNNs. **The problem with these networks were that they couldn't capture long term dependencies.** \n\nLSTMs and GRUs - variants of RNNs were capable of capturing the dependencies but it is also limited. \n\nSo, the main inspiration behind the transformer was to get rid of this recurrence and still end up capturing almost all the dependencies, to be precise global dependencies, yes the reference window of tranformers is full-range. This was achieved using a variant of attention mechanism called *self-attention* (multi-headed) which is very important for their success. One other advantage of Tranformer models are that they are highly parallelizable.  \n\n\n### Transformer Architecture\n**Note: The architecture diagrams are annotated with the corresponding step in the explanation.**\n\n![TranformerArchitecture](transformer-arch.png)\n\n- Transformer has two parts, the decoder which is on the left side on the above diagram and the encoder which is on the right. \n- Imagine we are doing machine translation for now.\n- The encoder takes the input data (sentence), and produces an intermediate representation of the input. \n- The decoder decodes this intermediate representation step by step and generates the output. The difference however is in how it is doing this. \n- Understanding the Encoder section is enough for ViT. \n\n> **Note: The explanations here are more about the intuition behind the architectures. For more mathematical details check out the respective research papers instead.**\n\n### Tranformers: Step by step overview\n**(1)** The input data first gets embedded into a vector. The embedding layer helps us grab a learned vector representation for each word.\n\n**(2)** In the next stage a positional encoding is injected into the input embeddings. This is because a transformer has no idea about the order of the sequence that is being passed as input - for example a sentence.\n\n**(3)** Now the multi-headed attention is where things get a little different. \n\n**Multi-headed-attention architecture:**\n![multi-headed-attn](multi-headed-attention.png)\n\n**(4)** Multi-Headed Attention consists of three learnable vectors. Query, Key and Value vectors. The motivation of this reportedly comes from information retrival where you search (query) and the search engine compares your query with a key and responds with a value. \n\n**(5)** The Q and K representations undergo a dot product matrix multiplication to produce a score matrix which represents how much a word has to attend to every other word. Higher score means more attention and vice-versa. \n\n**(6)** Then the Score matrix is scaled down according to the dimensions of the Q and K vectors. This is to ensure more stable gradients as multiplication can have exploding effects. \n\n(We'll discuss the mask part when we reach the decoder section)\n\n**(7)** Next the Score matrix is softmaxed to turn attention scores into probabilities. Obviously higher scores are heightened and lower scores are depressed. This ensures the model to be confident on which words to attend to. \n\n**(8)** Then the resultant matrix with probabilites is multiplied with the value vector. This will make the higher probaility scores the model has learned to be more important. The low scoring words will effectively drown out to become irrelevant. \n\n**(9)** Then, the concatenated output of QK and V vectors are fed into the Linear layer to process further. \n\n**(10)** Self-Attention is performed for each word in the sequence. Since one doesn't depend on the other a copy of the self attention module can be used to process everything simultaneously making this **multi-headed**. \n\n**(11)** Then the output value vectors are concatenated and added to the residual connection coming from the input layer and then the resultant respresentation is passed into a *LayerNorm* for normalization. (Residual connection help gradients flow through the network and LayernNorm helps reduce the training time by a small fraction and stabilize the network)\n\n**(12)** Further, the output is passed into a point-wise feed forward network to obtain an even richer representation.  \n\n**(13)** The outputs are again Layer-normed and residuals are added from the previous layer. \n\n\n**Note: This wraps up the encoder section and trust me this is enough to fully understand Vision Transformer. I'll be largely leaving it up to you to understand the decoder part as it is very similar to the encoding layer.**\n\n**(14)** The output from the encoder along with the inputs (if any) from the previous time steps/words are fed into the decoder where the outputs undergo masked-multi headed attention before being fed into the next attention layer along with the output from encoder. \n\n**(15)** Masked multi headed attention is necessary because the network shouldn't have any visibility into the words that are to come later in the sequence while decoding, to ensure there is no leak. This is done by masking the entries of words that come later in the series in the Score matrix. Current and previous words in the sequence are added with 1 and the future word scores are added with `-inf`. This ensures the future words in the series get drowned out into 0 when performing softmax to obtain the probabilities, while the rest are retained. \n\n**(16)** There are residual connections here as well, to improve the flow of gradients. Finally the output is sent to a Linear layer and softmaxed to obtain the outputs in probabilities. ","metadata":{}},{"cell_type":"markdown","source":"## How Vision Tranformers works?\n\nNow that we have covered transformers' internal working at a high level, we are finally ready to tackle Vision Tranformers. \n\nApplying Transformers on images was always going to be a challenge for the following reasons,\n- Unlike words/sentences/paragraphs, images contain much much more information in them basically in form of pixels. \n- It would be very hard, even with current hardware to attend to every other pixel in the image. \n- Instead, a popular alternative was to use localized attention. \n- In fact CNNs do something very similar through convolutions and the receptive field essentially grows bigger as we go deeper into the model's layers, but Tranformers were always going to be computationally more expensive than CNNs because of the' nature of Transformers. And of course, we know how incredibly much CNNs have contributed to the current advancements in Computer Vision.\n\nGoogle researchers have proposed something different in their paper than can possibly be the next big step in Computer Vision. They show that the reliance on CNNs may not be necessary anymore. So, let's dive right in and explore more about Vision Transformers.\n\n\n### Vision Transformer Architecture\n\n![vit-architecture](vit-arch.png)\n\n**(1)** They are only using the Encoder part of the transformer but the difference is in how they are feeding the images into the network. \n\n**(2)** They are breaking down the image into fixed size patches. So one of these patches can be of dimension 16x16 or 32x32 as proposed in the paper. More patches means more simpler it is to train these networks as the patches themselves get smaller. Hence we have that in the title - \"An Image is worth 16x16 words\". \n\n**(3)** The patches are then unrolled (flattened) and sent for further processing into the network.\n\n**(4)** Unlike NNs here the model has no idea whatsoever about the position of the samples in the sequence, here each sample is a patch from the input image. So the image is fed **along with a positional embedding vector** and into the encoder. One thing to note here is the positional embeddings are also learnable so you don't actually feed hard-coded vectors w.r.t to their positions. \n\n**(5)** There is also a special token at the start just like BERT.\n\n**(6)** So each image patch is first unrolled (flattened) into a big vector and gets multiplied with an embedding matrix which is also learnable, creating embedded patches. And these embedded patches are combined with the positional embedding vector and that gets fed into the Tranformer. \n\n> **Note: From here everything is just the same as a standard transformer**\n\n**(7)** With the only difference being, instead of a decoder the output from the encoder is passed directly into a Feed Forward Neural Network to obtain the classification output. \n\n### Things to note:\n- The paper ALMOST completely neglects Convolutions. \n- They are however using a couple of variants of ViT in which Convolutional embeddings of image patches are used. But that doesn't seem to impact performance much.  \n- Vision Transformers, at the time of writing this are topping Image Classification benchmarks on ImageNet. \n\n<img src=\"benchmarks-chart.png\" width=\"700\">\n<!-- ![BenchmarksChart](benchmarks-charpng) -->\n\n- There are a lot more interesting things in this paper but the one thing that stands out for me and potentially shows the power of transformers over CNNs is illustrated in the image below which shows the attention distance with respect to the layers. \n\n\n<img src=\"attn-distance.png\" width=\"300\" height=\"300\">\n<br>\n\n- The graph above suggests that the Transformers are already capable of paying attention to  regions that are far apart right from the starting layers of the network which is a pretty significant gain the Transformers bring over CNNs which has a finite receptive field at the start.","metadata":{}},{"cell_type":"markdown","source":"# <font size=4 color='blue'>If you find this notebook useful, leave an upvote, that motivates me to write more such notebooks.</font>","metadata":{}},{"cell_type":"markdown","source":"## Vision Transformer Implementation in PyTorch\nNow that you understand vision transformers, let's build a baseline model for [this competition](https://www.kaggle.com/c/cassava-leaf-disease-classification)\n\nFirst lets install torch-xla to be able to use the TPU and torch-image-models (timm). ","metadata":{}},{"cell_type":"code","source":"!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-21T17:02:29.037859Z","iopub.execute_input":"2024-04-21T17:02:29.038232Z","iopub.status.idle":"2024-04-21T17:02:30.280495Z","shell.execute_reply.started":"2024-04-21T17:02:29.038175Z","shell.execute_reply":"2024-04-21T17:02:30.279820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python pytorch-xla-env-setup.py --version 1.7","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:02:35.678091Z","iopub.execute_input":"2024-04-21T17:02:35.678462Z","iopub.status.idle":"2024-04-21T17:02:36.887939Z","shell.execute_reply.started":"2024-04-21T17:02:35.678427Z","shell.execute_reply":"2024-04-21T17:02:36.887098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:01:32.436075Z","iopub.execute_input":"2024-04-21T17:01:32.436751Z","iopub.status.idle":"2024-04-21T17:01:39.028266Z","shell.execute_reply.started":"2024-04-21T17:01:32.436713Z","shell.execute_reply":"2024-04-21T17:01:39.027216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nplt.style.use(\"ggplot\")\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\n\nimport timm\n\nimport gc\nimport os\nimport time\nimport random\nfrom datetime import datetime\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom sklearn import model_selection, metrics","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:01:42.678693Z","iopub.execute_input":"2024-04-21T17:01:42.679005Z","iopub.status.idle":"2024-04-21T17:01:43.817236Z","shell.execute_reply.started":"2024-04-21T17:01:42.678974Z","shell.execute_reply":"2024-04-21T17:01:43.816559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install --upgrade torch-xla","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:03:37.313892Z","iopub.execute_input":"2024-04-21T17:03:37.314631Z","iopub.status.idle":"2024-04-21T17:03:43.980546Z","shell.execute_reply.started":"2024-04-21T17:03:37.314594Z","shell.execute_reply":"2024-04-21T17:03:43.979695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.distributed.parallel_loader as pl\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:03:43.982402Z","iopub.execute_input":"2024-04-21T17:03:43.982643Z","iopub.status.idle":"2024-04-21T17:03:44.168040Z","shell.execute_reply.started":"2024-04-21T17:03:43.982616Z","shell.execute_reply":"2024-04-21T17:03:44.167008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport shutil\n\ndef split_folder(folder_path, train_ratio=0.8):\n    \"\"\"\n    This function splits a folder into train and test folders based on a specified ratio.\n\n    Args:\n        folder_path (str): The path to the folder containing the files.\n        train_ratio (float, optional): The proportion of files to be placed in the train folder. Defaults to 0.8.\n    \"\"\"\n    # Get all files in the folder\n    files = [f for f in os.listdir(folder_path) if os.path.isfile(os.path.join(folder_path, f))]\n\n    # Shuffle the list of files for random splitting\n    random.shuffle(files)\n\n    # Split the list based on train ratio\n    train_count = int(len(files) * train_ratio)\n    train_files = files[:train_count]\n    test_files = files[train_count:]\n\n    # Create train and test folders if they don't exist\n    train_folder = os.path.join('/kaggle/working/', \"train\")\n    test_folder = os.path.join('/kaggle/working/', \"test\")\n    os.makedirs(train_folder, exist_ok=True)\n    os.makedirs(test_folder, exist_ok=True)\n\n    # Copy files to respective folders\n    for file in train_files:\n        shutil.copy(os.path.join(folder_path, file), os.path.join(train_folder, file))\n    for file in test_files:\n        shutil.copy(os.path.join(folder_path, file), os.path.join(test_folder, file))\n\n    print(f\"Split complete! Train files: {len(train_files)}, Test files: {len(test_files)}\")\n\n# Set the folder path (assuming you've already set it)\nfolder_path = \"/kaggle/input/crop-and-weed-detection-data-with-bounding-boxes/agri_data/data\"\n\n# Split the folder with a train ratio of 80% (optional, adjust as needed)\nsplit_folder(folder_path, train_ratio=0.8)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T16:59:15.068945Z","iopub.execute_input":"2024-04-21T16:59:15.069297Z","iopub.status.idle":"2024-04-21T16:59:19.165044Z","shell.execute_reply.started":"2024-04-21T16:59:15.069265Z","shell.execute_reply":"2024-04-21T16:59:19.164423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For parallelization in TPUs\nos.environ[\"XLA_USE_BF16\"] = \"1\"\nos.environ[\"XLA_TENSOR_ALLOCATOR_MAXSIZE\"] = \"100000000\"","metadata":{"execution":{"iopub.status.busy":"2024-04-21T16:59:19.166341Z","iopub.execute_input":"2024-04-21T16:59:19.166538Z","iopub.status.idle":"2024-04-21T16:59:19.170016Z","shell.execute_reply.started":"2024-04-21T16:59:19.166514Z","shell.execute_reply":"2024-04-21T16:59:19.169546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    \"\"\"\n    Seeds basic parameters for reproductibility of results\n    \n    Arguments:\n        seed {int} -- Number of the seed\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(1001)","metadata":{"execution":{"iopub.status.busy":"2024-04-21T16:59:19.170819Z","iopub.execute_input":"2024-04-21T16:59:19.171001Z","iopub.status.idle":"2024-04-21T16:59:19.185566Z","shell.execute_reply.started":"2024-04-21T16:59:19.170977Z","shell.execute_reply":"2024-04-21T16:59:19.185050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# general global variables\nDATA_PATH = \"/kaggle/input/crop-and-weed-detection-data-with-bounding-boxes/agri_data/data\"\nTRAIN_PATH = \"/kaggle/working/train\"\nTEST_PATH = \"/kaggle/working/test\"\nMODEL_PATH = (\n    \"../input/vit-base-models-pretrained-pytorch/jx_vit_base_p16_224-80ecf9dd.pth\"\n)\n\n# model specific global variables\nIMG_SIZE = 224\nBATCH_SIZE = 16\nLR = 2e-05\nGAMMA = 0.7\nN_EPOCHS = 10","metadata":{"execution":{"iopub.status.busy":"2024-04-21T16:59:35.627014Z","iopub.execute_input":"2024-04-21T16:59:35.627346Z","iopub.status.idle":"2024-04-21T16:59:35.631528Z","shell.execute_reply.started":"2024-04-21T16:59:35.627311Z","shell.execute_reply":"2024-04-21T16:59:35.631002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport csv\n\ndef create_csv_with_class_and_filename(directory_path, csv_filename):\n    \"\"\"\n    Create a CSV file with class labels and corresponding filenames based on the contents of text files in a directory.\n\n    Args:\n        directory_path (str): The path to the directory containing the text files.\n        csv_filename (str): The filename for the CSV file.\n    \"\"\"\n    data = []  # List to store class labels and filenames\n\n    # Iterate through all files in the directory\n    for filename in os.listdir(directory_path):\n        if filename.endswith(\".txt\"):\n            file_path = os.path.join(directory_path, filename)\n            with open(file_path, 'r') as file:\n                for line in file:\n                    # Split the line and extract the class label and filename\n                    class_label = line.split()[0]\n                    data.append([class_label, directory_path + \"/\" + filename[:-3]+'jpeg'])\n\n    # Write data to CSV file\n    with open(csv_filename, mode='w', newline='') as file:\n        writer = csv.writer(file)\n        writer.writerow(['label', 'Filename'])\n        writer.writerows(data)\n\n# Example usage:\ndirectory_path = '/kaggle/input/crop-and-weed-detection-data-with-bounding-boxes/agri_data/data'\ncsv_filename = 'classes.csv'\ncreate_csv_with_class_and_filename(directory_path, csv_filename)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:35:21.042934Z","iopub.execute_input":"2024-04-21T17:35:21.043703Z","iopub.status.idle":"2024-04-21T17:35:21.658748Z","shell.execute_reply.started":"2024-04-21T17:35:21.043671Z","shell.execute_reply":"2024-04-21T17:35:21.657921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(\"classes.csv\"))\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:35:23.097504Z","iopub.execute_input":"2024-04-21T17:35:23.098211Z","iopub.status.idle":"2024-04-21T17:35:23.110245Z","shell.execute_reply.started":"2024-04-21T17:35:23.098178Z","shell.execute_reply":"2024-04-21T17:35:23.109650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.info()","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:35:23.899578Z","iopub.execute_input":"2024-04-21T17:35:23.899772Z","iopub.status.idle":"2024-04-21T17:35:23.909040Z","shell.execute_reply.started":"2024-04-21T17:35:23.899751Z","shell.execute_reply":"2024-04-21T17:35:23.908482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.label.value_counts().plot(kind=\"bar\")","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:35:24.828563Z","iopub.execute_input":"2024-04-21T17:35:24.829129Z","iopub.status.idle":"2024-04-21T17:35:24.935422Z","shell.execute_reply.started":"2024-04-21T17:35:24.829083Z","shell.execute_reply":"2024-04-21T17:35:24.934874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, valid_df = model_selection.train_test_split(\n    df, test_size=0.1, random_state=42, stratify=df['label'].values\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:35:25.814772Z","iopub.execute_input":"2024-04-21T17:35:25.815503Z","iopub.status.idle":"2024-04-21T17:35:25.822017Z","shell.execute_reply.started":"2024-04-21T17:35:25.815472Z","shell.execute_reply":"2024-04-21T17:35:25.821428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:35:26.970993Z","iopub.execute_input":"2024-04-21T17:35:26.971815Z","iopub.status.idle":"2024-04-21T17:35:26.983864Z","shell.execute_reply.started":"2024-04-21T17:35:26.971780Z","shell.execute_reply":"2024-04-21T17:35:26.983091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class MyDataset(torch.utils.data.Dataset):\n#     \"\"\"\n#     Helper Class to create the pytorch dataset\n#     \"\"\"\n\n#     def __init__(self, df, data_path=DATA_PATH, mode=\"train\", transforms=None):\n#         super().__init__()\n#         self.df_data = df.values\n#         self.data_path = data_path\n#         self.transforms = transforms\n#         self.mode = mode\n#         self.data_dir = \"train_images\" if mode == \"train\" else \"test_images\"\n\n#     def __len__(self):\n#         return len(self.df_data)\n\n#     def __getitem__(self, index):\n#         img_name, label = self.df_data[index]\n#         img_path = os.path.join(self.data_path, self.data_dir, img_name)\n#         img = Image.open(img_path).convert(\"RGB\")\n\n#         if self.transforms is not None:\n#             image = self.transforms(img)\n\n#         return image, label\n\nclass MyDataset(torch.utils.data.Dataset):\n    \"\"\"\n    Helper Class to create the pytorch dataset\n    \"\"\"\n\n    def __init__(self, df, data_path=DATA_PATH, mode=\"train\", transforms=None):\n        super().__init__()\n        self.df_data = df.values\n        self.data_path = data_path\n        self.transforms = transforms\n        self.mode = mode\n        self.data_dir = \"train_images\" if mode == \"train\" else \"test_images\"\n\n    def __len__(self):\n        return len(self.df_data)\n\n    def __getitem__(self, index):\n        label, img_name = self.df_data[index]\n#         print(\"img_name:==\" + img_name)\n        img_name = str(img_name)  # Convert img_name to string\n        img_path = os.path.join(self.data_path, self.data_dir, img_name)\n        img = Image.open(img_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            img = self.transforms(img)\n\n        return img, label\n\n    ","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:08.675277Z","iopub.execute_input":"2024-04-21T17:36:08.675575Z","iopub.status.idle":"2024-04-21T17:36:08.684640Z","shell.execute_reply.started":"2024-04-21T17:36:08.675544Z","shell.execute_reply":"2024-04-21T17:36:08.683971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# create image augmentations\ntransforms_train = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.RandomHorizontalFlip(p=0.3),\n        transforms.RandomVerticalFlip(p=0.3),\n        transforms.RandomResizedCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]\n)\n\ntransforms_valid = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:09.869339Z","iopub.execute_input":"2024-04-21T17:36:09.869596Z","iopub.status.idle":"2024-04-21T17:36:09.875531Z","shell.execute_reply.started":"2024-04-21T17:36:09.869569Z","shell.execute_reply":"2024-04-21T17:36:09.874943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Available Vision Transformer Models: \")\ntimm.list_models(\"vit*\")","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:10.281991Z","iopub.execute_input":"2024-04-21T17:36:10.282260Z","iopub.status.idle":"2024-04-21T17:36:10.289970Z","shell.execute_reply.started":"2024-04-21T17:36:10.282233Z","shell.execute_reply":"2024-04-21T17:36:10.289450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Uncomment for running first time\nclass ViTBase16(nn.Module):\n    def __init__(self, n_classes, pretrained=False):\n\n        super(ViTBase16, self).__init__()\n\n        self.model = timm.create_model(\"vit_base_patch16_224\", pretrained=False)\n        if pretrained:\n            self.model.load_state_dict(torch.load(MODEL_PATH))\n\n        self.model.head = nn.Linear(self.model.head.in_features, n_classes)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\n    def train_one_epoch(self, train_loader, criterion, optimizer, device):\n        # keep track of training loss\n        epoch_loss = 0.0\n        epoch_accuracy = 0.0\n\n        ###################\n        # train the model #\n        ###################\n        self.model.train()\n        for i, (data, target) in enumerate(train_loader):\n            # move tensors to GPU if CUDA is available\n            if device.type == \"cuda\":\n                data, target = data.cuda(), target.cuda()\n            elif device.type == \"xla\":\n                data = data.to(device, dtype=torch.float32)\n                target = target.to(device, dtype=torch.int64)\n\n            # clear the gradients of all optimized variables\n            optimizer.zero_grad()\n            # forward pass: compute predicted outputs by passing inputs to the model\n            output = self.forward(data)\n            # calculate the batch loss\n            loss = criterion(output, target)\n            # backward pass: compute gradient of the loss with respect to model parameters\n            loss.backward()\n            # Calculate Accuracy\n            accuracy = (output.argmax(dim=1) == target).float().mean()\n            # update training loss and accuracy\n            epoch_loss += loss\n            epoch_accuracy += accuracy\n\n            # perform a single optimization step (parameter update)\n            if device.type == \"xla\":\n                xm.optimizer_step(optimizer)\n\n                if i % 20 == 0:\n                    xm.master_print(f\"\\tBATCH {i+1}/{len(train_loader)} - LOSS: {loss}\")\n\n            else:\n                optimizer.step()\n\n        return epoch_loss / len(train_loader), epoch_accuracy / len(train_loader)\n\n    def validate_one_epoch(self, valid_loader, criterion, device):\n        # keep track of validation loss\n        valid_loss = 0.0\n        valid_accuracy = 0.0\n\n        ######################\n        # validate the model #\n        ######################\n        self.model.eval()\n        for data, target in valid_loader:\n            # move tensors to GPU if CUDA is available\n            if device.type == \"cuda\":\n                data, target = data.cuda(), target.cuda()\n            elif device.type == \"xla\":\n                data = data.to(device, dtype=torch.float32)\n                target = target.to(device, dtype=torch.int64)\n\n            with torch.no_grad():\n                # forward pass: compute predicted outputs by passing inputs to the model\n                output = self.model(data)\n                # calculate the batch loss\n                loss = criterion(output, target)\n                # Calculate Accuracy\n                accuracy = (output.argmax(dim=1) == target).float().mean()\n                # update average validation loss and accuracy\n                valid_loss += loss\n                valid_accuracy += accuracy\n\n        return valid_loss / len(valid_loader), valid_accuracy / len(valid_loader)","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:10.645784Z","iopub.execute_input":"2024-04-21T17:36:10.646051Z","iopub.status.idle":"2024-04-21T17:36:10.650889Z","shell.execute_reply.started":"2024-04-21T17:36:10.646022Z","shell.execute_reply":"2024-04-21T17:36:10.650398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"***Our Custom Vision Transformer***","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torch.utils.data import DataLoader, RandomSampler\nfrom datetime import datetime\n\nclass CustomViT(nn.Module):\n    def __init__(self, image_size, patch_size, num_classes, embedding_dim, pretrained=False):\n        super(CustomViT, self).__init__()\n\n        self.patch_embed = nn.Conv2d(in_channels=3, out_channels=embedding_dim, kernel_size=patch_size, stride=patch_size)\n        num_patches = (image_size // patch_size) ** 2\n        self.cls_token = nn.Parameter(torch.randn(1, 1, embedding_dim))\n        self.pos_embed = nn.Parameter(torch.randn(1, num_patches + 1, embedding_dim))\n        self.pos_dropout = nn.Dropout(p=0.1)\n\n        self.transformer_layers = nn.TransformerEncoderLayer(embedding_dim, nhead=8)\n        self.transformer = nn.TransformerEncoder(self.transformer_layers, num_layers=6)\n\n        self.fc = nn.Linear(embedding_dim, num_classes)\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        x = self.patch_embed(x)\n        x = x.flatten(2).transpose(1, 2)  # (B, N, C)\n        cls_tokens = self.cls_token.expand(B, -1, -1)\n        x = torch.cat((cls_tokens, x), dim=1)\n        x = x + self.pos_embed\n        x = self.pos_dropout(x)\n\n        x = self.transformer(x)\n\n        cls_output = x[:, 0, :]  # Extract class token output\n        output = self.fc(cls_output)\n        return output\n\ndef _run():\n    train_dataset = MyDataset(train_df, transforms=transforms_train)\n    valid_dataset = MyDataset(valid_df, transforms=transforms_valid)\n\n    train_loader = DataLoader(\n        dataset=train_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=RandomSampler(train_dataset),\n        drop_last=True,\n        num_workers=8,\n    )\n\n    valid_loader = DataLoader(\n        dataset=valid_dataset,\n        batch_size=BATCH_SIZE,\n        drop_last=True,\n        num_workers=8,\n    )\n\n    model = CustomViT(image_size=224, patch_size=16, num_classes=2, embedding_dim=768, pretrained=True)  # Adjust parameters as needed\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    lr = LR\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n\n    print(f\"INITIALIZING TRAINING ON {'GPU' if torch.cuda.is_available() else 'CPU'}\")\n    start_time = datetime.now()\n    print(f\"Start Time: {start_time}\")\n\n    for epoch in range(N_EPOCHS):\n        model.train()\n        epoch_loss = 0.0\n        for i, (data, target) in enumerate(train_loader):\n            data, target = data.to(device), target.to(device)\n\n            optimizer.zero_grad()\n            output = model(data)\n            loss = criterion(output, target)\n            loss.backward()\n            optimizer.step()\n\n            epoch_loss += loss.item()\n\n        print(f\"Epoch [{epoch+1}/{N_EPOCHS}], Train Loss: {epoch_loss / len(train_loader)}\")\n\n    print(f\"Execution time: {datetime.now() - start_time}\")\n\n    print(\"Saving Model\")\n    torch.save(\n        model.state_dict(), f'model_5e_{datetime.now().strftime(\"%Y%m%d-%H%M\")}.pth'\n    )\n\n_run()\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def fit_tpu(\n#     model, epochs, device, criterion, optimizer, train_loader, valid_loader=None\n# ):\n\n#     valid_loss_min = np.Inf  # track change in validation loss\n\n#     # keeping track of losses as it happen\n#     train_losses = []\n#     valid_losses = []\n#     train_accs = []\n#     valid_accs = []\n\n#     for epoch in range(1, epochs + 1):\n#         gc.collect()\n#         para_train_loader = pl.ParallelLoader(train_loader, [device])\n\n#         xm.master_print(f\"{'='*50}\")\n#         xm.master_print(f\"EPOCH {epoch} - TRAINING...\")\n#         train_loss, train_acc = model.train_one_epoch(\n#             para_train_loader.per_device_loader(device), criterion, optimizer, device\n#         )\n#         xm.master_print(\n#             f\"\\n\\t[TRAIN] EPOCH {epoch} - LOSS: {train_loss}, ACCURACY: {train_acc}\\n\"\n#         )\n#         train_losses.append(train_loss)\n#         train_accs.append(train_acc)\n#         gc.collect()\n\n#         if valid_loader is not None:\n#             gc.collect()\n#             para_valid_loader = pl.ParallelLoader(valid_loader, [device])\n#             xm.master_print(f\"EPOCH {epoch} - VALIDATING...\")\n#             valid_loss, valid_acc = model.validate_one_epoch(\n#                 para_valid_loader.per_device_loader(device), criterion, device\n#             )\n#             xm.master_print(f\"\\t[VALID] LOSS: {valid_loss}, ACCURACY: {valid_acc}\\n\")\n#             valid_losses.append(valid_loss)\n#             valid_accs.append(valid_acc)\n#             gc.collect()\n\n#             # save model if validation loss has decreased\n#             if valid_loss <= valid_loss_min and epoch != 1:\n#                 xm.master_print(\n#                     \"Validation loss decreased ({:.4f} --> {:.4f}).  Saving model ...\".format(\n#                         valid_loss_min, valid_loss\n#                     )\n#                 )\n#             #                 xm.save(model.state_dict(), 'best_model.pth')\n\n#             valid_loss_min = valid_loss\n\n#     return {\n#         \"train_loss\": train_losses,\n#         \"valid_losses\": valid_losses,\n#         \"train_acc\": train_accs,\n#         \"valid_acc\": valid_accs,\n#     }","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:11.033907Z","iopub.execute_input":"2024-04-21T17:36:11.034159Z","iopub.status.idle":"2024-04-21T17:36:11.038485Z","shell.execute_reply.started":"2024-04-21T17:36:11.034133Z","shell.execute_reply":"2024-04-21T17:36:11.037948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = ViTBase16(n_classes=2, pretrained=True)   \n# # Comment this after running this once. Or else it will take a lot of time to excute again","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:11.391429Z","iopub.execute_input":"2024-04-21T17:36:11.391661Z","iopub.status.idle":"2024-04-21T17:36:11.394565Z","shell.execute_reply.started":"2024-04-21T17:36:11.391636Z","shell.execute_reply":"2024-04-21T17:36:11.394007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# print(torch.__version__)\n\n# import torch_xla\n# print(torch_xla.__version__)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:11.733024Z","iopub.execute_input":"2024-04-21T17:36:11.733304Z","iopub.status.idle":"2024-04-21T17:36:11.737009Z","shell.execute_reply.started":"2024-04-21T17:36:11.733270Z","shell.execute_reply":"2024-04-21T17:36:11.736147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch_xla.core.xla_model as xm\n\n# # Now you can use xm functions\n# device = xm.xla_device()\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:36:12.053494Z","iopub.execute_input":"2024-04-21T17:36:12.053744Z","iopub.status.idle":"2024-04-21T17:36:12.056918Z","shell.execute_reply.started":"2024-04-21T17:36:12.053718Z","shell.execute_reply":"2024-04-21T17:36:12.056301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fit_cpu(model, epochs, device, criterion, optimizer, train_loader, valid_loader):\n    for epoch in range(1, epochs + 1):\n        model.train()\n        print('hello1')\n        for data, targets in train_loader:\n            print(\"hello2\")\n            data, targets = data.to(device), targets.to(device)\n#             print(data, targets)\n            optimizer.zero_grad()\n            outputs = model(data)\n            loss = criterion(outputs, targets)\n            loss.backward()\n            optimizer.step()\n\n        model.eval()\n        valid_loss = 0.0\n        correct = 0\n        total = 0\n        with torch.no_grad():\n            for data, targets in valid_loader:\n                data, targets = data.to(device), targets.to(device)\n                outputs = model(data)\n                loss = criterion(outputs, targets)\n                valid_loss += loss.item()\n                _, predicted = torch.max(outputs.data, 1)\n                total += targets.size(0)\n                correct += (predicted == targets).sum().item()\n\n        avg_valid_loss = valid_loss / len(valid_loader)\n        valid_accuracy = correct / total\n\n        print(f\"Epoch: {epoch}, Validation Loss: {avg_valid_loss:.4f}, Validation Accuracy: {valid_accuracy:.4f}\")\n\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:37:46.724903Z","iopub.execute_input":"2024-04-21T17:37:46.725730Z","iopub.status.idle":"2024-04-21T17:37:46.732940Z","shell.execute_reply.started":"2024-04-21T17:37:46.725692Z","shell.execute_reply":"2024-04-21T17:37:46.732375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def _run():\n#     # Initialize distributed training if available\n#     if torch.cuda.is_available():\n#         torch.cuda.set_device(torch.cuda.current_device())\n#         torch.distributed.init_process_group(backend='nccl', init_method='env://')\n    \n#     # Create datasets and dataloaders\n#     train_dataset = MyDataset(train_df, transforms=transforms_train)\n#     valid_dataset = MyDataset(valid_df, transforms=transforms_valid)\n\n#     train_sampler = torch.utils.data.distributed.DistributedSampler(\n#         train_dataset,\n#         shuffle=True if torch.distributed.get_rank() == 0 else False  # Shuffle only on the first rank\n#     )\n\n#     valid_sampler = torch.utils.data.distributed.DistributedSampler(\n#         valid_dataset,\n#         shuffle=False\n#     )\n\n#     train_loader = torch.utils.data.DataLoader(\n#         dataset=train_dataset,\n#         batch_size=BATCH_SIZE,\n#         sampler=train_sampler,\n#         drop_last=True,\n#         num_workers=8\n#     )\n\n#     valid_loader = torch.utils.data.DataLoader(\n#         dataset=valid_dataset,\n#         batch_size=BATCH_SIZE,\n#         sampler=valid_sampler,\n#         drop_last=True,\n#         num_workers=8\n#     )\n\n#     # Initialize model, criterion, and optimizer\n#     device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n#     model.to(device)\n#     criterion = nn.CrossEntropyLoss()\n#     optimizer = torch.optim.Adam(model.parameters(), lr=LR)\n\n#     # Start training\n#     start_time = datetime.now()\n#     print(f\"Start Time: {start_time}\")\n\n#     for epoch in range(N_EPOCHS):\n#         train_sampler.set_epoch(epoch)\n        \n#         model.train()\n#         train_loss = 0.0\n#         for inputs, labels in train_loader:\n#             inputs, labels = inputs.to(device), labels.to(device)\n#             optimizer.zero_grad()\n#             outputs = model(inputs)\n#             loss = criterion(outputs, labels)\n#             loss.backward()\n#             optimizer.step()\n#             train_loss += loss.item()\n\n#         # Validation\n#         model.eval()\n#         valid_loss = 0.0\n#         with torch.no_grad():\n#             for inputs, labels in valid_loader:\n#                 inputs, labels = inputs.to(device), labels.to(device)\n#                 outputs = model(inputs)\n#                 loss = criterion(outputs, labels)\n#                 valid_loss += loss.item()\n\n#         # Print training progress\n#         print(f\"Epoch [{epoch+1}/{N_EPOCHS}], \"\n#               f\"Train Loss: {train_loss/len(train_loader):.4f}, \"\n#               f\"Valid Loss: {valid_loss/len(valid_loader):.4f}\")\n\n#     print(f\"Execution time: {datetime.now() - start_time}\")\n\n#     # Save model\n#     print(\"Saving Model\")\n#     torch.save(model.state_dict(), f'model_5e_{datetime.now().strftime(\"%Y%m%d-%H%M\")}.pth')\n\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import DataLoader, RandomSampler\nfrom datetime import datetime\n\ndef _run():\n    train_dataset = MyDataset(train_df, transforms=transforms_train)\n    valid_dataset = MyDataset(valid_df, transforms=transforms_valid)\n\n    train_loader = DataLoader(\n        dataset=train_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=RandomSampler(train_dataset),\n        drop_last=True,\n        num_workers=8,\n    )\n\n    valid_loader = DataLoader(\n        dataset=valid_dataset,\n        batch_size=BATCH_SIZE,\n        drop_last=True,\n        num_workers=8,\n    )\n\n    criterion = nn.CrossEntropyLoss()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n    \n    lr = LR\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n\n    print(f\"INITIALIZING TRAINING ON {'GPU' if torch.cuda.is_available() else 'CPU'}\")\n    start_time = datetime.now()\n    print(f\"Start Time: {start_time}\")\n\n    logs = fit_cpu(\n        model=model,\n        epochs=N_EPOCHS,\n        device=device,\n        criterion=criterion,\n        optimizer=optimizer,\n        train_loader=train_loader,\n        valid_loader=valid_loader,\n    )\n\n\n    print(f\"Execution time: {datetime.now() - start_time}\")\n\n    print(\"Saving Model\")\n    torch.save(\n        model.state_dict(), f'model_5e_{datetime.now().strftime(\"%Y%m%d-%H%M\")}.pth'\n    )\n\n_run()\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T17:37:46.983027Z","iopub.execute_input":"2024-04-21T17:37:46.983331Z","iopub.status.idle":"2024-04-21T18:26:06.506698Z","shell.execute_reply.started":"2024-04-21T17:37:46.983300Z","shell.execute_reply":"2024-04-21T18:26:06.505664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Start training processes\n# def _mp_fn(rank, flags):\n#     torch.set_default_tensor_type(\"torch.FloatTensor\")\n#     a = _run()\n\n\n# # _run()\n# FLAGS = {}\n# xmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method=\"fork\")\n\n_run()","metadata":{"execution":{"iopub.status.busy":"2024-04-21T18:26:06.507715Z","iopub.status.idle":"2024-04-21T18:26:06.507966Z","shell.execute_reply.started":"2024-04-21T18:26:06.507837Z","shell.execute_reply":"2024-04-21T18:26:06.507851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Validation loss and accuracy data (replace these with your actual values)\nepochs = range(1, 10)\nvalidation_loss = [0.1311, 0.1483, 0.1310, 0.1269, 0.1261, 0.1424, 0.1355, 0.1390, 0.1441]\nvalidation_accuracy = [0.9519, 0.9519, 0.9663, 0.9663, 0.9663, 0.9663, 0.9663, 0.9615, 0.9663]\n\n# Plotting validation loss\nplt.figure(figsize=(10, 5))\nplt.plot(epochs, validation_loss, label='Validation Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Validation Loss')\nplt.title('Validation Loss over Epochs')\nplt.legend()\nplt.grid(True)\nplt.show()\n\n# Plotting validation accuracy\nplt.figure(figsize=(10, 5))\nplt.plot(epochs, validation_accuracy, label='Validation Accuracy')\nplt.xlabel('Epoch')\nplt.ylabel('Validation Accuracy')\nplt.title('Validation Accuracy over Epochs')\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-04-21T18:30:53.556262Z","iopub.execute_input":"2024-04-21T18:30:53.556541Z","iopub.status.idle":"2024-04-21T18:30:53.995827Z","shell.execute_reply.started":"2024-04-21T18:30:53.556513Z","shell.execute_reply":"2024-04-21T18:30:53.995264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _run():\n    # Initialize distributed training if available\n    if torch.cuda.is_available():\n        torch.cuda.set_device(torch.cuda.current_device())\n        torch.distributed.init_process_group(backend='nccl', init_method='env://')\n    \n    # Create datasets and dataloaders\n    train_dataset = MyDataset(train_df, transforms=transforms_train)\n    valid_dataset = MyDataset(valid_df, transforms=transforms_valid)\n\n    # Initialize process group\n    if torch.distributed.is_initialized():\n        train_sampler = torch.utils.data.distributed.DistributedSampler(\n            train_dataset,\n            shuffle=True if torch.distributed.get_rank() == 0 else False  # Shuffle only on the first rank\n        )\n\n        valid_sampler = torch.utils.data.distributed.DistributedSampler(\n            valid_dataset,\n            shuffle=False\n        )\n\n        train_loader = torch.utils.data.DataLoader(\n            dataset=train_dataset,\n            batch_size=BATCH_SIZE,\n            sampler=train_sampler,\n            drop_last=True,\n            num_workers=8\n        )\n\n        valid_loader = torch.utils.data.DataLoader(\n            dataset=valid_dataset,\n            batch_size=BATCH_SIZE,\n            sampler=valid_sampler,\n            drop_last=True,\n            num_workers=8\n        )\n\n        # Initialize model, criterion, and optimizer\n        device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        model.to(device)\n        criterion = nn.CrossEntropyLoss()\n        optimizer = torch.optim.Adam(model.parameters(), lr=LR)\n\n        # Start training\n        start_time = datetime.now()\n        print(f\"Start Time: {start_time}\")\n\n        for epoch in range(N_EPOCHS):\n            train_sampler.set_epoch(epoch)\n            \n            model.train()\n            train_loss = 0.0\n            for inputs, labels in train_loader:\n                inputs, labels = inputs.to(device), labels.to(device)\n                optimizer.zero_grad()\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n                train_loss += loss.item()\n\n            # Validation\n            model.eval()\n            valid_loss = 0.0\n            with torch.no_grad():\n                for inputs, labels in valid_loader:\n                    inputs, labels = inputs.to(device), labels.to(device)\n                    outputs = model(inputs)\n                    loss = criterion(outputs, labels)\n                    valid_loss += loss.item()\n\n            # Print training progress\n            print(f\"Epoch [{epoch+1}/{N_EPOCHS}], \"\n                  f\"Train Loss: {train_loss/len(train_loader):.4f}, \"\n                  f\"Valid Loss: {valid_loss/len(valid_loader):.4f}\")\n\n        print(f\"Execution time: {datetime.now() - start_time}\")\n\n        # Save model\n        print(\"Saving Model\")\n        torch.save(model.state_dict(), f'model_5e_{datetime.now().strftime(\"%Y%m%d-%H%M\")}.pth')\n\n    else:\n        print(\"Failed to initialize distributed training.\")\n_run()","metadata":{"execution":{"iopub.status.busy":"2024-04-21T18:30:42.848393Z","iopub.execute_input":"2024-04-21T18:30:42.849244Z","iopub.status.idle":"2024-04-21T18:30:42.863495Z","shell.execute_reply.started":"2024-04-21T18:30:42.849204Z","shell.execute_reply":"2024-04-21T18:30:42.862826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}