{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nfrom torchvision import transforms, models\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim as optim\n\nimport matplotlib.pyplot as plt\n\nimport cv2\n\nfrom time import perf_counter\n\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-24T20:57:16.833068Z","iopub.execute_input":"2024-04-24T20:57:16.833335Z","iopub.status.idle":"2024-04-24T20:57:23.933419Z","shell.execute_reply.started":"2024-04-24T20:57:16.833311Z","shell.execute_reply":"2024-04-24T20:57:23.932401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DIRECTORY_PATH = \"/kaggle/input/cassava-leaf-disease-classification/train_images/\"\nTRAIN_SPLIT = 0.9\n\ndf = pd.read_csv(\"/kaggle/input/cassava-leaf-disease-classification/train.csv\")\ndf[\"image_id\"] = df[\"image_id\"].apply(lambda x: DIRECTORY_PATH + x)\ntrain_df = df.sample(frac=TRAIN_SPLIT, random_state=42)\ntest_df = df.drop(train_df.index)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:23.935193Z","iopub.execute_input":"2024-04-24T20:57:23.935709Z","iopub.status.idle":"2024-04-24T20:57:24.000313Z","shell.execute_reply.started":"2024-04-24T20:57:23.935676Z","shell.execute_reply":"2024-04-24T20:57:23.999366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeafDataset(Dataset):\n    def __init__(self, image_paths, labels, num_labels, transform=None):\n        self.one_hot_labels = list()\n        self.image_paths = image_paths\n        self.transform = transform\n        \n        for label in labels:\n            one_hot = torch.eye(num_labels)[label]\n            self.one_hot_labels.append(one_hot)\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image = cv2.imread(self.image_paths[idx])[:, :, ::-1]\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, self.one_hot_labels[idx]","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:24.005316Z","iopub.execute_input":"2024-04-24T20:57:24.005571Z","iopub.status.idle":"2024-04-24T20:57:24.012853Z","shell.execute_reply.started":"2024-04-24T20:57:24.005549Z","shell.execute_reply":"2024-04-24T20:57:24.011858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\nIMAGE_SHAPE = (224, 224)\n\ntransform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Resize(IMAGE_SHAPE, antialias=True),\n    transforms.Normalize(mean=[0.4314, 0.4990, 0.3129], std=[0.2048, 0.2090, 0.1842]),\n])\ntrain_dataset = LeafDataset(train_df[\"image_id\"].to_list(), train_df[\"label\"].to_list(), 5, transform=transform)\ntest_dataset = LeafDataset(test_df[\"image_id\"].to_list(), test_df[\"label\"].to_list(), 5, transform=transform)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:24.013828Z","iopub.execute_input":"2024-04-24T20:57:24.014108Z","iopub.status.idle":"2024-04-24T20:57:24.323649Z","shell.execute_reply.started":"2024-04-24T20:57:24.014085Z","shell.execute_reply":"2024-04-24T20:57:24.322726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LeafDiseaseViT(nn.Module):\n    def __init__(self, num_classes):\n        super(LeafDiseaseViT, self).__init__()\n        self.vit = models.vit_b_16(weights=models.ViT_B_16_Weights.DEFAULT)\n        self.vit.heads = nn.Linear(self.vit.hidden_dim, num_classes)\n        \n        with torch.no_grad():\n            self.vit.heads.weight.zero_()\n            self.vit.heads.bias.zero_()\n\n    def forward(self, x):\n        x = self.vit(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:24.325037Z","iopub.execute_input":"2024-04-24T20:57:24.325399Z","iopub.status.idle":"2024-04-24T20:57:24.331895Z","shell.execute_reply.started":"2024-04-24T20:57:24.325366Z","shell.execute_reply":"2024-04-24T20:57:24.330991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 64\nLEARNING_RATE = 1e-3\nMOMENTUM = 0.9\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = LeafDiseaseViT(5)\nmodel.to(device)\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\ntest_dataloader = DataLoader(test_dataset, batch_size=1, shuffle=True)\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.SGD(params=model.parameters(), lr=LEARNING_RATE, momentum=MOMENTUM)","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:24.332916Z","iopub.execute_input":"2024-04-24T20:57:24.333251Z","iopub.status.idle":"2024-04-24T20:57:30.263102Z","shell.execute_reply.started":"2024-04-24T20:57:24.333220Z","shell.execute_reply":"2024-04-24T20:57:30.262128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_model(model, dataloader):\n    correct = 0\n    total = 0\n    model.eval()\n    with torch.no_grad():\n        print(\"---Testing---\")\n        start = perf_counter()\n        for image, label in dataloader:\n            image = image.to(device)\n            \n            prediction = torch.argmax(model(image), dim=-1)\n            actual = torch.argmax(label, dim=-1).to(device)\n            \n            correct += torch.sum(prediction == actual).item()\n            total += prediction.shape[-1]\n        end = perf_counter()\n    print(f\"Test time: {end - start:.2f}\")\n    print(f\"Percentage correct: {correct / total * 100:.2f}%\")\n    return correct / total","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:30.264358Z","iopub.execute_input":"2024-04-24T20:57:30.264699Z","iopub.status.idle":"2024-04-24T20:57:30.272179Z","shell.execute_reply.started":"2024-04-24T20:57:30.264669Z","shell.execute_reply":"2024-04-24T20:57:30.271319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 10\nSAVE_PATH = \"/kaggle/working/model1.pt\"\nval_scores = []\n\nfor epoch in range(EPOCHS):\n    val_score = test_model(model, test_dataloader)\n    val_scores.append(val_score)\n    \n    print(f\"---Training Epoch {epoch+1}---\")\n    model.train()\n    start = perf_counter()\n    for image, label in train_dataloader:\n        image, label = image.to(device), label.to(device)\n        optimizer.zero_grad()\n\n        prediction = model(image)\n        loss = criterion(prediction, label)\n        loss.backward()\n\n        optimizer.step()\n\n    end = perf_counter()\n    print(f\"Train time: {end - start:.2f}\")\n    print(f\"Saving to {SAVE_PATH}\")\n    torch.save(model.state_dict(), SAVE_PATH)\n\n    print(f\"Epoch {epoch + 1} Loss: {loss}\")","metadata":{"execution":{"iopub.status.busy":"2024-04-24T20:57:30.273396Z","iopub.execute_input":"2024-04-24T20:57:30.273722Z","iopub.status.idle":"2024-04-24T22:10:27.683601Z","shell.execute_reply.started":"2024-04-24T20:57:30.273692Z","shell.execute_reply":"2024-04-24T22:10:27.682595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot([i+1 for i in range(10)], val_scores, marker=\"o\", linestyle=\"-\")\n\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Val percentage\")\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-24T22:10:27.686518Z","iopub.execute_input":"2024-04-24T22:10:27.686790Z","iopub.status.idle":"2024-04-24T22:10:27.939624Z","shell.execute_reply.started":"2024-04-24T22:10:27.686767Z","shell.execute_reply":"2024-04-24T22:10:27.938865Z"},"trusted":true},"execution_count":null,"outputs":[]}]}