{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":5048,"databundleVersionId":868335},{"sourceType":"modelInstanceVersion","sourceId":208033,"databundleVersionId":10581945,"modelInstanceId":149492,"modelId":172002}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport tensorflow as tf\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-24T08:02:34.798870Z","iopub.execute_input":"2026-04-24T08:02:34.799123Z","iopub.status.idle":"2026-04-24T08:03:18.583589Z","shell.execute_reply.started":"2026-04-24T08:02:34.799096Z","shell.execute_reply":"2026-04-24T08:03:18.582658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\n\nprint(\"--- 1. Loading and Analyzing Driver Data ---\")\ncsv_path = '/kaggle/input/competitions/state-farm-distracted-driver-detection/driver_imgs_list.csv'\ntrain_dir = '/kaggle/input/competitions/state-farm-distracted-driver-detection/imgs/train'\n\ndf = pd.read_csv(csv_path)\n\n# استخراج قائمة السائقين الفريدين (26 سائق في الداتاسيت)\nunique_drivers = df['subject'].unique()\nprint(f\"Total Unique Drivers: {len(unique_drivers)}\")\n\nprint(\"\\n--- 2. Subject-Based Splitting (The PRO Way) ---\")\nnp.random.seed(42)\nnp.random.shuffle(unique_drivers)\n\ntrain_drivers = unique_drivers[:18]\nval_drivers = unique_drivers[18:22]\ntest_drivers = unique_drivers[22:]\n\nprint(f\"Train Drivers: {train_drivers}\")\nprint(f\"Val Drivers:   {val_drivers}\")\nprint(f\"Test Drivers:  {test_drivers}\")\n\ntrain_df = df[df['subject'].isin(train_drivers)].copy()\nval_df = df[df['subject'].isin(val_drivers)].copy()\ntest_df = df[df['subject'].isin(test_drivers)].copy()\ndef create_filepath(row):\n    return os.path.join(train_dir, row['classname'], row['img'])\n\ntrain_df['filepath'] = train_df.apply(create_filepath, axis=1)\nval_df['filepath'] = val_df.apply(create_filepath, axis=1)\ntest_df['filepath'] = test_df.apply(create_filepath, axis=1)\n\ntrain_df = train_df.rename(columns={'classname': 'label'})\nval_df = val_df.rename(columns={'classname': 'label'})\ntest_df = test_df.rename(columns={'classname': 'label'})\n\nprint(f\"\\nImages in Train: {len(train_df)}\")\nprint(f\"Images in Val:   {len(val_df)}\")\nprint(f\"Images in Test:  {len(test_df)}\")\n\nprint(\"\\n--- 3. Creating Generators ---\")\nIMG_HEIGHT, IMG_WIDTH = 224, 224\nBATCH_SIZE = 32\n\ndatagen = ImageDataGenerator(rescale=1./255)\n\ntrain_data = datagen.flow_from_dataframe(\n    dataframe=train_df, x_col='filepath', y_col='label',\n    target_size=(IMG_HEIGHT, IMG_WIDTH), batch_size=BATCH_SIZE, class_mode='categorical', shuffle=True\n)\n\nval_data = datagen.flow_from_dataframe(\n    dataframe=val_df, x_col='filepath', y_col='label',\n    target_size=(IMG_HEIGHT, IMG_WIDTH), batch_size=BATCH_SIZE, class_mode='categorical', shuffle=False\n)\n\ntest_data = datagen.flow_from_dataframe(\n    dataframe=test_df, x_col='filepath', y_col='label',\n    target_size=(IMG_HEIGHT, IMG_WIDTH), batch_size=BATCH_SIZE, class_mode='categorical', shuffle=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T08:03:27.300494Z","iopub.execute_input":"2026-04-24T08:03:27.301201Z","iopub.status.idle":"2026-04-24T08:05:27.144551Z","shell.execute_reply.started":"2026-04-24T08:03:27.301162Z","shell.execute_reply":"2026-04-24T08:05:27.143689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images, labels = next(train_data)\n\nplt.figure(figsize=(10, 10))\nfor i in range(9):\n    plt.subplot(3, 3, i + 1)\n    plt.imshow(images[i])\n    # استخراج رقم الفئة من المصفوفة\n    label_index = tf.argmax(labels[i]).numpy()\n    plt.title(f\"Class: c{label_index}\")\n    plt.axis(\"off\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:34:28.088379Z","iopub.execute_input":"2026-04-24T07:34:28.089175Z","iopub.status.idle":"2026-04-24T07:34:29.943366Z","shell.execute_reply.started":"2026-04-24T07:34:28.089142Z","shell.execute_reply":"2026-04-24T07:34:29.942554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout, BatchNormalization\n\nprint(\"--- Building Custom CNN Architecture from Scratch ---\")\n\nmodel = Sequential([\n    Conv2D(32, (3, 3), activation='relu', input_shape=(IMG_HEIGHT, IMG_WIDTH, 3), padding='same'),\n    BatchNormalization(),\n    MaxPooling2D(pool_size=(2, 2)),\n    \n    Conv2D(64, (3, 3), activation='relu', padding='same'),\n    BatchNormalization(),\n    MaxPooling2D(pool_size=(2, 2)),\n    \n    # ------------------ Block 3 ------------------\n    Conv2D(128, (3, 3), activation='relu', padding='same'),\n    BatchNormalization(),\n    MaxPooling2D(pool_size=(2, 2)),\n    \n    Conv2D(256, (3, 3), activation='relu', padding='same'),\n    BatchNormalization(),\n    MaxPooling2D(pool_size=(2, 2)),\n    \n \n    Flatten(),\n    \n    Dense(512, activation='relu'),\n    BatchNormalization(),\n    Dropout(0.5), \n    \n    Dense(128, activation='relu'),\n    BatchNormalization(),\n    Dropout(0.5),\n    \n    # طبقة الخرج النهائي (10 فئات)\n    Dense(10, activation='softmax')\n])\n\nmodel.compile(optimizer='adam', \n              loss='categorical_crossentropy', \n              metrics=['accuracy'])\n\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:34:47.883610Z","iopub.execute_input":"2026-04-24T07:34:47.884508Z","iopub.status.idle":"2026-04-24T07:34:49.261160Z","shell.execute_reply.started":"2026-04-24T07:34:47.884472Z","shell.execute_reply":"2026-04-24T07:34:49.260459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping\n\nprint(\"--- Setting up Callbacks ---\")\n\ncheckpoint = ModelCheckpoint(\"best_driver_model.keras\", \n                             monitor='val_accuracy',   \n                             save_best_only=True,      \n                             mode='max', \n                             verbose=1)\n\nearly_stop = EarlyStopping(monitor='val_accuracy', \n                           patience=5,                \n                           restore_best_weights=True,  \n                           verbose=1)\n\nprint(\"--- Starting Training Process ---\")\n\nEPOCHS = 20  \n\n# بدء عملية التعلم الفعلي (Fit)\nhistory = model.fit(\n    train_data,\n    validation_data=val_data,\n    epochs=EPOCHS,\n    callbacks=[checkpoint, early_stop]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:34:53.215701Z","iopub.execute_input":"2026-04-24T07:34:53.216390Z","iopub.status.idle":"2026-04-24T07:50:36.354032Z","shell.execute_reply.started":"2026-04-24T07:34:53.216360Z","shell.execute_reply":"2026-04-24T07:50:36.353277Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(history.history['accuracy'], label='Training Accuracy', linewidth=2)\nplt.plot(history.history['val_accuracy'], label='Validation Accuracy', linewidth=2)\nplt.title('Model Accuracy')\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\nplt.legend()\nplt.grid(True)\n\nplt.subplot(1, 2, 2)\nplt.plot(history.history['loss'], label='Training Loss', linewidth=2)\nplt.plot(history.history['val_loss'], label='Validation Loss', linewidth=2)\nplt.title('Model Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:59:50.049835Z","iopub.execute_input":"2026-04-24T07:59:50.050705Z","iopub.status.idle":"2026-04-24T07:59:50.370966Z","shell.execute_reply.started":"2026-04-24T07:59:50.050663Z","shell.execute_reply":"2026-04-24T07:59:50.370187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, classification_report\nimport seaborn as sns\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tensorflow.keras.models import load_model\n\nbest_model = load_model('best_driver_model.keras')\nprint(\"--- Best Model Loaded Successfully ---\")\n\n# 2. قياس الدقة الكلية على بيانات الاختبار (Test Set)\nprint(\"\\n--- Evaluating on the UNSEEN Test Set (15%) ---\")\ntest_loss, test_accuracy = best_model.evaluate(test_data)\nprint(f\"🎉 Final Test Accuracy: {test_accuracy * 100:.2f}%\")\nprint(f\"📉 Final Test Loss: {test_loss:.4f}\")\n\nprint(\"\\nGenerating predictions for the Confusion Matrix...\")\ntest_data.reset() \nY_pred = best_model.predict(test_data)\ny_pred = np.argmax(Y_pred, axis=1)\n\ny_true = test_data.classes\n\n# قاموس الأسماء للرسم البياني\nclass_labels_dict = {\n    'c0': 'Safe driving',\n    'c1': 'Texting - right',\n    'c2': 'Talking on phone - right',\n    'c3': 'Texting - left',\n    'c4': 'Talking on phone - left',\n    'c5': 'Operating the radio',\n    'c6': 'Drinking',\n    'c7': 'Reaching behind',\n    'c8': 'Hair and makeup',\n    'c9': 'Talking to passenger'\n}\n\nclass_names = [class_labels_dict[k] for k in test_data.class_indices.keys()]\n\ncm = confusion_matrix(y_true, y_pred)\n\nplt.figure(figsize=(12, 10))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n            xticklabels=class_names, yticklabels=class_names)\nplt.title('Confusion Matrix - Final Test Set Evaluation', fontsize=16, fontweight='bold')\nplt.ylabel('True Label', fontsize=14)\nplt.xlabel('Predicted Label', fontsize=14)\nplt.xticks(rotation=45, ha='right')\nplt.yticks(rotation=0)\nplt.tight_layout()\nplt.show()\n\n# 6. طباعة تقرير التصنيف المفصل\nprint(\"\\n--- Detailed Classification Report ---\")\nprint(classification_report(y_true, y_pred, target_names=class_names, digits=4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:59:01.360327Z","iopub.execute_input":"2026-04-24T07:59:01.361172Z","iopub.status.idle":"2026-04-24T07:59:28.017031Z","shell.execute_reply.started":"2026-04-24T07:59:01.361138Z","shell.execute_reply":"2026-04-24T07:59:28.016270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nfrom tensorflow.keras.preprocessing import image\n\nidx_to_class = {v: k for k, v in test_data.class_indices.items()}\n\nprint(\"---  Visualizing Model Predictions on UNSEEN Test Images ---\")\n\nplt.figure(figsize=(16, 10))\n\nsample_test_images = test_df.sample(6)\n\nfor i, (index, row) in enumerate(sample_test_images.iterrows()):\n    img_path = row['filepath']\n    true_label_code = row['label'] \n    true_label_name = class_labels_dict[true_label_code]\n    \n    img = image.load_img(img_path, target_size=(224, 224))\n    img_array = image.img_to_array(img)\n    img_array_expanded = np.expand_dims(img_array, axis=0) / 255.0  # Normalization\n    \n    prediction = best_model.predict(img_array_expanded, verbose=0)\n    predicted_idx = np.argmax(prediction)\n    predicted_label_code = idx_to_class[predicted_idx]\n    predicted_label_name = class_labels_dict[predicted_label_code]\n    confidence = np.max(prediction) * 100\n    \n    plt.subplot(2, 3, i + 1)\n    plt.imshow(img)\n    \n    if true_label_code == predicted_label_code:\n        color = 'green'\n        title_text = f\"✅ Pred: {predicted_label_name}\\nTrue: {true_label_name}\\nConf: {confidence:.2f}%\"\n    else:\n        color = 'red'\n        title_text = f\"❌ Pred: {predicted_label_name}\\nTrue: {true_label_name}\\nConf: {confidence:.2f}%\"\n        \n    plt.title(title_text, color=color, fontsize=13, fontweight='bold')\n    plt.axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:53:05.119260Z","iopub.execute_input":"2026-04-24T07:53:05.119856Z","iopub.status.idle":"2026-04-24T07:53:07.398603Z","shell.execute_reply.started":"2026-04-24T07:53:05.119826Z","shell.execute_reply":"2026-04-24T07:53:07.397472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T07:56:01.436433Z","iopub.execute_input":"2026-04-24T07:56:01.437187Z","iopub.status.idle":"2026-04-24T07:56:10.984107Z","shell.execute_reply.started":"2026-04-24T07:56:01.437149Z","shell.execute_reply":"2026-04-24T07:56:10.983508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nfrom tqdm import tqdm\n\nprint(\"--- 1. INITIALIZING DATASET CLASS ---\")\n\n# تعريف الكلاس الذي كان مفقوداً\nclass DriverDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.classes = sorted(df['label'].unique())\n        self.class_to_idx = {c: i for i, c in enumerate(self.classes)}\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_path = self.df.loc[idx, 'filepath']\n        label = self.class_to_idx[self.df.loc[idx, 'label']]\n        image = Image.open(img_path).convert('RGB')\n        \n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\nprint(\"--- 2. APPLYING SAFE RANDOM ERASING ---\")\n\n# التجهيز مع الـ Random Erasing الآمن\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomRotation(10), \n    transforms.RandomResizedCrop(224, scale=(0.85, 1.0)), \n    transforms.ColorJitter(brightness=0.2, contrast=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    transforms.RandomErasing(p=0.5, scale=(0.02, 0.10), ratio=(0.3, 3.3), value=0)\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# إنشاء الـ Loaders\ntrain_dataset = DriverDataset(train_df, transform=train_transform)\nval_dataset = DriverDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)\n\nprint(\"--- 3. UNLEASHING THE HEAVYWEIGHT CHAMPION (ViT-Base) ---\")\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ترقية الموديل إلى Base (أقوى 4 مرات من Small)\nmodel = timm.create_model('vit_base_patch16_224', \n                          pretrained=True, \n                          num_classes=10,\n                          drop_rate=0.4,       \n                          attn_drop_rate=0.2)  \n\nfor param in model.parameters():\n    param.requires_grad = True\n\nmodel = model.to(device)\n\n# إعداد التدريب (LR صغير للحفاظ على الأوزان الذهبية)\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.2)\noptimizer = optim.AdamW(model.parameters(), lr=1e-5, weight_decay=0.05)\nscheduler = CosineAnnealingLR(optimizer, T_max=15)\n\nEPOCHS = 15\nbest_val_acc = 0.0\n\nprint(\"\\n--- Starting Full Fine-Tuning ---\")\nfor epoch in range(EPOCHS):\n    model.train()\n    train_loss, train_correct, train_total = 0.0, 0, 0\n    \n    for images, labels in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS} [Train]\"):\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        \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        \n        train_loss += loss.item() * images.size(0)\n        _, predicted = torch.max(outputs, 1)\n        train_total += labels.size(0)\n        train_correct += (predicted == labels).sum().item()\n        \n    train_acc = train_correct / train_total\n    scheduler.step()\n    current_lr = scheduler.get_last_lr()[0]\n    \n    model.eval()\n    val_loss, val_correct, val_total = 0.0, 0, 0\n    with torch.no_grad():\n        for images, labels in tqdm(val_loader, desc=f\"Epoch {epoch+1}/{EPOCHS} [Val]\"):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            val_loss += loss.item() * images.size(0)\n            _, predicted = torch.max(outputs, 1)\n            val_total += labels.size(0)\n            val_correct += (predicted == labels).sum().item()\n            \n    val_acc = val_correct / val_total\n    \n    print(f\"Epoch {epoch+1} | LR: {current_lr:.6f} | Train Acc: {train_acc:.4f} | Val Acc: {val_acc:.4f}\")\n    \n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        torch.save(model.state_dict(), 'best_vit_base_finetuned.pth')\n        print(\"✅ New Best Model Saved!\")\n\nprint(f\"\\n--- Training Complete. Best Val Accuracy: {best_val_acc:.4f} ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T08:05:36.414166Z","iopub.execute_input":"2026-04-24T08:05:36.414857Z","iopub.status.idle":"2026-04-24T10:40:43.336305Z","shell.execute_reply.started":"2026-04-24T08:05:36.414808Z","shell.execute_reply":"2026-04-24T10:40:43.335495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\nimport torch.nn.functional as F\n\nmodel.load_state_dict(torch.load('best_vit_base_finetuned.pth'))\nmodel.eval()\n\nall_preds = []\nall_labels =[]\n\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Evaluating on Validation Set\"):\n        images, labels = images.to(device), labels.to(device)\n        outputs = model(images)\n        _, predicted = torch.max(outputs, 1)\n        \n        all_preds.extend(predicted.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\nclass_labels_dict = {\n    'c0': 'Safe driving',\n    'c1': 'Texting - right',\n    'c2': 'Talking on phone - right',\n    'c3': 'Texting - left',\n    'c4': 'Talking on phone - left',\n    'c5': 'Operating the radio',\n    'c6': 'Drinking',\n    'c7': 'Reaching behind',\n    'c8': 'Hair and makeup',\n    'c9': 'Talking to passenger'\n}\n\nclass_names_original = val_dataset.classes\nclass_names_descriptive = [class_labels_dict[c] for c in class_names_original]\n\n# 3. رسم مصفوفة الارتباك\ncm = confusion_matrix(all_labels, all_preds)\n\nplt.figure(figsize=(12, 10))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=class_names_descriptive, \n            yticklabels=class_names_descriptive)\n\nplt.title('Confusion Matrix - Validation Set', fontsize=16, fontweight='bold')\nplt.ylabel('True Label', fontsize=14)\nplt.xlabel('Predicted Label', fontsize=14)\nplt.xticks(rotation=45, ha='right')\nplt.yticks(rotation=0)\nplt.tight_layout()\nplt.show()\n\n# 4. طباعة تقرير التصنيف الشامل\nprint(\"\\n--- Classification Report ---\")\nprint(classification_report(all_labels, all_preds, target_names=class_names_descriptive, digits=4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T10:44:55.974702Z","iopub.execute_input":"2026-04-24T10:44:55.975021Z","iopub.status.idle":"2026-04-24T10:45:44.467216Z","shell.execute_reply.started":"2026-04-24T10:44:55.974994Z","shell.execute_reply":"2026-04-24T10:45:44.466298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\n\nidx_to_class = {i: c for i, c in enumerate(class_names_original)}\n\nplt.figure(figsize=(16, 10))\nindices = random.sample(range(len(val_dataset)), 6)\n\nfor i, idx in enumerate(indices):\n    image, label = val_dataset[idx]\n    \n    img_tensor = image.unsqueeze(0).to(device)\n    \n    with torch.no_grad():\n        output = model(img_tensor)\n        prob = F.softmax(output, dim=1)\n        confidence, predicted = torch.max(prob, 1)\n        \n    predicted_class = idx_to_class[predicted.item()]\n    true_class = idx_to_class[label]\n    \n    pred_name = class_labels_dict[predicted_class]\n    true_name = class_labels_dict[true_class]\n    conf_score = confidence.item() * 100\n    \n    img_vis = image.permute(1, 2, 0).cpu().numpy()\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    img_vis = std * img_vis + mean\n    img_vis = np.clip(img_vis, 0, 1)\n    \n    plt.subplot(2, 3, i + 1)\n    plt.imshow(img_vis)\n    \n    if predicted_class == true_class:\n        color = 'green'\n        title = f\"Pred: {pred_name}\\nTrue: {true_name}\\nConf: {conf_score:.2f}%\"\n    else:\n        color = 'red'\n        title = f\"Pred: {pred_name}\\nTrue: {true_name}\\nConf: {conf_score:.2f}%\"\n        \n    plt.title(title, color=color, fontsize=12, fontweight='bold')\n    plt.axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T10:55:37.390696Z","iopub.execute_input":"2026-04-24T10:55:37.391013Z","iopub.status.idle":"2026-04-24T10:55:38.490935Z","shell.execute_reply.started":"2026-04-24T10:55:37.390984Z","shell.execute_reply":"2026-04-24T10:55:38.489823Z"}},"outputs":[],"execution_count":null}]}