{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30635,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install einops\n!pip install tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-29T14:55:38.161690Z","iopub.execute_input":"2024-01-29T14:55:38.162560Z","iopub.status.idle":"2024-01-29T14:56:03.807582Z","shell.execute_reply.started":"2024-01-29T14:55:38.162525Z","shell.execute_reply":"2024-01-29T14:56:03.806523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom torch import nn\nfrom torch import Tensor\nfrom PIL import Image\nfrom torchvision.transforms import Compose, Resize, ToTensor\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import SGD\nimport torchvision\nfrom einops import rearrange, reduce, repeat\nfrom einops.layers.torch import Rearrange, Reduce\n# from torchsummary import summary\nimport json\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nfrom accelerate import Accelerator, notebook_launcher","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:03.810018Z","iopub.execute_input":"2024-01-29T14:56:03.810917Z","iopub.status.idle":"2024-01-29T14:56:08.223371Z","shell.execute_reply.started":"2024-01-29T14:56:03.810874Z","shell.execute_reply":"2024-01-29T14:56:08.222584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# patch_size = 16 # 16 pixels\n# x = torch.rand((1, 3, 144, 144)) # (h s1) --> (9, 16)\n# patches = rearrange(x, 'b c (h s1) (w s2) -> b (h w) (s1 s2 c)', s1=patch_size, s2=patch_size)\n# patches.shape","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:08.224661Z","iopub.execute_input":"2024-01-29T14:56:08.225477Z","iopub.status.idle":"2024-01-29T14:56:08.229644Z","shell.execute_reply.started":"2024-01-29T14:56:08.225440Z","shell.execute_reply":"2024-01-29T14:56:08.228793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.read_csv(\"/kaggle/input/cassava-leaf-disease-classification/train.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:08.231992Z","iopub.execute_input":"2024-01-29T14:56:08.232468Z","iopub.status.idle":"2024-01-29T14:56:08.286060Z","shell.execute_reply.started":"2024-01-29T14:56:08.232431Z","shell.execute_reply":"2024-01-29T14:56:08.285367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train, val = train_test_split(data, test_size=0.3, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:08.286991Z","iopub.execute_input":"2024-01-29T14:56:08.287293Z","iopub.status.idle":"2024-01-29T14:56:08.303912Z","shell.execute_reply.started":"2024-01-29T14:56:08.287267Z","shell.execute_reply":"2024-01-29T14:56:08.303032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(\"/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json\", 'r') as file:\n    label = json.load(file)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:08.305111Z","iopub.execute_input":"2024-01-29T14:56:08.305471Z","iopub.status.idle":"2024-01-29T14:56:08.321077Z","shell.execute_reply.started":"2024-01-29T14:56:08.305426Z","shell.execute_reply":"2024-01-29T14:56:08.320329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def PositionEmbedding(seq_len, emb_size):\n    embeddings = torch.ones(seq_len, emb_size)\n    for i in range(seq_len):\n        for j in range(emb_size):\n            embeddings[i][j] = np.sin(i / (pow(10000, j / emb_size))) if j % 2 == 0 else np.cos(i / (pow(10000, (j - 1) / emb_size)))\n    return torch.tensor(embeddings)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:08.322188Z","iopub.execute_input":"2024-01-29T14:56:08.322458Z","iopub.status.idle":"2024-01-29T14:56:08.330107Z","shell.execute_reply.started":"2024-01-29T14:56:08.322434Z","shell.execute_reply":"2024-01-29T14:56:08.329238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PatchEmbedding(nn.Module):\n    def __init__(self, in_channels: int = 3, patch_size: int = 16, emb_size: int = 768, img_size=224):\n        self.patch_size = patch_size\n        super().__init__()\n        self.projection = nn.Sequential(\n            # using a conv layer instead of a linear one -> performance gains\n            nn.Conv2d(in_channels, emb_size, kernel_size=patch_size, stride=patch_size),\n            Rearrange('b e (h) (w) -> b (h w) e'),\n        )\n\n        self.cls_token = nn.Parameter(torch.rand(1, 1, emb_size))\n        self.pos_embed = nn.Parameter(PositionEmbedding((img_size // patch_size)**2 + 1, emb_size))\n    \n    def forward(self, x: Tensor) -> Tensor:\n        b, _, _, _ = x.shape\n        x = self.projection(x)    \n\n        cls_token = repeat(self.cls_token, ' () s e -> b s e', b=b)\n\n        x = torch.cat([cls_token, x], dim=1)\n\n        x = x + self.pos_embed\n        return x\nx = torch.rand(1, 3, 224, 224)\npemb = PatchEmbedding()\npemb(x).shape","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:11.837440Z","iopub.execute_input":"2024-01-29T14:56:11.837804Z","iopub.status.idle":"2024-01-29T14:56:13.921339Z","shell.execute_reply.started":"2024-01-29T14:56:11.837775Z","shell.execute_reply":"2024-01-29T14:56:13.920406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# conv = nn.Conv2d(2, 4, kernel_size = 3, stride=3)\n# x = torch.rand(1, 2, 111, 111) \n# print(x)\n# rearrange(conv(x), 'b e (h) (w) -> b (h w) e').shape","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:28.425754Z","iopub.execute_input":"2024-01-29T14:56:28.426513Z","iopub.status.idle":"2024-01-29T14:56:28.430487Z","shell.execute_reply.started":"2024-01-29T14:56:28.426479Z","shell.execute_reply":"2024-01-29T14:56:28.429570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"w = torch.rand((1, 4, 5, 9))  \nk = torch.rand((1, 4, 5, 9))","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:28.630927Z","iopub.execute_input":"2024-01-29T14:56:28.631217Z","iopub.status.idle":"2024-01-29T14:56:28.636161Z","shell.execute_reply.started":"2024-01-29T14:56:28.631192Z","shell.execute_reply":"2024-01-29T14:56:28.635264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# q = w@k.transpose(3, 2)\n# F.softmax(q)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:29.040703Z","iopub.execute_input":"2024-01-29T14:56:29.040984Z","iopub.status.idle":"2024-01-29T14:56:29.045637Z","shell.execute_reply.started":"2024-01-29T14:56:29.040961Z","shell.execute_reply":"2024-01-29T14:56:29.044922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultiHead(nn.Module):\n  def __init__(self, emb_size, num_head):\n    super().__init__()\n    self.emb_size = emb_size\n    self.num_head = num_head\n    self.key = nn.Linear(emb_size, emb_size)\n    self.value = nn.Linear(emb_size, emb_size)\n    self.query = nn.Linear(emb_size, emb_size)  # rm bias=True\n    self.att_dr = nn.Dropout(0.1)\n  def forward(self, x):\n    k = rearrange(self.key(x), 'b n (h e) -> b h n e', h=self.num_head)\n    q = rearrange(self.key(x), 'b n (h e) -> b h n e', h=self.num_head)\n    v = rearrange(self.key(x), 'b n (h e) -> b h n e', h=self.num_head)\n\n\n    wei = q@k.transpose(3,2)/self.num_head ** 0.5    \n    wei = F.softmax(wei, dim=2)\n    wei = self.att_dr(wei)\n\n    out = wei@v\n\n    out = rearrange(out, 'b h n e -> b n (h e)')\n    return out\n\npatches_embedded = PatchEmbedding()(x)\nprint(patches_embedded.shape)\nMultiHead(768, 4)(patches_embedded).shape\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:29.655855Z","iopub.execute_input":"2024-01-29T14:56:29.656311Z","iopub.status.idle":"2024-01-29T14:56:31.701976Z","shell.execute_reply.started":"2024-01-29T14:56:29.656275Z","shell.execute_reply":"2024-01-29T14:56:31.701044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FeedForward(nn.Module):\n  def __init__(self, emb_size):\n    super().__init__()\n    self.ff = nn.Sequential(\n        nn.Linear(emb_size, 4*emb_size),\n        nn.Linear(4*emb_size, emb_size)\n    )\n  def forward(self, x):\n    return self.ff(x)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:31.703884Z","iopub.execute_input":"2024-01-29T14:56:31.704333Z","iopub.status.idle":"2024-01-29T14:56:31.710312Z","shell.execute_reply.started":"2024-01-29T14:56:31.704298Z","shell.execute_reply":"2024-01-29T14:56:31.709342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Block(nn.Module):\n  def __init__(self,emb_size, num_head):\n    super().__init__()\n    self.att = MultiHead(emb_size, num_head)\n    self.ll =   nn.LayerNorm(emb_size)\n    self.dropout = nn.Dropout(0.1)\n    self.ff = FeedForward(emb_size)\n  def forward(self, x):\n    x = x + self.dropout(self.att(self.ll(x)))  # self.att(x): x -> (b , n, emb_size) \n    x = x + self.dropout(self.ff(self.ll(x)))\n    return x","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:31.711338Z","iopub.execute_input":"2024-01-29T14:56:31.711656Z","iopub.status.idle":"2024-01-29T14:56:31.723691Z","shell.execute_reply.started":"2024-01-29T14:56:31.711620Z","shell.execute_reply":"2024-01-29T14:56:31.722972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VissionTransformer(nn.Module):\n  def __init__(self, _layers, emb_size, num_head, num_class):\n    super().__init__()\n    self.attention = nn.Sequential(*[Block(emb_size, num_head) for _ in range(_layers)])\n    self.patchemb = PatchEmbedding()\n    self.ff = nn.Linear(emb_size, num_class)\n\n  def forward(self, x):     # x -> (b, c, h, w)\n    embeddings = self.patchemb(x)    \n    x = self.attention(embeddings)    # embed -> (b, 197, emb_size)\n    x = self.ff(x[:, 0, :])\n    return x","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:32.008732Z","iopub.execute_input":"2024-01-29T14:56:32.009045Z","iopub.status.idle":"2024-01-29T14:56:32.015399Z","shell.execute_reply.started":"2024-01-29T14:56:32.009019Z","shell.execute_reply":"2024-01-29T14:56:32.014328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = torch.rand(1, 3, 224, 224)\nm = VissionTransformer(8, 768, 4, 10)\nm(x).shape","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:32.729182Z","iopub.execute_input":"2024-01-29T14:56:32.729828Z","iopub.status.idle":"2024-01-29T14:56:35.274435Z","shell.execute_reply.started":"2024-01-29T14:56:32.729795Z","shell.execute_reply":"2024-01-29T14:56:35.273461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform_train = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomRotation(10),\n    transforms.RandomHorizontalFlip(),\n    ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std = [0.229, 0.224, 0.225])\n])\n\ntransform_test = transforms.Compose([\n    transforms.Resize((224, 224)),\n    ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std = [0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:35.275892Z","iopub.execute_input":"2024-01-29T14:56:35.276209Z","iopub.status.idle":"2024-01-29T14:56:35.284884Z","shell.execute_reply.started":"2024-01-29T14:56:35.276182Z","shell.execute_reply":"2024-01-29T14:56:35.283956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, file, transform):\n        self.file = file\n        self.transform = transform\n    def __len__(self):\n        return len(self.file)\n    \n    def __getitem__(self, idx):\n        img, label = self.file.iloc[idx]\n        img = Image.open(\"/kaggle/input/cassava-leaf-disease-classification/train_images/\" + img)\n\n        if self.transform:\n            img = self.transform(img)\n        return img, label","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:35.593744Z","iopub.execute_input":"2024-01-29T14:56:35.594344Z","iopub.status.idle":"2024-01-29T14:56:35.599903Z","shell.execute_reply.started":"2024-01-29T14:56:35.594316Z","shell.execute_reply":"2024-01-29T14:56:35.599107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = CustomDataset(train, transform_train)\nvalid_data = CustomDataset(val, transform_test)\ntrain_loader = DataLoader(train_data, batch_size=32, shuffle=True)\nvalid_loader = DataLoader(valid_data, batch_size = 32, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:36.618809Z","iopub.execute_input":"2024-01-29T14:56:36.619191Z","iopub.status.idle":"2024-01-29T14:56:36.624653Z","shell.execute_reply.started":"2024-01-29T14:56:36.619160Z","shell.execute_reply":"2024-01-29T14:56:36.623749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nnum_layers = 8\nemb_size = 768\nnum_head = 6\nnum_class=5\nmodel = VissionTransformer(num_layers, emb_size=emb_size,num_head=num_head,num_class=num_class).to(device)\n# model = ViT(image_size=(224, 224), patch_size=16, num_classes=5, dim=768, depth=8, heads=4, mlp_dim=4*768).to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = SGD(model.parameters(), lr=1e-3,momentum=0.9, weight_decay=1e-4)","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:37.239401Z","iopub.execute_input":"2024-01-29T14:56:37.239698Z","iopub.status.idle":"2024-01-29T14:56:40.132691Z","shell.execute_reply.started":"2024-01-29T14:56:37.239672Z","shell.execute_reply":"2024-01-29T14:56:40.131895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_epoch(model, loader, criterion, optimizer, device):\n    final_loss = 0\n    final_acc = 0\n    model.train()\n    for i, (x, y) in enumerate(tqdm(loader)):\n        \n        image = x.type(torch.FloatTensor).to(device)\n        y = y.to(device)\n        \n        predict = model(image)\n#         predict = predict.type(torch.LongTensor).to(device)\n        loss = criterion(predict, y)\n        final_loss += loss.item()\n        \n        _, predict = predict.max(1)\n        accuracy = (predict == y).float()\n        final_acc += accuracy.sum()/len(accuracy) * 100\n        \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n    final_loss /= len(loader)\n    final_acc /= len(loader)\n    \n    return final_loss , final_acc\n\ndef valid_epoch(model, loader, criterion, optimizer, device):\n    final_loss = 0\n    final_acc = 0\n    model.eval()\n    with torch.no_grad():\n        \n        for i, (x, y) in enumerate(tqdm(loader)):\n            image = x.type(torch.FloatTensor).to(device)\n            y = y.to(device)\n\n            predict = model(image)\n#             predict = predict.type(torch.LongTensor).to(device)\n            loss = criterion(predict, y)\n            final_loss += loss.item()\n\n            _, predict = predict.max(1)\n            accuracy = (predict == y).float()\n            final_acc += accuracy.sum()/len(accuracy) * 100\n\n        final_loss /= len(loader)\n        final_acc /= len(loader)\n    \n    return final_loss , final_acc ","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:40.134123Z","iopub.execute_input":"2024-01-29T14:56:40.134413Z","iopub.status.idle":"2024-01-29T14:56:40.145100Z","shell.execute_reply.started":"2024-01-29T14:56:40.134388Z","shell.execute_reply":"2024-01-29T14:56:40.144262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 60\nbest_acc =   0\n\nfor i in range(EPOCHS):\n    print(\"EPOCHS: {}/{}\".format(i+1, EPOCHS))\n\n    train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n    valid_loss, valid_acc = valid_epoch(model, valid_loader, criterion, optimizer, device)\n    print(f\"Training: loss {train_loss}, accuracy {train_acc}\")\n    print(f\"Validation: loss {valid_loss}, accuracy {valid_acc}\")    \n\n    if valid_acc > best_acc:\n        best_acc = valid_acc\n        torch.save(model.state_dict(), \"vit_model.pth\")\n","metadata":{"execution":{"iopub.status.busy":"2024-01-29T14:56:40.146244Z","iopub.execute_input":"2024-01-29T14:56:40.146555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, \"vit.pth\")\ntorch.save(model.state_dict(), \"vit_state_dict.pth\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}