{"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":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Exploaring data\n### just to know the structure of directory that contains our data","metadata":{}},{"cell_type":"code","source":"import os\nDir = '../input/cassava-leaf-disease-classification'\nos.listdir(Dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:19.896278Z","iopub.execute_input":"2025-12-11T17:53:19.896504Z","iopub.status.idle":"2025-12-11T17:53:19.906215Z","shell.execute_reply.started":"2025-12-11T17:53:19.896486Z","shell.execute_reply":"2025-12-11T17:53:19.905505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(os.listdir('../input/cassava-leaf-disease-classification/train_images')))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:32.203712Z","iopub.execute_input":"2025-12-11T17:53:32.204455Z","iopub.status.idle":"2025-12-11T17:53:32.389762Z","shell.execute_reply.started":"2025-12-11T17:53:32.204427Z","shell.execute_reply":"2025-12-11T17:53:32.389177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\nprint(len(os.listdir('../input/cassava-leaf-disease-classification/test_images')))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:34.021857Z","iopub.execute_input":"2025-12-11T17:53:34.022478Z","iopub.status.idle":"2025-12-11T17:53:34.029993Z","shell.execute_reply.started":"2025-12-11T17:53:34.022453Z","shell.execute_reply":"2025-12-11T17:53:34.029316Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# importing some libs to help in data inspection ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:36.421456Z","iopub.execute_input":"2025-12-11T17:53:36.421723Z","iopub.status.idle":"2025-12-11T17:53:39.555302Z","shell.execute_reply.started":"2025-12-11T17:53:36.421703Z","shell.execute_reply":"2025-12-11T17:53:39.554685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ntrain_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:39.556387Z","iopub.execute_input":"2025-12-11T17:53:39.557154Z","iopub.status.idle":"2025-12-11T17:53:39.666147Z","shell.execute_reply.started":"2025-12-11T17:53:39.557104Z","shell.execute_reply":"2025-12-11T17:53:39.665507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sns.countplot(train_df , x = 'label')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:42.310323Z","iopub.execute_input":"2025-12-11T17:53:42.310963Z","iopub.status.idle":"2025-12-11T17:53:42.521359Z","shell.execute_reply.started":"2025-12-11T17:53:42.310938Z","shell.execute_reply":"2025-12-11T17:53:42.520632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df['label'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:45.189337Z","iopub.execute_input":"2025-12-11T17:53:45.189934Z","iopub.status.idle":"2025-12-11T17:53:45.199223Z","shell.execute_reply.started":"2025-12-11T17:53:45.189906Z","shell.execute_reply":"2025-12-11T17:53:45.198511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nnp.round((train_df['label'].value_counts()/len(train_df['label']))*100, 2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:47.233776Z","iopub.execute_input":"2025-12-11T17:53:47.234516Z","iopub.status.idle":"2025-12-11T17:53:47.241845Z","shell.execute_reply.started":"2025-12-11T17:53:47.234488Z","shell.execute_reply":"2025-12-11T17:53:47.240984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nwith open('../input/cassava-leaf-disease-classification/label_num_to_disease_map.json') as file:\n    print(json.dumps(json.loads(file.read()), indent=4))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:49.635751Z","iopub.execute_input":"2025-12-11T17:53:49.636393Z","iopub.status.idle":"2025-12-11T17:53:49.646073Z","shell.execute_reply.started":"2025-12-11T17:53:49.636368Z","shell.execute_reply":"2025-12-11T17:53:49.645406Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# show some random samples of class 1\n# which is  \"Cassava Bacterial Blight (CBB)\"","metadata":{}},{"cell_type":"code","source":"import cv2\nsample = train_df[train_df.label == 0].sample(9)\nplt.figure(figsize=(12,12))\nfor ind, (image_id, label) in enumerate(zip(sample.image_id, sample.label)):\n    plt.subplot(3, 3, ind + 1)\n    image = cv2.imread(os.path.join(\"../input/cassava-leaf-disease-classification/train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    plt.axis(\"off\")\n    \nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:52.899039Z","iopub.execute_input":"2025-12-11T17:53:52.899671Z","iopub.status.idle":"2025-12-11T17:53:54.184921Z","shell.execute_reply.started":"2025-12-11T17:53:52.899645Z","shell.execute_reply":"2025-12-11T17:53:54.183864Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# show some random samples of class 0\n# which is \"Cassava Brown Streak Disease (CBSD)\"\n","metadata":{}},{"cell_type":"code","source":"\n\nsample = train_df[train_df.label == 1].sample(9)\nplt.figure(figsize=(12,12))\nfor ind, (image_id, label) in enumerate(zip(sample.image_id, sample.label)):\n    plt.subplot(4, 3, ind + 1)\n    image = cv2.imread(os.path.join(\"../input/cassava-leaf-disease-classification/train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    plt.axis(\"off\")\n    \nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:53:59.722343Z","iopub.execute_input":"2025-12-11T17:53:59.722650Z","iopub.status.idle":"2025-12-11T17:54:00.671478Z","shell.execute_reply.started":"2025-12-11T17:53:59.722627Z","shell.execute_reply":"2025-12-11T17:54:00.670407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = train_df[train_df.label == 2].sample(9)\nplt.figure(figsize=(12,12))\nfor ind, (image_id, label) in enumerate(zip(sample.image_id, sample.label)):\n    plt.subplot(3, 3, ind + 1)\n    image = cv2.imread(os.path.join(\"../input/cassava-leaf-disease-classification/train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    plt.axis(\"off\")\n    \nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:04.081734Z","iopub.execute_input":"2025-12-11T17:54:04.082547Z","iopub.status.idle":"2025-12-11T17:54:05.055375Z","shell.execute_reply.started":"2025-12-11T17:54:04.082519Z","shell.execute_reply":"2025-12-11T17:54:05.054321Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = train_df[train_df.label == 4].sample(9)\nplt.figure(figsize=(12,12))\nfor ind, (img_id, lab) in enumerate(zip(sample.image_id, sample.label)):\n    plt.subplot(3,3,ind+1)\n    image = cv2.imread(os.path.join(\"../input/cassava-leaf-disease-classification/train_images\", img_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    plt.axis(\"off\")\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:07.225287Z","iopub.execute_input":"2025-12-11T17:54:07.225571Z","iopub.status.idle":"2025-12-11T17:54:08.171165Z","shell.execute_reply.started":"2025-12-11T17:54:07.225551Z","shell.execute_reply":"2025-12-11T17:54:08.170313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score\ny_pred = [3] * len(train_df.label)\nprint(f\"The baseline accuracy is {accuracy_score(y_pred, train_df.label)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:13.269296Z","iopub.execute_input":"2025-12-11T17:54:13.269842Z","iopub.status.idle":"2025-12-11T17:54:13.604338Z","shell.execute_reply.started":"2025-12-11T17:54:13.269818Z","shell.execute_reply":"2025-12-11T17:54:13.603635Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluating the model ","metadata":{}},{"cell_type":"code","source":"Batch_size = 16\nimg_height, img_width = 250, 250","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:16.830751Z","iopub.execute_input":"2025-12-11T17:54:16.831409Z","iopub.status.idle":"2025-12-11T17:54:16.835185Z","shell.execute_reply.started":"2025-12-11T17:54:16.831384Z","shell.execute_reply":"2025-12-11T17:54:16.834246Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df['label'].dtype","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:18.352734Z","iopub.execute_input":"2025-12-11T17:54:18.353000Z","iopub.status.idle":"2025-12-11T17:54:18.357751Z","shell.execute_reply.started":"2025-12-11T17:54:18.352981Z","shell.execute_reply":"2025-12-11T17:54:18.357154Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Augmentation ","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom torchvision import transforms, models\nfrom PIL import Image\n# import pandas as pd\nimport timm  # For EfficientNetB3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:44.002959Z","iopub.execute_input":"2025-12-11T17:54:44.003334Z","iopub.status.idle":"2025-12-11T17:54:44.007487Z","shell.execute_reply.started":"2025-12-11T17:54:44.003309Z","shell.execute_reply":"2025-12-11T17:54:44.006785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Settings\nBatch_size = 16\nimg_height, img_width = 300, 300\nepochs = 5\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nDir = '../input/cassava-leaf-disease-classification'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:49.721709Z","iopub.execute_input":"2025-12-11T17:54:49.721975Z","iopub.status.idle":"2025-12-11T17:54:49.725832Z","shell.execute_reply.started":"2025-12-11T17:54:49.721956Z","shell.execute_reply":"2025-12-11T17:54:49.725172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Custom Dataset\nclass CassavaDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_name = self.df.loc[idx, \"image_id\"]\n        label = int(self.df.loc[idx, \"label\"])\n        image_path = os.path.join(self.img_dir, img_name)\n        image = Image.open(image_path).convert(\"RGB\")\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:52.869263Z","iopub.execute_input":"2025-12-11T17:54:52.869840Z","iopub.status.idle":"2025-12-11T17:54:52.874962Z","shell.execute_reply.started":"2025-12-11T17:54:52.869817Z","shell.execute_reply":"2025-12-11T17:54:52.874261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3. Transforms / Augmentations\n# -------------------------------\ntrain_transforms = transforms.Compose([\n    transforms.Resize((img_height, img_width)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406],\n                         [0.229, 0.224, 0.225])  # ImageNet normalization\n])\n\nval_transforms = transforms.Compose([\n    transforms.Resize((img_height, img_width)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406],\n                         [0.229, 0.224, 0.225])\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:54.757180Z","iopub.execute_input":"2025-12-11T17:54:54.757830Z","iopub.status.idle":"2025-12-11T17:54:54.762643Z","shell.execute_reply.started":"2025-12-11T17:54:54.757804Z","shell.execute_reply":"2025-12-11T17:54:54.762005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4. Train/Validation Split\n# -------------------------------\ntrain_size = int(0.8 * len(train_df))\nval_size = len(train_df) - train_size\ntrain_df_split, val_df_split = random_split(train_df, [train_size, val_size])\n\nimg_dir = os.path.join(Dir, \"train_images\")\ntrain_dataset = CassavaDataset(train_df.iloc[train_df_split.indices], img_dir, transform=train_transforms)\nval_dataset   = CassavaDataset(train_df.iloc[val_df_split.indices],   img_dir, transform=val_transforms)\n\ntrain_loader = DataLoader(train_dataset, batch_size=Batch_size, shuffle=True, num_workers=2)\nval_loader   = DataLoader(val_dataset,   batch_size=Batch_size, shuffle=False, num_workers=2)\n\n# -------------------------------","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:56.901179Z","iopub.execute_input":"2025-12-11T17:54:56.901433Z","iopub.status.idle":"2025-12-11T17:54:56.914100Z","shell.execute_reply.started":"2025-12-11T17:54:56.901417Z","shell.execute_reply":"2025-12-11T17:54:56.913583Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5. Model (EfficientNetB3)\n# -------------------------------\nmodel = timm.create_model(\"efficientnet_b3\", pretrained=True, num_classes=5)\nmodel = model.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:54:58.923906Z","iopub.execute_input":"2025-12-11T17:54:58.924588Z","iopub.status.idle":"2025-12-11T17:55:00.744055Z","shell.execute_reply.started":"2025-12-11T17:54:58.924565Z","shell.execute_reply":"2025-12-11T17:55:00.743462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 6. Loss, Optimizer, Scheduler\n# -------------------------------\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.0001)\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.2, patience=2, min_lr=1e-6, verbose=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:55:22.120251Z","iopub.execute_input":"2025-12-11T17:55:22.120527Z","iopub.status.idle":"2025-12-11T17:55:22.126171Z","shell.execute_reply.started":"2025-12-11T17:55:22.120507Z","shell.execute_reply":"2025-12-11T17:55:22.125435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 7. Early stopping\n# -------------------------------\nclass EarlyStopping:\n    def __init__(self, patience=3, verbose=True):\n        self.patience = patience\n        self.verbose = verbose\n        self.counter = 0\n        self.best_loss = None\n        self.early_stop = False\n        self.best_state = None\n\n    def __call__(self, val_loss, model):\n        if self.best_loss is None or val_loss < self.best_loss:\n            self.best_loss = val_loss\n            self.best_state = model.state_dict()\n            self.counter = 0\n        else:\n            self.counter += 1\n            if self.verbose:\n                print(f\"EarlyStopping counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n                model.load_state_dict(self.best_state)\n\nearly_stopping = EarlyStopping(patience=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:55:24.564036Z","iopub.execute_input":"2025-12-11T17:55:24.564350Z","iopub.status.idle":"2025-12-11T17:55:24.570695Z","shell.execute_reply.started":"2025-12-11T17:55:24.564328Z","shell.execute_reply":"2025-12-11T17:55:24.569877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------\n# 8. Training Loop\n# -------------------------------\nfor epoch in range(epochs):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n\n    train_loss = running_loss / total\n    train_acc = correct / total  \n\n    # ---------------- Validation ----------------\n    model.eval()\n    val_loss = 0.0\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            val_loss += loss.item() * images.size(0)\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n    val_loss /= total\n    val_acc = correct / total\n\n    print(f\"Epoch [{epoch+1}/{epochs}]  Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}  Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}\")\n\n    scheduler.step(val_loss)\n    early_stopping(val_loss, model)\n    if early_stopping.early_stop:\n        print(\"Early stopping triggered.\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T17:55:26.869006Z","iopub.execute_input":"2025-12-11T17:55:26.869512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9. Save Model\n# -------------------------------\ntorch.save(model.state_dict(), \"Cassava_Model.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5. Custom CNN Model\n# -------------------------------\nclass SimpleCNN(nn.Module):\n    \n    def __init__(self, num_classes=5):\n        \n        super(SimpleCNN, self).__init__()\n        \n        self.features = nn.Sequential(\n            nn.Conv2d(3, 32, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2,2),\n\n            nn.Conv2d(32, 64, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2,2),\n\n            nn.Conv2d(64, 128, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2,2),\n\n            nn.Conv2d(128, 256, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2,2)\n        )\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(256 * 18 * 18, 256),  # 300x300 after 4 poolings -> 300/16 ~18\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, num_classes)\n        )\n        \n    def forward(self, x):\n        x = self.features(x)\n        x = self.classifier(x)\n        return x\n\nmodel = SimpleCNN(num_classes=5).to(device)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy\n\nclass EarlyStopping:\n    def __init__(self, patience=3, verbose=True):\n        self.patience = patience\n        self.verbose = verbose\n        self.counter = 0\n        self.best_loss = float('inf')  # start very high\n        self.early_stop = False\n        self.best_state = None\n\n    def __call__(self, val_loss, model):\n        if val_loss < self.best_loss:\n            self.best_loss = val_loss\n            self.best_state = copy.deepcopy(model.state_dict())\n            self.counter = 0\n        else:\n            self.counter += 1\n            if self.verbose:\n                print(f\"EarlyStopping counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                if self.best_state is not None:\n                    model.load_state_dict(self.best_state)\n                self.early_stop = True\n\n# For CNN model\nearly_stopping = EarlyStopping(patience=3, verbose=True)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 8. Training Loop\n# -------------------------------\nfor epoch in range(epochs):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n\n    train_loss = running_loss / total\n    train_acc = correct / total\n     # Validation\n    model.eval()\n    val_loss = 0.0\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            val_loss += loss.item() * images.size(0)\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n    val_loss /= total\n    val_acc = correct / total\n\n    print(f\"Epoch [{epoch+1}/{epochs}]  Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}  Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}\")\n\n    scheduler.step(val_loss)\n    early_stopping(val_loss, model)\n    if early_stopping.early_stop:\n        print(\"Early stopping triggered.\")\n        break\n\n# -------------------------------","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9. Save Model\n# -------------------------------\ntorch.save(model.state_dict(), \"Cassava_CNN_Model.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}