{"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":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport cv2\nimport os \nimport torch.nn as nn\nfrom torch.utils.data import Dataset ,DataLoader, Subset\nimport torch.optim as optim\nimport numpy as np\nimport torch\nfrom torchvision import transforms\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom torch.utils.data import random_split\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-15T15:34:37.959926Z","iopub.execute_input":"2024-05-15T15:34:37.960377Z","iopub.status.idle":"2024-05-15T15:34:43.210029Z","shell.execute_reply.started":"2024-05-15T15:34:37.960319Z","shell.execute_reply":"2024-05-15T15:34:43.209137Z"},"trusted":true},"execution_count":2,"outputs":[]},{"cell_type":"code","source":"data_dir = '/kaggle/input/cassava-leaf-disease-classification'\nbatch_size = 32\nepochs = 30\nnum_classes = 5\nlr_rate = 0.05\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:34:43.212075Z","iopub.execute_input":"2024-05-15T15:34:43.212866Z","iopub.status.idle":"2024-05-15T15:34:43.239743Z","shell.execute_reply.started":"2024-05-15T15:34:43.212833Z","shell.execute_reply":"2024-05-15T15:34:43.23847Z"},"trusted":true},"execution_count":3,"outputs":[]},{"cell_type":"code","source":"os.listdir(data_dir)\nprint('Train images: %d' %len(os.listdir(os.path.join(data_dir, \"train_images\"))))","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:34:43.240921Z","iopub.execute_input":"2024-05-15T15:34:43.241231Z","iopub.status.idle":"2024-05-15T15:34:43.696971Z","shell.execute_reply.started":"2024-05-15T15:34:43.241206Z","shell.execute_reply":"2024-05-15T15:34:43.69585Z"},"trusted":true},"execution_count":4,"outputs":[{"name":"stdout","text":"Train images: 21397\n","output_type":"stream"}]},{"cell_type":"code","source":"train_labels = pd.read_csv(os.path.join(data_dir, \"train.csv\"))\ntrain_labels.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:34:43.699943Z","iopub.execute_input":"2024-05-15T15:34:43.700776Z","iopub.status.idle":"2024-05-15T15:34:43.745982Z","shell.execute_reply.started":"2024-05-15T15:34:43.700736Z","shell.execute_reply":"2024-05-15T15:34:43.744842Z"},"trusted":true},"execution_count":5,"outputs":[{"execution_count":5,"output_type":"execute_result","data":{"text/plain":"         image_id  label\n0  1000015157.jpg      0\n1  1000201771.jpg      3\n2   100042118.jpg      1\n3  1000723321.jpg      1\n4  1000812911.jpg      3","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>image_id</th>\n      <th>label</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>1000015157.jpg</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>1000201771.jpg</td>\n      <td>3</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>100042118.jpg</td>\n      <td>1</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>1000723321.jpg</td>\n      <td>1</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>1000812911.jpg</td>\n      <td>3</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}]},{"cell_type":"code","source":"# Design a specific Dataset to read local data\nclass CustomDataset(Dataset):\n    # Initialize the dataset with the CSV file, root directory, and optional transform\n    def __init__(self, csv_file, root_dir, transform=None):\n        self.data_frame = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.transform = transform\n        \n    # Return the number of samples in the dataset\n    def __len__(self):\n        return len(self.data_frame)\n\n    # Get the file name and label for the specified index\n    def __getitem__(self, idx):\n        img_name = self.data_frame.iloc[idx, 0]\n        img_path = os.path.join(self.root_dir, img_name)\n        image = Image.open(img_path).convert('RGB')\n\n        label = int(self.data_frame.iloc[idx, 1])\n        # Apply the specified transform to the image if it exists\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n# Define data transform\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n])\n# Read images from local file\ndataset = CustomDataset(csv_file=os.path.join(data_dir, \"train.csv\"), root_dir=os.path.join(data_dir, \"train_images\"), transform=transform)\n\n# Calculate the dataset size\ndataset_size = len(dataset)\n\n# Define the ratio\ntrain_ratio = 0.8\nval_ratio = 0.1\ntest_ratio = 0.1\n\n# Calculate the size of every dataset\ntrain_size = int(train_ratio * dataset_size)\nval_size = int(val_ratio * dataset_size)\ntest_size = dataset_size - train_size - val_size\n\n# Split the dataset\ntrain_dataset, val_dataset, test_dataset = random_split(dataset, [train_size, val_size, test_size])\n\n# Use DataLoader to load data\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n\n# Print the size of dataset\nprint(f'Train dataset size: {len(train_dataset)}')\nprint(f'Validation dataset size: {len(val_dataset)}')\nprint(f'Test dataset size: {len(test_dataset)}')","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:34:43.747287Z","iopub.execute_input":"2024-05-15T15:34:43.74767Z","iopub.status.idle":"2024-05-15T15:34:43.798259Z","shell.execute_reply.started":"2024-05-15T15:34:43.747641Z","shell.execute_reply":"2024-05-15T15:34:43.797273Z"},"trusted":true},"execution_count":6,"outputs":[{"name":"stdout","text":"Train dataset size: 17117\nValidation dataset size: 2139\nTest dataset size: 2141\n","output_type":"stream"}]},{"cell_type":"code","source":"# Define a simple Convolutional Neural Network (CNN) class\nclass SimpleCNN(nn.Module):\n    # init()：Start initialize\n    def __init__(self, num_classes):\n        super(SimpleCNN, self).__init__()\n        # First convolutional layer\n        self.conv1 = nn.Conv2d(3, 32, kernel_size=3, padding=1)\n        # Use ReLU as the activation function\n        self.relu = nn.ReLU(inplace=True)\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n        # Second convolutional layer\n        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)\n        # Fully connected layers for classification\n        self.fc1 = nn.Linear(64 * 56 * 56, 64)\n        self.fc2 = nn.Linear(64, num_classes)\n\n    def forward(self, x):\n        # Move the input tensor to the device (GPU or CPU)\n        x = x.to(device)\n        # First convolutional layer\n        x = self.conv1(x)\n        x = self.relu(x)\n        x = self.pool(x)\n        # Second convolutional layer\n        x = self.conv2(x)\n        x = self.relu(x)\n        x = self.pool(x)\n        # Flatten the output for the fully connected layers\n        x = x.view(x.size(0), -1)\n        x = self.fc1(x)\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:34:43.799441Z","iopub.execute_input":"2024-05-15T15:34:43.799759Z","iopub.status.idle":"2024-05-15T15:34:43.809652Z","shell.execute_reply.started":"2024-05-15T15:34:43.799733Z","shell.execute_reply":"2024-05-15T15:34:43.808561Z"},"trusted":true},"execution_count":7,"outputs":[]},{"cell_type":"code","source":"# Design ResNet\nclass BasicBlock(nn.Module):\n    # expansion refers to the multiple of decreasing the scale to increase the dimension in each small residual block\n    expansion = 1\n \n    # init()：Start initialize\n    def __init__(self, in_channel, out_channel, stride=1, downsample=None, **kwargs):\n        super(BasicBlock, self).__init__()\n        self.conv1 = nn.Conv2d(in_channels=in_channel, out_channels=out_channel,\n                               kernel_size=3, stride=stride, padding=1, bias=False)\n        # Use batch normalization\n        self.bn1 = nn.BatchNorm2d(out_channel)\n        # Use ReLU as the activation function\n        self.relu = nn.ReLU()\n        self.conv2 = nn.Conv2d(in_channels=out_channel, out_channels=out_channel,\n                               kernel_size=3, stride=1, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channel)\n        self.downsample = downsample\n \n    # forward()：The forward propagation process is defined and the connections between the layers are described\n    def forward(self, x):\n        # The residual block retains the original input\n        identity = x\n        # In the case of a dashed residual structure, downsampling is performed\n        if self.downsample is not None:\n            identity = self.downsample(x)\n \n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n        # -----------------------------------------\n        out = self.conv2(out)\n        out = self.bn2(out)\n        # The main branch and the shortcut branch data are added\n        out += identity\n        out = self.relu(out)\n \n        return out\n \n \n# Define the residual structure of ResNet50/101/152\nclass Bottleneck(nn.Module):\n    # expansion refers to the multiple of decreasing the scale to increase the dimension in each small residual block\n    expansion = 4\n \n    # init()：Start initialize\n    def __init__(self, in_channel, out_channel, stride=1, downsample=None,\n                 groups=1, width_per_group=64):\n        super(Bottleneck, self).__init__()\n \n        width = int(out_channel * (width_per_group / 64.)) * groups\n \n        self.conv1 = nn.Conv2d(in_channels=in_channel, out_channels=width,\n                               kernel_size=1, stride=1, bias=False)\n        # Use batch normalization\n        self.bn1 = nn.BatchNorm2d(width)\n        # -----------------------------------------\n        self.conv2 = nn.Conv2d(in_channels=width, out_channels=width, groups=groups,\n                               kernel_size=3, stride=stride, bias=False, padding=1)\n        self.bn2 = nn.BatchNorm2d(width)\n        # -----------------------------------------\n        self.conv3 = nn.Conv2d(in_channels=width, out_channels=out_channel * self.expansion,\n                               kernel_size=1, stride=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(out_channel * self.expansion)\n        # Use ReLU as the activation function\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n \n    # forward()：The forward propagation process is defined and the connections between the layers are described\n    def forward(self, x):\n        # The residual block retains the original input\n        identity = x\n        # In the case of a dashed residual structure, downsampling is performed\n        if self.downsample is not None:\n            identity = self.downsample(x)\n \n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n \n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n \n        out = self.conv3(out)\n        out = self.bn3(out)\n        # The main branch and the shortcut branch data are added\n        out += identity\n        out = self.relu(out)\n \n        return out\n \n \n# Define ResNet class\nclass ResNet(nn.Module):\n    # initialize the function\n    def __init__(self,\n                 block,\n                 blocks_num,\n                 num_classes=1000,\n                 include_top=True,\n                 groups=1,\n                 width_per_group=64):\n        super(ResNet, self).__init__()\n        self.include_top = include_top\n        # The maxpool has 64 output channels and 64 residual structure input channels\n        self.in_channel = 64\n \n        self.groups = groups\n        self.width_per_group = width_per_group\n \n        self.conv1 = nn.Conv2d(3, self.in_channel, kernel_size=7, stride=2,\n                               padding=3, bias=False)\n        self.bn1 = nn.BatchNorm2d(self.in_channel)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        # Shallow stride=1, deep stride=2\n        # block：Two types of residual modules are defined\n        # block_num：The number of residual blocks in the module\n        self.layer1 = self._make_layer(block, 64, blocks_num[0])\n        self.layer2 = self._make_layer(block, 128, blocks_num[1], stride=2)\n        self.layer3 = self._make_layer(block, 256, blocks_num[2], stride=2)\n        self.layer4 = self._make_layer(block, 512, blocks_num[3], stride=2)\n        if self.include_top:\n            # Adaptive average pooling, with specified outputs (H, W) and no change in the number of channels\n            self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n            # Fully connected layer\n            self.fc = nn.Linear(512 * block.expansion, num_classes)\n        # Inherit nn. Module class, self.modules(), which returns all modules in the network\n        for m in self.modules():\n            # if convolution layer\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n \n    # Define the Residuals module, which consists of several residual blocks\n    def _make_layer(self, block, channel, block_num, stride=1):\n        downsample = None\n        if stride != 1 or self.in_channel != channel * block.expansion:\n            downsample = nn.Sequential(\n                nn.Conv2d(self.in_channel, channel * block.expansion, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(channel * block.expansion))\n \n        layers = []\n        layers.append(block(self.in_channel,\n                            channel,\n                            downsample=downsample,\n                            stride=stride,\n                            groups=self.groups,\n                            width_per_group=self.width_per_group))\n        self.in_channel = channel * block.expansion\n \n        for _ in range(1, block_num):\n            layers.append(block(self.in_channel,\n                                channel,\n                                groups=self.groups,\n                                width_per_group=self.width_per_group))\n        # Sequential：Custom sequences are connected into models to generate network structures\n        return nn.Sequential(*layers)\n \n    # forward()：The forward propagation process is defined and the connections between the layers are described\n    def forward(self, x):\n        # Static layer\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n        # Dynamic layers\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n \n        if self.include_top:\n            x = self.avgpool(x)\n            x = torch.flatten(x, 1)\n            x = self.fc(x)\n \n        return x\n\n# ResNet50\ndef resnet50(num_classes=num_classes, include_top=True):\n    return ResNet(Bottleneck, [3, 4, 6, 3], num_classes=num_classes, include_top=include_top)\n\n# ResNet101\ndef resnet101(num_classes=num_classes, include_top=True):\n    return ResNet(Bottleneck, [3, 4, 23, 3], num_classes=num_classes, include_top=include_top)\n\n# ResNet152\ndef resnet152(num_classes=num_classes, include_top=True):\n    return ResNet(Bottleneck, [3, 8, 36, 3], num_classes=num_classes, include_top=include_top)","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:35:43.693044Z","iopub.execute_input":"2024-05-15T15:35:43.693411Z","iopub.status.idle":"2024-05-15T15:35:43.73166Z","shell.execute_reply.started":"2024-05-15T15:35:43.693385Z","shell.execute_reply":"2024-05-15T15:35:43.73028Z"},"trusted":true},"execution_count":10,"outputs":[]},{"cell_type":"code","source":"# choose the model you want\nmodel = resnet50(num_classes).to(device)\n# design the optimizers\noptimizer1 = optim.Adam(model.parameters(), lr=lr_rate)\noptimizer2 = optim.SGD(model.parameters(), lr=lr_rate, momentum=0.9)\n\n# Validate step\ndef validate(model, val_loader):\n    # Set the model to evaluation mode\n    model.eval()\n    correct = 0\n    total = 0\n    # Use tqdm for a progress bar during training\n    progress_bar = tqdm(enumerate(val_loader), total=len(val_loader))\n    with torch.no_grad():\n        for batch_idx, (images, labels) in progress_bar:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images)\n            _, predicted = torch.max(outputs, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n            \n    # Calculate accuracy and print the result\n    accuracy = correct / total\n    print(\"\\n Evaluation accuracy: {}\".format(accuracy))\n    return accuracy\n\n# Test step\ndef test(model, test_loader):\n    # Set the model to evaluation mode\n    model.eval()\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        for images, labels in test_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images)\n            _, predicted = torch.max(outputs, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n            \n    # Calculate accuracy and print the result\n    accuracy = correct / total\n    print(f'Test Accuracy: {accuracy * 100:.2f}%')\n\n# Train step\ndef train(model, train_loader, val_loader, test_loader, num_epochs=5, optim = optimizer1):\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim\n    # Use a cyclic learning rate scheduler for adaptive learning rates\n    scheduler = torch.optim.lr_scheduler.CyclicLR(optimizer, base_lr=0.005, max_lr=0.05)\n    trn_acc_hist = []\n    val_acc_hist = []\n    \n    for epoch in range(num_epochs):\n        total_loss = 0.0\n        # Use tqdm for a progress bar during training\n        progress_bar = tqdm(enumerate(train_loader), total=len(train_loader), desc=f'Epoch {epoch+1}/{num_epochs}')\n        for batch_idx, (images, labels) in progress_bar:\n            images = images.to(device)\n            labels = labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            total_loss += loss.item()\n            progress_bar.set_postfix({'Loss': total_loss / (batch_idx + 1)})\n        \n        # Print average loss at the end of each epoch\n        print(f'Epoch {epoch+1}/{num_epochs}, Average Loss: {total_loss/len(train_loader)}')\n        # Record training and validation accuracy history\n        trn_acc_hist.append(validate(model, train_loader))\n        # Validation is performed at the end of each epoch\n        print(\"\\n Evaluate on validation set...\")\n        val_acc_hist.append(validate(model, val_loader))\n\n    # Test at the end of the training session\n    print(\"\\nTraining completed. Starting test evaluation.\")\n    test(model, test_loader)\n    return trn_acc_hist, val_acc_hist\n\n# Start training\nprint(\"Training started...\")\ntrn_acc_hist, val_acc_hist = train(model, train_loader, val_loader, test_loader, epochs, optimizer2)","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:34:43.851891Z","iopub.execute_input":"2024-05-15T15:34:43.852644Z","iopub.status.idle":"2024-05-15T15:35:01.096207Z","shell.execute_reply.started":"2024-05-15T15:34:43.852617Z","shell.execute_reply":"2024-05-15T15:35:01.094343Z"},"trusted":true},"execution_count":9,"outputs":[{"name":"stdout","text":"Training started...\n","output_type":"stream"},{"name":"stderr","text":"Epoch 1/30:   4%|▎         | 19/535 [00:15<07:03,  1.22it/s, Loss=2.6] \n","output_type":"stream"},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mKeyboardInterrupt\u001b[0m                         Traceback (most recent call last)","Cell \u001b[0;32mIn[9], line 89\u001b[0m\n\u001b[1;32m     87\u001b[0m \u001b[38;5;66;03m# Start training\u001b[39;00m\n\u001b[1;32m     88\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mTraining started...\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n\u001b[0;32m---> 89\u001b[0m trn_acc_hist, val_acc_hist \u001b[38;5;241m=\u001b[39m \u001b[43mtrain\u001b[49m\u001b[43m(\u001b[49m\u001b[43mmodel\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mtrain_loader\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mval_loader\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mtest_loader\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mepochs\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43moptimizer2\u001b[49m\u001b[43m)\u001b[49m\n","Cell \u001b[0;32mIn[9], line 62\u001b[0m, in \u001b[0;36mtrain\u001b[0;34m(model, train_loader, val_loader, test_loader, num_epochs, optim)\u001b[0m\n\u001b[1;32m     60\u001b[0m \u001b[38;5;66;03m# Use tqdm for a progress bar during training\u001b[39;00m\n\u001b[1;32m     61\u001b[0m progress_bar \u001b[38;5;241m=\u001b[39m tqdm(\u001b[38;5;28menumerate\u001b[39m(train_loader), total\u001b[38;5;241m=\u001b[39m\u001b[38;5;28mlen\u001b[39m(train_loader), desc\u001b[38;5;241m=\u001b[39m\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mEpoch \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mepoch\u001b[38;5;241m+\u001b[39m\u001b[38;5;241m1\u001b[39m\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m/\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mnum_epochs\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m)\n\u001b[0;32m---> 62\u001b[0m \u001b[38;5;28;01mfor\u001b[39;00m batch_idx, (images, labels) \u001b[38;5;129;01min\u001b[39;00m progress_bar:\n\u001b[1;32m     63\u001b[0m     images \u001b[38;5;241m=\u001b[39m images\u001b[38;5;241m.\u001b[39mto(device)\n\u001b[1;32m     64\u001b[0m     labels \u001b[38;5;241m=\u001b[39m labels\u001b[38;5;241m.\u001b[39mto(device)\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/tqdm/std.py:1182\u001b[0m, in \u001b[0;36mtqdm.__iter__\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m   1179\u001b[0m time \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_time\n\u001b[1;32m   1181\u001b[0m \u001b[38;5;28;01mtry\u001b[39;00m:\n\u001b[0;32m-> 1182\u001b[0m     \u001b[38;5;28;01mfor\u001b[39;00m obj \u001b[38;5;129;01min\u001b[39;00m iterable:\n\u001b[1;32m   1183\u001b[0m         \u001b[38;5;28;01myield\u001b[39;00m obj\n\u001b[1;32m   1184\u001b[0m         \u001b[38;5;66;03m# Update and possibly print the progressbar.\u001b[39;00m\n\u001b[1;32m   1185\u001b[0m         \u001b[38;5;66;03m# Note: does not call self.update(1) for speed optimisation.\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataloader.py:630\u001b[0m, in \u001b[0;36m_BaseDataLoaderIter.__next__\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    627\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_sampler_iter \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m:\n\u001b[1;32m    628\u001b[0m     \u001b[38;5;66;03m# TODO(https://github.com/pytorch/pytorch/issues/76750)\u001b[39;00m\n\u001b[1;32m    629\u001b[0m     \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_reset()  \u001b[38;5;66;03m# type: ignore[call-arg]\u001b[39;00m\n\u001b[0;32m--> 630\u001b[0m data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_next_data\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m    631\u001b[0m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_num_yielded \u001b[38;5;241m+\u001b[39m\u001b[38;5;241m=\u001b[39m \u001b[38;5;241m1\u001b[39m\n\u001b[1;32m    632\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_dataset_kind \u001b[38;5;241m==\u001b[39m _DatasetKind\u001b[38;5;241m.\u001b[39mIterable \u001b[38;5;129;01mand\u001b[39;00m \\\n\u001b[1;32m    633\u001b[0m         \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_IterableDataset_len_called \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;129;01mnot\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m \u001b[38;5;129;01mand\u001b[39;00m \\\n\u001b[1;32m    634\u001b[0m         \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_num_yielded \u001b[38;5;241m>\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_IterableDataset_len_called:\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataloader.py:674\u001b[0m, in \u001b[0;36m_SingleProcessDataLoaderIter._next_data\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    672\u001b[0m \u001b[38;5;28;01mdef\u001b[39;00m \u001b[38;5;21m_next_data\u001b[39m(\u001b[38;5;28mself\u001b[39m):\n\u001b[1;32m    673\u001b[0m     index \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_next_index()  \u001b[38;5;66;03m# may raise StopIteration\u001b[39;00m\n\u001b[0;32m--> 674\u001b[0m     data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_dataset_fetcher\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mfetch\u001b[49m\u001b[43m(\u001b[49m\u001b[43mindex\u001b[49m\u001b[43m)\u001b[49m  \u001b[38;5;66;03m# may raise StopIteration\u001b[39;00m\n\u001b[1;32m    675\u001b[0m     \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_pin_memory:\n\u001b[1;32m    676\u001b[0m         data \u001b[38;5;241m=\u001b[39m _utils\u001b[38;5;241m.\u001b[39mpin_memory\u001b[38;5;241m.\u001b[39mpin_memory(data, \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_pin_memory_device)\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/_utils/fetch.py:49\u001b[0m, in \u001b[0;36m_MapDatasetFetcher.fetch\u001b[0;34m(self, possibly_batched_index)\u001b[0m\n\u001b[1;32m     47\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mauto_collation:\n\u001b[1;32m     48\u001b[0m     \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mhasattr\u001b[39m(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdataset, \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124m__getitems__\u001b[39m\u001b[38;5;124m\"\u001b[39m) \u001b[38;5;129;01mand\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdataset\u001b[38;5;241m.\u001b[39m__getitems__:\n\u001b[0;32m---> 49\u001b[0m         data \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mdataset\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m__getitems__\u001b[49m\u001b[43m(\u001b[49m\u001b[43mpossibly_batched_index\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m     50\u001b[0m     \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[1;32m     51\u001b[0m         data \u001b[38;5;241m=\u001b[39m [\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdataset[idx] \u001b[38;5;28;01mfor\u001b[39;00m idx \u001b[38;5;129;01min\u001b[39;00m possibly_batched_index]\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataset.py:364\u001b[0m, in \u001b[0;36mSubset.__getitems__\u001b[0;34m(self, indices)\u001b[0m\n\u001b[1;32m    362\u001b[0m     \u001b[38;5;28;01mreturn\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdataset\u001b[38;5;241m.\u001b[39m__getitems__([\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mindices[idx] \u001b[38;5;28;01mfor\u001b[39;00m idx \u001b[38;5;129;01min\u001b[39;00m indices])  \u001b[38;5;66;03m# type: ignore[attr-defined]\u001b[39;00m\n\u001b[1;32m    363\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[0;32m--> 364\u001b[0m     \u001b[38;5;28;01mreturn\u001b[39;00m [\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdataset[\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mindices[idx]] \u001b[38;5;28;01mfor\u001b[39;00m idx \u001b[38;5;129;01min\u001b[39;00m indices]\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/torch/utils/data/dataset.py:364\u001b[0m, in \u001b[0;36m<listcomp>\u001b[0;34m(.0)\u001b[0m\n\u001b[1;32m    362\u001b[0m     \u001b[38;5;28;01mreturn\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdataset\u001b[38;5;241m.\u001b[39m__getitems__([\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mindices[idx] \u001b[38;5;28;01mfor\u001b[39;00m idx \u001b[38;5;129;01min\u001b[39;00m indices])  \u001b[38;5;66;03m# type: ignore[attr-defined]\u001b[39;00m\n\u001b[1;32m    363\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[0;32m--> 364\u001b[0m     \u001b[38;5;28;01mreturn\u001b[39;00m [\u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mdataset\u001b[49m\u001b[43m[\u001b[49m\u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mindices\u001b[49m\u001b[43m[\u001b[49m\u001b[43midx\u001b[49m\u001b[43m]\u001b[49m\u001b[43m]\u001b[49m \u001b[38;5;28;01mfor\u001b[39;00m idx \u001b[38;5;129;01min\u001b[39;00m indices]\n","Cell \u001b[0;32mIn[6], line 17\u001b[0m, in \u001b[0;36mCustomDataset.__getitem__\u001b[0;34m(self, idx)\u001b[0m\n\u001b[1;32m     15\u001b[0m img_name \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdata_frame\u001b[38;5;241m.\u001b[39miloc[idx, \u001b[38;5;241m0\u001b[39m]\n\u001b[1;32m     16\u001b[0m img_path \u001b[38;5;241m=\u001b[39m os\u001b[38;5;241m.\u001b[39mpath\u001b[38;5;241m.\u001b[39mjoin(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mroot_dir, img_name)\n\u001b[0;32m---> 17\u001b[0m image \u001b[38;5;241m=\u001b[39m \u001b[43mImage\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mopen\u001b[49m\u001b[43m(\u001b[49m\u001b[43mimg_path\u001b[49m\u001b[43m)\u001b[49m\u001b[38;5;241m.\u001b[39mconvert(\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mRGB\u001b[39m\u001b[38;5;124m'\u001b[39m)\n\u001b[1;32m     19\u001b[0m label \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mint\u001b[39m(\u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mdata_frame\u001b[38;5;241m.\u001b[39miloc[idx, \u001b[38;5;241m1\u001b[39m])\n\u001b[1;32m     20\u001b[0m \u001b[38;5;66;03m# Apply the specified transform to the image if it exists\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/PIL/Image.py:3245\u001b[0m, in \u001b[0;36mopen\u001b[0;34m(fp, mode, formats)\u001b[0m\n\u001b[1;32m   3242\u001b[0m     fp \u001b[38;5;241m=\u001b[39m io\u001b[38;5;241m.\u001b[39mBytesIO(fp\u001b[38;5;241m.\u001b[39mread())\n\u001b[1;32m   3243\u001b[0m     exclusive_fp \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;01mTrue\u001b[39;00m\n\u001b[0;32m-> 3245\u001b[0m prefix \u001b[38;5;241m=\u001b[39m \u001b[43mfp\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mread\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;241;43m16\u001b[39;49m\u001b[43m)\u001b[49m\n\u001b[1;32m   3247\u001b[0m preinit()\n\u001b[1;32m   3249\u001b[0m accept_warnings \u001b[38;5;241m=\u001b[39m []\n","\u001b[0;31mKeyboardInterrupt\u001b[0m: "],"ename":"KeyboardInterrupt","evalue":"","output_type":"error"}]},{"cell_type":"code","source":"# Svae the trained model\ntorch.save(model.state_dict(),'./model_best.pth')","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:35:01.09717Z","iopub.status.idle":"2024-05-15T15:35:01.0976Z","shell.execute_reply.started":"2024-05-15T15:35:01.097399Z","shell.execute_reply":"2024-05-15T15:35:01.097415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ei part ta important \n\n# Plot the training and validation accuracy histories\n\nx = np.arange(epochs)\nplt.figure()\nplt.plot(x, trn_acc_hist)\nplt.plot(x, val_acc_hist)\nplt.legend(['Training', 'Validation'])\nplt.xticks(x)\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Classification')\nplt.gcf().set_size_inches(10, 5)\n# Save the plot as an image (PNG) with a specified DPI (dots per inch)\nplt.savefig('classify.png', dpi=300)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:35:01.099279Z","iopub.status.idle":"2024-05-15T15:35:01.099657Z","shell.execute_reply.started":"2024-05-15T15:35:01.099463Z","shell.execute_reply":"2024-05-15T15:35:01.099476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get an enumeration of the test loader\nexamples = enumerate(test_loader)\nbatch_idx, (example_data, example_targets) = next(examples)\n# Get predictions through the model\nwith torch.no_grad():\n    example_data = example_data.to(device)\n    output = model(example_data)\nfig = plt.figure()\n\n# Iterate over the first 9 examples in the batch\nfor i in range(9):\n    plt.subplot(3,3,i+1)\n    plt.tight_layout()\n    plt.imshow(example_data[i][0].cpu())\n    plt.title(\"Prediction: {}\".format(output.data.max(1, keepdim=True)[1][i].item()))\n    plt.xticks([])\n    plt.yticks([])\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-15T15:35:01.102421Z","iopub.status.idle":"2024-05-15T15:35:01.103112Z","shell.execute_reply.started":"2024-05-15T15:35:01.102834Z","shell.execute_reply":"2024-05-15T15:35:01.102855Z"},"trusted":true},"execution_count":null,"outputs":[]}]}