{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"databundleVersionId":3437841,"sourceId":34478,"sourceType":"competition"}],"dockerImageVersionId":31329,"isGpuEnabled":true,"isInternetEnabled":true,"language":"python","sourceType":"notebook"},"papermill":{"default_parameters":{},"duration":24000.207563,"end_time":"2026-03-28T18:54:10.973063+00:00","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-03-28T12:14:10.7655+00:00","version":"2.7.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"2d6b9543867c4c8a996ffb565bb5ffa3":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"ProgressView","bar_style":"danger","description":"","description_allow_html":false,"layout":"IPY_MODEL_c678ea80191749b2aaebc41017634f99","max":1217502758,"min":0,"orientation":"horizontal","style":"IPY_MODEL_ec3166c8f17c49bf95ee53a560db70e2","tabbable":null,"tooltip":null,"value":536539590}},"426c57220d0144bdb64fcabc8f219902":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_a7f61b4a6a2b4ded90c61cde5455491e","IPY_MODEL_2d6b9543867c4c8a996ffb565bb5ffa3","IPY_MODEL_d1a547f9f3994eb7ae49ae2b62aa9808"],"layout":"IPY_MODEL_75a68992f9544c7fb6a1a4d0173000fa","tabbable":null,"tooltip":null}},"508f951d0cff4fd89428a01cdb060b69":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"5287bb18c0e94751980906cd58e74f7b":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"75a68992f9544c7fb6a1a4d0173000fa":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"8633955752b54531817c362cc9ecd4b1":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"a3b3a3ff4cb44d41a44475a986d4b959":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","background":null,"description_width":"","font_size":null,"text_color":null}},"a7f61b4a6a2b4ded90c61cde5455491e":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_8633955752b54531817c362cc9ecd4b1","placeholder":"​","style":"IPY_MODEL_a3b3a3ff4cb44d41a44475a986d4b959","tabbable":null,"tooltip":null,"value":"model.safetensors:  44%"}},"c678ea80191749b2aaebc41017634f99":{"model_module":"@jupyter-widgets/base","model_module_version":"2.0.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"2.0.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border_bottom":null,"border_left":null,"border_right":null,"border_top":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"d1a547f9f3994eb7ae49ae2b62aa9808":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"2.0.0","_view_name":"HTMLView","description":"","description_allow_html":false,"layout":"IPY_MODEL_5287bb18c0e94751980906cd58e74f7b","placeholder":"​","style":"IPY_MODEL_508f951d0cff4fd89428a01cdb060b69","tabbable":null,"tooltip":null,"value":" 537M/1.22G [00:05&lt;00:23, 28.4MB/s]"}},"ec3166c8f17c49bf95ee53a560db70e2":{"model_module":"@jupyter-widgets/controls","model_module_version":"2.0.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"2.0.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"2.0.0","_view_name":"StyleView","bar_color":null,"description_width":""}}},"version_major":2,"version_minor":0}}},"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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2026-06-16T15:10:29.240645Z","iopub.execute_input":"2026-06-16T15:10:29.241461Z","iopub.status.idle":"2026-06-16T15:10:29.245912Z","shell.execute_reply.started":"2026-06-16T15:10:29.241428Z","shell.execute_reply":"2026-06-16T15:10:29.244848Z"},"papermill":{"duration":0.012068,"end_time":"2026-03-28T12:14:13.247682+00:00","exception":false,"start_time":"2026-03-28T12:14:13.235614+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n# Clean up old checkpoints to ensure a truly fresh start\nfor f in [\"/kaggle/working/checkpoint.pth\", \"/kaggle/working/checkpoint_v2.pth\"]:\n    if os.path.exists(f):\n        os.remove(f)\n        print(f\"Removed old checkpoint: {f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:10:29.247482Z","iopub.execute_input":"2026-06-16T15:10:29.247791Z","iopub.status.idle":"2026-06-16T15:10:29.278088Z","shell.execute_reply.started":"2026-06-16T15:10:29.247770Z","shell.execute_reply":"2026-06-16T15:10:29.277486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport gc\nimport torch\nimport timm\nimport os\nimport gc\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\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":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.278790Z","iopub.execute_input":"2026-06-16T15:10:29.279026Z","iopub.status.idle":"2026-06-16T15:10:29.284361Z","shell.execute_reply.started":"2026-06-16T15:10:29.279005Z","shell.execute_reply":"2026-06-16T15:10:29.283579Z"},"papermill":{"duration":15.925034,"end_time":"2026-03-28T12:14:29.177913+00:00","exception":false,"start_time":"2026-03-28T12:14:13.252879+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import shutil\n\n# shutil.rmtree(\"/kaggle/working/snakeclef\", ignore_errors=True)","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.286083Z","iopub.execute_input":"2026-06-16T15:10:29.286297Z","iopub.status.idle":"2026-06-16T15:10:29.295860Z","shell.execute_reply.started":"2026-06-16T15:10:29.286277Z","shell.execute_reply":"2026-06-16T15:10:29.295198Z"},"papermill":{"duration":0.010755,"end_time":"2026-03-28T12:14:29.194159+00:00","exception":false,"start_time":"2026-03-28T12:14:29.183404+00:00","status":"completed"},"tags":[],"trusted":true},"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":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.296616Z","iopub.execute_input":"2026-06-16T15:10:29.296911Z","iopub.status.idle":"2026-06-16T15:10:29.308233Z","shell.execute_reply.started":"2026-06-16T15:10:29.296879Z","shell.execute_reply":"2026-06-16T15:10:29.307556Z"},"papermill":{"duration":0.009886,"end_time":"2026-03-28T12:14:29.209041+00:00","exception":false,"start_time":"2026-03-28T12:14:29.199155+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"Total VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.308956Z","iopub.execute_input":"2026-06-16T15:10:29.309263Z","iopub.status.idle":"2026-06-16T15:10:29.321308Z","shell.execute_reply.started":"2026-06-16T15:10:29.309230Z","shell.execute_reply":"2026-06-16T15:10:29.320568Z"},"papermill":{"duration":0.261293,"end_time":"2026-03-28T12:14:29.475595+00:00","exception":false,"start_time":"2026-03-28T12:14:29.214302+00:00","status":"completed"},"tags":[],"trusted":true},"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":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.322318Z","iopub.execute_input":"2026-06-16T15:10:29.322701Z","iopub.status.idle":"2026-06-16T15:10:29.336143Z","shell.execute_reply.started":"2026-06-16T15:10:29.322672Z","shell.execute_reply":"2026-06-16T15:10:29.335453Z"},"papermill":{"duration":0.01071,"end_time":"2026-03-28T12:14:29.4915+00:00","exception":false,"start_time":"2026-03-28T12:14:29.48079+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(TRAIN_METADATA)\nprint(\"Total samples:\", len(df))\nprint(\"Total classes:\", df[\"class_id\"].nunique())","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.336989Z","iopub.execute_input":"2026-06-16T15:10:29.337821Z","iopub.status.idle":"2026-06-16T15:10:29.753713Z","shell.execute_reply.started":"2026-06-16T15:10:29.337788Z","shell.execute_reply":"2026-06-16T15:10:29.753000Z"},"papermill":{"duration":0.703548,"end_time":"2026-03-28T12:14:30.200288+00:00","exception":false,"start_time":"2026-03-28T12:14:29.49674+00:00","status":"completed"},"tags":[],"trusted":true},"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 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":{"execution":{"iopub.status.busy":"2026-06-16T15:10:29.756230Z","iopub.execute_input":"2026-06-16T15:10:29.756568Z","iopub.status.idle":"2026-06-16T15:10:30.221443Z","shell.execute_reply.started":"2026-06-16T15:10:29.756543Z","shell.execute_reply":"2026-06-16T15:10:30.220768Z"},"papermill":{"duration":0.491965,"end_time":"2026-03-28T12:14:30.697984+00:00","exception":false,"start_time":"2026-03-28T12:14:30.206019+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# __getitem__ is defined inside SnakeDataset below — this cell is intentionally left empty","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:10:30.222308Z","iopub.execute_input":"2026-06-16T15:10:30.222580Z","iopub.status.idle":"2026-06-16T15:10:30.226635Z","shell.execute_reply.started":"2026-06-16T15:10:30.222556Z","shell.execute_reply":"2026-06-16T15:10:30.225809Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Train/Validation Split","metadata":{"papermill":{"duration":0.005128,"end_time":"2026-03-28T12:14:30.708534+00:00","exception":false,"start_time":"2026-03-28T12:14:30.703406+00:00","status":"completed"},"tags":[]}},{"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":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.227727Z","iopub.execute_input":"2026-06-16T15:10:30.228072Z","iopub.status.idle":"2026-06-16T15:10:30.250151Z","shell.execute_reply.started":"2026-06-16T15:10:30.228021Z","shell.execute_reply":"2026-06-16T15:10:30.249627Z"},"papermill":{"duration":0.019981,"end_time":"2026-03-28T12:14:30.733512+00:00","exception":false,"start_time":"2026-03-28T12:14:30.713531+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Fix Labels","metadata":{"papermill":{"duration":0.005124,"end_time":"2026-03-28T12:14:30.74397+00:00","exception":false,"start_time":"2026-03-28T12:14:30.738846+00:00","status":"completed"},"tags":[]}},{"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 = train_df.copy()\nval_df   = val_df.copy()\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)\nprint(\"Num classes:\", num_classes)","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.251092Z","iopub.execute_input":"2026-06-16T15:10:30.251356Z","iopub.status.idle":"2026-06-16T15:10:30.260866Z","shell.execute_reply.started":"2026-06-16T15:10:30.251336Z","shell.execute_reply":"2026-06-16T15:10:30.260020Z"},"papermill":{"duration":0.017036,"end_time":"2026-03-28T12:14:30.76704+00:00","exception":false,"start_time":"2026-03-28T12:14:30.750004+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Transforms","metadata":{"papermill":{"duration":0.005261,"end_time":"2026-03-28T12:14:30.77788+00:00","exception":false,"start_time":"2026-03-28T12:14:30.772619+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# DenseNet was pretrained at 224×224. Using 518×518 wastes memory\n# with no accuracy benefit for this architecture.\n# Standard ImageNet normalisation is correct for timm DenseNet weights.\n\ntrain_transform = transforms.Compose([\n    transforms.Resize((448, 448)),  # Double the size before cropping\n    transforms.RandAugment(num_ops=2, magnitude=10),\n    transforms.RandomResizedCrop(224), \n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\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":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.262020Z","iopub.execute_input":"2026-06-16T15:10:30.262358Z","iopub.status.idle":"2026-06-16T15:10:30.272839Z","shell.execute_reply.started":"2026-06-16T15:10:30.262325Z","shell.execute_reply":"2026-06-16T15:10:30.272265Z"},"papermill":{"duration":0.012056,"end_time":"2026-03-28T12:14:30.795114+00:00","exception":false,"start_time":"2026-03-28T12:14:30.783058+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset Class","metadata":{"papermill":{"duration":0.005051,"end_time":"2026-03-28T12:14:30.805364+00:00","exception":false,"start_time":"2026-03-28T12:14:30.800313+00:00","status":"completed"},"tags":[]}},{"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.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            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 Exception:\n            return self.__getitem__((idx + 1) % len(self.df))\n\n        label = int(self.df.iloc[idx][\"class_id\"])\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.273875Z","iopub.execute_input":"2026-06-16T15:10:30.274206Z","iopub.status.idle":"2026-06-16T15:10:30.287871Z","shell.execute_reply.started":"2026-06-16T15:10:30.274172Z","shell.execute_reply":"2026-06-16T15:10:30.286829Z"},"papermill":{"duration":0.012848,"end_time":"2026-03-28T12:14:30.823316+00:00","exception":false,"start_time":"2026-03-28T12:14:30.810468+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### DataLoaders","metadata":{"papermill":{"duration":0.004975,"end_time":"2026-03-28T12:14:30.833461+00:00","exception":false,"start_time":"2026-03-28T12:14:30.828486+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from torch.utils.data import WeightedRandomSampler\nimport torch\n\n# --- 1. Calculate Sample Weights for the Training Set ---\n# Get class counts from the training dataframe\nclass_counts = train_df[\"class_id\"].value_counts().sort_index().values\n\n# Calculate weights: inversely proportional to class frequencies\nclass_weights = 1.0 / class_counts\nclass_weights = torch.FloatTensor(class_weights)\n\n# Assign a weight to each individual sample in the training set based on its class\nsample_weights = [class_weights[label] for label in train_df[\"class_id\"]]\nsample_weights = torch.DoubleTensor(sample_weights)\n\n# --- 2. Create the Sampler ---\n# replacement=True allows the sampler to oversample minority classes\nsampler = WeightedRandomSampler(\n    weights=sample_weights, \n    num_samples=len(sample_weights), \n    replacement=True\n)\n\n# --- 3. Update DataLoaders ---\nBATCH_SIZE = 256\nNUM_WORKERS = 2\n\ntrain_dataset = SnakeDataset(train_df, TRAIN_IMG_DIR, train_transform)\nval_dataset = SnakeDataset(val_df, TRAIN_IMG_DIR, val_transform)\n\n# IMPORTANT: Remove 'shuffle=True' because the sampler handles shuffling!\ntrain_loader = DataLoader(\n    train_dataset, \n    batch_size=64, \n    sampler=sampler, \n    num_workers=2,         # KEEP AT 2\n    pin_memory=True,       # ESSENTIAL\n    prefetch_factor=2      # KEEP THIS\n)\n\nval_loader = DataLoader(\n    val_dataset, \n    batch_size=BATCH_SIZE, \n    shuffle=False, \n    num_workers=NUM_WORKERS, \n    pin_memory=True, \n    persistent_workers=True\n)\n\nprint(f\"Train batches : {len(train_loader)}\")\nprint(f\"Val   batches : {len(val_loader)}\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.288974Z","iopub.execute_input":"2026-06-16T15:10:30.289268Z","iopub.status.idle":"2026-06-16T15:10:30.702361Z","shell.execute_reply.started":"2026-06-16T15:10:30.289233Z","shell.execute_reply":"2026-06-16T15:10:30.701634Z"},"papermill":{"duration":0.011768,"end_time":"2026-03-28T12:14:30.85058+00:00","exception":false,"start_time":"2026-03-28T12:14:30.838812+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{"papermill":{"duration":0.005276,"end_time":"2026-03-28T12:14:30.861422+00:00","exception":false,"start_time":"2026-03-28T12:14:30.856146+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# print(\"Available Vision Transformer Models: \")\n# timm.list_models(\"vit*\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.703476Z","iopub.execute_input":"2026-06-16T15:10:30.703859Z","iopub.status.idle":"2026-06-16T15:10:30.707237Z","shell.execute_reply.started":"2026-06-16T15:10:30.703835Z","shell.execute_reply":"2026-06-16T15:10:30.706485Z"},"papermill":{"duration":0.00975,"end_time":"2026-03-28T12:14:30.876372+00:00","exception":false,"start_time":"2026-03-28T12:14:30.866622+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport timm\nimport numpy as np\nimport pandas as pd\nimport torchvision.transforms as transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:10:30.708151Z","iopub.execute_input":"2026-06-16T15:10:30.708465Z","iopub.status.idle":"2026-06-16T15:10:30.719143Z","shell.execute_reply.started":"2026-06-16T15:10:30.708421Z","shell.execute_reply":"2026-06-16T15:10:30.718329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DENSENET_VARIANT = \"densenet121\"\n\nmodel = timm.create_model(\n    DENSENET_VARIANT, \n    pretrained=True, \n    num_classes=num_classes, \n    drop_rate=0.3\n)\nmodel = model.to(device)\n\n# 1. Turn on checkpointing FIRST\nmodel.set_grad_checkpointing(True)\n\n# 2. THEN wrap it in DataParallel\n\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"Model        : {DENSENET_VARIANT}\")\nprint(f\"Total params : {total_params:,}\")\nprint(f\"Trainable    : {trainable_params:,}\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:30.720052Z","iopub.execute_input":"2026-06-16T15:10:30.720285Z","iopub.status.idle":"2026-06-16T15:10:31.028840Z","shell.execute_reply.started":"2026-06-16T15:10:30.720264Z","shell.execute_reply":"2026-06-16T15:10:31.027967Z"},"papermill":{"duration":10.5501,"end_time":"2026-03-28T12:14:41.431713+00:00","exception":false,"start_time":"2026-03-28T12:14:30.881613+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Helper: works regardless of DataParallel wrapper\ndef get_base_model(m):\n    return m.module if isinstance(m, torch.nn.DataParallel) else m\n\n# Freeze everything except the classifier head initially\ndef freeze_backbone(m):\n    for param in m.parameters():\n        param.requires_grad = False\n    for param in get_base_model(m).classifier.parameters():\n        param.requires_grad = True\n\ndef unfreeze_backbone(m):\n    for param in m.parameters():\n        param.requires_grad = True\n\nfreeze_backbone(model)\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Trainable params after freeze: {trainable:,}  (classifier head only)\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:31.031761Z","iopub.execute_input":"2026-06-16T15:10:31.032094Z","iopub.status.idle":"2026-06-16T15:10:31.040982Z","shell.execute_reply.started":"2026-06-16T15:10:31.032072Z","shell.execute_reply":"2026-06-16T15:10:31.040272Z"},"papermill":{"duration":0.012882,"end_time":"2026-03-28T12:14:41.450632+00:00","exception":false,"start_time":"2026-03-28T12:14:41.43775+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loss + Optimizer","metadata":{"papermill":{"duration":0.005392,"end_time":"2026-03-28T12:14:41.461448+00:00","exception":false,"start_time":"2026-03-28T12:14:41.456056+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Keep label_smoothing=0.15, it acts as a regularizer that prevents the model from being overconfident\ncriterion = torch.nn.CrossEntropyLoss(label_smoothing=0.15)\n\n# Only optimise unfrozen (classifier) params initially\noptimizer = torch.optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()), \n    lr=3e-4, \n    weight_decay=1e-3  # Increased from 1e-4\n)\n\n# CosineAnnealingWarmRestarts restarts every T_0 epochs — helps escape plateaus\n# Use a more aggressive factor for the final stage of training\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, \n    mode='min', \n    factor=0.3,    # Steeper reduction to settle into the local minimum\n    patience=3,    # Wait a bit longer to be sure we are plateaued\n    min_lr=1e-7\n)\n\nscaler = torch.amp.GradScaler(\"cuda\")\nprint(\"Criterion, optimizer, scheduler and scaler initialised.\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:31.041962Z","iopub.execute_input":"2026-06-16T15:10:31.042646Z","iopub.status.idle":"2026-06-16T15:10:31.064596Z","shell.execute_reply.started":"2026-06-16T15:10:31.042622Z","shell.execute_reply":"2026-06-16T15:10:31.064014Z"},"papermill":{"duration":0.012959,"end_time":"2026-03-28T12:14:41.479591+00:00","exception":false,"start_time":"2026-03-28T12:14:41.466632+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Added \"_v2\" to force a completely fresh start!\nCKPT_PATH = \"/kaggle/working/checkpoint_v2.pth\" \nBEST_PATH = \"/kaggle/working/best_model_v2.pth\"\n\ndef save_checkpoint(epoch, batch_idx, model, optimizer, scheduler, scaler, 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        \"scaler_state\": scaler.state_dict(),\n        \"best_acc\": best_acc\n    }, CKPT_PATH)","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:31.065550Z","iopub.execute_input":"2026-06-16T15:10:31.066409Z","iopub.status.idle":"2026-06-16T15:10:31.071351Z","shell.execute_reply.started":"2026-06-16T15:10:31.066376Z","shell.execute_reply":"2026-06-16T15:10:31.070377Z"},"papermill":{"duration":0.010645,"end_time":"2026-03-28T12:14:41.495964+00:00","exception":false,"start_time":"2026-03-28T12:14:41.485319+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_epoch = 0\nstart_batch = 0\nbest_acc    = 0.0\n\nif os.path.exists(CKPT_PATH):\n    checkpoint = torch.load(CKPT_PATH, 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    if \"scaler_state\" in checkpoint:\n        scaler.load_state_dict(checkpoint[\"scaler_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 + 1}, batch {start_batch}\")\nelse:\n    print(\"No checkpoint found — starting fresh.\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:31.072376Z","iopub.execute_input":"2026-06-16T15:10:31.073131Z","iopub.status.idle":"2026-06-16T15:10:31.084175Z","shell.execute_reply.started":"2026-06-16T15:10:31.073106Z","shell.execute_reply":"2026-06-16T15:10:31.083395Z"},"papermill":{"duration":0.011819,"end_time":"2026-03-28T12:14:41.513285+00:00","exception":false,"start_time":"2026-03-28T12:14:41.501466+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EarlyStopping:\n    def __init__(self, patience=5, min_delta=0.0):\n        \"\"\"\n        Args:\n            patience (int): How many epochs to wait after last time validation loss improved.\n            min_delta (float): Minimum change in the monitored quantity to qualify as an improvement.\n        \"\"\"\n        self.patience = patience\n        self.min_delta = min_delta\n        self.counter = 0\n        self.best_loss = None\n        self.early_stop = False\n\n    def __call__(self, val_loss):\n        if self.best_loss is None:\n            self.best_loss = val_loss\n        elif val_loss > self.best_loss - self.min_delta:\n            self.counter += 1\n            print(f\"  [EarlyStopping] Patience counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_loss = val_loss\n            self.counter = 0\n            print(f\"  [EarlyStopping] Validation loss improved.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:10:31.085137Z","iopub.execute_input":"2026-06-16T15:10:31.085829Z","iopub.status.idle":"2026-06-16T15:10:31.099987Z","shell.execute_reply.started":"2026-06-16T15:10:31.085798Z","shell.execute_reply":"2026-06-16T15:10:31.099247Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training Loop","metadata":{"papermill":{"duration":0.005131,"end_time":"2026-03-28T12:14:41.541463+00:00","exception":false,"start_time":"2026-03-28T12:14:41.536332+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_losses = []\nval_losses   = []\nval_accs     = []\n\nUNFREEZE_EPOCH   = 3     # Unfreeze backbone after N warm-up epochs\nFINETUNE_LR      = 5e-5  # Lower LR for full fine-tuning\nCHECKPOINT_EVERY = 300   # Save checkpoint every N batches\nNUM_EPOCHS       = 25    # Reduced from 30; with larger batch + proper LR, converges faster\n\nbackbone_unfrozen = (start_epoch >= UNFREEZE_EPOCH)\n\n# If resuming past unfreeze point, make sure backbone is unfrozen\nif backbone_unfrozen:\n    unfreeze_backbone(model)\n    for pg in optimizer.param_groups:\n        pg['lr'] = FINETUNE_LR\nearly_stopping = EarlyStopping(patience=5, min_delta=0.001)\nema_model = timm.utils.ModelEmaV2(model, decay=0.999)\nfor epoch in range(start_epoch, NUM_EPOCHS):\n    model.train()\n    # ── Unfreeze backbone ──────────────────────────────────────────────────\n    if epoch == UNFREEZE_EPOCH:\n        print(f\"Epoch {epoch}: Backbone unfrozen | LR = {FINETUNE_LR}\")\n        unfreeze_backbone(model)\n        \n        # MUST re-initialize the optimizer to include all 7 million parameters!\n        optimizer = torch.optim.AdamW(model.parameters(), lr=FINETUNE_LR, weight_decay=1e-4)\n        \n        # Re-initialize the scheduler for the new optimizer\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n            optimizer, T_0=5, T_mult=2, eta_min=1e-6\n        )\n        print(f\"Epoch {epoch+1}: Backbone unfrozen | LR = {FINETUNE_LR}\")\n\n    # ── Train ──────────────────────────────────────────────────────────────\n    \n    running_loss = 0.0\n    batches_run  = 0\n\n    for batch_idx, (images, labels) in enumerate(train_loader):\n\n        # Resume skip\n        if epoch == start_epoch and batch_idx < start_batch:\n            continue\n\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\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        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        ema_model.update(model)\n        running_loss += loss.item()\n        batches_run  += 1\n\n        if batch_idx % 100 == 0:\n            current_lr = optimizer.param_groups[0]['lr']\n            print(f\"Epoch {epoch+1} | Batch {batch_idx}/{len(train_loader)} \"\n                  f\"| Loss {loss.item():.4f} | LR {current_lr:.2e}\")\n\n        if batch_idx % CHECKPOINT_EVERY == 0 and batch_idx != 0:\n            save_checkpoint(epoch, batch_idx, model, optimizer, scheduler, scaler, best_acc)\n            print(f\"  Checkpoint saved at batch {batch_idx}\")\n\n        # Free batch tensors immediately\n        del images, labels, outputs, loss\n\n    # Reset start_batch after first epoch resumes\n    if epoch == start_epoch:\n        start_batch = 0\n\n\n    epoch_loss = running_loss / max(batches_run, 1)\n    train_losses.append(epoch_loss)\n\n    # ── Validate ───────────────────────────────────────────────────────────\n    model.eval()\n\n    all_preds   = []\n    all_targets = []\n    val_running_loss = 0.0\n\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            with torch.amp.autocast(\"cuda\"):\n                outputs  = model(images)\n                val_loss = criterion(outputs, labels)\n\n            val_running_loss += val_loss.item()\n            _, predicted = torch.max(outputs, 1)\n\n            all_preds.extend(predicted.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n\n            del images, labels, outputs, val_loss\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    # Pass the validation loss variable here\n    scheduler.step(val_epoch_loss)\n    # ── Early Stopping Check ───────────────────────────────────────────────\n    early_stopping(val_epoch_loss)\n    if early_stopping.early_stop:\n        print(\"Early stopping triggered. Training halted.\")\n        break\n\n    # ── Save best ──────────────────────────────────────────────────────────\n    if acc > best_acc:\n        best_acc = acc\n        torch.save(model.state_dict(), BEST_PATH)\n        print(f\"  Best model saved! ({best_acc:.4f})\")\n\n    # ── End-of-epoch checkpoint ────────────────────────────────────────────\n    save_checkpoint(epoch + 1, 0, model, optimizer, scheduler, scaler, best_acc)\n\n    # ── Memory cleanup ─────────────────────────────────────────────────────\n    del all_preds, all_targets\n    gc.collect()\n    torch.cuda.empty_cache()\n    print(f\"  VRAM after cleanup: \"\n          f\"{torch.cuda.memory_allocated() / 1e9:.2f} GB / \"\n          f\"{torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\n\nprint(f\"\\nTraining complete. Best val accuracy: {best_acc:.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:10:31.100964Z","iopub.execute_input":"2026-06-16T15:10:31.101247Z","iopub.status.idle":"2026-06-16T15:48:11.358856Z","shell.execute_reply.started":"2026-06-16T15:10:31.101217Z","shell.execute_reply":"2026-06-16T15:48:11.358092Z"},"papermill":{"duration":23959.399229,"end_time":"2026-03-28T18:54:00.945952+00:00","exception":false,"start_time":"2026-03-28T12:14:41.546723+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"EVALUATION","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load(BEST_PATH, weights_only=False))\nmodel.eval()\n\nall_preds   = []\nall_targets = []\nall_probs   = []\n\nwith torch.no_grad():\n    for images, labels in val_loader:\n        images = images.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        with torch.amp.autocast(\"cuda\"):\n            outputs = model(images)\n\n        prob      = torch.softmax(outputs.float(), dim=1)   # float32 for softmax accuracy\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        del images, labels, outputs, prob\n\nprint(\"Evaluation done.\")\nprint(f\"Val samples : {len(all_preds)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:11.361100Z","iopub.execute_input":"2026-06-16T15:48:11.361526Z","iopub.status.idle":"2026-06-16T15:48:18.714818Z","shell.execute_reply.started":"2026-06-16T15:48:11.361492Z","shell.execute_reply":"2026-06-16T15:48:18.713921Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Save Model","metadata":{"papermill":{"duration":0.014226,"end_time":"2026-03-28T18:54:00.983303+00:00","exception":false,"start_time":"2026-03-28T18:54:00.969077+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/snake_densenet_final.pth\")\nprint(\"Model saved.\")","metadata":{"execution":{"iopub.status.busy":"2026-06-16T15:48:18.715959Z","iopub.execute_input":"2026-06-16T15:48:18.716286Z","iopub.status.idle":"2026-06-16T15:48:18.835298Z","shell.execute_reply.started":"2026-06-16T15:48:18.716252Z","shell.execute_reply":"2026-06-16T15:48:18.834460Z"},"papermill":{"duration":1.725122,"end_time":"2026-03-28T18:54:02.722731+00:00","exception":false,"start_time":"2026-03-28T18:54:00.997609+00:00","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualizations\n## Group A — Training Curves","metadata":{}},{"cell_type":"markdown","source":"### 1. Learning Curves","metadata":{}},{"cell_type":"code","source":"fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))\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()\nax1.grid(True, alpha=0.3)\n\nax2.plot(val_accs, marker='o', color='seagreen')\nax2.axhline(y=best_acc, color='red', linestyle='--', label=f'Best: {best_acc:.4f}')\nax2.set_title(\"Val Accuracy per Epoch\")\nax2.set_xlabel(\"Epoch\")\nax2.set_ylabel(\"Accuracy\")\nax2.legend()\nax2.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/learning_curves.png\", dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:18.836293Z","iopub.execute_input":"2026-06-16T15:48:18.836666Z","iopub.status.idle":"2026-06-16T15:48:19.532366Z","shell.execute_reply.started":"2026-06-16T15:48:18.836641Z","shell.execute_reply":"2026-06-16T15:48:19.531544Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2. Confidence Distribution","metadata":{}},{"cell_type":"code","source":"confidences = np.max(all_probs, axis=1)\n\nplt.figure(figsize=(8, 5))\nplt.hist(confidences, bins=40, color='steelblue', edgecolor='white')\nplt.title(\"Prediction Confidence Distribution\")\nplt.xlabel(\"Confidence\")\nplt.ylabel(\"Frequency\")\nplt.axvline(x=confidences.mean(), color='red', linestyle='--',\n            label=f'Mean: {confidences.mean():.2f}')\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:19.533565Z","iopub.execute_input":"2026-06-16T15:48:19.533897Z","iopub.status.idle":"2026-06-16T15:48:19.708878Z","shell.execute_reply.started":"2026-06-16T15:48:19.533870Z","shell.execute_reply":"2026-06-16T15:48:19.708248Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3. CLASSIFICATION DISTRIBLUTION","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\nprint(classification_report(all_targets, all_preds, zero_division=0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:19.710028Z","iopub.execute_input":"2026-06-16T15:48:19.710391Z","iopub.status.idle":"2026-06-16T15:48:19.728191Z","shell.execute_reply.started":"2026-06-16T15:48:19.710352Z","shell.execute_reply":"2026-06-16T15:48:19.727466Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4. 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=(12, 10))\nsns.heatmap(cm, annot=(num_classes <= 20), cmap=\"Blues\", fmt='d')\nplt.title(\"Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/confusion_matrix.png\", dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:19.729074Z","iopub.execute_input":"2026-06-16T15:48:19.729325Z","iopub.status.idle":"2026-06-16T15:48:20.858459Z","shell.execute_reply.started":"2026-06-16T15:48:19.729302Z","shell.execute_reply":"2026-06-16T15:48:20.857978Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Group B — Evaluation Metrics","metadata":{}},{"cell_type":"markdown","source":"### 5. Classification Report","metadata":{}},{"cell_type":"markdown","source":"### 6. Confusion Matrix","metadata":{}},{"cell_type":"markdown","source":"### 7. Per-Class Accuracy Bar Chart","metadata":{}},{"cell_type":"markdown","source":"### 8. Classification Report Heatmap","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import classification_report\nreport = classification_report(all_targets, all_preds, output_dict=True)\nreport_df = pd.DataFrame(report).transpose()\nreport_df = report_df.drop(columns=[\"support\"], errors=\"ignore\")\nreport_df = report_df[:-3]  # drop macro/weighted avg rows for clarity\n\nplt.figure(figsize=(8, max(8, len(report_df) // 3)))\nsns.heatmap(report_df.astype(float), annot=True, fmt=\".2f\", cmap=\"YlGnBu\",\n            linewidths=0.5, cbar=True)\nplt.title(\"Classification Report Heatmap (Precision / Recall / F1)\")\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/classification_report_heatmap.png\", dpi=150)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:20.859396Z","iopub.execute_input":"2026-06-16T15:48:20.860036Z","iopub.status.idle":"2026-06-16T15:48:22.096660Z","shell.execute_reply.started":"2026-06-16T15:48:20.860010Z","shell.execute_reply":"2026-06-16T15:48:22.095825Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9. 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))\ny_true_bin  = label_binarize(all_targets, classes=range(n_classes))\ny_score     = np.array(all_probs)\n\n# Compute macro-average ROC to avoid 100-line legend\nfpr_all, tpr_all, roc_auc_all = {}, {}, {}\nfor i in range(n_classes):\n    fpr_all[i], tpr_all[i], _ = roc_curve(y_true_bin[:, i], y_score[:, i])\n    roc_auc_all[i] = auc(fpr_all[i], tpr_all[i])\n\nmean_auc = np.mean(list(roc_auc_all.values()))\n\n# Plot individual curves (lighter) + mean AUC annotation\nplt.figure(figsize=(8, 6))\nfor i in range(n_classes):\n    plt.plot(fpr_all[i], tpr_all[i], alpha=0.2, linewidth=0.8)\nplt.plot([0,1],[0,1],'k--', linewidth=1.5)\nplt.title(f\"ROC Curves — Macro-avg AUC = {mean_auc:.3f}\")\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/roc_curves.png\", dpi=150, bbox_inches='tight')\nplt.show()\n\nprint(f\"Macro-average AUC: {mean_auc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:22.097833Z","iopub.execute_input":"2026-06-16T15:48:22.098255Z","iopub.status.idle":"2026-06-16T15:48:22.506093Z","shell.execute_reply.started":"2026-06-16T15:48:22.098227Z","shell.execute_reply":"2026-06-16T15:48:22.505436Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 10. Precision-Recall Curve","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import precision_recall_curve\nplt.figure(figsize=(10, 6))\nfor i in range(n_classes):\n    prec, rec, _ = precision_recall_curve(y_true_bin[:, i], y_score[:, i])\n    plt.plot(rec, prec, linewidth=0.6, alpha=0.5, label=f\"Class {i}\")\n\nplt.title(\"Precision-Recall Curve (per class)\")\nplt.xlabel(\"Recall\")\nplt.ylabel(\"Precision\")\nplt.legend(loc=\"lower left\", fontsize=6, ncol=3)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/pr_curve.png\", dpi=150)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:22.507200Z","iopub.execute_input":"2026-06-16T15:48:22.507512Z","iopub.status.idle":"2026-06-16T15:48:23.567117Z","shell.execute_reply.started":"2026-06-16T15:48:22.507488Z","shell.execute_reply":"2026-06-16T15:48:23.566267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 11. Top-K Accuracy (Top-1, Top-3, Top-5)","metadata":{}},{"cell_type":"code","source":"y_score_t = np.array(all_probs)\ny_true_t  = np.array(all_targets)\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(y_true_t, y_score_t, k) for k in ks]\n\nplt.figure(figsize=(6, 4))\nbars = plt.bar([f\"Top-{k}\" for k in ks], accs, color=['steelblue','seagreen','tomato'])\nfor bar, acc in zip(bars, accs):\n    plt.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.005,\n             f\"{acc:.4f}\", ha='center', fontsize=11)\nplt.title(\"Top-K Accuracy\")\nplt.ylabel(\"Accuracy\")\nplt.ylim(0, 1.05)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/topk_accuracy.png\", dpi=150)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:23.568285Z","iopub.execute_input":"2026-06-16T15:48:23.568652Z","iopub.status.idle":"2026-06-16T15:48:23.795741Z","shell.execute_reply.started":"2026-06-16T15:48:23.568626Z","shell.execute_reply":"2026-06-16T15:48:23.795113Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 12. Confidence Score Distribution (Correct vs Incorrect)","metadata":{}},{"cell_type":"markdown","source":"### 13. Calibration Curve (Reliability Diagram)","metadata":{}},{"cell_type":"code","source":"from sklearn.calibration import calibration_curve\n\n# Use max confidence as the probability for the predicted class\nconfidences  = np.max(all_probs, axis=1)\ncorrect_mask = (np.array(all_preds) == np.array(all_targets)).astype(int)\n\nfraction_of_positives, mean_predicted_value = calibration_curve(\n    correct_mask, confidences, n_bins=10\n)\n\nplt.figure(figsize=(7, 5))\nplt.plot(mean_predicted_value, fraction_of_positives, marker='o',\n         color='steelblue', label='Model')\nplt.plot([0, 1], [0, 1], 'k--', label='Perfect calibration')\nplt.title(\"Calibration Curve (Reliability Diagram)\")\nplt.xlabel(\"Mean Predicted Confidence\")\nplt.ylabel(\"Fraction Correct\")\nplt.legend()\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/calibration_curve.png\", dpi=150)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:23.796649Z","iopub.execute_input":"2026-06-16T15:48:23.797013Z","iopub.status.idle":"2026-06-16T15:48:24.099616Z","shell.execute_reply.started":"2026-06-16T15:48:23.796986Z","shell.execute_reply":"2026-06-16T15:48:24.098988Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import optuna\n\ndef objective(trial):\n    # 1. Suggest hyperparameters\n    lr = trial.suggest_float(\"lr\", 1e-5, 1e-3, log=True)\n    drop_rate = trial.suggest_float(\"drop_rate\", 0.1, 0.5)\n    \n    # 2. Build model with suggested params\n    # Change this line in your Optuna cell:\n    model = timm.create_model(\"densenet121\", pretrained=True, num_classes=48, drop_rate=drop_rate)\n    model = model.to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr)\n    \n    # 3. Use a mini-loop to evaluate this trial\n    # We only train for 3 epochs to see if these parameters are \"promising\"\n    for epoch in range(3):\n        model.train()\n        for images, labels in train_loader:\n            images, labels = images.to(device), labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n        \n        # 4. Validation\n        model.eval()\n        correct = 0\n        with torch.no_grad():\n            for images, labels in val_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                _, predicted = torch.max(outputs, 1)\n                correct += (predicted == labels).sum().item()\n        \n        val_accuracy = correct / len(val_dataset)\n        \n        # 5. Report to Optuna\n        trial.report(val_accuracy, epoch)\n        if trial.should_prune():\n            raise optuna.exceptions.TrialPruned()\n            \n    return val_accuracy\n\n# Run the study\nstudy = optuna.create_study(direction=\"maximize\")\nstudy.optimize(objective, n_trials=10)\n\nprint(\"Best hyperparameters:\", study.best_params)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-16T15:48:24.100603Z","iopub.execute_input":"2026-06-16T15:48:24.100909Z","iopub.status.idle":"2026-06-16T15:48:35.701441Z","shell.execute_reply.started":"2026-06-16T15:48:24.100884Z","shell.execute_reply":"2026-06-16T15:48:35.699913Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Group C — Error Analysis","metadata":{}},{"cell_type":"markdown","source":"### 14. Most Confused Class Pairs","metadata":{}},{"cell_type":"markdown","source":"### 15. High Confidence Wrong Predictions","metadata":{}},{"cell_type":"markdown","source":"### 16. Low Confidence Correct Predictions","metadata":{}},{"cell_type":"markdown","source":"### 17. Per-Class Confidence Box Plot","metadata":{}},{"cell_type":"markdown","source":"## Group D — Dataset Analysis","metadata":{}},{"cell_type":"markdown","source":"### 18. No-Snake vs Snake Balance","metadata":{}},{"cell_type":"markdown","source":"### 19. Train vs Val Class Distribution","metadata":{}},{"cell_type":"markdown","source":"## Group E — Feature Embeddings","metadata":{}},{"cell_type":"markdown","source":"### 20. PCA of Feature Embeddings","metadata":{}},{"cell_type":"markdown","source":"### 21. t-SNE of Feature Embeddings","metadata":{}},{"cell_type":"markdown","source":"### 22. PCA Explained Variance","metadata":{}},{"cell_type":"markdown","source":"# Save Model","metadata":{}}]}