{"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":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30747,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os                                         # To work with OS commands\nimport pandas as pd\nfrom PIL import Image                             # To read Images\nfrom termcolor import colored                     # To Colorfull output\nfrom datetime import datetime                     # To calculate durations\n\nimport numpy as np                                # To work with Numpy arrays\nimport matplotlib.pyplot as plt                   # To Visualization\nimport seaborn as sns                             # To Visualization\n\nimport torch                                      # To work with TORCH framework\nimport torch.nn as nn                             # To work with Neural Networks\nimport torchvision                                # To work with image datasets\nimport torchvision.transforms as transforms       # To create data transforms\nfrom torch.utils.data import Dataset, DataLoader,random_split\n\nfrom torchvision.models import efficientnet_b5, EfficientNet_B5_Weights # Pretrained model with its weights\n\nfrom sklearn.metrics import confusion_matrix, classification_report # To Evaluate the result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:30.691276Z","iopub.execute_input":"2025-05-07T05:15:30.691958Z","iopub.status.idle":"2025-05-07T05:15:36.738870Z","shell.execute_reply.started":"2025-05-07T05:15:30.691918Z","shell.execute_reply":"2025-05-07T05:15:36.737830Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_path = '/kaggle/input/cassava-leaf-disease-classification/train.csv'\ntrain_csv = pd.read_csv(train_path);\ntrain_csv.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.740675Z","iopub.execute_input":"2025-05-07T05:15:36.741118Z","iopub.status.idle":"2025-05-07T05:15:36.792619Z","shell.execute_reply.started":"2025-05-07T05:15:36.741091Z","shell.execute_reply":"2025-05-07T05:15:36.791555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"base_dir = '/kaggle/input/cassava-leaf-disease-classification/train_images'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.793689Z","iopub.execute_input":"2025-05-07T05:15:36.793936Z","iopub.status.idle":"2025-05-07T05:15:36.798271Z","shell.execute_reply.started":"2025-05-07T05:15:36.793915Z","shell.execute_reply":"2025-05-07T05:15:36.797392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose(\n    [        \n        transforms.ToTensor(),                         # Convert images to tensor\n        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))  # Normalize R, G, B channels\n    ]\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.799689Z","iopub.execute_input":"2025-05-07T05:15:36.800025Z","iopub.status.idle":"2025-05-07T05:15:36.811717Z","shell.execute_reply.started":"2025-05-07T05:15:36.800002Z","shell.execute_reply":"2025-05-07T05:15:36.810788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_path = '/kaggle/input/cassava-leaf-disease-classification/train_images/1005200906.jpg'\n\n# Open the image file\nwith Image.open(image_path) as img:\n    # Get the size of the image\n    width, height = img.size\n    print(f\"Image size: {width} x {height}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.814770Z","iopub.execute_input":"2025-05-07T05:15:36.815156Z","iopub.status.idle":"2025-05-07T05:15:36.839443Z","shell.execute_reply.started":"2025-05-07T05:15:36.815129Z","shell.execute_reply":"2025-05-07T05:15:36.838334Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, csv_file, root_dir, transform=None):\n        # Read the CSV file\n        self.data = pd.read_csv(csv_file)\n        # Set the directory where images are stored\n        self.root_dir = root_dir\n        # Apply any image transformations if provided\n        self.transform = transform\n    \n    def __len__(self):\n        # Return the total number of images in the dataset\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        # Get the image name from the CSV file\n        img_name = os.path.join(self.root_dir, self.data.iloc[idx, 0])\n        # Open the image using PIL\n        image = Image.open(img_name).convert('RGB')\n        # Get the corresponding label from the CSV file\n        label = self.data.iloc[idx, 1]\n        \n        # Apply any transformations to the image\n        if self.transform:\n            image = self.transform(image)\n        \n        # Return the image and its label\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.840846Z","iopub.execute_input":"2025-05-07T05:15:36.841960Z","iopub.status.idle":"2025-05-07T05:15:36.848145Z","shell.execute_reply.started":"2025-05-07T05:15:36.841931Z","shell.execute_reply":"2025-05-07T05:15:36.847198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(type(train_path))  # Should be <class 'str'>","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.849478Z","iopub.execute_input":"2025-05-07T05:15:36.850187Z","iopub.status.idle":"2025-05-07T05:15:36.863330Z","shell.execute_reply.started":"2025-05-07T05:15:36.850148Z","shell.execute_reply":"2025-05-07T05:15:36.862139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_path)       # Should be the file path as a string","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.864597Z","iopub.execute_input":"2025-05-07T05:15:36.864939Z","iopub.status.idle":"2025-05-07T05:15:36.876540Z","shell.execute_reply.started":"2025-05-07T05:15:36.864908Z","shell.execute_reply":"2025-05-07T05:15:36.875450Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, csv_file, root_dir, transform=None):\n        # Read the CSV file\n        self.data = pd.read_csv(csv_file)\n        # Set the directory where images are stored\n        self.root_dir = root_dir\n        # Apply any image transformations if provided\n        self.transform = transform\n    \n    def __len__(self):\n        # Return the total number of images in the dataset\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        # Get the image name from the CSV file\n        img_name = os.path.join(self.root_dir, self.data.iloc[idx, 0])\n        # Open the image using PIL\n        image = Image.open(img_name).convert('RGB')\n        # Get the corresponding label from the CSV file\n        label = self.data.iloc[idx, 1]\n        \n        # Apply any transformations to the image\n        if self.transform:\n            image = self.transform(image)\n        \n        # Return the image and its label\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.877714Z","iopub.execute_input":"2025-05-07T05:15:36.878035Z","iopub.status.idle":"2025-05-07T05:15:36.888008Z","shell.execute_reply.started":"2025-05-07T05:15:36.878011Z","shell.execute_reply":"2025-05-07T05:15:36.887023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = CassavaDataset(csv_file=train_path, root_dir=base_dir, transform=transform)\n\nprint(\"Images added with labels Now creating new dataset\")\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.889104Z","iopub.execute_input":"2025-05-07T05:15:36.889388Z","iopub.status.idle":"2025-05-07T05:15:36.931981Z","shell.execute_reply.started":"2025-05-07T05:15:36.889336Z","shell.execute_reply":"2025-05-07T05:15:36.930930Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Lets divide into train and valid data","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE=32\n# Define the split ratio\ntrain_size = int(0.8 * len(dataset))  # 80% for training\nval_size = len(dataset) - train_size  # Remaining 20% for validation\n\n# Split the dataset\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)\n\nprint(\"Dataloaders are done\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.933134Z","iopub.execute_input":"2025-05-07T05:15:36.933488Z","iopub.status.idle":"2025-05-07T05:15:36.965522Z","shell.execute_reply.started":"2025-05-07T05:15:36.933461Z","shell.execute_reply":"2025-05-07T05:15:36.964465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image, label = dataset[3]\n# image is the tensor form and label is the index no of the species\nprint(image.shape)\nprint(label)\nimage","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:36.966734Z","iopub.execute_input":"2025-05-07T05:15:36.966987Z","iopub.status.idle":"2025-05-07T05:15:37.083433Z","shell.execute_reply.started":"2025-05-07T05:15:36.966966Z","shell.execute_reply":"2025-05-07T05:15:37.082425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# To visualize the image, convert the tensor to a numpy array and plot it\nimage_np = image.permute(1, 2, 0).numpy()  # Convert from (C, H, W) to (H, W, C) for visualization\n\nplt.imshow(image_np)\nplt.title(f\"Label: {label}\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:37.084698Z","iopub.execute_input":"2025-05-07T05:15:37.084993Z","iopub.status.idle":"2025-05-07T05:15:37.509607Z","shell.execute_reply.started":"2025-05-07T05:15:37.084971Z","shell.execute_reply":"2025-05-07T05:15:37.508655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\n\n# Load the label to disease mapping\nwith open('/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json', 'r') as file:\n    label_map = json.load(file)\n\n# Print to check the contents\nprint(label_map)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:37.513096Z","iopub.execute_input":"2025-05-07T05:15:37.513423Z","iopub.status.idle":"2025-05-07T05:15:37.522667Z","shell.execute_reply.started":"2025-05-07T05:15:37.513396Z","shell.execute_reply":"2025-05-07T05:15:37.521520Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# lets 1st Fetch 1 Batch from train_loader(contain 32 Images and Labels)\nfor images, labels in train_loader :\n    break\n    \n# Plot images and their labels in 1 batch (32 images and labels)\nplt.subplots(4, 8, figsize=(20, 10))\nplt.suptitle('Images and Labels in 1 batch', fontsize=20, fontweight='bold')\n\nfor i in range(BATCH_SIZE) :\n    ax = plt.subplot(4, 8, i+1)\n    img = torch.permute(images[i], (1, 2, 0))\n    #Many visualization libraries, such as matplotlib expect images to be in the format [Height, Width, Channels]\n    #But PyTorch’s default image tensor format is [Channels, Height, Width] (i.e., [C, H, W]),\n    plt.imshow(img)\n    plt.axis('off')\n    # Convert label to class name using label_map\n    label_name = label_map.get(str(labels[i].item()), 'Unknown')  # Ensure keys are strings\n    plt.title(label_name)\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:37.523962Z","iopub.execute_input":"2025-05-07T05:15:37.524358Z","iopub.status.idle":"2025-05-07T05:15:43.355817Z","shell.execute_reply.started":"2025-05-07T05:15:37.524252Z","shell.execute_reply":"2025-05-07T05:15:43.354456Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Using efficientnet_model.Then only uncomment these","metadata":{}},{"cell_type":"code","source":"# # Load efficientnet_b5 weight (best weight)\n# weights = EfficientNet_B5_Weights.DEFAULT   # I have imported this..Check the beginning of this file\n\n# # This function prepare images to feed into efficientnet_b5 model\n# preprocess = weights.transforms()\n\n# # Load efficientnet_b5 model with best weights\n# efficientnet_model = efficientnet_b5(weights=weights)\n\n# # Freezing parameters means that during training, these parameters will not be updated (i.e., their gradients will not be computed).\n# # This is typically done when you want to use a pretrained model as a feature extractor and do not want to modify the pretrained weights.\n# for param in efficientnet_model.parameters() :\n#     param.requires_grad = False\n\n# # Then Unfreeze some of last layers of model\n# # By unfreezing the last layers, you allow their weights to be updated during training.\n# # ONLY unfreeze LAST LAYER when using a pretrained model !!!!\n# for param in efficientnet_model.features[6].parameters() :   # refers to the seventh block\n#     param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:43.357137Z","iopub.execute_input":"2025-05-07T05:15:43.357442Z","iopub.status.idle":"2025-05-07T05:15:43.362042Z","shell.execute_reply.started":"2025-05-07T05:15:43.357413Z","shell.execute_reply":"2025-05-07T05:15:43.361164Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Using Vgg16 model ","metadata":{}},{"cell_type":"code","source":"from torchvision.models import vgg16, VGG16_Weights\n\n# Load the pre-trained VGG16 weights\nweights = VGG16_Weights.DEFAULT\n\n# This function prepares images to feed into VGG16\npreprocess = weights.transforms()\n\n# Load VGG16 model with pre-trained weights\nvgg16_model = vgg16(weights=weights)\n\nfor name, param in vgg16_model.named_parameters():\n    print(name)\nfor name, param in vgg16_model.named_parameters():\n    print(f\"{name}: requires_grad = {param.requires_grad}\")\n\n\n# Freeze all layers except the final fully connected layer\nfor name, param in vgg16_model.named_parameters():\n    if \"classifier.6\" in name:  # Unfreeze the final fully connected layer\n        param.requires_grad = True\n    else:\n        param.requires_grad = False\n\nfor name, param in vgg16_model.named_parameters():\n    print(f\"{name}: requires_grad = {param.requires_grad}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:43.363250Z","iopub.execute_input":"2025-05-07T05:15:43.363544Z","iopub.status.idle":"2025-05-07T05:15:48.837080Z","shell.execute_reply.started":"2025-05-07T05:15:43.363520Z","shell.execute_reply":"2025-05-07T05:15:48.836043Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### I named it as CNN model","metadata":{}},{"cell_type":"code","source":"# Number of classes\nnum_classes =5\n\nclass CNN(nn.Module) :\n    def __init__(self, num_classes) :\n        super(CNN, self).__init__()\n        ####  CONVs\n        self.conv_layers = vgg16_model   # UPDATE YOUR MODEL !!!!!!!!!!!!!!!!!!!!!!!\n        ####  Dense\n        self.dense_layers = nn.Sequential(\n            nn.Linear(1000, 256),  # UPDATE !efficientnet_model has 1000 classes in its layer thats why we take from there\n            nn.ReLU(),\n            nn.Linear(256, 64),\n            nn.ReLU(),\n            nn.Linear(64, 5)    # The number 5 corresponds to the number of output classes\n        )\n    def forward(self, X) :\n        out = self.conv_layers(X)\n         # the output tensor is typically multi-dimensional (e.g., [batch_size, channels, height, width]). out.size(0): This is the batch size\n        out = out.view(out.size(0), -1)    # used to flatten the output tensor\n        out = self.dense_layers(out)\n        return out\n\n\nmodel = CNN(num_classes)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters())  # default => torch.optim.Adam (Adaptive Moment Estimation): lr=0.001\n\n#print(model)\nprint(\"Model created !!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:48.838456Z","iopub.execute_input":"2025-05-07T05:15:48.838742Z","iopub.status.idle":"2025-05-07T05:15:48.850882Z","shell.execute_reply.started":"2025-05-07T05:15:48.838718Z","shell.execute_reply":"2025-05-07T05:15:48.849953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nprint(device)\n\nmodel.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:15:48.852129Z","iopub.execute_input":"2025-05-07T05:15:48.852533Z","iopub.status.idle":"2025-05-07T05:15:49.250247Z","shell.execute_reply.started":"2025-05-07T05:15:48.852490Z","shell.execute_reply":"2025-05-07T05:15:49.249215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.manual_seed(32)\n\nnum_epochs = 15\nprint( \"LETS START !! \")\n# below, train_losses is initialized as a NumPy array with zeros.\n#Predefined Size: This method creates a fixed-size array with a size of num_epochs,\n#which is useful if you want to use array indexing for storing loss values.\n\ntrain_losses = np.zeros(num_epochs)  # 0 0 0 0 0 ..stored in this format\nval_losses = np.zeros(num_epochs)\ntrain_accs = np.zeros(num_epochs)\nval_accs = np.zeros(num_epochs)\n# Initialize batch counter\nbatch_counter = 0\nval_counter = 0\n    \nfor it in range(num_epochs) :\n\n    # Define some variabels to store metrics in each epoch\n    train_loss = []\n    val_loss = []\n    n_correct = 0\n    n_total = 0\n    # Change mode to TRAIN\n    model.train()\n    t0 = datetime.now()\n\n    # Fetch images from DataLoader\n    for images, labels in train_loader :\n        batch_counter += 1 \n        # Move data to GPU\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        # Zero Grad optimizer\n        optimizer.zero_grad()\n\n        # Forward pass\n        y_pred = model(images)\n        loss = criterion(y_pred, labels)\n\n        # Backward pass\n        loss.backward()\n        optimizer.step()\n\n        # Train loss\n        train_loss.append(loss.item())\n\n        # Train Accuracy Calculate\n        _, prediction = torch.max(y_pred, 1)\n        n_correct += (prediction==labels).sum().item()\n        n_total += labels.shape[0]\n        \n        if batch_counter % 50 == 0:\n            print(f'Epoch [{it + 1}/{num_epochs}], Batch [{batch_counter}/{len(train_loader)}] -> '\n                  f'Train Loss: {np.mean(train_loss):.4f}, Train Acc: {n_correct / n_total:.4f}')\n\n    train_loss = np.mean(train_loss)\n    train_losses[it] = train_loss\n    train_accs[it] = n_correct / n_total\n    print(\"train_loader done..starting val_loader\")\n    # Validation\n    for images, labels in val_loader :\n        val_counter += 1 \n         # Move data to GPU\n        images = images.to(device)\n        labels = labels.to(device)\n        # Forward padd\n        y_pred = model(images)\n        loss = criterion(y_pred, labels)\n\n        # Validation Loss\n        val_loss.append(loss.item())\n\n        # Validation Accuracy\n        _, prediction = torch.max(y_pred, 1)\n        n_correct += (prediction==labels).sum().item()\n        n_total += labels.shape[0]\n        if val_counter % 30 == 0:\n            print(f'Epoch [{it + 1}/{num_epochs}], Batch [{val_counter}/{len(val_loader)}] -> '\n                  f'Val Loss: {np.mean(val_loss):.4f}, Val Acc: {n_correct / n_total:.4f}')\n\n    val_loss = np.mean(val_loss)\n    val_losses[it] = val_loss\n    val_accs[it] = n_correct / n_total\n\n    dt = datetime.now() - t0\n    # Print the result of  each epochs\n    print(f'Epoch [{it+1}/{num_epochs}] -> '\n          f' Train Loss:{train_loss:.4f}, Train Acc:{train_accs[it]:.4f} '  #  using indexing WOW !!\n          f' |Val Loss:{val_loss:.4f},   Val Acc:{val_accs[it]:.4f}  | | Duration : {dt}')\n    batch_counter=0\n    val_counter=0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T05:16:13.696290Z","iopub.execute_input":"2025-05-07T05:16:13.696960Z","iopub.status.idle":"2025-05-07T08:25:41.713108Z","shell.execute_reply.started":"2025-05-07T05:16:13.696932Z","shell.execute_reply":"2025-05-07T08:25:41.711928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# 训练循环后绘图：\n\n# 绘制训练和验证的损失曲线\nplt.figure(figsize=(12, 6))\n\n# 训练损失\nplt.subplot(1, 2, 1)\nplt.plot(train_losses, label='Train Loss', color='blue', marker='o')\nplt.plot(val_losses, label='Validation Loss', color='red', marker='x')\nplt.title('The change of loss function with training epochs')\nplt.xlabel('Round by round')\nplt.ylabel('Loss')\nplt.legend()\n\n# 训练准确率\nplt.subplot(1, 2, 2)\nplt.plot(train_accs, label='Train Accuracy', color='blue', marker='o')\nplt.plot(val_accs, label='Validation Accuracy', color='red', marker='x')\nplt.title('Accuracy varies with training epochs')\nplt.xlabel('Round by round')\nplt.ylabel('Acc')\nplt.legend()\n\n# 显示图表\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T08:29:54.540989Z","iopub.execute_input":"2025-05-07T08:29:54.541626Z","iopub.status.idle":"2025-05-07T08:29:55.261526Z","shell.execute_reply.started":"2025-05-07T08:29:54.541592Z","shell.execute_reply":"2025-05-07T08:29:55.260440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 导入必要的库\nfrom sklearn.metrics import confusion_matrix\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# 定义类别名称映射（直接使用你提供的字典）\nclass_mapping = {\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}\n\n# 将字典键转换为整数后排序，确保顺序正确\nsorted_classes = sorted(class_mapping.items(), key=lambda x: int(x[0]))\nclass_names = [item[1] for item in sorted_classes]  # 最终顺序：0-4对应的名称\n\n# 准备绘制混淆矩阵\nmodel.eval()\nall_labels = []\nall_predictions = []\n\nwith torch.no_grad():\n    for images, labels in val_loader:\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        outputs = model(images)\n        _, preds = torch.max(outputs, 1)\n        \n        all_labels.extend(labels.cpu().numpy())\n        all_predictions.extend(preds.cpu().numpy())\n\n# 计算混淆矩阵\ncm = confusion_matrix(all_labels, all_predictions)\n\n# 可视化（使用调整后的参数）\nplt.figure(figsize=(12, 10))  # 增大画布尺寸以适应长标签\nsns.set(font_scale=1.2)  # 增大字体大小\n\nax = sns.heatmap(\n    cm, \n    annot=True,\n    fmt='d',\n    cmap='Blues',\n    xticklabels=class_names,\n    yticklabels=class_names\n)\n\n# 调整标签显示\nax.set_xticklabels(class_names, rotation=45, ha='right')  # 旋转x轴标签\nax.set_yticklabels(class_names, rotation=0)  # 保持y轴标签水平\nplt.xlabel('Predicted Labels', fontsize=14)\nplt.ylabel('True Labels', fontsize=14)\nplt.title('Confusion Matrix with Class Names', pad=20, fontsize=16)  # 添加标题间距\n\n# 调整布局防止标签被截断\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T08:30:01.573224Z","iopub.execute_input":"2025-05-07T08:30:01.573652Z","iopub.status.idle":"2025-05-07T08:32:30.447396Z","shell.execute_reply.started":"2025-05-07T08:30:01.573625Z","shell.execute_reply":"2025-05-07T08:32:30.446434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"MODEL Trrained now lets predict !\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T08:38:55.497321Z","iopub.execute_input":"2025-05-07T08:38:55.498317Z","iopub.status.idle":"2025-05-07T08:38:55.503258Z","shell.execute_reply.started":"2025-05-07T08:38:55.498274Z","shell.execute_reply":"2025-05-07T08:38:55.502189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save only the model's state dictionary\n# torch.save(model.state_dict(), 'my_VGG16_model.pth')\n\ntorch.save(model, 'my_VGG16_full_model.pth')\ntorch.save(model.state_dict(), 'my_VGG16_full_model000.pth')\n\nprint(\"Model saved successfully after training!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T08:38:57.680584Z","iopub.execute_input":"2025-05-07T08:38:57.680942Z","iopub.status.idle":"2025-05-07T08:38:59.971498Z","shell.execute_reply.started":"2025-05-07T08:38:57.680915Z","shell.execute_reply":"2025-05-07T08:38:59.970403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.onnx\n\n# 定义模型结构\nclass CassavaLeafDiseaseModel(nn.Module):\n    def __init__(self):\n        super(CassavaLeafDiseaseModel, self).__init__()\n        # 根据实际模型结构定义你的网络层\n        # 例如：\n        self.conv1 = nn.Conv2d(3, 16, kernel_size=3, stride=1, padding=1)\n        self.relu = nn.ReLU()\n        self.maxpool = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.fc = nn.Linear(16 * 56 * 56, 5)  # 假设有5个类别\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n        return x\n\n# 加载训练好的模型权重\nmodel = torch.load('/kaggle/working/my_VGG16_full_model.pth')  # 替换为你的模型权重文件路径\nmodel.eval()\n\n# 确保模型在CPU上\nmodel = model.cpu()\n\n# 创建一个示例输入，并确保它在CPU上\ndummy_input = torch.randn(1, 3, 224, 224, device='cpu')  # 假设输入是224x224的RGB图像\n\n# 导出为ONNX格式\ntorch.onnx.export(\n    model,                        # 要导出的模型\n    dummy_input,                  # 模型的输入\n    \"cassava_leaf_disease_model.onnx\",  # 输出的ONNX文件名\n    input_names=['input'],         # 输入的名称\n    output_names=['output'],       # 输出的名称\n    dynamic_axes={'input': {0: 'batch_size'},    # 批次大小是动态的\n                  'output': {0: 'batch_size'}})\n\nprint(\"模型已成功导出为ONNX格式\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T08:39:15.806475Z","iopub.execute_input":"2025-05-07T08:39:15.806853Z","iopub.status.idle":"2025-05-07T08:39:23.380212Z","shell.execute_reply.started":"2025-05-07T08:39:15.806826Z","shell.execute_reply":"2025-05-07T08:39:23.379140Z"}},"outputs":[],"execution_count":null}]}