{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","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":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:03:26.358364Z","iopub.execute_input":"2025-11-29T12:03:26.358570Z","iopub.status.idle":"2025-11-29T12:03:53.721835Z","shell.execute_reply.started":"2025-11-29T12:03:26.358553Z","shell.execute_reply":"2025-11-29T12:03:53.720748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport pandas as pd\nfrom PIL import Image\nfrom torchvision import transforms\nfrom torch.cuda.amp import autocast, GradScaler\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:12:07.377839Z","iopub.execute_input":"2025-11-29T12:12:07.378494Z","iopub.status.idle":"2025-11-29T12:12:07.382400Z","shell.execute_reply.started":"2025-11-29T12:12:07.378468Z","shell.execute_reply":"2025-11-29T12:12:07.381761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_DIR = \"/kaggle/input/cassava-leaf-disease-classification/train_images\"\nCSV_PATH  = \"/kaggle/input/cassava-leaf-disease-classification/train.csv\"\n\nBATCH_SIZE = 32\nEPOCHS = 15\nLR = 3e-5\nIMG_SIZE = 224\nNUM_CLASSES = 5\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:12:21.617945Z","iopub.execute_input":"2025-11-29T12:12:21.618639Z","iopub.status.idle":"2025-11-29T12:12:21.622688Z","shell.execute_reply.started":"2025-11-29T12:12:21.618614Z","shell.execute_reply":"2025-11-29T12:12:21.621939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.RandomResizedCrop(IMG_SIZE),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.ColorJitter(0.2,0.2,0.2,0.1),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n])\n\nvalid_transform = transforms.Compose([\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:12:35.048922Z","iopub.execute_input":"2025-11-29T12:12:35.049445Z","iopub.status.idle":"2025-11-29T12:12:35.054298Z","shell.execute_reply.started":"2025-11-29T12:12:35.049423Z","shell.execute_reply":"2025-11-29T12:12:35.053628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, df, img_root, transform=None):\n        self.df = df\n        self.img_root = img_root\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.img_root, row.image_id)\n        img = Image.open(img_path).convert(\"RGB\")\n\n        if self.transform:\n            img = self.transform(img)\n\n        label = row.label\n        return img, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:12:45.989156Z","iopub.execute_input":"2025-11-29T12:12:45.989445Z","iopub.status.idle":"2025-11-29T12:12:45.994460Z","shell.execute_reply.started":"2025-11-29T12:12:45.989424Z","shell.execute_reply":"2025-11-29T12:12:45.993742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(CSV_PATH)\n\n# (Optional) You can split into train/valid. Kaggle dataset has no valid set.\nfrom sklearn.model_selection import train_test_split\ntrain_df, valid_df = train_test_split(df, test_size=0.15, random_state=42, stratify=df.label)\n\ntrain_ds = CassavaDataset(train_df, TRAIN_DIR, train_transform)\nvalid_ds = CassavaDataset(valid_df, TRAIN_DIR, valid_transform)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=4)\nvalid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:12:58.080412Z","iopub.execute_input":"2025-11-29T12:12:58.080986Z","iopub.status.idle":"2025-11-29T12:12:58.847870Z","shell.execute_reply.started":"2025-11-29T12:12:58.080963Z","shell.execute_reply":"2025-11-29T12:12:58.847287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = timm.create_model(\n    'vit_base_patch16_224',\n    pretrained=True,\n    num_classes=NUM_CLASSES\n)\n\nmodel.to(DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:13:09.638228Z","iopub.execute_input":"2025-11-29T12:13:09.639004Z","iopub.status.idle":"2025-11-29T12:13:12.982104Z","shell.execute_reply.started":"2025-11-29T12:13:09.638977Z","shell.execute_reply":"2025-11-29T12:13:12.981340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=0.05)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\nscaler = GradScaler()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:13:27.619835Z","iopub.execute_input":"2025-11-29T12:13:27.620119Z","iopub.status.idle":"2025-11-29T12:13:27.625489Z","shell.execute_reply.started":"2025-11-29T12:13:27.620101Z","shell.execute_reply":"2025-11-29T12:13:27.624688Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(epoch):\n    model.train()\n    running_loss = 0\n    total = 0\n    correct = 0\n\n    for imgs, labels in train_loader:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n\n        optimizer.zero_grad()\n\n        with autocast():\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item()\n        _, preds = outputs.max(1)\n        total += labels.size(0)\n        correct += preds.eq(labels).sum().item()\n\n    print(f\"Epoch {epoch} Train Loss: {running_loss/len(train_loader):.4f}  Acc: {100*correct/total:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:13:42.752563Z","iopub.execute_input":"2025-11-29T12:13:42.753255Z","iopub.status.idle":"2025-11-29T12:13:42.758240Z","shell.execute_reply.started":"2025-11-29T12:13:42.753229Z","shell.execute_reply":"2025-11-29T12:13:42.757705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(epoch):\n    model.eval()\n    running_loss = 0\n    total = 0\n    correct = 0\n\n    with torch.no_grad():\n        for imgs, labels in valid_loader:\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item()\n            _, preds = outputs.max(1)\n            total += labels.size(0)\n            correct += preds.eq(labels).sum().item()\n\n    acc = 100 * correct / total\n    print(f\"Epoch {epoch} Valid Loss: {running_loss/len(valid_loader):.4f}  Acc: {acc:.2f}%\")\n    return acc\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:13:53.040202Z","iopub.execute_input":"2025-11-29T12:13:53.040723Z","iopub.status.idle":"2025-11-29T12:13:53.045498Z","shell.execute_reply.started":"2025-11-29T12:13:53.040701Z","shell.execute_reply":"2025-11-29T12:13:53.044927Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_acc = 0\n\nfor epoch in range(1, EPOCHS+1):\n    train_one_epoch(epoch)\n    val_acc = validate(epoch)\n    scheduler.step()\n\n    if val_acc > best_acc:\n        best_acc = val_acc\n        torch.save(model.state_dict(), \"best_vit_b16.pth\")\n        print(\"✔ Saved New Best Model\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-29T12:14:38.913126Z","iopub.execute_input":"2025-11-29T12:14:38.913878Z","iopub.status.idle":"2025-11-29T13:09:37.147028Z","shell.execute_reply.started":"2025-11-29T12:14:38.913848Z","shell.execute_reply":"2025-11-29T13:09:37.145910Z"}},"outputs":[],"execution_count":null}]}