{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================\n# TEST SET EVALUATION\n# Run this in a NEW Kaggle cell after training\n# =============================================\n\nimport os\nimport torch\nimport timm\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    classification_report, confusion_matrix,\n    roc_curve, auc, precision_recall_curve\n)\nfrom sklearn.preprocessing import label_binarize\n\n","metadata":{"_uuid":"d31df85d-82ec-412f-8b3b-8191e12d6397","_cell_guid":"ca34b3b9-862d-41d1-9782-f27e6f71fd16","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-06-10T06:17:35.192168Z","iopub.execute_input":"2026-06-10T06:17:35.192804Z","iopub.status.idle":"2026-06-10T06:17:49.248574Z","shell.execute_reply.started":"2026-06-10T06:17:35.192768Z","shell.execute_reply":"2026-06-10T06:17:49.247534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 1 — Recreate the full dataset df\n# (same code as your original notebook)\n# =============================================\n\nBASE_PATH      = \"/kaggle/input/competitions/snakeclef2022\"\nTRAIN_METADATA = BASE_PATH + \"/SnakeCLEF2022-TrainMetadata.csv\"\nTRAIN_IMG_DIR  = BASE_PATH + \"/SnakeCLEF2022-medium_size/SnakeCLEF2022-medium_size\"\nISO_MAPPING    = BASE_PATH + \"/SnakeCLEF2022-ISOxSpeciesMapping.csv\"\nNO_SNAKE_DIR   = \"/kaggle/input/datasets/rounak221bs/snake-like-objects/no_snake_dataset/no_snake_training\"\n\niso_df         = pd.read_csv(ISO_MAPPING)\nindia_species  = iso_df[iso_df['india'] == 1]['binomial'].tolist()\nfull_df        = pd.read_csv(TRAIN_METADATA)\nindia_df       = full_df[full_df['binomial_name'].isin(india_species)]\nclass_counts   = india_df['binomial_name'].value_counts()\ntop_species    = class_counts[(class_counts > 100) & (class_counts < 600)].index\ndf             = india_df[india_df['binomial_name'].isin(top_species)].copy()\n\n# Add no-snake\nno_snake_files = [f for f in os.listdir(NO_SNAKE_DIR)\n                  if f.lower().endswith((\".jpg\", \".jpeg\", \".png\"))]\nno_snake_df    = pd.DataFrame({\n    \"file_path\":     [os.path.join(NO_SNAKE_DIR, f) for f in no_snake_files],\n    \"binomial_name\": \"no_snake\",\n    \"class_id\":      -1\n})\ndf = pd.concat([df, no_snake_df], ignore_index=True)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:17:52.633320Z","iopub.execute_input":"2026-06-10T06:17:52.634147Z","iopub.status.idle":"2026-06-10T06:17:53.416583Z","shell.execute_reply.started":"2026-06-10T06:17:52.634112Z","shell.execute_reply":"2026-06-10T06:17:53.415892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================\n# STEP 2 — Retrospective 3-way split\n# Same random_state=42 → train_df identical to training\n# =============================================\n\ntrain_df, temp_df = train_test_split(\n    df, test_size=0.30,\n    stratify=df[\"class_id\"],\n    random_state=42\n)\n\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.50,\n    stratify=temp_df[\"class_id\"],\n    random_state=42\n)\n\nprint(f\"Train : {len(train_df)}\")\nprint(f\"Val   : {len(val_df)}\")\nprint(f\"Test  : {len(test_df)}  ← never seen during training\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:17:55.434083Z","iopub.execute_input":"2026-06-10T06:17:55.434792Z","iopub.status.idle":"2026-06-10T06:17:55.460818Z","shell.execute_reply.started":"2026-06-10T06:17:55.434759Z","shell.execute_reply":"2026-06-10T06:17:55.460177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================\n# STEP 3 — Fix labels (same mapping as training)\n# =============================================\n\nunique_classes = sorted(train_df[\"class_id\"].unique())\nclass_to_idx   = {cls: idx for idx, cls in enumerate(unique_classes)}\n\ntrain_df[\"class_id\"] = train_df[\"class_id\"].map(class_to_idx)\nval_df[\"class_id\"]   = val_df[\"class_id\"].map(class_to_idx)\ntest_df[\"class_id\"]  = test_df[\"class_id\"].map(class_to_idx)\n\nnum_classes = len(unique_classes)\nprint(f\"Num classes: {num_classes}\")\n\n# Build reverse mapping: idx → binomial_name\nidx_to_name = {}\nfor name, orig_id in df.set_index(\"binomial_name\")[\"class_id\"].items():\n    if orig_id in class_to_idx:\n        idx_to_name[class_to_idx[orig_id]] = name\nidx_to_name = dict(sorted(idx_to_name.items()))\nclass_names = [idx_to_name[i] for i in range(num_classes)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:17:57.582004Z","iopub.execute_input":"2026-06-10T06:17:57.582773Z","iopub.status.idle":"2026-06-10T06:17:57.610105Z","shell.execute_reply.started":"2026-06-10T06:17:57.582739Z","shell.execute_reply":"2026-06-10T06:17:57.609486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 4 — Dataset + DataLoader\n# =============================================\n\nImage.LOAD_TRUNCATED_IMAGES = True\n\nclass SnakeDataset(Dataset):\n    def __init__(self, df, root_dir, transform=None):\n        self.df       = df.reset_index(drop=True)\n        self.root_dir = root_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        try:\n            row = self.df.iloc[idx]\n            img_path = (row[\"file_path\"] if row[\"binomial_name\"] == \"no_snake\"\n                        else os.path.join(self.root_dir, row[\"file_path\"]))\n            image = Image.open(img_path).convert(\"RGB\")\n        except:\n            return self.__getitem__((idx + 1) % len(self.df))\n        label = self.df.iloc[idx][\"class_id\"]\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\ntest_transform = transforms.Compose([\n    transforms.Resize((518, 518)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\ntest_dataset = SnakeDataset(test_df, TRAIN_IMG_DIR, test_transform)\ntest_loader  = DataLoader(test_dataset, batch_size=8,\n                          shuffle=False, num_workers=4, pin_memory=True)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:18:01.309277Z","iopub.execute_input":"2026-06-10T06:18:01.309872Z","iopub.status.idle":"2026-06-10T06:18:01.318878Z","shell.execute_reply.started":"2026-06-10T06:18:01.309841Z","shell.execute_reply":"2026-06-10T06:18:01.318223Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 5 — Load best model\n# =============================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = timm.create_model(\n    \"vit_large_patch14_dinov2\",\n    pretrained=False,\n    num_classes=num_classes\n)\n\nif torch.cuda.device_count() > 1:\n    model = torch.nn.DataParallel(model)\n\nmodel = model.to(device)\n\n# Handle DataParallel prefix mismatch\ncheckpoint = torch.load(\n    \"/kaggle/input/models/rounak221bs/vit-large-patch14-dinov2-95/pytorch/default/1/snake_model.pth\",\n    map_location=device\n)\n\n# Strip 'module.' prefix if saved with DataParallel but loading on single GPU\nfrom collections import OrderedDict\nnew_state = OrderedDict()\nfor k, v in checkpoint.items():\n    new_key = k.replace(\"module.\", \"\") if k.startswith(\"module.\") else k\n    new_state[new_key] = v\n\n# Load into underlying model (handles both single and multi-GPU)\nif isinstance(model, torch.nn.DataParallel):\n    model.module.load_state_dict(new_state)\nelse:\n    model.load_state_dict(new_state)\n\nmodel.eval()\nprint(\"Model loaded.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:18:03.610352Z","iopub.execute_input":"2026-06-10T06:18:03.611067Z","iopub.status.idle":"2026-06-10T06:18:15.738309Z","shell.execute_reply.started":"2026-06-10T06:18:03.611032Z","shell.execute_reply":"2026-06-10T06:18:15.737625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 6 — Inference on test set\n# =============================================\n\nall_preds, all_targets, all_probs = [], [], []\n\nwith torch.no_grad():\n    for images, labels in test_loader:\n        images  = images.to(device)\n        outputs = model(images)\n        prob    = torch.softmax(outputs, dim=1)\n        _, predicted = torch.max(outputs, 1)\n        all_preds.extend(predicted.cpu().numpy())\n        all_targets.extend(labels.cpu().numpy())\n        all_probs.extend(prob.cpu().numpy())\n\nall_preds   = np.array(all_preds)\nall_targets = np.array(all_targets)\nall_probs   = np.array(all_probs)\n\ntest_acc = (all_preds == all_targets).mean()\nprint(f\"\\n{'='*40}\")\nprint(f\"  TEST ACCURACY : {test_acc*100:.2f}%\")\nprint(f\"{'='*40}\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:18:18.240666Z","iopub.execute_input":"2026-06-10T06:18:18.241537Z","iopub.status.idle":"2026-06-10T06:24:46.621901Z","shell.execute_reply.started":"2026-06-10T06:18:18.241481Z","shell.execute_reply":"2026-06-10T06:24:46.621056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================\n# STEP 7 — Classification Report\n# =============================================\n\nprint(classification_report(all_targets, all_preds,\n                             target_names=class_names, digits=4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:14.753479Z","iopub.execute_input":"2026-06-10T06:25:14.754182Z","iopub.status.idle":"2026-06-10T06:25:14.773591Z","shell.execute_reply.started":"2026-06-10T06:25:14.754147Z","shell.execute_reply":"2026-06-10T06:25:14.772822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 8 — Confusion Matrix\n# =============================================\n\ncm = confusion_matrix(all_targets, all_preds)\nplt.figure(figsize=(16, 14))\nsns.heatmap(cm, annot=False, cmap=\"Blues\",\n            xticklabels=class_names, yticklabels=class_names)\nplt.title(f\"Confusion Matrix — Test Set (Acc: {test_acc*100:.2f}%)\",\n          fontsize=14, fontweight='bold')\nplt.xlabel(\"Predicted\", fontsize=11)\nplt.ylabel(\"Actual\", fontsize=11)\nplt.xticks(rotation=90, fontsize=6)\nplt.yticks(rotation=0,  fontsize=6)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/test_confusion_matrix.png\", dpi=150)\nplt.show()\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:25.392994Z","iopub.execute_input":"2026-06-10T06:25:25.393994Z","iopub.status.idle":"2026-06-10T06:25:27.111280Z","shell.execute_reply.started":"2026-06-10T06:25:25.393947Z","shell.execute_reply":"2026-06-10T06:25:27.110670Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 9 — Per-Class Accuracy\n# =============================================\n\nper_class_acc = []\nfor c in range(num_classes):\n    mask = all_targets == c\n    per_class_acc.append((all_preds[mask] == c).mean() if mask.sum() > 0 else 0.0)\n\nplt.figure(figsize=(18, 5))\nbars = plt.bar(range(num_classes), [a*100 for a in per_class_acc],\n               color='steelblue', edgecolor='white', linewidth=0.4)\nplt.axhline(y=np.mean(per_class_acc)*100, color='tomato', linestyle='--',\n            linewidth=1.2, label=f'Mean: {np.mean(per_class_acc)*100:.2f}%')\nplt.xticks(range(num_classes), class_names, rotation=90, fontsize=6)\nplt.title(\"Per-Class Accuracy — Test Set\", fontsize=13, fontweight='bold')\nplt.xlabel(\"Species\")\nplt.ylabel(\"Accuracy (%)\")\nplt.ylim(0, 105)\nplt.legend()\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/test_per_class_accuracy.png\", dpi=150)\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:39.838687Z","iopub.execute_input":"2026-06-10T06:25:39.839134Z","iopub.status.idle":"2026-06-10T06:25:40.769333Z","shell.execute_reply.started":"2026-06-10T06:25:39.839090Z","shell.execute_reply":"2026-06-10T06:25:40.768634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 10 — Top-K Accuracy\n# =============================================\n\ndef topk_acc(y_true, y_score, k):\n    topk = np.argsort(y_score, axis=1)[:, -k:]\n    return np.mean([y_true[i] in topk[i] for i in range(len(y_true))])\n\nks   = [1, 3, 5]\naccs = [topk_acc(all_targets, all_probs, k) for k in ks]\n\nplt.figure(figsize=(6, 4))\nbars = plt.bar([f\"Top-{k}\" for k in ks], [a*100 for a in accs],\n               color=['steelblue', 'seagreen', 'tomato'])\nfor bar, acc in zip(bars, accs):\n    plt.text(bar.get_x() + bar.get_width()/2,\n             bar.get_height() + 0.3,\n             f\"{acc*100:.2f}%\", ha='center', fontsize=12, fontweight='bold')\nplt.title(\"Top-K Accuracy — Test Set\", fontsize=13, fontweight='bold')\nplt.ylabel(\"Accuracy (%)\")\nplt.ylim(0, 105)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/test_topk_accuracy.png\", dpi=150)\nplt.show()\n\nprint(f\"Top-1 : {accs[0]*100:.2f}%\")\nprint(f\"Top-3 : {accs[1]*100:.2f}%\")\nprint(f\"Top-5 : {accs[2]*100:.2f}%\")\n\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:44.625920Z","iopub.execute_input":"2026-06-10T06:25:44.626179Z","iopub.status.idle":"2026-06-10T06:25:44.861943Z","shell.execute_reply.started":"2026-06-10T06:25:44.626157Z","shell.execute_reply":"2026-06-10T06:25:44.861206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 11 — Confidence Distribution\n# =============================================\n\nconfidences  = np.max(all_probs, axis=1)\ncorrect_mask = all_preds == all_targets\n\nplt.figure(figsize=(10, 4))\nplt.hist(confidences[correct_mask],  bins=40, alpha=0.6,\n         color='seagreen', label='Correct')\nplt.hist(confidences[~correct_mask], bins=40, alpha=0.6,\n         color='tomato',   label='Incorrect')\nplt.title(\"Confidence Distribution — Correct vs Incorrect (Test Set)\",\n          fontsize=13, fontweight='bold')\nplt.xlabel(\"Confidence\")\nplt.ylabel(\"Count\")\nplt.legend()\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/test_confidence_dist.png\", dpi=150)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:50.329493Z","iopub.execute_input":"2026-06-10T06:25:50.329829Z","iopub.status.idle":"2026-06-10T06:25:50.750501Z","shell.execute_reply.started":"2026-06-10T06:25:50.329801Z","shell.execute_reply":"2026-06-10T06:25:50.749851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# STEP 12 — ROC Curve (macro avg)\n# =============================================\n\nn_classes  = num_classes\ny_true_bin = label_binarize(all_targets, classes=range(n_classes))\n\nplt.figure(figsize=(10, 6))\nmacro_fpr = np.linspace(0, 1, 300)\nmacro_tpr = np.zeros_like(macro_fpr)\n\nfor i in range(n_classes):\n    fpr, tpr, _ = roc_curve(y_true_bin[:, i], all_probs[:, i])\n    roc_auc     = auc(fpr, tpr)\n    plt.plot(fpr, tpr, linewidth=0.5, alpha=0.35)\n    macro_tpr  += np.interp(macro_fpr, fpr, tpr)\n\nmacro_tpr /= n_classes\nmacro_auc  = auc(macro_fpr, macro_tpr)\nplt.plot(macro_fpr, macro_tpr, color='navy', linewidth=2.5,\n         label=f'Macro-avg ROC (AUC = {macro_auc:.4f})')\nplt.plot([0,1],[0,1],'k--', linewidth=1)\nplt.title(\"ROC Curve — Test Set\", fontsize=13, fontweight='bold')\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.legend(fontsize=11)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/test_roc_curve.png\", dpi=150)\nplt.show()\n\nprint(f\"\\nMacro-avg AUC: {macro_auc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:54.673006Z","iopub.execute_input":"2026-06-10T06:25:54.673848Z","iopub.status.idle":"2026-06-10T06:25:55.424196Z","shell.execute_reply.started":"2026-06-10T06:25:54.673814Z","shell.execute_reply":"2026-06-10T06:25:55.423491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# SUMMARY\n# =============================================\n\nprint(\"\\n\" + \"=\"*40)\nprint(\"        FINAL TEST SET SUMMARY\")\nprint(\"=\"*40)\nprint(f\"  Test samples : {len(test_df)}\")\nprint(f\"  Top-1 Acc    : {accs[0]*100:.2f}%\")\nprint(f\"  Top-3 Acc    : {accs[1]*100:.2f}%\")\nprint(f\"  Top-5 Acc    : {accs[2]*100:.2f}%\")\nprint(f\"  Macro AUC    : {macro_auc:.4f}\")\nprint(\"=\"*40)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-10T06:25:59.793309Z","iopub.execute_input":"2026-06-10T06:25:59.793778Z","iopub.status.idle":"2026-06-10T06:25:59.801591Z","shell.execute_reply.started":"2026-06-10T06:25:59.793734Z","shell.execute_reply":"2026-06-10T06:25:59.800728Z"}},"outputs":[],"execution_count":null}]}