{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":34478,"databundleVersionId":3437841}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.719806Z","iopub.execute_input":"2026-03-28T12:09:10.720255Z","iopub.status.idle":"2026-03-28T12:09:10.724441Z","shell.execute_reply.started":"2026-03-28T12:09:10.720225Z","shell.execute_reply":"2026-03-28T12:09:10.723528Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\nimport os\nimport torch\nimport timm\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.72585Z","iopub.execute_input":"2026-03-28T12:09:10.72621Z","iopub.status.idle":"2026-03-28T12:09:10.738951Z","shell.execute_reply.started":"2026-03-28T12:09:10.726175Z","shell.execute_reply":"2026-03-28T12:09:10.738271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nif os.path.exists(\"/kaggle/working/checkpoint.pth\"):\n    os.remove(\"/kaggle/working/checkpoint.pth\")\n    print(\"Checkpoint deleted!\")\nelse:\n    print(\"No checkpoint found - starting fresh!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import shutil\n\n# shutil.rmtree(\"/kaggle/working/snakeclef\", ignore_errors=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.740007Z","iopub.execute_input":"2026-03-28T12:09:10.740366Z","iopub.status.idle":"2026-03-28T12:09:10.75023Z","shell.execute_reply.started":"2026-03-28T12:09:10.740343Z","shell.execute_reply":"2026-03-28T12:09:10.749386Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import shutil\n\n# shutil.copytree(\"/kaggle/input/competitions/snakeclef2022/SnakeCLEF2022-medium_size\", \"/kaggle/working/snakeclef\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.765488Z","iopub.execute_input":"2026-03-28T12:09:10.766225Z","iopub.status.idle":"2026-03-28T12:09:10.770322Z","shell.execute_reply.started":"2026-03-28T12:09:10.766186Z","shell.execute_reply":"2026-03-28T12:09:10.769657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.771592Z","iopub.execute_input":"2026-03-28T12:09:10.771938Z","iopub.status.idle":"2026-03-28T12:09:10.783319Z","shell.execute_reply.started":"2026-03-28T12:09:10.771915Z","shell.execute_reply":"2026-03-28T12:09:10.782601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_PATH = \"/kaggle/input/competitions/snakeclef2022\"\n\nTRAIN_METADATA = BASE_PATH + \"/SnakeCLEF2022-TrainMetadata.csv\"\nTRAIN_IMG_DIR  = BASE_PATH + \"/SnakeCLEF2022-medium_size/SnakeCLEF2022-medium_size\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.784169Z","iopub.execute_input":"2026-03-28T12:09:10.784477Z","iopub.status.idle":"2026-03-28T12:09:10.793979Z","shell.execute_reply.started":"2026-03-28T12:09:10.784445Z","shell.execute_reply":"2026-03-28T12:09:10.793174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(TRAIN_METADATA)\n\nprint(\"Total samples:\", len(df))\nprint(\"Total classes:\", df[\"class_id\"].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:10.794868Z","iopub.execute_input":"2026-03-28T12:09:10.795267Z","iopub.status.idle":"2026-03-28T12:09:11.185412Z","shell.execute_reply.started":"2026-03-28T12:09:10.795244Z","shell.execute_reply":"2026-03-28T12:09:11.184559Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# FILTER: Indian snake species (via ISO mapping)\n# =============================================\n\nISO_MAPPING = BASE_PATH + \"/SnakeCLEF2022-ISOxSpeciesMapping.csv\"\niso_df = pd.read_csv(ISO_MAPPING)\n\n# Step 1 — Get all species native to India\nindia_species = iso_df[iso_df['india'] == 1]['binomial'].tolist()\nprint(f\"Total species native to India: {len(india_species)}\")\n\n# Step 2 — Filter FULL TrainMetadata (df from Cell 6, not pre-filtered)\n# Re-read to make sure we're using the full dataset\nfull_df = pd.read_csv(TRAIN_METADATA)\nindia_df = full_df[full_df['binomial_name'].isin(india_species)]\nprint(f\"Total rows: {len(india_df)}\")\nprint(f\"Unique species: {india_df['binomial_name'].nunique()}\")\n\n# Step 3 — Keep species with >= 100 images\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\nprint(f\"\\nAfter >= 100 filter:\")\nprint(f\"  Dataset size : {len(df)}\")\nprint(f\"  Classes      : {df['binomial_name'].nunique()}\")\nprint()\nprint(df.groupby('binomial_name').size()\n        .reset_index(name='count')\n        .sort_values('count', ascending=False)\n        .to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.187826Z","iopub.execute_input":"2026-03-28T12:09:11.188069Z","iopub.status.idle":"2026-03-28T12:09:11.625539Z","shell.execute_reply.started":"2026-03-28T12:09:11.188046Z","shell.execute_reply":"2026-03-28T12:09:11.624864Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Train/Validation Split","metadata":{}},{"cell_type":"code","source":"train_df, val_df = train_test_split(\n    df,\n    test_size=0.15,\n    stratify=df[\"class_id\"],\n    random_state=42\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.626686Z","iopub.execute_input":"2026-03-28T12:09:11.627009Z","iopub.status.idle":"2026-03-28T12:09:11.641492Z","shell.execute_reply.started":"2026-03-28T12:09:11.626971Z","shell.execute_reply":"2026-03-28T12:09:11.640904Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Fix Labels","metadata":{}},{"cell_type":"code","source":"unique_classes = sorted(train_df[\"class_id\"].unique())\n\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)\n\nnum_classes = len(unique_classes)\n\nprint(\"Num classes:\", num_classes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.642258Z","iopub.execute_input":"2026-03-28T12:09:11.642579Z","iopub.status.idle":"2026-03-28T12:09:11.65392Z","shell.execute_reply.started":"2026-03-28T12:09:11.642547Z","shell.execute_reply":"2026-03-28T12:09:11.653188Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Transforms","metadata":{}},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ColorJitter(0.2,0.2,0.2,0.1),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],\n                         [0.229,0.224,0.225])\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],\n                         [0.229,0.224,0.225])\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.655325Z","iopub.execute_input":"2026-03-28T12:09:11.655616Z","iopub.status.idle":"2026-03-28T12:09:11.662266Z","shell.execute_reply.started":"2026-03-28T12:09:11.655581Z","shell.execute_reply":"2026-03-28T12:09:11.661456Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset Class","metadata":{}},{"cell_type":"code","source":"Image.LOAD_TRUNCATED_IMAGES = True\n\nclass SnakeDataset(Dataset):\n\n    def __init__(self, df, root_dir, transform=None):\n        self.df = df\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\n        try:\n            img_path = os.path.join(\n                self.root_dir,\n                self.df.iloc[idx][\"file_path\"]\n            )\n            image = Image.open(img_path).convert(\"RGB\")\n        except:\n            return self.__getitem__((idx+1) % len(self.df))\n\n        label = self.df.iloc[idx][\"class_id\"]\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.663352Z","iopub.execute_input":"2026-03-28T12:09:11.664171Z","iopub.status.idle":"2026-03-28T12:09:11.676751Z","shell.execute_reply.started":"2026-03-28T12:09:11.664105Z","shell.execute_reply":"2026-03-28T12:09:11.676184Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### DataLoaders","metadata":{}},{"cell_type":"code","source":"train_dataset = SnakeDataset(train_df, TRAIN_IMG_DIR, train_transform)\nval_dataset   = SnakeDataset(val_df, TRAIN_IMG_DIR, val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=0, pin_memory=True)\nval_loader   = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=0, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.677822Z","iopub.execute_input":"2026-03-28T12:09:11.678186Z","iopub.status.idle":"2026-03-28T12:09:11.688641Z","shell.execute_reply.started":"2026-03-28T12:09:11.678151Z","shell.execute_reply":"2026-03-28T12:09:11.687889Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# print(\"Available Vision Transformer Models: \")\n# timm.list_models(\"vit*\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.689417Z","iopub.execute_input":"2026-03-28T12:09:11.689683Z","iopub.status.idle":"2026-03-28T12:09:11.698903Z","shell.execute_reply.started":"2026-03-28T12:09:11.689648Z","shell.execute_reply":"2026-03-28T12:09:11.698191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = timm.create_model(\n    \"resnext50_32x4d\",\n    pretrained=True,\n    num_classes=num_classes\n)\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs!\")\n    model = torch.nn.DataParallel(model)\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:11.700614Z","iopub.execute_input":"2026-03-28T12:09:11.701277Z","iopub.status.idle":"2026-03-28T12:09:25.110594Z","shell.execute_reply.started":"2026-03-28T12:09:11.701254Z","shell.execute_reply":"2026-03-28T12:09:25.109915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False\n\nfor param in model.module.fc.parameters():\n    param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:25.111517Z","iopub.execute_input":"2026-03-28T12:09:25.112135Z","iopub.status.idle":"2026-03-28T12:09:25.117207Z","shell.execute_reply.started":"2026-03-28T12:09:25.112097Z","shell.execute_reply":"2026-03-28T12:09:25.11632Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loss + Optimizer","metadata":{}},{"cell_type":"code","source":"criterion = torch.nn.CrossEntropyLoss(label_smoothing=0.1)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=2e-5,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)\n\nscaler = torch.amp.GradScaler(\"cuda\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:25.120188Z","iopub.execute_input":"2026-03-28T12:09:25.120466Z","iopub.status.idle":"2026-03-28T12:09:25.131051Z","shell.execute_reply.started":"2026-03-28T12:09:25.120444Z","shell.execute_reply":"2026-03-28T12:09:25.130372Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import os\n\n# def file_exists(row):\n#     path = os.path.join(TRAIN_IMG_DIR, row[\"file_path\"])\n#     return os.path.exists(path)\n\n# df = df[df.apply(file_exists, axis=1)]\n\n# print(\"After filtering:\", len(df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:25.131956Z","iopub.execute_input":"2026-03-28T12:09:25.132784Z","iopub.status.idle":"2026-03-28T12:09:25.143299Z","shell.execute_reply.started":"2026-03-28T12:09:25.13275Z","shell.execute_reply":"2026-03-28T12:09:25.142595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_checkpoint(epoch, batch_idx, model, optimizer, scheduler, best_acc):\n    torch.save({\n        \"epoch\": epoch,\n        \"batch_idx\": batch_idx,\n        \"model_state\": model.state_dict(),\n        \"optimizer_state\": optimizer.state_dict(),\n        \"scheduler_state\": scheduler.state_dict(),\n        \"best_acc\": best_acc\n    }, \"/kaggle/working/checkpoint.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:25.144212Z","iopub.execute_input":"2026-03-28T12:09:25.144481Z","iopub.status.idle":"2026-03-28T12:09:25.155056Z","shell.execute_reply.started":"2026-03-28T12:09:25.144459Z","shell.execute_reply":"2026-03-28T12:09:25.154421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_epoch = 0\nstart_batch = 0\nbest_acc = 0\n\nif os.path.exists(\"/kaggle/working/checkpoint.pth\"):\n    checkpoint = torch.load(\"/kaggle/working/checkpoint.pth\",weights_only=False)\n\n    model.load_state_dict(checkpoint[\"model_state\"])\n    optimizer.load_state_dict(checkpoint[\"optimizer_state\"])\n    scheduler.load_state_dict(checkpoint[\"scheduler_state\"])\n\n    start_epoch = checkpoint[\"epoch\"]\n    start_batch = checkpoint[\"batch_idx\"]\n    best_acc = checkpoint[\"best_acc\"]\n\n    print(f\"Resuming from epoch {start_epoch}, batch {start_batch}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:25.155954Z","iopub.execute_input":"2026-03-28T12:09:25.156283Z","iopub.status.idle":"2026-03-28T12:09:25.166908Z","shell.execute_reply.started":"2026-03-28T12:09:25.156246Z","shell.execute_reply":"2026-03-28T12:09:25.166178Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training Loop","metadata":{}},{"cell_type":"code","source":"train_losses = []\nval_losses = []\nval_accs = []\n\n# =========================\n# FREEZE BACKBONE INITIALLY\n# =========================\nfor param in model.parameters():\n    param.requires_grad = False\n\nfor param in model.module.fc.parameters():\n    param.requires_grad = True\n\n# =========================\n# OPTIMIZER + SCHEDULER\n# =========================\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=2e-5,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=12\n)\n\n# =========================\n# TRAINING LOOP\n# =========================\nfor epoch in range(start_epoch, 12):\n\n    # Unfreeze after 3 epochs\n    if epoch == 3:\n        for param in model.parameters():\n            param.requires_grad = True\n        print(\"Backbone unfrozen\")\n\n    model.train()\n    running_loss = 0\n\n    for batch_idx, (images, labels) in enumerate(train_loader):\n\n        # Resume support\n        if epoch == start_epoch and batch_idx < start_batch:\n            continue\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.amp.autocast(\"cuda\"):\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n        scaler.scale(loss).backward()\n\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item()\n\n        if batch_idx % 500 == 0 and batch_idx != 0:\n            save_checkpoint(epoch, batch_idx, model, optimizer, scheduler, best_acc)\n            print(f\"Checkpoint saved at batch {batch_idx}\")\n\n        if batch_idx % 100 == 0:\n            print(f\"Epoch {epoch+1} | Batch {batch_idx}/{len(train_loader)} | Loss {loss.item():.4f}\")\n\n    scheduler.step()\n\n    num_batches_trained = batch_idx + 1 - (start_batch if epoch == start_epoch else 0)\n    epoch_loss = running_loss / num_batches_trained\n    train_losses.append(epoch_loss)\n\n    # =========================\n    # VALIDATION\n    # =========================\n    model.eval()\n\n    all_preds = []\n    all_targets = []\n    all_probs = []\n    val_running_loss = 0\n\n    with torch.no_grad():\n        for images, labels in val_loader:\n\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n\n            val_loss = criterion(outputs, labels)\n            val_running_loss += val_loss.item()\n\n            prob = torch.softmax(outputs, dim=1)\n            _, predicted = torch.max(outputs, 1)\n\n            all_preds.extend(predicted.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            all_probs.extend(prob.cpu().numpy())\n\n    acc = (np.array(all_preds) == np.array(all_targets)).mean()\n    val_epoch_loss = val_running_loss / len(val_loader)\n    val_accs.append(acc)\n    val_losses.append(val_epoch_loss)\n\n    print(f\"\\nEpoch {epoch+1} Completed\")\n    print(f\"  Train Loss : {epoch_loss:.4f}\")\n    print(f\"  Val Loss   : {val_epoch_loss:.4f}\")\n    print(f\"  Val Acc    : {acc:.4f}\\n\")\n\n    # =========================\n    # SAVE BEST MODEL\n    # =========================\n    if acc > best_acc:\n        best_acc = acc\n        torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n        print(\"Best model saved!\")\n\n    # =========================\n    # SAVE END-OF-EPOCH CHECKPOINT\n    # =========================\n    save_checkpoint(epoch, len(train_loader), model, optimizer, scheduler, best_acc)\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:49.468294Z","iopub.execute_input":"2026-03-28T12:09:49.468861Z","iopub.status.idle":"2026-03-28T12:09:50.643903Z","shell.execute_reply.started":"2026-03-28T12:09:49.468826Z","shell.execute_reply":"2026-03-28T12:09:50.642794Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Save Model","metadata":{}},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/snake_model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.703813Z","iopub.status.idle":"2026-03-28T12:09:48.704159Z","shell.execute_reply.started":"2026-03-28T12:09:48.703988Z","shell.execute_reply":"2026-03-28T12:09:48.704004Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization Curves","metadata":{}},{"cell_type":"markdown","source":"### Classification Report","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\nprint(classification_report(all_targets, all_preds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.706669Z","iopub.status.idle":"2026-03-28T12:09:48.707309Z","shell.execute_reply.started":"2026-03-28T12:09:48.707048Z","shell.execute_reply":"2026-03-28T12:09:48.707073Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Confusion Matrix","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns\n\ncm = confusion_matrix(all_targets, all_preds)\n\nplt.figure(figsize=(10,8))\nsns.heatmap(cm, annot=False, cmap=\"Blues\")\nplt.title(\"Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.708645Z","iopub.status.idle":"2026-03-28T12:09:48.709006Z","shell.execute_reply.started":"2026-03-28T12:09:48.708821Z","shell.execute_reply":"2026-03-28T12:09:48.708843Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ROC Curve","metadata":{}},{"cell_type":"code","source":"from sklearn.preprocessing import label_binarize\nfrom sklearn.metrics import roc_curve, auc\n\nn_classes = len(set(all_targets))\n\ny_true_bin = label_binarize(all_targets, classes=range(n_classes))\ny_score = np.array(all_probs)\n\nfor i in range(n_classes):\n    fpr, tpr, _ = roc_curve(y_true_bin[:, i], y_score[:, i])\n    roc_auc = auc(fpr, tpr)\n    plt.plot(fpr, tpr, label=f\"Class {i} (AUC={roc_auc:.2f})\")\n\nplt.plot([0,1],[0,1],'k--')\nplt.title(\"ROC Curve\")\nplt.xlabel(\"FPR\")\nplt.ylabel(\"TPR\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.710807Z","iopub.status.idle":"2026-03-28T12:09:48.711173Z","shell.execute_reply.started":"2026-03-28T12:09:48.71099Z","shell.execute_reply":"2026-03-28T12:09:48.711021Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Learning Curve","metadata":{}},{"cell_type":"code","source":"fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))\n\nax1.plot(train_losses, marker='o', color='steelblue', label='Train Loss')\nax1.plot(val_losses, marker='o', color='tomato', label='Val Loss')\nax1.set_title(\"Loss per Epoch\")\nax1.set_xlabel(\"Epoch\")\nax1.set_ylabel(\"Loss\")\nax1.legend()\n\nax2.plot(val_accs, marker='o', color='seagreen')\nax2.set_title(\"Val Accuracy per Epoch\")\nax2.set_xlabel(\"Epoch\")\nax2.set_ylabel(\"Accuracy\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.712171Z","iopub.status.idle":"2026-03-28T12:09:48.712476Z","shell.execute_reply.started":"2026-03-28T12:09:48.712327Z","shell.execute_reply":"2026-03-28T12:09:48.712342Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Confidence Score Distribution","metadata":{}},{"cell_type":"code","source":"confidences = np.max(all_probs, axis=1)\n\nplt.hist(confidences, bins=30)\nplt.title(\"Confidence Distribution\")\nplt.xlabel(\"Confidence\")\nplt.ylabel(\"Frequency\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.713911Z","iopub.status.idle":"2026-03-28T12:09:48.714205Z","shell.execute_reply.started":"2026-03-28T12:09:48.714032Z","shell.execute_reply":"2026-03-28T12:09:48.714047Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### False-Positive / False-Negative Visualization","metadata":{}},{"cell_type":"markdown","source":"### Correlation Matrix","metadata":{}},{"cell_type":"code","source":"corr = np.corrcoef(np.array(all_probs).T)\n\nplt.figure(figsize=(10,8))\nsns.heatmap(corr, cmap=\"coolwarm\")\nplt.title(\"Correlation Matrix (Classes)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.716018Z","iopub.status.idle":"2026-03-28T12:09:48.716357Z","shell.execute_reply.started":"2026-03-28T12:09:48.716222Z","shell.execute_reply":"2026-03-28T12:09:48.716246Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Class Activation Map","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\n\ndef show_cam(model, image, label):\n\n    model.eval()\n    image = image.unsqueeze(0).to(device)\n\n    features = []\n\n    def hook(module, input, output):\n        features.append(output)\n\n    handle = model.features[-1].register_forward_hook(hook)\n\n    output = model(image)\n    pred = output.argmax(dim=1)\n\n    feature_map = features[0][0].detach().cpu()\n\n    heatmap = torch.mean(feature_map, dim=0)\n    heatmap = F.relu(heatmap)\n    heatmap /= heatmap.max()\n\n    plt.imshow(heatmap, cmap='jet')\n    plt.title(f\"CAM | Pred: {pred.item()} | True: {label}\")\n    plt.show()\n\n    handle.remove()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.717468Z","iopub.status.idle":"2026-03-28T12:09:48.717853Z","shell.execute_reply.started":"2026-03-28T12:09:48.717655Z","shell.execute_reply":"2026-03-28T12:09:48.717677Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Precision vs Recall Curve","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import precision_recall_curve\n\nprecision = dict()\nrecall = dict()\n\nfor i in range(n_classes):\n    precision[i], recall[i], _ = precision_recall_curve(\n        y_true_bin[:, i], y_score[:, i]\n    )\n    plt.plot(recall[i], precision[i], label=f\"Class {i}\")\n\nplt.xlabel(\"Recall\")\nplt.ylabel(\"Precision\")\nplt.title(\"Precision-Recall Curve\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-28T12:09:48.718906Z","iopub.status.idle":"2026-03-28T12:09:48.71932Z","shell.execute_reply.started":"2026-03-28T12:09:48.719103Z","shell.execute_reply":"2026-03-28T12:09:48.719149Z"}},"outputs":[],"execution_count":null}]}