{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":9646355,"sourceType":"datasetVersion","datasetId":5891175}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Import Libraries and Data","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom torchvision import datasets\nfrom sklearn.model_selection import train_test_split\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\n\nimport os\nimport time\nimport random\nimport shutil","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:07.907546Z","iopub.execute_input":"2024-10-23T08:33:07.907852Z","iopub.status.idle":"2024-10-23T08:33:13.184770Z","shell.execute_reply.started":"2024-10-23T08:33:07.907821Z","shell.execute_reply":"2024-10-23T08:33:13.183744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:13.189337Z","iopub.execute_input":"2024-10-23T08:33:13.189608Z","iopub.status.idle":"2024-10-23T08:33:13.248657Z","shell.execute_reply.started":"2024-10-23T08:33:13.189578Z","shell.execute_reply":"2024-10-23T08:33:13.247644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# general global variables\ndir_dataset = \"/kaggle/input/cassava-leaf/\"\ndir_trainset = \"/kaggle/input/cassava-leaf/train_images/\"\ndir_testset = \"/kaggle/input/cassava-leaf/test_images/\"\ndir_model = \"/kaggle/input/cassava-leaf-model/\"\n\ncassava_diseases = {\n    \"0\": \"Cassava Bacterial Blight (CBB)\",\n    \"1\": \"Cassava Brown Streak Disease (CBSD)\",\n    \"2\": \"Cassava Green Mottle (CGM)\",\n    \"3\": \"Cassava Mosaic Disease (CMD)\",\n    \"4\": \"Healthy\"\n}","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:13.250143Z","iopub.execute_input":"2024-10-23T08:33:13.250550Z","iopub.status.idle":"2024-10-23T08:33:13.260163Z","shell.execute_reply.started":"2024-10-23T08:33:13.250496Z","shell.execute_reply":"2024-10-23T08:33:13.259242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\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  # Set seed for PyTorch (both CPU and GPU)\n    \nseed_everything()","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:13.262764Z","iopub.execute_input":"2024-10-23T08:33:13.263140Z","iopub.status.idle":"2024-10-23T08:33:13.277097Z","shell.execute_reply.started":"2024-10-23T08:33:13.263106Z","shell.execute_reply":"2024-10-23T08:33:13.276302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Data Analysis & Preparation","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(\"/kaggle/input/cassava-leaf/train.csv\")\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:13.278086Z","iopub.execute_input":"2024-10-23T08:33:13.278384Z","iopub.status.idle":"2024-10-23T08:33:13.308392Z","shell.execute_reply.started":"2024-10-23T08:33:13.278354Z","shell.execute_reply":"2024-10-23T08:33:13.307434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.set_style(\"whitegrid\")\nplt.figure(figsize=(6,6))\nsns.countplot(x=\"label\", data=df_train, edgecolor=\"black\", palette=\"mako\", hue=\"label\")","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:13.309467Z","iopub.execute_input":"2024-10-23T08:33:13.309768Z","iopub.status.idle":"2024-10-23T08:33:13.747710Z","shell.execute_reply.started":"2024-10-23T08:33:13.309734Z","shell.execute_reply":"2024-10-23T08:33:13.746604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Sample Images of the Various Diseases","metadata":{}},{"cell_type":"code","source":"for i in range(0, df_train[\"label\"].nunique()):\n    df_disease = df_train[df_train[\"label\"] == i]\n    files = df_disease[\"image_id\"].sample(3).tolist()\n    \n    print(cassava_diseases[str(i)])\n    plt.figure(figsize=(15,3))\n    index = 0\n    for file in files:\n        image = Image.open(dir_trainset + file)\n        plt.subplot(1, 3, index + 1)\n        plt.imshow(image)\n        plt.axis(\"off\")\n        index += 1\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:13.748867Z","iopub.execute_input":"2024-10-23T08:33:13.749189Z","iopub.status.idle":"2024-10-23T08:33:17.135132Z","shell.execute_reply.started":"2024-10-23T08:33:13.749156Z","shell.execute_reply":"2024-10-23T08:33:17.134232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Testing Image Augumentation Using Albumentations","metadata":{}},{"cell_type":"code","source":"image_size = 224\n\n# Declare an augmentation pipeline\ntransform = A.Compose([\n    A.Resize(height=image_size, width=image_size),  # Resize to 299x299\n    A.HorizontalFlip(p=0.5),  # Random horizontal flip\n    A.VerticalFlip(p=0.5),    # Random vertical flip\n    A.Rotate(limit=30, p=0.5),  # Random rotation\n    #A.RandomBrightnessContrast(p=0.1),  # Random brightness/contrast adjustment\n    #A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),  # Shift, scale, rotate\n    #A.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),  # Normalize to range [-1, 1] (for Inception)\n])\n\nfiles = df_train[\"image_id\"].sample(3).tolist()\n\nplt.figure(figsize=(8, 8))\nindex = 0 \n\nfor file in files:\n    original_image = cv2.imread(dir_trainset + file)\n    image = cv2.cvtColor(original_image, cv2.COLOR_BGR2RGB)\n\n    # Augment the image\n    transformed = transform(image=image)\n    transformed_image = transformed[\"image\"]\n\n    # Display original image\n    plt.subplot(3, 2, index + 1)\n    plt.imshow(image)  # Use `image` to show RGB format original image\n    plt.axis(\"off\")\n    plt.title(\"Original\")\n    index += 1\n\n    # Display transformed image\n    plt.subplot(3, 2, index + 1)\n    plt.imshow(transformed_image)\n    plt.axis(\"off\")\n    plt.title(\"Transformed\")\n    index += 1\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:17.136285Z","iopub.execute_input":"2024-10-23T08:33:17.136601Z","iopub.status.idle":"2024-10-23T08:33:18.504703Z","shell.execute_reply.started":"2024-10-23T08:33:17.136568Z","shell.execute_reply":"2024-10-23T08:33:18.503792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model: ViT (Fine-Tune)","metadata":{}},{"cell_type":"code","source":"class CassavaDataset(torch.utils.data.Dataset):\n    def __init__(self, dataframe, img_dir, transform=None):\n        self.dataframe = dataframe\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        img_name = os.path.join(self.img_dir, self.dataframe.iloc[idx, 0])\n        image = Image.open(img_name).convert(\"RGB\")\n        label = int(self.dataframe.iloc[idx, 1])\n\n        # Convert the PIL image to a NumPy array\n        image = np.array(image)\n\n        if self.transform:\n            # Pass the image as a named argument\n            image = self.transform(image=image)[\"image\"]\n\n        return image, label\n\nimage_size = 384\n    \ntrain_transform = A.Compose([\n    A.RandomResizedCrop(image_size, image_size),\n    A.Transpose(p=0.5),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.ShiftScaleRotate(p=0.5),\n    A.HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n    A.RandomBrightnessContrast(brightness_limit=(-0.1,0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n    A.CoarseDropout(p=0.5),\n    ToTensorV2(p=1.0)], p=1.)\n\nvalid_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.CenterCrop(image_size, image_size, p=1.),\n    A.Resize(384, 384),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n    ToTensorV2(p=1.0)], p=1.)\n\ndf_train, df_val = train_test_split(df_train, test_size=0.2, stratify=df_train['label'])\n\ntrain_dataset = CassavaDataset(df_train, dir_trainset, transform=train_transform)\nval_dataset = CassavaDataset(df_val, dir_trainset, transform=valid_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, num_workers=4, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=32, num_workers=4, shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:18.505901Z","iopub.execute_input":"2024-10-23T08:33:18.506238Z","iopub.status.idle":"2024-10-23T08:33:18.537452Z","shell.execute_reply.started":"2024-10-23T08:33:18.506202Z","shell.execute_reply":"2024-10-23T08:33:18.536361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = timm.create_model('vit_base_patch16_384', pretrained=True)\nmodel.head = nn.Linear(model.head.in_features, 5)  # 5 classes for cassava diseases\nmodel = model.to('cuda' if torch.cuda.is_available() else 'cpu')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:18.539464Z","iopub.execute_input":"2024-10-23T08:33:18.539924Z","iopub.status.idle":"2024-10-23T08:33:20.333649Z","shell.execute_reply.started":"2024-10-23T08:33:18.539877Z","shell.execute_reply":"2024-10-23T08:33:20.332794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:20.334763Z","iopub.execute_input":"2024-10-23T08:33:20.335087Z","iopub.status.idle":"2024-10-23T08:33:20.341391Z","shell.execute_reply.started":"2024-10-23T08:33:20.335054Z","shell.execute_reply":"2024-10-23T08:33:20.340450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grad_accum_steps = 2\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=10, grad_accum_steps=1):\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    \n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch+1}/{num_epochs}')\n        print('-' * 20)\n        \n        # Training phase\n        model.train()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n\n        train_loader_tqdm = tqdm(train_loader, desc='Training', leave=False)\n        \n        optimizer.zero_grad()\n        for batch_idx, (images, labels) in enumerate(train_loader_tqdm):\n            images, labels = images.to(device), labels.to(device)\n            \n            outputs = model(images)\n            loss = criterion(outputs, labels) / grad_accum_steps  # Normalize loss for gradient accumulation\n            \n            loss.backward()\n            \n            if (batch_idx + 1) % grad_accum_steps == 0:  # Update weights after `grad_accum_steps`\n                optimizer.step()\n                optimizer.zero_grad()\n            \n            running_loss += loss.item() * grad_accum_steps  # Undo normalization for display\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n            \n            train_loader_tqdm.set_postfix({'loss': running_loss / (total // labels.size(0)), \n                                           'accuracy': 100. * correct / total})\n\n        train_loss = running_loss / len(train_loader)\n        train_acc = 100. * correct / total\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        correct = 0\n        total = 0\n        \n        val_loader_tqdm = tqdm(val_loader, desc='Validating', leave=False)\n        \n        with torch.no_grad():\n            for images, labels in val_loader_tqdm:\n                images, labels = images.to(device), labels.to(device)\n                \n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                \n                val_loss += loss.item()\n                _, predicted = outputs.max(1)\n                total += labels.size(0)\n                correct += predicted.eq(labels).sum().item()\n                \n                val_loader_tqdm.set_postfix({'val_loss': val_loss / (total // labels.size(0)), \n                                             'val_accuracy': 100. * correct / total})\n        \n        val_loss /= len(val_loader)\n        val_acc = 100. * correct / total\n        \n        print(f'Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%, '\n              f'Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%')\n        \n        scheduler.step()\n","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:20.342903Z","iopub.execute_input":"2024-10-23T08:33:20.343450Z","iopub.status.idle":"2024-10-23T08:33:20.358484Z","shell.execute_reply.started":"2024-10-23T08:33:20.343417Z","shell.execute_reply":"2024-10-23T08:33:20.357621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=10)","metadata":{"execution":{"iopub.status.busy":"2024-10-23T08:33:20.361209Z","iopub.execute_input":"2024-10-23T08:33:20.361518Z","iopub.status.idle":"2024-10-23T14:07:52.941833Z","shell.execute_reply.started":"2024-10-23T08:33:20.361487Z","shell.execute_reply":"2024-10-23T14:07:52.940493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save the trained model's state dictionary\ntorch.save(model.state_dict(), 'base_patch16_384_trained_885.pth')","metadata":{"execution":{"iopub.status.busy":"2024-10-23T14:11:03.964837Z","iopub.execute_input":"2024-10-23T14:11:03.965903Z","iopub.status.idle":"2024-10-23T14:11:04.581764Z","shell.execute_reply.started":"2024-10-23T14:11:03.965858Z","shell.execute_reply":"2024-10-23T14:11:04.580873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}