{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":33679,"databundleVersionId":3212216,"sourceType":"competition"},{"sourceId":14117759,"sourceType":"datasetVersion","datasetId":8990132}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms, models\nfrom PIL import Image\nfrom sklearn.metrics import roc_curve, auc, precision_recall_curve, confusion_matrix, f1_score\nimport os\nimport json","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(f\"Using device: {device}\")\n\n# used 0.5 for the old ones, and 0.98 for your final tuned version\nMODELS_CONFIG = [\n    (\"First (20 Rotations)\", \"/kaggle/input/presentedtoclasscheckpoints/balanced_do_not_touch_model_rotating_safe_20rotations_rich.pth\", 0.5),\n    (\"Middle (40 Rotations)\", \"/kaggle/input/presentedtoclasscheckpoints/balanced_do_not_touch_model_rotating_safe_40rotations_On_20_rotations.pth\", 0.5),\n    (\"Final (60 Rotations)\", \"/kaggle/input/presentedtoclasscheckpoints/another_60_runs.pth\", 0.98) \n]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PlantDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row[\"image_path\"]\n        try:\n            img = Image.open(img_path).convert(\"RGB\")\n        except:\n            img = Image.new(\"RGB\", (224, 224), (0, 0, 0))\n        if self.transform:\n            img = self.transform(img)\n        return img, torch.tensor(int(row[\"do_not_touch\"]), dtype=torch.long)\n\ntest_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/herbarium-2022-fgvc9\"\nprint(\"Loading Metadata...\")\nwith open(os.path.join(DATA_DIR, \"train_metadata.json\"), \"r\") as f:\n    train_meta = json.load(f)\n\nann_df = pd.DataFrame(train_meta[\"annotations\"])\nimg_df = pd.DataFrame(train_meta[\"images\"])\ncat_df = pd.DataFrame(train_meta[\"categories\"])\nmerged = ann_df.merge(img_df, on=\"image_id\", how=\"left\").merge(cat_df, on=\"category_id\", how=\"left\")\nmerged[\"image_path\"] = merged[\"file_name\"].apply(lambda fn: os.path.join(DATA_DIR, \"train_images\", fn))\n\n\ntoxic_genus_list = [\"Toxicodendron\", \"Euphorbia\", \"Urtica\", \"Cicuta\", \"Conium\", \"Heracleum\"]\nmerged[\"do_not_touch\"] = merged[\"genus\"].isin(toxic_genus_list).astype(int)\n\n\nfrom sklearn.model_selection import train_test_split\n_, test_df = train_test_split(merged, test_size=0.05, random_state=42, stratify=merged[\"do_not_touch\"])\n\nprint(f\"Test Set Size: {len(test_df)}\")\ntest_dataset = PlantDataset(test_df, transform=test_transforms)\ntest_loader = DataLoader(test_dataset, batch_size=128, shuffle=False, num_workers=2)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = {} \nfor name, path, thresh in MODELS_CONFIG:\n    print(f\"\\nProcessing: {name}...\")\n    \n    # Init Model\n    model = models.efficientnet_b0(weights=None)\n    model.classifier = nn.Sequential(nn.Dropout(0.2), nn.Linear(1280, 1))\n    model = model.to(device)\n    \n    # Load Weights\n    if os.path.exists(path):\n        ckpt = torch.load(path, map_location=device)\n        if \"model_state_dict\" in ckpt: model.load_state_dict(ckpt[\"model_state_dict\"])\n        else: model.load_state_dict(ckpt)\n    else:\n        print(f\"!! WARNING: File not found {path}\")\n        continue\n        \n    model.eval()\n    all_probs = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for images, labels in test_loader:\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.sigmoid(outputs).squeeze(1).cpu().numpy()\n            all_probs.append(probs)\n            all_labels.append(labels.numpy())\n            \n    y_true = np.concatenate(all_labels)\n    y_scores = np.concatenate(all_probs)\n    y_pred = (y_scores >= thresh).astype(int)\n    \n    results[name] = {\n        \"y_true\": y_true,\n        \"y_scores\": y_scores,\n        \"y_pred\": y_pred,\n        \"threshold\": thresh\n    }","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.style.use('seaborn-v0_8-whitegrid')\n\n# Plot 1: ROC Curve Overlay\nplt.figure(figsize=(10, 6))\nfor name, data in results.items():\n    fpr, tpr, _ = roc_curve(data[\"y_true\"], data[\"y_scores\"])\n    roc_auc = auc(fpr, tpr)\n    plt.plot(fpr, tpr, label=f'{name} (AUC = {roc_auc:.4f})', linewidth=2)\n\nplt.plot([0, 1], [0, 1], 'k--', linestyle='--')\nplt.xlabel('False Positive Rate')\nplt.ylabel('True Positive Rate')\nplt.title('Improvement in ROC Curve over Training Stages')\nplt.legend(loc=\"lower right\")\nplt.show()\n\n# Plot 2: Precision-Recall Curve Overlay \nplt.figure(figsize=(10, 6))\nfor name, data in results.items():\n    precision, recall, _ = precision_recall_curve(data[\"y_true\"], data[\"y_scores\"])\n    pr_auc = auc(recall, precision)\n    plt.plot(recall, precision, label=f'{name} (PR AUC = {pr_auc:.4f})', linewidth=2)\n\nplt.xlabel('Recall (Sensitivity)')\nplt.ylabel('Precision (Confidence)')\nplt.title('Precision-Recall Curve: The Real Improvement')\nplt.legend(loc=\"lower left\")\nplt.show()\n\n# plot 3: Confusion Matrices Side-by-Side \nfig, axes = plt.subplots(1, 3, figsize=(20, 5))\n\nfor ax, (name, data) in zip(axes, results.items()):\n    cm = confusion_matrix(data[\"y_true\"], data[\"y_pred\"])\n    \n    # Custom annotations with \"Safe\" and \"Toxic\" labels\n    group_names = ['True Safe','False Toxic','False Safe','True Toxic']\n    group_counts = [\"{0:0.0f}\".format(value) for value in cm.flatten()]\n    labels = [f\"{v1}\\n{v2}\" for v1, v2 in zip(group_names, group_counts)]\n    labels = np.asarray(labels).reshape(2,2)\n    \n    sns.heatmap(cm, annot=labels, fmt='', cmap='Blues', cbar=False, ax=ax, annot_kws={\"size\": 14})\n    ax.set_title(f\"{name}\\nThreshold: {data['threshold']}\", fontsize=14, fontweight='bold')\n    ax.set_xlabel('Predicted')\n    ax.set_ylabel('Actual')\n    ax.set_xticklabels(['Safe', 'Toxic'])\n    ax.set_yticklabels(['Safe', 'Toxic'])\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}