{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":4521,"databundleVersionId":326986,"sourceType":"competition"}],"dockerImageVersionId":31193,"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\nimport numpy as np # linear algebra\nimport 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\nimport os\nfor 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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# 1. Load the Data\n# NOTE: On Kaggle, your path is likely \"../input/[competition-name]/train.csv\"\n# You may need to update this string depending on where the file is located.\nfile_path = '/kaggle/input/noaa-right-whale-recognition/train.csv' \n\ntry:\n    df = pd.read_csv(file_path)\n    print(\"✅ Data loaded successfully!\")\nexcept FileNotFoundError:\n    print(f\"❌ File not found at {file_path}. Please check your path.\")\n    # Creating dummy data so the code runs for demonstration purposes if file is missing\n    data = {'Image': [f'img_{i}.jpg' for i in range(100)], \n            'whaleID': ['whale_001']*20 + ['whale_002']*15 + ['whale_003']*10 + [f'whale_{i}' for i in range(4, 59)]}\n    df = pd.DataFrame(data)\n\n# 2. Basic Statistics\nnum_images = len(df)\nnum_whales = df['whaleID'].nunique()\nimages_per_whale = df['whaleID'].value_counts()\n\nprint(\"-\" * 30)\nprint(f\"Total Images: {num_images}\")\nprint(f\"Total Unique Whales: {num_whales}\")\nprint(\"-\" * 30)\nprint(\"Top 5 most frequent whales:\")\nprint(images_per_whale.head())\nprint(\"-\" * 30)\nprint(\"Stats on images per whale:\")\nprint(images_per_whale.describe())\nprint(\"-\" * 30)\n\n# 3. Visualize the Top 20 Whales\nplt.figure(figsize=(12, 8))\ntop_20_whales = images_per_whale.head(20)\n\nsns.barplot(x=top_20_whales.index, y=top_20_whales.values, palette=\"viridis\")\n\nplt.title('Top 20 Most Frequent Whales', fontsize=16)\nplt.xlabel('Whale ID', fontsize=12)\nplt.ylabel('Number of Images', fontsize=12)\nplt.xticks(rotation=45, ha='right') # Rotate labels for readability\nplt.tight_layout()\n\nplt.show()\n\n# 4. Visualize the \"Long Tail\" (Optional but recommended)\n# This shows how many whales have very few images (imbalance check)\nplt.figure(figsize=(10, 6))\nplt.hist(images_per_whale.values, bins=50, color='teal', edgecolor='black')\nplt.yscale('log') # Log scale because the disparity is usually huge\nplt.title('Distribution of Images per Whale (Log Scale)', fontsize=16)\nplt.xlabel('Number of Images', fontsize=12)\nplt.ylabel('Count of Whales (Log Scale)', fontsize=12)\nplt.grid(axis='y', alpha=0.5)\nplt.tight_layout()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T12:59:08.946101Z","iopub.execute_input":"2025-11-23T12:59:08.946677Z","iopub.status.idle":"2025-11-23T12:59:09.614402Z","shell.execute_reply.started":"2025-11-23T12:59:08.946642Z","shell.execute_reply":"2025-11-23T12:59:09.613677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport os\n\n# --- CONFIGURATION ---\nCSV_PATH = '/kaggle/input/noaa-right-whale-recognition/train.csv'\n\ndef run_simple_analysis():\n    print(\"📊 Loading dataset...\")\n    try:\n        df = pd.read_csv(CSV_PATH)\n        # Remove known bad image if present\n        df = df[df['Image'] != 'w_7489.jpg']\n    except FileNotFoundError:\n        print(\"⚠️ CSV file not found. Please check the path.\")\n        return\n\n    # --- INSIGHTS WITHOUT GRAPHS (Text Analysis) ---\n    total_images = len(df)\n    total_whales = df['whaleID'].nunique()\n    counts = df['whaleID'].value_counts()\n    \n    single_image_whales = counts[counts == 1].count()\n    rare_whales = counts[counts < 5].count()\n    \n    print(\"\\n\" + \"=\"*30)\n    print(\"🧐 DATASET INSIGHTS SUMMARY\")\n    print(\"=\"*30)\n    print(f\"• Total Images:      {total_images}\")\n    print(f\"• Unique Whales:     {total_whales}\")\n    print(f\"• Average Images/Whale: {total_images / total_whales:.2f}\")\n    print(\"-\" * 30)\n    print(f\"• Most Popular Whale:   {counts.index[0]} ({counts.iloc[0]} images)\")\n    print(f\"• 'One-Shot' Whales:    {single_image_whales} whales have exactly 1 image.\")\n    print(f\"• Rare Whales (<5 imgs): {rare_whales} whales ({rare_whales/total_whales*100:.1f}% of total).\")\n    print(\"=\"*30 + \"\\n\")\n\n    # --- GRAPH 1: THE \"CELEBRITY\" WHALES (Bar Chart) ---\n    # Shows who the most common whales are.\n    plt.figure(figsize=(12, 6))\n    top_20 = counts.head(20)\n    sns.barplot(x=top_20.index, y=top_20.values, palette=\"viridis\")\n    plt.title('Top 20 Most Photographed Whales', fontsize=15)\n    plt.xlabel('Whale ID')\n    plt.ylabel('Number of Images')\n    plt.xticks(rotation=45, ha='right')\n    plt.grid(axis='y', alpha=0.3)\n    plt.tight_layout()\n    plt.savefig('graph_1_top_whales.png')\n    plt.show()\n    \n    # --- GRAPH 2: THE \"LONG TAIL\" (Histogram) ---\n    # Shows the extreme class imbalance.\n    plt.figure(figsize=(10, 6))\n    plt.hist(counts.values, bins=range(1, 50), color='#34495e', edgecolor='white')\n    plt.title('Distribution of Images per Whale (Class Imbalance)', fontsize=15)\n    plt.xlabel('Number of Images Available')\n    plt.ylabel('Count of Whales')\n    plt.axvline(x=5, color='red', linestyle='--', label='Rare Threshold (5 images)')\n    plt.legend()\n    plt.grid(axis='y', alpha=0.3)\n    plt.tight_layout()\n    plt.savefig('graph_2_imbalance_hist.png')\n    plt.show()\n\n    # --- GRAPH 3: DIFFICULTY BREAKDOWN (Pie Chart) ---\n    # Visualizes the challenge level.\n    labels = ['One-Shot (1 Img)', 'Few-Shot (2-10 Imgs)', 'Frequent (>10 Imgs)']\n    \n    c1 = counts[counts == 1].count()\n    c2 = counts[(counts >= 2) & (counts <= 10)].count()\n    c3 = counts[counts > 10].count()\n    \n    sizes = [c1, c2, c3]\n    colors = ['#e74c3c', '#f39c12', '#2ecc71'] # Red (Hard), Orange (Medium), Green (Easy)\n    \n    plt.figure(figsize=(8, 8))\n    plt.pie(sizes, labels=labels, autopct='%1.1f%%', colors=colors, startangle=140, explode=(0.05, 0, 0))\n    plt.title('Whale Rarity Breakdown (Difficulty Level)', fontsize=15)\n    plt.savefig('graph_3_rarity_pie.png')\n    plt.show()\n\n    print(\"✅ Analysis complete. Graphs saved as png files.\")\n\nif __name__ == \"__main__\":\n    run_simple_analysis()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T15:53:23.768111Z","iopub.execute_input":"2025-11-23T15:53:23.768967Z","iopub.status.idle":"2025-11-23T15:53:24.800224Z","shell.execute_reply.started":"2025-11-23T15:53:23.768932Z","shell.execute_reply":"2025-11-23T15:53:24.799415Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport zipfile\nimport glob\n\n# --- CONFIGURATION ---\n# Note: On Kaggle, input is read-only. We will extract to /kaggle/working/\nCONFIG = {\n    # Point this to the ZIP file if that's what you have\n    'IMG_DIR': '/kaggle/input/noaa-right-whale-recognition/imgs.zip', \n    'CSV_PATH': '/kaggle/input/noaa-right-whale-recognition/train.csv',\n    'IMG_SIZE': 256,\n    'BATCH_SIZE': 32,\n    'NUM_WORKERS': 2,\n    'SEED': 42\n}\n\ndef handle_zip_extraction(path_from_config):\n    \"\"\"\n    Checks if IMG_DIR is a zip file. If so, extracts it to a writable directory\n    and returns the new path to the folder containing images.\n    \"\"\"\n    if not path_from_config.endswith('.zip'):\n        return path_from_config\n\n    # Define extraction target (Kaggle working dir)\n    extract_root = \"./extracted_data\"\n    \n    # Check if already extracted to avoid re-doing it on re-runs\n    if os.path.exists(extract_root) and len(os.listdir(extract_root)) > 0:\n        print(f\"✅ Zip already extracted at: {extract_root}\")\n    else:\n        print(f\"📦 Extracting {path_from_config} to {extract_root}...\")\n        os.makedirs(extract_root, exist_ok=True)\n        with zipfile.ZipFile(path_from_config, 'r') as zip_ref:\n            zip_ref.extractall(extract_root)\n        print(\"✅ Extraction complete.\")\n\n    # CRITICAL STEP: Find where the images actually are.\n    # Sometimes zip files have a folder inside them (e.g., imgs.zip -> imgs/ -> image.jpg)\n    # or they have images at the root (imgs.zip -> image.jpg).\n    \n    # Check for a subdirectory that matches the zip name (common Kaggle pattern)\n    subfolder_name = os.path.splitext(os.path.basename(path_from_config))[0] # 'imgs'\n    potential_subfolder = os.path.join(extract_root, subfolder_name)\n    \n    if os.path.exists(potential_subfolder):\n        return potential_subfolder\n    \n    # Otherwise, images are likely in the extract_root directly\n    return extract_root\n\ndef load_and_preprocess_df(csv_path):\n    \"\"\"\n    Loads CSV, encodes labels, and handles the train/val split logic\n    specifically for the rare whale problem.\n    \"\"\"\n    df = pd.read_csv(csv_path)\n    \n    # 1. Label Encoding\n    encoder = LabelEncoder()\n    df['label_idx'] = encoder.fit_transform(df['whaleID'])\n    classes = encoder.classes_\n    print(f\"✅ Label Encoding Complete. Found {len(classes)} unique whales.\")\n    \n    # 2. Stratified Split Logic\n    counts = df.whaleID.value_counts()\n    single_shot_whales = counts[counts == 1].index\n    \n    df_single = df[df.whaleID.isin(single_shot_whales)]\n    df_multi = df[~df.whaleID.isin(single_shot_whales)]\n    \n    print(f\"📊 Split Stats: {len(df_single)} whales have only 1 image (forced to Train).\")\n    \n    train_multi, val_multi = train_test_split(\n        df_multi, \n        test_size=0.1, \n        random_state=CONFIG['SEED'], \n        stratify=df_multi['whaleID']\n    )\n    \n    df_train = pd.concat([train_multi, df_single]).sample(frac=1).reset_index(drop=True)\n    df_val = val_multi.sample(frac=1).reset_index(drop=True)\n    \n    print(f\"✅ Final Split: Train: {len(df_train)} images, Val: {len(df_val)} images\")\n    \n    return df_train, df_val, classes\n\n# --- AUGMENTATION ---\ndef get_transforms(data='train'):\n    if data == 'train':\n        return A.Compose([\n            A.Resize(CONFIG['IMG_SIZE'], CONFIG['IMG_SIZE']),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Rotate(limit=30, p=0.7),\n            A.RandomBrightnessContrast(p=0.2),\n            A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=20, p=0.2),\n            A.GaussianBlur(p=0.1),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n    elif data == 'val':\n        return A.Compose([\n            A.Resize(CONFIG['IMG_SIZE'], CONFIG['IMG_SIZE']),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n\n# --- DATASET CLASS ---\nclass WhaleDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_name = row['Image']\n        label = row['label_idx']\n        \n        # Construct full path\n        img_path = os.path.join(self.img_dir, img_name)\n        \n        # Load Image\n        image = cv2.imread(img_path)\n        \n        if image is None:\n            # Fallback/Debug info if image isn't found\n            # print(f\"Warning: Could not load {img_path}\") \n            image = np.zeros((CONFIG['IMG_SIZE'], CONFIG['IMG_SIZE'], 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            \n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n            \n        return image, label\n\n# --- VISUALIZATION HELPER ---\ndef visualize_batch(dataloader, classes):\n    images, labels = next(iter(dataloader))\n    plt.figure(figsize=(16, 8))\n    for i in range(min(8, len(images))):\n        ax = plt.subplot(2, 4, i + 1)\n        img = images[i].permute(1, 2, 0).numpy()\n        img = img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406]\n        img = np.clip(img, 0, 1)\n        plt.imshow(img)\n        plt.title(classes[labels[i]])\n        plt.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n\n# --- MAIN EXECUTION ---\nif __name__ == \"__main__\":\n    \n    # 0. Handle Zip Extraction\n    # This automatically updates the IMG_DIR to the folder where we extracted the files\n    real_img_dir = handle_zip_extraction(CONFIG['IMG_DIR'])\n    print(f\"📂 Images will be loaded from: {real_img_dir}\")\n\n    # 1. Prepare DataFrames\n    if os.path.exists(CONFIG['CSV_PATH']):\n        train_df, val_df, class_names = load_and_preprocess_df(CONFIG['CSV_PATH'])\n        \n        # 2. Create Datasets using the REAL image directory\n        train_dataset = WhaleDataset(train_df, real_img_dir, transform=get_transforms('train'))\n        val_dataset = WhaleDataset(val_df, real_img_dir, transform=get_transforms('val'))\n        \n        # 3. Create DataLoaders\n        train_loader = DataLoader(\n            train_dataset, \n            batch_size=CONFIG['BATCH_SIZE'], \n            shuffle=True, \n            num_workers=CONFIG['NUM_WORKERS']\n        )\n        \n        val_loader = DataLoader(\n            val_dataset, \n            batch_size=CONFIG['BATCH_SIZE'], \n            shuffle=False, \n            num_workers=CONFIG['NUM_WORKERS']\n        )\n        \n        print(\"\\n🔎 Visualizing a batch of augmented training data...\")\n        visualize_batch(train_loader, class_names)\n        \n        print(\"\\n✅ Ready for model training!\")\n    else:\n        print(f\"⚠️ CSV file not found at {CONFIG['CSV_PATH']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T11:02:24.772869Z","iopub.execute_input":"2025-11-23T11:02:24.773423Z","iopub.status.idle":"2025-11-23T11:04:25.377222Z","shell.execute_reply.started":"2025-11-23T11:02:24.773385Z","shell.execute_reply":"2025-11-23T11:04:25.376244Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport zipfile\nimport glob\n\n# --- CONFIGURATION ---\nCONFIG = {\n    'IMG_DIR': '/kaggle/input/noaa-right-whale-recognition/imgs.zip', \n    'CSV_PATH': '/kaggle/input/noaa-right-whale-recognition/train.csv',\n    \n    # T4 x2 OPTIMIZATIONS\n    'IMG_SIZE': 384,         \n    'BATCH_SIZE': 32,\n    'NUM_WORKERS': 2,\n    'SEED': 42\n}\n\ndef handle_zip_extraction(path_from_config):\n    extract_root = \"./extracted_data\"\n    \n    if path_from_config.endswith('.zip'):\n        if os.path.exists(extract_root) and len(os.listdir(extract_root)) > 0:\n            print(f\"✅ Found existing data in: {extract_root}. Skipping extraction.\")\n        else:\n            print(f\"📦 Extracting {path_from_config} to {extract_root}...\")\n            os.makedirs(extract_root, exist_ok=True)\n            with zipfile.ZipFile(path_from_config, 'r') as zip_ref:\n                zip_ref.extractall(extract_root)\n            print(\"✅ Extraction complete.\")\n    else:\n        extract_root = path_from_config\n\n    potential_subfolder = os.path.join(extract_root, 'imgs')\n    if os.path.exists(potential_subfolder):\n        return potential_subfolder\n    \n    return extract_root\n\ndef load_and_preprocess_df(csv_path):\n    df = pd.read_csv(csv_path)\n    \n    # 1. REMOVE KNOWN MISSING IMAGES\n    # w_7489.jpg is known to be missing in this dataset\n    df = df[df['Image'] != 'w_7489.jpg'].copy()\n    \n    encoder = LabelEncoder()\n    df['label_idx'] = encoder.fit_transform(df['whaleID'])\n    classes = encoder.classes_\n    print(f\"✅ Label Encoding Complete. Found {len(classes)} unique whales.\")\n    \n    counts = df.whaleID.value_counts()\n    single_shot_whales = counts[counts == 1].index\n    \n    df_single = df[df.whaleID.isin(single_shot_whales)]\n    df_multi = df[~df.whaleID.isin(single_shot_whales)]\n    \n    print(f\"📊 Split Stats: {len(df_single)} whales have only 1 image (forced to Train).\")\n    \n    train_multi, val_multi = train_test_split(\n        df_multi, \n        test_size=0.1, \n        random_state=CONFIG['SEED'], \n        stratify=df_multi['whaleID']\n    )\n    \n    df_train = pd.concat([train_multi, df_single]).sample(frac=1).reset_index(drop=True)\n    df_val = val_multi.sample(frac=1).reset_index(drop=True)\n    \n    print(f\"✅ Final Split: Train: {len(df_train)} images, Val: {len(df_val)} images\")\n    \n    return df_train, df_val, classes\n\n# --- AUGMENTATION ---\ndef get_transforms(data='train'):\n    if data == 'train':\n        return A.Compose([\n            # FIX: Ensure image is at least 1200x1200 before cropping\n            # If image is smaller, it pads with zeros (black)\n            A.PadIfNeeded(min_height=1200, min_width=1200, border_mode=cv2.BORDER_CONSTANT, value=0),\n            \n            # Smart \"Zoom\" - Crop the center \n            A.CenterCrop(height=1200, width=1200, p=1.0), \n            \n            A.Resize(CONFIG['IMG_SIZE'], CONFIG['IMG_SIZE']),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Rotate(limit=30, p=0.7),\n            A.RandomBrightnessContrast(p=0.2),\n            A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=20, p=0.2),\n            A.GaussianBlur(p=0.1),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2(),\n        ])\n    elif data == 'val':\n        return A.Compose([\n            A.PadIfNeeded(min_height=1200, min_width=1200, border_mode=cv2.BORDER_CONSTANT, value=0),\n            A.CenterCrop(height=1200, width=1200, p=1.0),\n            A.Resize(CONFIG['IMG_SIZE'], CONFIG['IMG_SIZE']),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2(),\n        ])\n\n# --- DATASET CLASS ---\nclass WhaleDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_name = row['Image']\n        label = row['label_idx']\n        img_path = os.path.join(self.img_dir, img_name)\n        \n        image = cv2.imread(img_path)\n        if image is None:\n            # FIX: Fallback image must be large enough for the crop!\n            # Using 1500x1500x3 ensures CenterCrop(1200, 1200) works fine.\n            image = np.zeros((1500, 1500, 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            \n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n            \n        return image, label\n\n# --- VISUALIZATION HELPER ---\ndef visualize_batch(dataloader, classes):\n    try:\n        images, labels = next(iter(dataloader))\n    except StopIteration:\n        print(\"⚠️ DataLoader is empty. Check your paths.\")\n        return\n\n    plt.figure(figsize=(16, 8))\n    for i in range(min(8, len(images))):\n        ax = plt.subplot(2, 4, i + 1)\n        img = images[i].permute(1, 2, 0).numpy()\n        img = img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406]\n        img = np.clip(img, 0, 1)\n        plt.imshow(img)\n        plt.title(classes[labels[i]])\n        plt.axis(\"off\")\n    plt.tight_layout()\n    plt.show()\n\nif __name__ == \"__main__\":\n    real_img_dir = handle_zip_extraction(CONFIG['IMG_DIR'])\n    print(f\"📂 Images will be loaded from: {real_img_dir}\")\n\n    if os.path.exists(CONFIG['CSV_PATH']):\n        train_df, val_df, class_names = load_and_preprocess_df(CONFIG['CSV_PATH'])\n        \n        train_dataset = WhaleDataset(train_df, real_img_dir, transform=get_transforms('train'))\n        val_dataset = WhaleDataset(val_df, real_img_dir, transform=get_transforms('val'))\n        \n        train_loader = DataLoader(train_dataset, batch_size=CONFIG['BATCH_SIZE'], shuffle=True, num_workers=CONFIG['NUM_WORKERS'])\n        val_loader = DataLoader(val_dataset, batch_size=CONFIG['BATCH_SIZE'], shuffle=False, num_workers=CONFIG['NUM_WORKERS'])\n        \n        print(\"\\n🔎 Visualizing a batch of augmented training data...\")\n        visualize_batch(train_loader, class_names)\n        \n        print(\"\\n✅ Ready for model training!\")\n    else:\n        print(f\"⚠️ CSV file not found at {CONFIG['CSV_PATH']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T11:49:34.715921Z","iopub.execute_input":"2025-11-23T11:49:34.716290Z","iopub.status.idle":"2025-11-23T11:49:40.671696Z","shell.execute_reply.started":"2025-11-23T11:49:34.716242Z","shell.execute_reply":"2025-11-23T11:49:40.670943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nimport timm\nfrom tqdm import tqdm\nimport numpy as np\nimport time\nimport copy\nimport sys\n\ntry:\n    from preprocessing_pipeline import CONFIG, train_loader, val_loader, class_names\nexcept ImportError:\n    if 'train_loader' not in globals():\n        print(\"⚠️ Variables from preprocessing not found. Please run the preprocessing step first.\")\n        sys.exit(1)\n\n# --- CONFIGURATION ---\nTRAIN_CONFIG = {\n    'MODEL_NAME': 'resnet26d', \n    'NUM_EPOCHS': 10,          \n    'LEARNING_RATE': 3e-4,    \n    'WEIGHT_DECAY': 1e-4,\n    'DEVICE': torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\") # Changed to generic 'cuda'\n}\n\n# --- MODEL DEFINITION ---\nclass WhaleClassifier(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=True):\n        super(WhaleClassifier, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, num_classes=0)\n        in_features = self.model.num_features\n        \n        self.head = nn.Sequential(\n            nn.BatchNorm1d(in_features),\n            nn.Dropout(0.3),\n            nn.Linear(in_features, num_classes)\n        )\n        \n    def forward(self, x):\n        features = self.model(x)\n        output = self.head(features)\n        return output\n\n# --- TRAINING HELPER FUNCTIONS ---\ndef train_one_epoch(model, loader, criterion, optimizer, device, epoch):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc=f\"Epoch {epoch+1}/{TRAIN_CONFIG['NUM_EPOCHS']} [Train]\")\n    \n    for images, labels in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * images.size(0)\n        _, predicted = torch.max(outputs, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n        \n        pbar.set_postfix({'loss': loss.item(), 'acc': correct/total})\n        \n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = correct / total\n    return epoch_loss, epoch_acc\n\ndef validate(model, loader, criterion, device, epoch):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc=f\"Epoch {epoch+1}/{TRAIN_CONFIG['NUM_EPOCHS']} [Val]\")\n    \n    with torch.no_grad():\n        for images, labels in pbar:\n            images = images.to(device)\n            labels = labels.to(device)\n            \n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item() * images.size(0)\n            _, predicted = torch.max(outputs, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n            \n            pbar.set_postfix({'loss': loss.item(), 'acc': correct/total})\n            \n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc = correct / total\n    return epoch_loss, epoch_acc\n\n# --- MAIN TRAINING LOOP ---\nif __name__ == \"__main__\":\n    print(f\"🚀 Training on device: {TRAIN_CONFIG['DEVICE']}\")\n    print(f\"🐳 Number of classes: {len(class_names)}\")\n    \n    # 1. Initialize Model\n    model = WhaleClassifier(\n        model_name=TRAIN_CONFIG['MODEL_NAME'],\n        num_classes=len(class_names)\n    )\n    \n    # --- MULTI-GPU LOGIC ---\n    if torch.cuda.device_count() > 1:\n        print(f\"🔥 Found {torch.cuda.device_count()} GPUs! Using DataParallel.\")\n        model = nn.DataParallel(model)\n    \n    model = model.to(TRAIN_CONFIG['DEVICE'])\n    \n    # 2. Loss and Optimizer\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=TRAIN_CONFIG['LEARNING_RATE'], weight_decay=TRAIN_CONFIG['WEIGHT_DECAY'])\n    scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=TRAIN_CONFIG['NUM_EPOCHS'], eta_min=1e-6)\n    \n    # 3. Training Loop\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_loss = float('inf')\n    \n    start_time = time.time()\n    \n    for epoch in range(TRAIN_CONFIG['NUM_EPOCHS']):\n        train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, TRAIN_CONFIG['DEVICE'], epoch)\n        val_loss, val_acc = validate(model, val_loader, criterion, TRAIN_CONFIG['DEVICE'], epoch)\n        scheduler.step()\n        \n        print(f\"Epoch {epoch+1} Summary: Train Loss: {train_loss:.4f} Acc: {train_acc:.4f} | Val Loss: {val_loss:.4f} Acc: {val_acc:.4f}\")\n        \n        if val_loss < best_loss:\n            print(f\"⭐️ Validation Loss Improved ({best_loss:.4f} -> {val_loss:.4f}). Saving model...\")\n            best_loss = val_loss\n            \n            # Handle saving for DataParallel (unwrap 'module.')\n            if isinstance(model, nn.DataParallel):\n                best_model_wts = copy.deepcopy(model.module.state_dict())\n            else:\n                best_model_wts = copy.deepcopy(model.state_dict())\n                \n            torch.save(best_model_wts, \"best_whale_model.pth\")\n            \n    time_elapsed = time.time() - start_time\n    print(f\"\\n✅ Training complete in {time_elapsed // 60:.0f}m {time_elapsed % 60:.0f}s\")\n    print(f\"🏆 Best Validation Loss: {best_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T11:50:04.022603Z","iopub.execute_input":"2025-11-23T11:50:04.022988Z","iopub.status.idle":"2025-11-23T12:17:53.386533Z","shell.execute_reply.started":"2025-11-23T11:50:04.022955Z","shell.execute_reply":"2025-11-23T12:17:53.385674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nimport timm\nfrom tqdm import tqdm\nimport numpy as np\nimport time\nimport copy\nimport sys\n\ntry:\n    from preprocessing_pipeline import CONFIG, train_loader, val_loader, class_names\nexcept ImportError:\n    if 'train_loader' not in globals():\n        print(\"⚠️ Variables from preprocessing not found. Please run the preprocessing step first.\")\n        sys.exit(1)\n\n# --- CONFIGURATION ---\nTRAIN_CONFIG = {\n    'MODEL_NAME': 'resnet26d', \n    'NUM_EPOCHS': 31,          # INCREASED to 30 to allow convergence\n    'LEARNING_RATE': 3e-4,    \n    'WEIGHT_DECAY': 1e-4,\n    'DEVICE': torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n}\n\n# --- METRIC HELPER ---\ndef calculate_accuracy(output, target, topk=(1, 5)):\n    \"\"\"Computes the accuracy over the k top predictions for the specified values of k\"\"\"\n    with torch.no_grad():\n        maxk = max(topk)\n        batch_size = target.size(0)\n\n        _, pred = output.topk(maxk, 1, True, True)\n        pred = pred.t()\n        correct = pred.eq(target.view(1, -1).expand_as(pred))\n\n        res = []\n        for k in topk:\n            correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)\n            res.append(correct_k.mul_(1.0 / batch_size))\n        return res\n\n# --- MODEL DEFINITION ---\nclass WhaleClassifier(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=True):\n        super(WhaleClassifier, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, num_classes=0)\n        in_features = self.model.num_features\n        \n        self.head = nn.Sequential(\n            nn.BatchNorm1d(in_features),\n            nn.Dropout(0.3),\n            nn.Linear(in_features, num_classes)\n        )\n        \n    def forward(self, x):\n        features = self.model(x)\n        output = self.head(features)\n        return output\n\n# --- TRAINING HELPER FUNCTIONS ---\ndef train_one_epoch(model, loader, criterion, optimizer, device, epoch):\n    model.train()\n    running_loss = 0.0\n    correct_1 = 0\n    correct_5 = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc=f\"Epoch {epoch+1}/{TRAIN_CONFIG['NUM_EPOCHS']} [Train]\")\n    \n    for images, labels in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        loss.backward()\n        optimizer.step()\n        \n        # Metrics\n        running_loss += loss.item() * images.size(0)\n        total += labels.size(0)\n        \n        # Calculate Top-1 and Top-5\n        acc1, acc5 = calculate_accuracy(outputs, labels, topk=(1, 5))\n        correct_1 += acc1.item() * labels.size(0)\n        correct_5 += acc5.item() * labels.size(0)\n        \n        pbar.set_postfix({'loss': loss.item(), 'acc1': correct_1/total, 'acc5': correct_5/total})\n        \n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc1 = correct_1 / total\n    epoch_acc5 = correct_5 / total\n    return epoch_loss, epoch_acc1, epoch_acc5\n\ndef validate(model, loader, criterion, device, epoch):\n    model.eval()\n    running_loss = 0.0\n    correct_1 = 0\n    correct_5 = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc=f\"Epoch {epoch+1}/{TRAIN_CONFIG['NUM_EPOCHS']} [Val]\")\n    \n    with torch.no_grad():\n        for images, labels in pbar:\n            images = images.to(device)\n            labels = labels.to(device)\n            \n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item() * images.size(0)\n            total += labels.size(0)\n            \n            acc1, acc5 = calculate_accuracy(outputs, labels, topk=(1, 5))\n            correct_1 += acc1.item() * labels.size(0)\n            correct_5 += acc5.item() * labels.size(0)\n            \n            pbar.set_postfix({'loss': loss.item(), 'acc1': correct_1/total, 'acc5': correct_5/total})\n            \n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_acc1 = correct_1 / total\n    epoch_acc5 = correct_5 / total\n    return epoch_loss, epoch_acc1, epoch_acc5\n\n# --- MAIN TRAINING LOOP ---\nif __name__ == \"__main__\":\n    print(f\"🚀 Training on device: {TRAIN_CONFIG['DEVICE']}\")\n    print(f\"🐳 Number of classes: {len(class_names)}\")\n    \n    # 1. Initialize Model\n    model = WhaleClassifier(\n        model_name=TRAIN_CONFIG['MODEL_NAME'],\n        num_classes=len(class_names)\n    )\n    \n    # --- MULTI-GPU LOGIC ---\n    if torch.cuda.device_count() > 1:\n        print(f\"🔥 Found {torch.cuda.device_count()} GPUs! Using DataParallel.\")\n        model = nn.DataParallel(model)\n    \n    model = model.to(TRAIN_CONFIG['DEVICE'])\n    \n    # 2. Loss and Optimizer\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.AdamW(model.parameters(), lr=TRAIN_CONFIG['LEARNING_RATE'], weight_decay=TRAIN_CONFIG['WEIGHT_DECAY'])\n    scheduler = lr_scheduler.CosineAnnealingLR(optimizer, T_max=TRAIN_CONFIG['NUM_EPOCHS'], eta_min=1e-6)\n    \n    # 3. Training Loop\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_loss = float('inf')\n    \n    start_time = time.time()\n    \n    for epoch in range(TRAIN_CONFIG['NUM_EPOCHS']):\n        # Train\n        train_loss, train_acc1, train_acc5 = train_one_epoch(model, train_loader, criterion, optimizer, TRAIN_CONFIG['DEVICE'], epoch)\n        \n        # Validate\n        val_loss, val_acc1, val_acc5 = validate(model, val_loader, criterion, TRAIN_CONFIG['DEVICE'], epoch)\n        \n        # Scheduler Step\n        scheduler.step()\n        \n        print(f\"Epoch {epoch+1}: Train Loss: {train_loss:.4f} Acc1: {train_acc1:.4f} Acc5: {train_acc5:.4f}\")\n        print(f\"          Val Loss:   {val_loss:.4f} Acc1: {val_acc1:.4f} Acc5: {val_acc5:.4f}\")\n        \n        # Save Best Model\n        if val_loss < best_loss:\n            print(f\"⭐️ Val Loss Improved ({best_loss:.4f} -> {val_loss:.4f}). Saving...\")\n            best_loss = val_loss\n            \n            if isinstance(model, nn.DataParallel):\n                best_model_wts = copy.deepcopy(model.module.state_dict())\n            else:\n                best_model_wts = copy.deepcopy(model.state_dict())\n                \n            torch.save(best_model_wts, \"best_whale_model.pth\")\n            \n    time_elapsed = time.time() - start_time\n    print(f\"\\n✅ Training complete in {time_elapsed // 60:.0f}m {time_elapsed % 60:.0f}s\")\n    print(f\"🏆 Best Validation Loss: {best_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T13:00:13.946939Z","iopub.execute_input":"2025-11-23T13:00:13.947664Z","iopub.status.idle":"2025-11-23T14:30:18.949151Z","shell.execute_reply.started":"2025-11-23T13:00:13.947629Z","shell.execute_reply":"2025-11-23T14:30:18.948291Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport timm\nimport math\n\n# Import necessary variables from previous steps\n# If running in a notebook, these should already be in memory.\ntry:\n    from preprocessing_pipeline import val_loader, class_names, CONFIG\nexcept ImportError:\n    print(\"⚠️ properties not found. Make sure you ran the preprocessing step!\")\n\n# --- REDEFINE MODEL (Must match training exactly) ---\nclass WhaleClassifier(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=False):\n        super(WhaleClassifier, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, num_classes=0)\n        in_features = self.model.num_features\n        \n        self.head = nn.Sequential(\n            nn.BatchNorm1d(in_features),\n            nn.Dropout(0.3),\n            nn.Linear(in_features, num_classes)\n        )\n        \n    def forward(self, x):\n        features = self.model(x)\n        output = self.head(features)\n        return output\n\ndef visualize_model_predictions(model_path, num_images=12):\n    \"\"\"\n    Loads the best model and visualizes predictions on the validation set.\n    \"\"\"\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"🔍 Running inference on: {device}\")\n    \n    # 1. Re-initialize the model structure\n    model = WhaleClassifier(model_name='resnet26d', num_classes=len(class_names), pretrained=False)\n    \n    # 2. Load the trained weights\n    try:\n        state_dict = torch.load(model_path, map_location=device)\n        model.load_state_dict(state_dict)\n        print(\"✅ Model weights loaded successfully.\")\n    except Exception as e:\n        print(f\"❌ Failed to load model weights: {e}\")\n        return\n\n    model = model.to(device)\n    model.eval() # Set to evaluation mode\n    \n    # 3. Get a batch of validation data\n    images, labels = next(iter(val_loader))\n    images = images.to(device)\n    labels = labels.to(device)\n    \n    # 4. Predict\n    with torch.no_grad():\n        outputs = model(images)\n        # Get probabilities using Softmax\n        probs = torch.nn.functional.softmax(outputs, dim=1)\n        # Get top 1 prediction\n        confidences, preds = torch.max(probs, 1)\n        \n    # 5. Plotting\n    # Calculate grid size (e.g., 3x4 for 12 images)\n    cols = 4\n    rows = math.ceil(min(num_images, len(images)) / cols)\n    \n    plt.figure(figsize=(20, 5 * rows))\n    \n    for i in range(min(num_images, len(images))):\n        ax = plt.subplot(rows, cols, i + 1)\n        \n        # Un-normalize image for display\n        img = images[i].cpu().permute(1, 2, 0).numpy()\n        img = img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406] # ImageNet stats\n        img = np.clip(img, 0, 1)\n        \n        true_label = class_names[labels[i].cpu().item()]\n        pred_label = class_names[preds[i].cpu().item()]\n        confidence = confidences[i].cpu().item() * 100\n        \n        plt.imshow(img)\n        plt.axis(\"off\")\n        \n        # Color code the title\n        if true_label == pred_label:\n            color = 'green'\n            title_text = f\"✅ {true_label}\\nConf: {confidence:.1f}%\"\n        else:\n            color = 'red'\n            title_text = f\"❌ True: {true_label}\\nPred: {pred_label} ({confidence:.1f}%)\"\n            \n        plt.title(title_text, color=color, fontsize=12, fontweight='bold')\n        \n    plt.tight_layout()\n    plt.show()\n\nif __name__ == \"__main__\":\n    # Check if the weight file exists\n    import os\n    if os.path.exists(\"best_whale_model.pth\"):\n        visualize_model_predictions(\"best_whale_model.pth\", num_images=16)\n    else:\n        print(\"⚠️ 'best_whale_model.pth' not found. Did you finish training?\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T12:20:32.021588Z","iopub.execute_input":"2025-11-23T12:20:32.022025Z","iopub.status.idle":"2025-11-23T12:20:41.530438Z","shell.execute_reply.started":"2025-11-23T12:20:32.021993Z","shell.execute_reply":"2025-11-23T12:20:41.529537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport timm\nimport math\nfrom tqdm import tqdm\n\n# Import necessary variables from previous steps\ntry:\n    from preprocessing_pipeline import val_loader, class_names, CONFIG\nexcept ImportError:\n    print(\"⚠️ properties not found. Make sure you ran the preprocessing step!\")\n\n# --- CONFIGURATION ---\n# Ensure this matches what you used in training!\nMODEL_NAME = 'resnet26d' \n\n# --- REDEFINE MODEL (Must match training exactly) ---\nclass WhaleClassifier(nn.Module):\n    def __init__(self, model_name, num_classes, pretrained=False):\n        super(WhaleClassifier, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained, num_classes=0)\n        in_features = self.model.num_features\n        \n        self.head = nn.Sequential(\n            nn.BatchNorm1d(in_features),\n            nn.Dropout(0.3),\n            nn.Linear(in_features, num_classes)\n        )\n        \n    def forward(self, x):\n        features = self.model(x)\n        output = self.head(features)\n        return output\n\ndef calculate_accuracy(output, target, topk=(1, 5)):\n    \"\"\"Computes the accuracy over the k top predictions\"\"\"\n    with torch.no_grad():\n        maxk = max(topk)\n        batch_size = target.size(0)\n\n        _, pred = output.topk(maxk, 1, True, True)\n        pred = pred.t()\n        correct = pred.eq(target.view(1, -1).expand_as(pred))\n\n        res = []\n        for k in topk:\n            correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)\n            res.append(correct_k.item())\n        return res\n\ndef evaluate_and_visualize(model_path):\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"🔍 Loading model from: {model_path}\")\n    print(f\"⚙️  Using Device: {device}\")\n    \n    # 1. Initialize Model\n    model = WhaleClassifier(model_name=MODEL_NAME, num_classes=len(class_names), pretrained=False)\n    \n    # 2. Load Weights\n    try:\n        if torch.cuda.device_count() > 1:\n            print(\"Note: Training used multiple GPUs. Adjusting keys...\")\n        \n        state_dict = torch.load(model_path, map_location=device)\n        model.load_state_dict(state_dict)\n        print(\"✅ Weights loaded successfully.\")\n    except Exception as e:\n        print(f\"❌ Error loading weights: {e}\")\n        return\n\n    model = model.to(device)\n    model.eval()\n    \n    # --- PART 1: CALCULATE OVERALL ACCURACY ---\n    print(\"\\n📊 Calculating Accuracy on ENTIRE Validation Set...\")\n    total_correct_1 = 0\n    total_correct_5 = 0\n    total_samples = 0\n    \n    # We iterate through the whole validation loader to get the real score\n    with torch.no_grad():\n        for images, labels in tqdm(val_loader, desc=\"Evaluating\"):\n            images = images.to(device)\n            labels = labels.to(device)\n            \n            outputs = model(images)\n            \n            # Calculate batch accuracy\n            acc1, acc5 = calculate_accuracy(outputs, labels, topk=(1, 5))\n            \n            total_correct_1 += acc1\n            total_correct_5 += acc5\n            total_samples += labels.size(0)\n            \n    avg_acc1 = (total_correct_1 / total_samples) * 100\n    avg_acc5 = (total_correct_5 / total_samples) * 100\n    \n    print(\"-\" * 40)\n    print(f\"🏆 Final Results for {MODEL_NAME}:\")\n    print(f\"   Top-1 Accuracy: {avg_acc1:.2f}% (Exact Match)\")\n    print(f\"   Top-5 Accuracy: {avg_acc5:.2f}% (Correct whale is in top 5 guesses)\")\n    print(\"-\" * 40)\n\n    # --- PART 2: VISUALIZE A BATCH ---\n    print(\"\\n🖼️  Visualizing Random Batch...\")\n    \n    # Get a batch\n    images, labels = next(iter(val_loader))\n    images = images.to(device)\n    labels = labels.to(device)\n    \n    with torch.no_grad():\n        outputs = model(images)\n        probs = torch.nn.functional.softmax(outputs, dim=1)\n        confidences, preds = torch.max(probs, 1)\n        \n    # Plot 16 images\n    num_images = min(16, len(images))\n    cols = 4\n    rows = math.ceil(num_images / cols)\n    \n    plt.figure(figsize=(20, 5 * rows))\n    \n    for i in range(num_images):\n        ax = plt.subplot(rows, cols, i + 1)\n        \n        # Un-normalize\n        img = images[i].cpu().permute(1, 2, 0).numpy()\n        img = img * [0.229, 0.224, 0.225] + [0.485, 0.456, 0.406]\n        img = np.clip(img, 0, 1)\n        \n        true_label = class_names[labels[i].cpu().item()]\n        pred_label = class_names[preds[i].cpu().item()]\n        confidence = confidences[i].cpu().item() * 100\n        \n        plt.imshow(img)\n        plt.axis(\"off\")\n        \n        if true_label == pred_label:\n            color = 'green'\n            title_text = f\"✅ {true_label}\\nConf: {confidence:.1f}%\"\n        else:\n            color = 'red'\n            title_text = f\"❌ True: {true_label}\\nPred: {pred_label} ({confidence:.1f}%)\"\n            \n        plt.title(title_text, color=color, fontsize=12, fontweight='bold')\n        \n    plt.tight_layout()\n    plt.show()\n\nif __name__ == \"__main__\":\n    import os\n    if os.path.exists(\"best_whale_model.pth\"):\n        evaluate_and_visualize(\"best_whale_model.pth\")\n    else:\n        print(\"⚠️ 'best_whale_model.pth' not found. Did you finish training?\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T14:34:43.053378Z","iopub.execute_input":"2025-11-23T14:34:43.053806Z","iopub.status.idle":"2025-11-23T14:35:11.792893Z","shell.execute_reply.started":"2025-11-23T14:34:43.053777Z","shell.execute_reply":"2025-11-23T14:35:11.791569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport os\nimport timm\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import log_loss, precision_recall_fscore_support, confusion_matrix\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# --- COMPACT CONFIG & CLASSES ---\nCONFIG = {'CSV': '/kaggle/input/noaa-right-whale-recognition/train.csv', 'IMG': './extracted_data/imgs', 'SIZE': 384, 'BATCH': 32}\n\nclass WhaleDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df, self.img_dir, self.transform = df, img_dir, transform\n    def __len__(self): return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = cv2.imread(os.path.join(self.img_dir, row['Image']))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if img is not None else np.zeros((1500, 1500, 3), np.uint8)\n        return self.transform(image=img)['image'] if self.transform else img, row['label_idx']\n\nclass WhaleClassifier(nn.Module):\n    def __init__(self, model_name, num_classes):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=False, num_classes=0)\n        self.head = nn.Sequential(nn.BatchNorm1d(self.model.num_features), nn.Dropout(0.3), nn.Linear(self.model.num_features, num_classes))\n    def forward(self, x): return self.head(self.model(x))\n\n# --- MAIN EVALUATION FUNCTION ---\ndef run_evaluation():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"⚙️  Running Fast Evaluation on {device}...\")\n\n    # 1. Quick Data Setup\n    df = pd.read_csv(CONFIG['CSV'])\n    df = df[df['Image'] != 'w_7489.jpg'].copy() # Remove bad image\n    \n    encoder = LabelEncoder()\n    df['label_idx'] = encoder.fit_transform(df['whaleID'])\n    classes = encoder.classes_\n    \n    # Handle Single Image Whales\n    counts = df.whaleID.value_counts()\n    single_shot_whales = counts[counts == 1].index\n    df_multi = df[~df.whaleID.isin(single_shot_whales)]\n    \n    # Stratified Split\n    _, val_df = train_test_split(df_multi, test_size=0.1, random_state=42, stratify=df_multi['whaleID'])\n    \n    transforms = A.Compose([\n        A.PadIfNeeded(min_height=1200, min_width=1200, border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.CenterCrop(1200, 1200), A.Resize(CONFIG['SIZE'], CONFIG['SIZE']),\n        A.Normalize(), ToTensorV2()\n    ])\n    loader = DataLoader(WhaleDataset(val_df, CONFIG['IMG'], transforms), batch_size=CONFIG['BATCH'], shuffle=False, num_workers=2)\n\n    # 2. Load Model\n    model = WhaleClassifier('resnet26d', len(classes))\n    state_dict = torch.load(\"best_whale_model.pth\", map_location=device)\n    model.load_state_dict(state_dict)\n    model.to(device).eval()\n\n    # 3. Get Predictions\n    all_preds, all_targets, all_probs = [], [], []\n    with torch.no_grad():\n        for img, label in tqdm(loader, desc=\"Predicting\"):\n            out = model(img.to(device))\n            prob = torch.softmax(out, dim=1)\n            all_preds.extend(prob.argmax(1).cpu().numpy())\n            all_targets.extend(label.numpy())\n            all_probs.extend(prob.cpu().numpy())\n\n    # 4. Calculate Metrics\n    all_targets, all_probs = np.array(all_targets), np.array(all_probs)\n    all_preds = np.array(all_preds)\n    \n    acc1 = np.mean(all_preds == all_targets)\n    top5 = sum([all_targets[i] in np.argsort(all_probs[i])[::-1][:5] for i in range(len(all_targets))]) / len(all_targets)\n    p, r, f1, _ = precision_recall_fscore_support(all_targets, all_preds, average='weighted', zero_division=0)\n    try: loss = log_loss(all_targets, all_probs, labels=list(range(len(classes))))\n    except: loss = 99.9\n\n    # 5. Print Report\n    print(f\"\\n{'='*30}\\n🏆 FINAL EVALUATION RESULTS\\n{'='*30}\")\n    print(f\"✅ Top-1 Accuracy:  {acc1:.2%}\")\n    print(f\"✅ Top-5 Accuracy:  {top5:.2%}\")\n    print(f\"📉 Log Loss:        {loss:.4f}\")\n    print(f\"{'-'*30}\")\n    print(f\"🎯 Precision:       {p:.4f}\")\n    print(f\"📡 Recall:          {r:.4f}\")\n    print(f\"⚖️  F1-Score:        {f1:.4f}\")\n    print(f\"{'='*30}\")\n\n    # --- 6. CONFUSION MATRIX (TOP 20) ---\n    print(\"\\n🎨 Generating Confusion Matrix for Top 20 Whales...\")\n    from collections import Counter\n    \n    # Get the 20 most frequent whales in the VALIDATION set\n    counts = Counter(all_targets)\n    top_20_indices = [k for k, v in counts.most_common(20)]\n    top_20_names = [classes[i] for i in top_20_indices]\n    \n    # Filter targets and preds to only include these 20 whales\n    mask = [i for i, t in enumerate(all_targets) if t in top_20_indices]\n    filtered_targets = all_targets[mask]\n    filtered_preds = all_preds[mask]\n    \n    # Generate Matrix\n    cm = confusion_matrix(filtered_targets, filtered_preds, labels=top_20_indices)\n    \n    # Plot\n    plt.figure(figsize=(14, 12))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                xticklabels=top_20_names, yticklabels=top_20_names)\n    plt.xlabel('Predicted Label')\n    plt.ylabel('True Label')\n    plt.title('Confusion Matrix (Top 20 Most Frequent Whales)')\n    plt.xticks(rotation=45, ha='right')\n    plt.tight_layout()\n    \n    # --- SAVE THE IMAGE ---\n    plt.savefig('confusion_matrix.png', bbox_inches='tight', dpi=300)\n    print(\"✅ Saved plot to 'confusion_matrix.png'\")\n    \n    plt.show()\n\nif __name__ == \"__main__\":\n    if os.path.exists(\"best_whale_model.pth\"): run_evaluation()\n    else: print(\"❌ Model not found.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-23T15:11:35.047280Z","iopub.execute_input":"2025-11-23T15:11:35.048042Z","iopub.status.idle":"2025-11-23T15:11:55.383875Z","shell.execute_reply.started":"2025-11-23T15:11:35.048008Z","shell.execute_reply":"2025-11-23T15:11:55.383031Z"}},"outputs":[],"execution_count":null}]}