{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Importing Required Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport copy\nimport timm\nimport random\nimport time\nimport torch\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.optimizer\nimport concurrent.futures\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nfrom collections import OrderedDict\nfrom torch.utils.data import Dataset, DataLoader, Subset, random_split\nfrom torch.cuda import amp\nfrom torchvision import transforms as T\nfrom torchvision.io import read_image\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.metrics import accuracy_score, confusion_matrix, f1_score, classification_report\nfrom tqdm import tqdm\n\nprint(torch.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:52:15.062753Z","iopub.execute_input":"2025-06-10T14:52:15.063060Z","iopub.status.idle":"2025-06-10T14:52:25.459250Z","shell.execute_reply.started":"2025-06-10T14:52:15.063038Z","shell.execute_reply":"2025-06-10T14:52:25.458293Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Seeds for Reproducibility","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed):\n    \"\"\"\n    Sets seeds for reproducibility in training.\n\n    Args:\n        seed (int): Seed value to ensure determinism.\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)  # Seed for hash-based operations\n    np.random.seed(seed)  # Seed for NumPy\n    torch.manual_seed(seed)  # Seed for PyTorch (CPU)\n    torch.cuda.manual_seed(seed)  # Seed for PyTorch (GPU)\n    torch.backends.cudnn.deterministic = True  # Make CuDNN deterministic\n    torch.backends.cudnn.benchmark = False  # Enable benchmark mode for CuDNN","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:52:57.767776Z","iopub.execute_input":"2025-06-10T14:52:57.768305Z","iopub.status.idle":"2025-06-10T14:52:57.773031Z","shell.execute_reply.started":"2025-06-10T14:52:57.768276Z","shell.execute_reply":"2025-06-10T14:52:57.772044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:53:06.145708Z","iopub.execute_input":"2025-06-10T14:53:06.146048Z","iopub.status.idle":"2025-06-10T14:53:06.151813Z","shell.execute_reply.started":"2025-06-10T14:53:06.146019Z","shell.execute_reply":"2025-06-10T14:53:06.150989Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"data = pd.read_csv('../input/aptos2019-blindness-detection/train.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:53:14.483319Z","iopub.execute_input":"2025-06-10T14:53:14.483745Z","iopub.status.idle":"2025-06-10T14:53:14.500812Z","shell.execute_reply.started":"2025-06-10T14:53:14.483704Z","shell.execute_reply":"2025-06-10T14:53:14.499979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Number of samples: ', data.shape[0])\ndisplay(data.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:53:18.988122Z","iopub.execute_input":"2025-06-10T14:53:18.988440Z","iopub.status.idle":"2025-06-10T14:53:19.011378Z","shell.execute_reply.started":"2025-06-10T14:53:18.988414Z","shell.execute_reply":"2025-06-10T14:53:19.010711Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data['diagnosis'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:53:24.370430Z","iopub.execute_input":"2025-06-10T14:53:24.370868Z","iopub.status.idle":"2025-06-10T14:53:24.384449Z","shell.execute_reply.started":"2025-06-10T14:53:24.370831Z","shell.execute_reply":"2025-06-10T14:53:24.383540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"f, ax = plt.subplots(figsize=(14, 8.7))\nax = sns.countplot(x=\"diagnosis\", data=data, palette=\"GnBu_d\")\nsns.despine()\nplt.savefig('Distruption class', dpi=300, transparent=True)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:53:29.350895Z","iopub.execute_input":"2025-06-10T14:53:29.351263Z","iopub.status.idle":"2025-06-10T14:53:30.204244Z","shell.execute_reply.started":"2025-06-10T14:53:29.351236Z","shell.execute_reply":"2025-06-10T14:53:30.203053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Setting the style for the plot\nsns.set_style(\"white\")\n\n# Mapping class labels to their corresponding categories\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Plotting the first 15 images along with their labels\ncount = 1\nplt.figure(figsize=[20, 20])\n\nfor img_name in data['id_code'][:15]:  # Assuming 'train' contains the dataset\n    img = cv2.imread(f\"../input/aptos2019-blindness-detection/train_images/{img_name}.png\")[..., [2, 1, 0]]  # Reading the image\n    \n    # Getting the label (class) for the image\n    label = data[data['id_code'] == img_name]['diagnosis'].values[0]  # Assuming 'diagnosis' is the label column\n    \n    # Setting up the subplot with image and label\n    plt.subplot(5, 5, count)\n    plt.imshow(img)\n    plt.title(f\"Image {count}: {level_to_category[label]}\")  # Display the class label\n    count += 1\n    \n# Display the plot\nplt.savefig('/kaggle/working/imagebeforepreprecssing.png')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-10T14:53:47.792966Z","iopub.execute_input":"2025-06-10T14:53:47.793259Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"# Function to crop the image based on grayscale threshold\ndef crop_image_from_gray(img, tol=7):\n    if img.ndim == 2:\n        mask = img > tol\n        return img[np.ix_(mask.any(1), mask.any(0))]\n    elif img.ndim == 3:\n        gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        mask = gray_img > tol\n        \n        check_shape = img[:,:,0][np.ix_(mask.any(1), mask.any(0))].shape[0]\n        if check_shape == 0:  # Image is too dark so that we crop out everything\n            return img  # Return original image\n        else:\n            img1 = img[:,:,0][np.ix_(mask.any(1), mask.any(0))]\n            img2 = img[:,:,1][np.ix_(mask.any(1), mask.any(0))]\n            img3 = img[:,:,2][np.ix_(mask.any(1), mask.any(0))]\n            img = np.stack([img1, img2, img3], axis=-1)\n        return img","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T04:31:04.348242Z","iopub.execute_input":"2025-05-27T04:31:04.348468Z","iopub.status.idle":"2025-05-27T04:31:04.354602Z","shell.execute_reply.started":"2025-05-27T04:31:04.348447Z","shell.execute_reply":"2025-05-27T04:31:04.353838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set input and output directories\ninput_dir = '/kaggle/input/aptos2019-blindness-detection/train_images/'\noutput_dir = '/kaggle/working/processed_images/'\n\n# Ensure the output directory exists\nos.makedirs(output_dir, exist_ok=True)\n\n# Load the CSV containing image names and labels\ncsv_path = '/kaggle/input/aptos2019-blindness-detection/train.csv'\ndf = pd.read_csv(csv_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T04:31:04.356345Z","iopub.execute_input":"2025-05-27T04:31:04.35662Z","iopub.status.idle":"2025-05-27T04:31:04.374662Z","shell.execute_reply.started":"2025-05-27T04:31:04.356597Z","shell.execute_reply":"2025-05-27T04:31:04.373646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_image(row, sigmaX=10):\n    sample_image_id = row['id_code']\n    sample_image_file = sample_image_id + '.png'\n    sample_image_path = os.path.join(input_dir, sample_image_file)\n    \n    if os.path.exists(sample_image_path):\n        # Ben Graham's preprocessing\n        image = cv2.imread(sample_image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = crop_image_from_gray(image)\n        image = cv2.resize(image, (384, 384))\n        #image = cv2.addWeighted(image, 4, cv2.GaussianBlur(image, (0, 0), sigmaX), -4, 128)\n        \n        # Save the processed image to the output directory\n        output_path = os.path.join(output_dir, sample_image_file)\n        cv2.imwrite(output_path, cv2.cvtColor(image, cv2.COLOR_RGB2BGR))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T04:34:35.684945Z","iopub.execute_input":"2025-05-27T04:34:35.685357Z","iopub.status.idle":"2025-05-27T04:34:35.690399Z","shell.execute_reply.started":"2025-05-27T04:34:35.685325Z","shell.execute_reply":"2025-05-27T04:34:35.689569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Using ThreadPoolExecutor to process images in parallel\nwith concurrent.futures.ThreadPoolExecutor(max_workers=8) as executor:\n    list(tqdm(executor.map(process_image, [row for _, row in df.iterrows()]), total=df.shape[0], desc=\"Processing images\", unit=\"image\"))\n\nprint(\"Processing complete for all images.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T04:34:35.704077Z","iopub.execute_input":"2025-05-27T04:34:35.704408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Setting the style for the plot\nsns.set_style(\"white\")\n\n# Mapping class labels to their corresponding categories`\nlevel_to_category = {\n    0: \"No_DR\",\n    1: \"Mild\",\n    2: \"Moderate\",\n    3: \"Severe\",\n    4: \"Proliferate_DR\"\n}\n\n# Plotting the first 15 images along with their labels\ncount = 1\nplt.figure(figsize=[20, 20])\n\nfor img_name in data['id_code'][:15]:  # Assuming 'train' contains the dataset\n    img = cv2.imread(f\"/kaggle/working/processed_images/{img_name}.png\")[..., [2, 1, 0]]  # Reading the image\n    \n    # Getting the label (class) for the image\n    label = data[data['id_code'] == img_name]['diagnosis'].values[0]  # Assuming 'diagnosis' is the label column\n    \n    # Setting up the subplot with image and label\n    plt.subplot(5, 5, count)\n    plt.imshow(img)\n    plt.title(f\"Image {count}: {level_to_category[label]}\")  # Display the class label\n    count += 1\n\n# Display the plot\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Splitting and Transformation","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/working/processed_images\"\nCSV_PATH = \"/kaggle/input/aptos2019-blindness-detection/train.csv\"\nMODEL_PATH = \"./kaggle/working/\"\nLEARNING_RATE = 1e-4\nTRAIN_BATCH_SIZE = 16\nVALID_BATCH_SIZE = 16\nTEST_BATCH_SIZE = 16\nTRAIN_SPLIT = 0.7\nVAL_SPLIT = 0.2\nTEST_SPLIT = 0.1\nNUM_WORKERS = 4\nUSE_AMP = True\nEPOCHS = 20","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:12.272705Z","iopub.execute_input":"2025-05-19T13:01:12.273082Z","iopub.status.idle":"2025-05-19T13:01:12.278938Z","shell.execute_reply.started":"2025-05-19T13:01:12.273039Z","shell.execute_reply":"2025-05-19T13:01:12.27811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RetinopathyDataset(Dataset):\n    def __init__(self, image_dir, csv_file, transforms=None):\n        self.data = pd.read_csv(csv_file)\n        self.transforms = transforms\n        self.image_dir = image_dir\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        img_name = os.path.join(self.image_dir, self.data.loc[idx, 'id_code'] + '.png')\n\n        tensor_image = read_image(img_name)\n        label = torch.tensor(self.data.loc[idx, 'diagnosis'], dtype=torch.long)\n\n        if self.transforms is not None:\n            tensor_image = self.transforms(tensor_image)\n\n        return (tensor_image, label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:12.279816Z","iopub.execute_input":"2025-05-19T13:01:12.280024Z","iopub.status.idle":"2025-05-19T13:01:12.297539Z","shell.execute_reply.started":"2025-05-19T13:01:12.280006Z","shell.execute_reply":"2025-05-19T13:01:12.296754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Augmentation\ndata_transforms = T.Compose([\n    T.RandomResizedCrop(384, scale=(0.8, 1.0)),\n    T.RandomHorizontalFlip(p=0.5),\n    T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),\n    T.ConvertImageDtype(torch.float32),\n    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\n# Load Dataset\nfull_dataset = RetinopathyDataset(DATA_DIR, CSV_PATH, transforms=data_transforms)\nlabels = full_dataset.data['diagnosis'].values\n\n# Dataset size for train, val, test\ntotal_size = len(full_dataset)\ntrain_size = int(TRAIN_SPLIT * total_size)\nval_size = int(VAL_SPLIT * total_size)\ntest_size = total_size - train_size - val_size  # Sisa dataset untuk test\n\n# Split dataset\ntrain_idx, test_idx = train_test_split(np.arange(len(full_dataset)), test_size=TEST_SPLIT, \n                                       stratify=labels, random_state=42)\ntrain_idx, val_idx = train_test_split(train_idx, test_size=VAL_SPLIT/(1-TEST_SPLIT), \n                                      stratify=labels[train_idx], random_state=42)\n\n# Use Subset to divide the dataset\ntrain_dataset = Subset(full_dataset, train_idx)\nval_dataset = Subset(full_dataset, val_idx)\ntest_dataset = Subset(full_dataset, test_idx)\n\n# DataLoader for each dataset\ntrain_loader = DataLoader(train_dataset, batch_size=TRAIN_BATCH_SIZE, shuffle=True, \n                          num_workers=NUM_WORKERS, drop_last=True, pin_memory=True,\n                          prefetch_factor=2)\n\nval_loader = DataLoader(val_dataset, batch_size=VALID_BATCH_SIZE, shuffle=False, \n                        num_workers=NUM_WORKERS, drop_last=False, pin_memory=True,\n                          prefetch_factor=2)\n\ntest_loader = DataLoader(test_dataset, batch_size=TEST_BATCH_SIZE, shuffle=False, \n                         num_workers=NUM_WORKERS, drop_last=False, pin_memory=True,\n                          prefetch_factor=2)\n\nprint(f\"Total Dataset: {total_size}\")\nprint(f\"Train Set: {train_size} samples\")\nprint(f\"Validation Set: {val_size} samples\")\nprint(f\"Test Set: {test_size} samples\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:12.298364Z","iopub.execute_input":"2025-05-19T13:01:12.29863Z","iopub.status.idle":"2025-05-19T13:01:12.328236Z","shell.execute_reply.started":"2025-05-19T13:01:12.298605Z","shell.execute_reply":"2025-05-19T13:01:12.327531Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Fine Tune the Model","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:12.32893Z","iopub.execute_input":"2025-05-19T13:01:12.329139Z","iopub.status.idle":"2025-05-19T13:01:12.423775Z","shell.execute_reply.started":"2025-05-19T13:01:12.329121Z","shell.execute_reply":"2025-05-19T13:01:12.423094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# List of model\nMODEL_NAME = \"convnextv2_base.fcmae_ft_in22k_in1k_384\"\nMODEL_SAVE = \"/kaggle/working/deit_training_results.csv\"\nCHECKPOINT_PATH = f\"/kaggle/working/{MODEL_NAME}_best.pt\"\n\n# Load Model\nmodel = timm.create_model(MODEL_NAME, pretrained=True, num_classes=5)\nmodel.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:12.424707Z","iopub.execute_input":"2025-05-19T13:01:12.424993Z","iopub.status.idle":"2025-05-19T13:01:15.948481Z","shell.execute_reply.started":"2025-05-19T13:01:12.42494Z","shell.execute_reply":"2025-05-19T13:01:15.947632Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Handle Imbalance","metadata":{}},{"cell_type":"code","source":"# Ambil label\nlabels = data['diagnosis'].values\n\n# Hitung class weights\nclasses = sorted(data['diagnosis'].unique())\nclass_weights = compute_class_weight(class_weight='balanced', classes=classes, y=labels)\n\n# Ubah ke tensor\nclass_weights_tensor = torch.tensor(class_weights, dtype=torch.float).to(device)\n\nprint(f\"Class Weights: {class_weights_tensor}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:15.949404Z","iopub.execute_input":"2025-05-19T13:01:15.949681Z","iopub.status.idle":"2025-05-19T13:01:16.221884Z","shell.execute_reply.started":"2025-05-19T13:01:15.949651Z","shell.execute_reply":"2025-05-19T13:01:16.221172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(weight=class_weights_tensor)\noptimizer = optim.AdamW(model.parameters(), lr=1e-4)\n\n# Mixed Precision Training tetap sama\nscaler = torch.cuda.amp.GradScaler() if torch.cuda.is_available() else None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:16.222536Z","iopub.execute_input":"2025-05-19T13:01:16.222739Z","iopub.status.idle":"2025-05-19T13:01:16.228736Z","shell.execute_reply.started":"2025-05-19T13:01:16.222721Z","shell.execute_reply":"2025-05-19T13:01:16.227823Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"# Training Function\ndef train_step(model, train_loader, criterion, optimizer, device, scaler):\n    model.train()\n    total_loss, correct, total_samples = 0, 0, 0\n\n    with tqdm(train_loader, desc=\"Training\", unit=\"batch\") as pbar:\n        for inputs, target in pbar:\n            inputs, target = inputs.to(device), target.to(device)\n            optimizer.zero_grad()\n\n            if scaler:\n                with torch.cuda.amp.autocast():\n                    output = model(inputs)\n                    loss = criterion(output, target)\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                output = model(inputs)\n                loss = criterion(output, target)\n                loss.backward()\n                optimizer.step()\n\n            _, predicted = output.max(1)\n            total_loss += loss.item() * inputs.size(0)\n            correct += predicted.eq(target).sum().item()\n            total_samples += inputs.size(0)\n\n            pbar.set_postfix({\"Loss\": f\"{total_loss / total_samples:.4f}\", \"Accuracy\": f\"{100.0 * correct / total_samples:.2f}%\"})\n\n    return {\"loss\": total_loss / total_samples, \"accuracy\": 100.0 * correct / total_samples}\n\n# Validation Function\n@torch.no_grad()\ndef val_step(model, val_loader, criterion, device):\n    model.eval()\n    total_loss, correct, total_samples = 0, 0, 0\n\n    with tqdm(val_loader, desc=\"Validation\", unit=\"batch\") as pbar:\n        for inputs, target in pbar:\n            inputs, target = inputs.to(device), target.to(device)\n            output = model(inputs)\n            loss = criterion(output, target)\n\n            _, predicted = output.max(1)\n            total_loss += loss.item() * inputs.size(0)\n            correct += predicted.eq(target).sum().item()\n            total_samples += inputs.size(0)\n\n            pbar.set_postfix({\"Loss\": f\"{total_loss / total_samples:.4f}\", \"Accuracy\": f\"{100.0 * correct / total_samples:.2f}%\"})\n\n    return {\"loss\": total_loss / total_samples, \"accuracy\": 100.0 * correct / total_samples}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:16.22967Z","iopub.execute_input":"2025-05-19T13:01:16.229894Z","iopub.status.idle":"2025-05-19T13:01:16.247772Z","shell.execute_reply.started":"2025-05-19T13:01:16.229875Z","shell.execute_reply":"2025-05-19T13:01:16.246989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Training Loop\nbest_val_acc = 0\nbest_model_wts = copy.deepcopy(model.state_dict())\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\ntrain_loss, train_acc, val_loss, val_acc = [], [], [], []\n\nfor epoch in range(EPOCHS):\n    print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n\n    train_metrics = train_step(model, train_loader, criterion, optimizer, device, scaler)\n    val_metrics = val_step(model, val_loader, criterion, device)\n\n    scheduler.step()\n\n    train_loss.append(train_metrics[\"loss\"])\n    train_acc.append(train_metrics[\"accuracy\"])\n    val_loss.append(val_metrics[\"loss\"])\n    val_acc.append(val_metrics[\"accuracy\"])\n\n    if val_metrics[\"accuracy\"] > best_val_acc:\n        best_val_acc = val_metrics[\"accuracy\"]\n        best_model_wts = copy.deepcopy(model.state_dict())\n        torch.save(best_model_wts, CHECKPOINT_PATH)\n        print(f\"=> Model saved at {CHECKPOINT_PATH}\")\n\n# Save Metrics\nmetrics_df = pd.DataFrame({\n    \"epoch\": range(1, len(train_loss) + 1),\n    \"train_loss\": train_loss,\n    \"train_accuracy\": train_acc,\n    \"val_loss\": val_loss,\n    \"val_accuracy\": val_acc\n})\nmetrics_df.to_csv(MODEL_SAVE, index=False)\n\n# Plot Training Results\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nsns.lineplot(x='epoch', y='train_loss', data=metrics_df, label='Train Loss')\nsns.lineplot(x='epoch', y='val_loss', data=metrics_df, label='Validation Loss')\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Loss per Epoch\")\nplt.legend()\n\nplt.subplot(1, 2, 2)\nsns.lineplot(x='epoch', y='train_accuracy', data=metrics_df, label='Train Accuracy')\nsns.lineplot(x='epoch', y='val_accuracy', data=metrics_df, label='Validation Accuracy')\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Accuracy per Epoch\")\nplt.legend()\nplt.show()\n\n# Load Best Model for Testing\nmodel.load_state_dict(torch.load(CHECKPOINT_PATH, map_location=device))\nmodel.to(device)\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T13:01:16.248573Z","iopub.execute_input":"2025-05-19T13:01:16.248874Z","iopub.status.idle":"2025-05-19T14:08:12.791261Z","shell.execute_reply.started":"2025-05-19T13:01:16.248838Z","shell.execute_reply":"2025-05-19T14:08:12.790091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation","metadata":{}},{"cell_type":"code","source":"# Function to Evaluate Test Set\n@torch.no_grad()\ndef evaluate_test_set(model, test_loader, device):\n    y_true, y_pred = [], []\n\n    for inputs, labels in tqdm(test_loader, desc=\"Testing\"):\n        inputs, labels = inputs.to(device), labels.to(device)\n        outputs = model(inputs)\n        _, preds = torch.max(outputs, 1)\n\n        y_true.extend(labels.cpu().numpy())\n        y_pred.extend(preds.cpu().numpy())\n\n    return np.array(y_true), np.array(y_pred)\n\n# Evaluate on Test Set\ny_true_test, y_pred_test = evaluate_test_set(model, test_loader, device)\n\n# Calculate Test Accuracy\ntest_accuracy = accuracy_score(y_true_test, y_pred_test) * 100\nprint(f\"\\n Test Accuracy: {test_accuracy:.3f}%\")\n\n# Generate Classification Report\nclass_report = classification_report(y_true_test, y_pred_test, digits=3)\nprint(\"\\n Classification Report:\\n\", class_report)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T14:08:12.792603Z","iopub.execute_input":"2025-05-19T14:08:12.793043Z","iopub.status.idle":"2025-05-19T14:08:29.555386Z","shell.execute_reply.started":"2025-05-19T14:08:12.792999Z","shell.execute_reply":"2025-05-19T14:08:29.554425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save Model\ntorch.save(model, CHECKPOINT_PATH)\nprint(f\"Model saved to {CHECKPOINT_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-19T14:08:29.556391Z","iopub.execute_input":"2025-05-19T14:08:29.556642Z","iopub.status.idle":"2025-05-19T14:08:30.460966Z","shell.execute_reply.started":"2025-05-19T14:08:29.55662Z","shell.execute_reply":"2025-05-19T14:08:30.459979Z"}},"outputs":[],"execution_count":null}]}