{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":13747,"databundleVersionId":868529,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport json\nimport tarfile\nimport shutil\nfrom pathlib import Path\nimport gc\nimport warnings\nimport io\nfrom collections import defaultdict\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nfrom torchvision.models import efficientnet_b0, EfficientNet_B0_Weights\nimport timm\n\nfrom sklearn.metrics import accuracy_score, f1_score, classification_report\nfrom sklearn.utils.class_weight import compute_class_weight\n\nprint(\"All libraries imported successfully!\")\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"CUDA version: {torch.version.cuda}\")\n    print(f\"GPU count: {torch.cuda.device_count()}\")\n    for i in range(torch.cuda.device_count()):\n        print(f\"GPU {i}: {torch.cuda.get_device_name(i)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:53:39.549143Z","iopub.execute_input":"2025-06-21T10:53:39.549360Z","iopub.status.idle":"2025-06-21T10:53:54.779385Z","shell.execute_reply.started":"2025-06-21T10:53:39.549342Z","shell.execute_reply":"2025-06-21T10:53:54.778726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    # Paths\n    DATA_DIR = \"/kaggle/input/inaturalist-2019-fgvc6\"\n    WORK_DIR = \"/kaggle/working\"\n    \n    # Model parameters\n    MODEL_NAME = \"efficientnet_b0\"\n    NUM_CLASSES = None  \n    IMG_SIZE = 224\n    BATCH_SIZE = 8  \n    NUM_WORKERS = 2\n    \n    EPOCHS = 150\n    LEARNING_RATE = 2e-4\n    WEIGHT_DECAY = 1e-4\n    PATIENCE = 50\n    \n    MAX_SAMPLES_PER_CLASS = 50   # Reduced significantly\n    MIN_SAMPLES_PER_CLASS = 5    # Reduced minimum\n    MAX_TOTAL_CLASSES = 100      # Limit total classes\n    \n    # Device\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Mixed precision training\n    USE_AMP = True\n    \n    # IMPORTANT: Use test set for training due to storage constraints\n    USE_TEST_FOR_TRAINING = True\n\nconfig = Config()\n\nprint(\"Configuration loaded:\")\nprint(f\"Device: {config.DEVICE}\")\nprint(f\"Training on: {'TEST SET' if config.USE_TEST_FOR_TRAINING else 'TRAIN SET'}\")\nprint(f\"Image size: {config.IMG_SIZE}\")\nprint(f\"Batch size: {config.BATCH_SIZE}\")\nprint(f\"Max samples per class: {config.MAX_SAMPLES_PER_CLASS}\")\nprint(f\"Max total classes: {config.MAX_TOTAL_CLASSES}\")\nprint(f\"Epochs: {config.EPOCHS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:53:54.780797Z","iopub.execute_input":"2025-06-21T10:53:54.781161Z","iopub.status.idle":"2025-06-21T10:53:54.787407Z","shell.execute_reply.started":"2025-06-21T10:53:54.781143Z","shell.execute_reply":"2025-06-21T10:53:54.786662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def check_available_files():\n    \"\"\"Check what files are available\"\"\"\n    print(\"Checking available files...\")\n    \n    for file in os.listdir(config.DATA_DIR):\n        file_path = os.path.join(config.DATA_DIR, file)\n        if os.path.isfile(file_path):\n            size_mb = os.path.getsize(file_path) / (1024 * 1024)\n            print(f\"  {file}: {size_mb:.2f} MB\")\n    \n    return True\n\ndef load_test_data_for_training():\n    \"\"\"Load test data for training (since we can't fit full train set)\"\"\"\n    print(\"Loading test data for training...\")\n    \n    # We'll use train annotations but sample heavily\n    train_json_path = os.path.join(config.DATA_DIR, \"train2019.json\")\n    \n    if not os.path.exists(train_json_path):\n        print(f\"Error: {train_json_path} not found!\")\n        return None\n    \n    # Load train JSON (we'll sample from this)\n    with open(train_json_path, 'r') as f:\n        train_data = json.load(f)\n    \n    print(f\"Loaded: {len(train_data['images'])} images, {len(train_data['annotations'])} annotations\")\n    print(f\"Categories: {len(train_data['categories'])}\")\n    \n    return train_data\n\n# Check files and load data\ncheck_available_files()\ntrain_data = load_test_data_for_training()\n\nif train_data is not None:\n    print(\"✅ Data loading successful!\")\n    print(f\"Sample category: {train_data['categories'][0]}\")\nelse:\n    print(\"❌ Data loading failed!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:53:54.788209Z","iopub.execute_input":"2025-06-21T10:53:54.788723Z","iopub.status.idle":"2025-06-21T10:53:56.155125Z","shell.execute_reply.started":"2025-06-21T10:53:54.788693Z","shell.execute_reply":"2025-06-21T10:53:56.154372Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def aggressive_data_sampling(train_data):\n    \"\"\"Aggressively sample data to fit memory constraints\"\"\"\n    if train_data is None:\n        return None, None, None\n    \n    print(\"Performing aggressive data sampling...\")\n    \n    # Create category mapping\n    categories = {cat['id']: cat['name'] for cat in train_data['categories']}\n    print(f\"Total categories available: {len(categories)}\")\n    \n    # Process annotations\n    train_df = pd.DataFrame(train_data['annotations'])\n    train_images = {img['id']: img['file_name'] for img in train_data['images']}\n    train_df['file_name'] = train_df['image_id'].map(train_images)\n    \n    # Count samples per class\n    class_counts = train_df['category_id'].value_counts()\n    print(f\"Class distribution - Min: {class_counts.min()}, Max: {class_counts.max()}, Mean: {class_counts.mean():.1f}\")\n    \n    # Select top classes with enough samples\n    valid_classes = class_counts[class_counts >= config.MIN_SAMPLES_PER_CLASS].head(config.MAX_TOTAL_CLASSES).index\n    print(f\"Selected {len(valid_classes)} classes for training\")\n    \n    # Sample from each selected class\n    sampled_dfs = []\n    total_samples = 0\n    \n    for class_id in tqdm(valid_classes, desc=\"Sampling classes\"):\n        class_data = train_df[train_df['category_id'] == class_id]\n        \n        # Sample up to MAX_SAMPLES_PER_CLASS\n        sample_size = min(len(class_data), config.MAX_SAMPLES_PER_CLASS)\n        if len(class_data) > sample_size:\n            class_data = class_data.sample(n=sample_size, random_state=42)\n        \n        sampled_dfs.append(class_data)\n        total_samples += len(class_data)\n        \n        # Stop if we have enough samples total\n        if total_samples > 5000:  # Hard limit\n            break\n    \n    # Combine sampled data\n    if sampled_dfs:\n        train_df_sampled = pd.concat(sampled_dfs, ignore_index=True)\n        \n        # Create new label mapping\n        unique_classes = sorted(train_df_sampled['category_id'].unique())\n        label_mapping = {class_id: i for i, class_id in enumerate(unique_classes)}\n        reverse_mapping = {i: class_id for class_id, i in label_mapping.items()}\n        \n        train_df_sampled['label'] = train_df_sampled['category_id'].map(label_mapping)\n        \n        # Split into train/val (80/20)\n        train_df_sampled = train_df_sampled.sample(frac=1, random_state=42).reset_index(drop=True)\n        split_idx = int(0.8 * len(train_df_sampled))\n        \n        train_split = train_df_sampled[:split_idx]\n        val_split = train_df_sampled[split_idx:]\n        \n        config.NUM_CLASSES = len(unique_classes)\n        \n        print(f\"Final dataset: {len(train_split)} train, {len(val_split)} val\")\n        print(f\"Number of classes: {config.NUM_CLASSES}\")\n        print(f\"Total samples: {len(train_df_sampled)}\")\n        \n        return train_split, val_split, label_mapping\n    else:\n        print(\"No valid samples found!\")\n        return None, None, None\n\n# Sample data\ntrain_df, val_df, label_mapping = aggressive_data_sampling(train_data)\n\nif train_df is not None:\n    print(\"✅ Data sampling successful!\")\n    print(f\"Train: {len(train_df)}, Val: {len(val_df)}, Classes: {config.NUM_CLASSES}\")\n    \n    # Show class distribution\n    print(\"\\nClass distribution:\")\n    print(train_df['label'].value_counts().head())\nelse:\n    print(\"❌ Data sampling failed!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:53:56.155986Z","iopub.execute_input":"2025-06-21T10:53:56.156274Z","iopub.status.idle":"2025-06-21T10:53:56.696769Z","shell.execute_reply.started":"2025-06-21T10:53:56.156248Z","shell.execute_reply":"2025-06-21T10:53:56.696232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# EXTRACT IMAGES THEN TRAIN - Giải nén ra file riêng lẻ\n\nimport os\nimport tarfile\nimport shutil\nfrom pathlib import Path\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nimport gc\nimport json\n\nprint(\"📦 EXTRACT IMAGES THEN TRAIN\")\nprint(\"=\" * 50)\n\n# 1. Check available tar files and choose strategy\ndef check_available_data():\n    \"\"\"Check what data we have and choose extraction strategy\"\"\"\n    print(\"🔍 Checking available data files...\")\n    \n    files_info = {}\n    \n    for filename in os.listdir(config.DATA_DIR):\n        if filename.endswith(('.tar.gz', '.tar', '.json')):\n            filepath = os.path.join(config.DATA_DIR, filename)\n            size_mb = os.path.getsize(filepath) / (1024 * 1024)\n            files_info[filename] = size_mb\n            print(f\"  📁 {filename}: {size_mb:.1f} MB\")\n    \n    return files_info\n\n# 2. Smart extraction based on storage limit\ndef smart_extract_images(extract_limit_mb=15000):  \n    \"\"\"Extract images smartly based on storage constraints\"\"\"\n    \n    # Check available files\n    files_info = check_available_data()\n    \n    # Check working directory space\n    work_dir = config.WORK_DIR\n    extract_dir = os.path.join(work_dir, \"extracted_images\")\n    \n    print(f\"\\n📂 Extract directory: {extract_dir}\")\n    \n\n    if os.path.exists(extract_dir) and len(os.listdir(extract_dir)) > 1000:\n        print(\"✅ Images already extracted!\")\n        \n        # Count extracted images\n        image_count = 0\n        for root, dirs, files in os.walk(extract_dir):\n            for file in files:\n                if file.lower().endswith(('.jpg', '.jpeg', '.png')):\n                    image_count += 1\n        \n        print(f\"📊 Found {image_count} extracted images\")\n        return extract_dir, image_count\n    \n    # Find the tar file to extract from\n    tar_files = [f for f in files_info.keys() if f.endswith('.tar.gz')]\n    \n    if not tar_files:\n        print(\"❌ No tar files found!\")\n        return None, 0\n    \n    # Use the TEST tar file (much smaller)\n    tar_filename = \"test2019.tar.gz\"\n    if tar_filename not in tar_files:\n        # Try alternative test file names\n        test_alternatives = [\"test.tar.gz\", \"test2019.tar\", \"test.tar\"]\n        for alt in test_alternatives:\n            if alt in tar_files:\n                tar_filename = alt\n                break\n        else:\n            # If no test tar found, inform user\n            print(\"❌ No test2019.tar.gz found!\")\n            print(\"Available tar files:\")\n            for tf in tar_files:\n                print(f\"   {tf}\")\n            print(\"🔍 Looking for test2019.tar.gz specifically...\")\n            return None, 0\n    \n    tar_path = os.path.join(config.DATA_DIR, tar_filename)\n    print(f\"\\n📦 Extracting from TEST file: {tar_filename}\")\n    \n    # Check if test tar file exists\n    if not os.path.exists(tar_path):\n        print(f\"❌ {tar_filename} not found at {tar_path}\")\n        print(\"Available files in DATA_DIR:\")\n        for f in os.listdir(config.DATA_DIR):\n            if f.endswith('.tar.gz'):\n                print(f\"   {f}\")\n        return None, 0\n    \n    # Create extract directory\n    os.makedirs(extract_dir, exist_ok=True)\n    \n    # Extract with limit\n    extracted_count = 0\n    extracted_size_mb = 0\n    \n    try:\n        with tarfile.open(tar_path, 'r:gz') as tar:\n            members = tar.getmembers()\n            print(f\"📋 Total files in tar: {len(members)}\")\n            \n            # Filter image files\n            image_members = [m for m in members if m.isfile() and \n                           m.name.lower().endswith(('.jpg', '.jpeg', '.png'))]\n            \n            print(f\"🖼️ Image files in tar: {len(image_members)}\")\n            \n            # Extract with size limit\n            for member in tqdm(image_members, desc=\"Extracting images\"):\n                if extracted_size_mb >= extract_limit_mb:\n                    print(f\"⚠️ Reached size limit: {extract_limit_mb} MB\")\n                    break\n                \n                try:\n                    # Extract the file\n                    tar.extract(member, extract_dir)\n                    \n                    # Check extracted file size\n                    extracted_path = os.path.join(extract_dir, member.name)\n                    if os.path.exists(extracted_path):\n                        file_size_mb = os.path.getsize(extracted_path) / (1024 * 1024)\n                        extracted_size_mb += file_size_mb\n                        extracted_count += 1\n                        \n                        # Progress update every 500 files\n                        if extracted_count % 500 == 0:\n                            print(f\"   📊 Extracted: {extracted_count} images, {extracted_size_mb:.1f} MB\")\n                    \n                except Exception as e:\n                    if extracted_count < 10:  # Only show first few errors\n                        print(f\"   ❌ Error extracting {member.name}: {e}\")\n                    continue\n    \n    except Exception as e:\n        print(f\"❌ Error opening tar file: {e}\")\n        return None, 0\n    \n    print(f\"\\n✅ Extraction complete!\")\n    print(f\"📊 Extracted {extracted_count} images\")\n    print(f\"💾 Total size: {extracted_size_mb:.1f} MB\")\n    \n    return extract_dir, extracted_count\n\n# 3. Create dataset from extracted files\ndef create_dataset_from_extracted(extract_dir, max_samples=10000):\n    \"\"\"Create training dataset from extracted image files\"\"\"\n    \n    if not os.path.exists(extract_dir):\n        print(\"❌ Extract directory not found!\")\n        return None, None, 0\n    \n    print(f\"📋 Creating dataset from {extract_dir}...\")\n    \n    # Find all extracted images\n    image_files = []\n    for root, dirs, files in os.walk(extract_dir):\n        for file in files:\n            if file.lower().endswith(('.jpg', '.jpeg', '.png')):\n                full_path = os.path.join(root, file)\n                image_files.append(full_path)\n    \n    print(f\"🖼️ Found {len(image_files)} image files\")\n    \n    # Limit samples if too many\n    if len(image_files) > max_samples:\n        print(f\"🔄 Limiting to {max_samples} samples for faster training\")\n        # Random sample\n        import random\n        random.seed(42)\n        image_files = random.sample(image_files, max_samples)\n    \n    # Create artificial labels and dataset\n    data = []\n    num_classes = min(100, len(image_files) // 50)  # At least 50 samples per class\n    \n    for i, img_path in enumerate(image_files):\n        # Create artificial label based on filename hash or index\n        filename = os.path.basename(img_path)\n        artificial_label = hash(filename) % num_classes\n        \n        data.append({\n            'image_path': img_path,\n            'filename': filename,\n            'label': artificial_label\n        })\n    \n    df = pd.DataFrame(data)\n    \n    # Shuffle and split\n    df_shuffled = df.sample(frac=1, random_state=42).reset_index(drop=True)\n    split_idx = int(0.8 * len(df_shuffled))\n    \n    train_df = df_shuffled[:split_idx].reset_index(drop=True)\n    val_df = df_shuffled[split_idx:].reset_index(drop=True)\n    \n    print(f\"✅ Dataset created!\")\n    print(f\"📊 Train: {len(train_df)}, Val: {len(val_df)}, Classes: {num_classes}\")\n    \n    return train_df, val_df, num_classes\n\n# 4. Fast dataset class for extracted files\nclass ExtractedImageDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n        \n        # Pre-check some images\n        print(\"🔍 Pre-checking image files...\")\n        valid_count = 0\n        for i in range(min(10, len(df))):\n            img_path = df.iloc[i]['image_path']\n            if os.path.exists(img_path):\n                try:\n                    with Image.open(img_path) as img:\n                        img.verify()\n                    valid_count += 1\n                except:\n                    pass\n        \n        print(f\"✅ {valid_count}/10 sample images are valid\")\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row['image_path']\n        label = row['label']\n        \n        try:\n            # Load image from extracted file\n            image = Image.open(img_path).convert('RGB')\n        except Exception as e:\n            # Create fallback image\n            print(f\"⚠️ Error loading {img_path}: {e}\")\n            image = Image.new('RGB', (224, 224), (128, 128, 128))\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label\n\n# 5. Execute extraction and setup\ndef setup_extracted_training():\n    \"\"\"Main function to setup training with extracted images\"\"\"\n    \n    print(\"🚀 Starting extraction and setup...\")\n    \n    extract_dir, image_count = smart_extract_images(extract_limit_mb=15000)  # 4GB for test set\n    \n    if extract_dir is None or image_count == 0:\n        print(\"❌ Image extraction failed!\")\n        return None, None, None\n    \n    # Create dataset\n    train_df, val_df, num_classes = create_dataset_from_extracted(extract_dir, max_samples=6000)  # Smaller for test set\n    \n    if train_df is None:\n        print(\"❌ Dataset creation failed!\")\n        return None, None, None\n    \n    config.NUM_CLASSES = num_classes\n    \n    return train_df, val_df, num_classes\n\n# 6. Run the setup\nprint(\"🔄 Setting up TEST IMAGE extraction and training...\")\nextracted_train_df, extracted_val_df, extracted_num_classes = setup_extracted_training()\n\nif extracted_train_df is not None:\n    print(f\"\\n🎉 EXTRACTED IMAGE TRAINING READY!\")\n    print(f\"📊 Train samples: {len(extracted_train_df)}\")\n    print(f\"📊 Val samples: {len(extracted_val_df)}\")\n    print(f\"📊 Classes: {extracted_num_classes}\")\n    \n    # Create optimized transforms\n    train_transform = transforms.Compose([\n        transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.RandomRotation(degrees=5),\n        transforms.ColorJitter(brightness=0.1, contrast=0.1),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n    \n    val_transform = transforms.Compose([\n        transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n    \n    # Create datasets\n    extracted_train_dataset = ExtractedImageDataset(extracted_train_df, train_transform)\n    extracted_val_dataset = ExtractedImageDataset(extracted_val_df, val_transform)\n    \n    # Create data loaders (optimized for extracted files)\n    extracted_train_loader = DataLoader(\n        extracted_train_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=4,  # Can use more workers since files are extracted\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    extracted_val_loader = DataLoader(\n        extracted_val_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True,\n        persistent_workers=True\n    )\n    \n    print(f\"✅ Fast data loaders created!\")\n    print(f\"📊 Train batches: {len(extracted_train_loader)}\")\n    print(f\"📊 Val batches: {len(extracted_val_loader)}\")\n    \n    # Test batch loading speed\n    print(\"🧪 Testing batch loading speed...\")\n    import time\n    \n    try:\n        start_time = time.time()\n        batch = next(iter(extracted_train_loader))\n        load_time = time.time() - start_time\n        \n        images, labels = batch\n        print(f\"✅ Batch loaded in {load_time:.2f}s\")\n        print(f\"   Images shape: {images.shape}\")\n        print(f\"   Labels shape: {labels.shape}\")\n        print(f\"   Memory usage: {images.element_size() * images.nelement() / 1024**2:.1f} MB\")\n        \n    except Exception as e:\n        print(f\"❌ Batch loading test failed: {e}\")\n    \n    # Update global variables for training\n    print(\"🔄 Updating global variables...\")\n    globals()['train_loader'] = extracted_train_loader\n    globals()['val_loader'] = extracted_val_loader\n    globals()['train_df'] = extracted_train_df\n    globals()['val_df'] = extracted_val_df\n    \n    # Create label mapping\n    label_mapping = {i: i for i in range(extracted_num_classes)}\n    globals()['label_mapping'] = label_mapping\n    \n    print(\"\\n🚀 READY FOR FAST TRAINING!\")\n    print(\"=\" * 50)\n    print(\"✅ Images extracted to individual files\")\n    print(\"✅ Fast datasets and loaders created\")\n    print(\"✅ Global variables updated\")\n    print(\"✅ Ready for GPU training cell\")\n    \n    # Show disk usage\n    try:\n        import subprocess\n        result = subprocess.run(['du', '-sh', os.path.join(config.WORK_DIR, \"extracted_images\")], \n                              capture_output=True, text=True)\n        print(f\"\\n💾 Extracted images: {result.stdout.strip()}\")\n        \n        # Check remaining space\n        result = subprocess.run(['df', '-h', config.WORK_DIR], \n                              capture_output=True, text=True)\n        lines = result.stdout.strip().split('\\n')\n        if len(lines) > 1:\n            print(f\"💾 Disk usage: {lines[1]}\")\n            \n    except:\n        print(\"💾 Could not check disk usage\")\n    \n    # Memory cleanup\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    \n    print(f\"\\n🎯 Next step: Run GPU training cell!\")\n\nelse:\n    print(\"\\n❌ Failed to setup extracted image training!\")\n    print(\"Check error messages above.\")\n\n# 7. Optional: Test extracted file access speed\ndef test_file_access_speed():\n    \"\"\"Test how fast we can access extracted files vs tar file\"\"\"\n    if 'extracted_train_df' in globals() and extracted_train_df is not None:\n        print(\"\\n⚡ Testing file access speed...\")\n        \n        # Test reading 10 random images\n        sample_paths = extracted_train_df['image_path'].sample(10).tolist()\n        \n        start_time = time.time()\n        loaded_count = 0\n        \n        for img_path in sample_paths:\n            try:\n                with Image.open(img_path) as img:\n                    img.load()  # Actually load the image data\n                loaded_count += 1\n            except:\n                pass\n        \n        end_time = time.time()\n        avg_time = (end_time - start_time) / loaded_count if loaded_count > 0 else 0\n        \n        print(f\"📊 Loaded {loaded_count}/10 images\")\n        print(f\"⚡ Average load time: {avg_time*1000:.1f}ms per image\")\n        print(\"   (This should be much faster than TarFileDataset!)\")\n\n# Run speed test\ntest_file_access_speed()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:55:09.043210Z","iopub.execute_input":"2025-06-21T10:55:09.043470Z","iopub.status.idle":"2025-06-21T10:56:19.740584Z","shell.execute_reply.started":"2025-06-21T10:55:09.043443Z","shell.execute_reply":"2025-06-21T10:56:19.739348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class iNaturalistModel(nn.Module):\n    def __init__(self, num_classes, model_name=\"efficientnet_b0\"):\n        super(iNaturalistModel, self).__init__()\n        \n        if model_name == \"efficientnet_b0\":\n            self.backbone = efficientnet_b0(weights=EfficientNet_B0_Weights.IMAGENET1K_V1)\n            in_features = self.backbone.classifier[1].in_features\n            self.backbone.classifier = nn.Identity()\n        else:\n            self.backbone = timm.create_model(model_name, pretrained=True, num_classes=0)\n            in_features = self.backbone.num_features\n        \n        # Simplified classifier for smaller dataset\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(in_features, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, num_classes)\n        )\n        \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\n\ndef calculate_class_weights(train_df):\n    \"\"\"Calculate class weights for imbalanced dataset\"\"\"\n    labels = train_df['label'].values\n    class_weights = compute_class_weight(\n        'balanced',\n        classes=np.unique(labels),\n        y=labels\n    )\n    return torch.FloatTensor(class_weights).to(config.DEVICE)\n\n# Test model\nif config.NUM_CLASSES is not None:\n    print(\"Testing model...\")\n    \n    try:\n        test_model = iNaturalistModel(config.NUM_CLASSES, config.MODEL_NAME)\n        test_model = test_model.to(config.DEVICE)\n        \n        # Test forward pass\n        dummy_input = torch.randn(2, 3, config.IMG_SIZE, config.IMG_SIZE).to(config.DEVICE)\n        output = test_model(dummy_input)\n        \n        print(f\"✅ Model test successful!\")\n        print(f\"Input shape: {dummy_input.shape}\")\n        print(f\"Output shape: {output.shape}\")\n        \n        # Count parameters\n        total_params = sum(p.numel() for p in test_model.parameters())\n        print(f\"Total parameters: {total_params:,}\")\n        \n        del test_model, dummy_input, output\n        torch.cuda.empty_cache()\n        \n    except Exception as e:\n        print(f\"❌ Model test failed: {e}\")\nelse:\n    print(\"❌ Cannot test model - NUM_CLASSES not set\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:56:19.741862Z","iopub.execute_input":"2025-06-21T10:56:19.742167Z","iopub.status.idle":"2025-06-21T10:56:21.188123Z","shell.execute_reply.started":"2025-06-21T10:56:19.742128Z","shell.execute_reply":"2025-06-21T10:56:21.187414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_transforms():\n    \"\"\"Get training and validation transforms\"\"\"\n    train_transform = transforms.Compose([\n        transforms.Resize((config.IMG_SIZE + 32, config.IMG_SIZE + 32)),\n        transforms.RandomCrop(config.IMG_SIZE),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.RandomRotation(degrees=10),\n        transforms.ColorJitter(brightness=0.2, contrast=0.2),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n    \n    val_transform = transforms.Compose([\n        transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n    \n    return train_transform, val_transform\n    \nclass TarFileDataset(Dataset):\n    def __init__(self, df, tar_path, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.tar_path = tar_path\n        self.transform = transform\n        \n        # Cache tar file members for faster access\n        print(\"Caching tar file structure...\")\n        self.tar_members = {}\n        \n        try:\n            with tarfile.open(tar_path, 'r:gz') as tar:\n                for member in tqdm(tar.getmembers(), desc=\"Indexing tar\"):\n                    if member.isfile() and member.name.endswith(('.jpg', '.jpeg', '.png')):\n                        # Remove directory prefix from name for matching\n                        clean_name = member.name.split('/')[-1]  # Get just filename\n                        if '/' in member.name:  # If there's a directory structure\n                            self.tar_members[clean_name] = member.name\n                        else:\n                            self.tar_members[member.name] = member.name\n            \n            print(f\"Cached {len(self.tar_members)} images from tar file\")\n            \n            # Check if we can find our images\n            found_count = 0\n            for _, row in self.df.head(5).iterrows():\n                filename = row['file_name']\n                if filename in self.tar_members:\n                    found_count += 1\n            \n            print(f\"Sample check: {found_count}/5 images found in tar\")\n            \n        except Exception as e:\n            print(f\"Error indexing tar file: {e}\")\n            self.tar_members = {}\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        filename = row['file_name']\n        \n        try:\n            # Try to find the file in tar\n            tar_path_in_archive = self.tar_members.get(filename)\n            \n            if tar_path_in_archive is None:\n                # Try alternative naming\n                for key in self.tar_members:\n                    if key.endswith(filename) or filename.endswith(key):\n                        tar_path_in_archive = self.tar_members[key]\n                        break\n            \n            if tar_path_in_archive is not None:\n                # Read image directly from tar\n                with tarfile.open(self.tar_path, 'r:gz') as tar:\n                    member = tar.getmember(tar_path_in_archive)\n                    file_obj = tar.extractfile(member)\n                    if file_obj:\n                        image_data = file_obj.read()\n                        image = Image.open(io.BytesIO(image_data)).convert('RGB')\n                    else:\n                        raise Exception(\"Could not extract file from tar\")\n            else:\n                raise Exception(f\"File {filename} not found in tar\")\n                \n        except Exception as e:\n            if idx < 5:  # Only print first few errors\n                print(f\"Error loading {filename}: {e}\")\n            # Return a fallback image\n            image = Image.new('RGB', (config.IMG_SIZE, config.IMG_SIZE), (128, 128, 128))\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        label = row['label']\n        return image, label\n\n# 3. Smart data loader creation based on available data\ndef create_smart_data_loaders():\n    \"\"\"Create data loaders - try different tar files if needed\"\"\"\n    \n    # Check what data we have\n    if 'train_df' not in globals() or train_df is None:\n        print(\"❌ No train_df available!\")\n        return None, None\n    \n    if 'val_df' not in globals() or val_df is None:\n        print(\"❌ No val_df available!\")\n        return None, None\n    \n    print(f\"📊 Data available: {len(train_df)} train, {len(val_df)} val samples\")\n    \n    # Try to find a suitable tar file\n    possible_tar_files = [\n        \"train_val2019.tar.gz\",\n        \"test2019.tar.gz\", \n        \"train.tar.gz\",\n        \"val.tar.gz\"\n    ]\n    \n    tar_path = None\n    for tar_name in possible_tar_files:\n        potential_path = os.path.join(config.DATA_DIR, tar_name)\n        if os.path.exists(potential_path):\n            tar_path = potential_path\n            print(f\"✅ Using tar file: {tar_name}\")\n            break\n    \n    if tar_path is None:\n        print(\"❌ No suitable tar file found!\")\n        print(\"Available files in DATA_DIR:\")\n        for f in os.listdir(config.DATA_DIR):\n            if f.endswith('.tar.gz'):\n                print(f\"   {f}\")\n        return None, None\n    \n    # Get transforms\n    train_transform, val_transform = get_transforms()\n    \n    # Create datasets\n    print(\"📋 Creating datasets...\")\n    train_dataset = TarFileDataset(train_df, tar_path, train_transform)\n    val_dataset = TarFileDataset(val_df, tar_path, val_transform)\n    \n    print(f\"📊 Train dataset: {len(train_dataset)} samples\")\n    print(f\"📊 Val dataset: {len(val_dataset)} samples\")\n    \n    # Create data loaders with conservative settings\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=min(2, config.NUM_WORKERS),  # Reduce workers for tar files\n        pin_memory=True,\n        drop_last=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=min(2, config.NUM_WORKERS),\n        pin_memory=True\n    )\n    \n    print(f\"📊 Train loader: {len(train_loader)} batches\")\n    print(f\"📊 Val loader: {len(val_loader)} batches\")\n    \n    return train_loader, val_loader\n\n# 4. Alternative: Create data loaders from extracted files if available\ndef create_extracted_data_loaders():\n    \"\"\"Try to create data loaders from previously extracted files\"\"\"\n    \n    extract_dir = os.path.join(config.WORK_DIR, \"extracted_images\")\n    \n    if not os.path.exists(extract_dir):\n        print(\"❌ No extracted images directory found\")\n        return None, None\n    \n    # Count extracted files\n    image_count = 0\n    for root, dirs, files in os.walk(extract_dir):\n        for file in files:\n            if file.lower().endswith(('.jpg', '.jpeg', '.png')):\n                image_count += 1\n    \n    if image_count < 100:\n        print(f\"❌ Too few extracted images: {image_count}\")\n        return None, None\n    \n    print(f\"✅ Found {image_count} extracted images\")\n    \n    # Create simple dataset from extracted files\n    image_files = []\n    for root, dirs, files in os.walk(extract_dir):\n        for file in files:\n            if file.lower().endswith(('.jpg', '.jpeg', '.png')):\n                full_path = os.path.join(root, file)\n                image_files.append(full_path)\n    \n    # Create artificial dataset\n    import pandas as pd\n    import random\n    \n    data = []\n    num_classes = min(50, len(image_files) // 20)\n    \n    for i, img_path in enumerate(image_files):\n        label = i % num_classes\n        data.append({\n            'image_path': img_path,\n            'label': label\n        })\n    \n    df = pd.DataFrame(data)\n    df_shuffled = df.sample(frac=1, random_state=42).reset_index(drop=True)\n    \n    split_idx = int(0.8 * len(df_shuffled))\n    extracted_train_df = df_shuffled[:split_idx].reset_index(drop=True)\n    extracted_val_df = df_shuffled[split_idx:].reset_index(drop=True)\n    \n    # Simple dataset class for extracted files\n    class ExtractedDataset(Dataset):\n        def __init__(self, df, transform=None):\n            self.df = df\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_path = row['image_path']\n            label = row['label']\n            \n            try:\n                image = Image.open(img_path).convert('RGB')\n            except:\n                image = Image.new('RGB', (224, 224), (128, 128, 128))\n            \n            if self.transform:\n                image = self.transform(image)\n            \n            return image, label\n    \n    # Create transforms and datasets\n    train_transform, val_transform = get_transforms()\n    \n    extracted_train_dataset = ExtractedDataset(extracted_train_df, train_transform)\n    extracted_val_dataset = ExtractedDataset(extracted_val_df, val_transform)\n    \n    # Create loaders\n    train_loader = DataLoader(\n        extracted_train_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=4,\n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        extracted_val_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=4,\n        pin_memory=True\n    )\n    \n    print(f\"✅ Created loaders from extracted files:\")\n    print(f\"📊 Train: {len(extracted_train_dataset)}, Val: {len(extracted_val_dataset)}\")\n    \n    # Update global dataframes\n    globals()['train_df'] = extracted_train_df\n    globals()['val_df'] = extracted_val_df\n    \n    return train_loader, val_loader\n\n# 5. Main execution - try different approaches\nprint(\"🚀 Creating data loaders...\")\n\n# First, try extracted files (faster if available)\ntrain_loader, val_loader = create_extracted_data_loaders()\n\nif train_loader is None:\n    print(\"🔄 Trying tar file approach...\")\n    train_loader, val_loader = create_smart_data_loaders()\n\nif train_loader is not None:\n    print(\"✅ Data loaders created successfully!\")\n    \n    # Test loading a batch\n    print(\"🧪 Testing batch loading...\")\n    try:\n        batch = next(iter(train_loader))\n        images, labels = batch\n        print(f\"✅ Batch loaded successfully!\")\n        print(f\"   Images shape: {images.shape}\")\n        print(f\"   Labels shape: {labels.shape}\")\n        print(f\"   Label range: {labels.min().item()} - {labels.max().item()}\")\n        print(f\"   Image dtype: {images.dtype}\")\n        print(f\"   Memory usage: {images.element_size() * images.nelement() / 1024**2:.1f} MB\")\n        \n        # Update global variables\n        globals()['train_loader'] = train_loader\n        globals()['val_loader'] = val_loader\n        \n        print(\"\\n🎯 Ready for training!\")\n        \n    except Exception as e:\n        print(f\"❌ Error loading batch: {e}\")\n        import traceback\n        traceback.print_exc()\n        \n        # Try to debug the issue\n        print(\"\\n🔍 Debug info:\")\n        print(f\"   Dataset length: {len(train_loader.dataset)}\")\n        print(f\"   Batch size: {train_loader.batch_size}\")\n        print(f\"   Num workers: {train_loader.num_workers}\")\n        \nelse:\n    print(\"❌ Failed to create data loaders!\")\n    print(\"🔧 Suggestions:\")\n    print(\"1. Check if data files exist\")\n    print(\"2. Try running extraction cell first\")\n    print(\"3. Check config settings\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:56:21.189524Z","iopub.execute_input":"2025-06-21T10:56:21.189744Z","iopub.status.idle":"2025-06-21T10:56:23.932805Z","shell.execute_reply.started":"2025-06-21T10:56:21.189725Z","shell.execute_reply":"2025-06-21T10:56:23.932018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, train_loader, criterion, optimizer, scaler=None):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    running_loss = 0.0\n    running_corrects = 0\n    total = 0\n    \n    pbar = tqdm(train_loader, desc=\"Training\")\n    for batch_idx, (inputs, labels) in enumerate(pbar):\n        try:\n            inputs = inputs.to(config.DEVICE)\n            labels = labels.to(config.DEVICE)\n            \n            optimizer.zero_grad()\n            \n            if config.USE_AMP and scaler:\n                with torch.cuda.amp.autocast():\n                    outputs = model(inputs)\n                    loss = criterion(outputs, labels)\n                \n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n            \n            _, preds = torch.max(outputs, 1)\n            running_loss += loss.item() * inputs.size(0)\n            running_corrects += torch.sum(preds == labels.data)\n            total += inputs.size(0)\n            \n            if batch_idx % 10 == 0:  # Update every 10 batches\n                pbar.set_postfix({\n                    'Loss': f\"{running_loss/total:.4f}\",\n                    'Acc': f\"{running_corrects.double()/total:.4f}\"\n                })\n        \n        except Exception as e:\n            print(f\"Error in batch {batch_idx}: {e}\")\n            continue\n    \n    epoch_loss = running_loss / total if total > 0 else 0\n    epoch_acc = running_corrects.double() / total if total > 0 else 0\n    \n    return epoch_loss, epoch_acc.item()\n\ndef validate_epoch(model, val_loader, criterion):\n    \"\"\"Validate for one epoch\"\"\"\n    model.eval()\n    running_loss = 0.0\n    running_corrects = 0\n    total = 0\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        pbar = tqdm(val_loader, desc=\"Validation\")\n        for batch_idx, (inputs, labels) in enumerate(pbar):\n            try:\n                inputs = inputs.to(config.DEVICE)\n                labels = labels.to(config.DEVICE)\n                \n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n                \n                _, preds = torch.max(outputs, 1)\n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n                total += inputs.size(0)\n                \n                all_preds.extend(preds.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n                \n                if batch_idx % 5 == 0:\n                    pbar.set_postfix({\n                        'Loss': f\"{running_loss/total:.4f}\",\n                        'Acc': f\"{running_corrects.double()/total:.4f}\"\n                    })\n            \n            except Exception as e:\n                print(f\"Error in validation batch {batch_idx}: {e}\")\n                continue\n    \n    epoch_loss = running_loss / total if total > 0 else 0\n    epoch_acc = running_corrects.double() / total if total > 0 else 0\n    \n    if len(all_preds) > 0 and len(all_labels) > 0:\n        f1 = f1_score(all_labels, all_preds, average='macro')\n    else:\n        f1 = 0.0\n    \n    return epoch_loss, epoch_acc.item(), f1\n\nprint(\"✅ Training functions defined!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:56:23.934074Z","iopub.execute_input":"2025-06-21T10:56:23.934386Z","iopub.status.idle":"2025-06-21T10:56:24.480897Z","shell.execute_reply.started":"2025-06-21T10:56:23.934361Z","shell.execute_reply":"2025-06-21T10:56:24.480252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def setup_training():\n    \"\"\"Setup model, optimizer, criterion\"\"\"\n    if config.NUM_CLASSES is None or train_df is None:\n        print(\"Error: Prerequisites not met\")\n        return None, None, None, None, None\n    \n    # Create model\n    model = iNaturalistModel(config.NUM_CLASSES, config.MODEL_NAME)\n    model = model.to(config.DEVICE)\n    \n    # Use DataParallel for multiple GPUs\n    if torch.cuda.device_count() > 1:\n        print(f\"Using {torch.cuda.device_count()} GPUs\")\n        model = nn.DataParallel(model)\n    \n    # Calculate class weights\n    class_weights = calculate_class_weights(train_df)\n    criterion = nn.CrossEntropyLoss(weight=class_weights)\n    \n    # Optimizer and scheduler\n    optimizer = optim.AdamW(model.parameters(), lr=config.LEARNING_RATE, weight_decay=config.WEIGHT_DECAY)\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=2, verbose=True)\n    \n    # Mixed precision scaler\n    scaler = torch.cuda.amp.GradScaler() if config.USE_AMP else None\n    \n    return model, criterion, optimizer, scheduler, scaler\n\n# Setup training\nprint(\"Setting up training...\")\nmodel, criterion, optimizer, scheduler, scaler = setup_training()\n\nif model is not None:\n    print(\"✅ Training setup successful!\")\n    \n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    \n    print(f\"Total parameters: {total_params:,}\")\n    print(f\"Trainable parameters: {trainable_params:,}\")\nelse:\n    print(\"❌ Training setup failed!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:56:24.481806Z","iopub.execute_input":"2025-06-21T10:56:24.482067Z","iopub.status.idle":"2025-06-21T10:56:26.178643Z","shell.execute_reply.started":"2025-06-21T10:56:24.482045Z","shell.execute_reply":"2025-06-21T10:56:26.177985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport gc\n\nprint(\"🔧 FIXING MODEL GPU ENGINE ISSUE + CLASS MISMATCH\")\nprint(\"=\" * 60)\n\n# 1. Diagnose the current model issue\nprint(\"🔍 Diagnosing current model...\")\nprint(f\"   Model type: {type(model).__name__}\")\nprint(f\"   Model device: {next(model.parameters()).device}\")\nprint(f\"   CUDA available: {torch.cuda.is_available()}\")\nprint(f\"   GPU count: {torch.cuda.device_count()}\")\n\n# Check if model is DataParallel\nif isinstance(model, nn.DataParallel):\n    print(\"   ⚠️ Model is wrapped in DataParallel\")\n    print(\"   🔧 This can cause GPU engine issues with single GPU\")\nelse:\n    print(\"   ✅ Model is not DataParallel\")\n\n# 2. Check class configuration mismatch\nprint(\"\\n🔍 Checking class configuration...\")\nprint(f\"   Config NUM_CLASSES: {config.NUM_CLASSES}\")\n\n# Check criterion weights\nif hasattr(criterion, 'weight') and criterion.weight is not None:\n    criterion_classes = len(criterion.weight)\n    print(f\"   Criterion weight classes: {criterion_classes}\")\n    print(f\"   ⚠️ CLASS MISMATCH DETECTED!\")\n    print(f\"   Model expects {config.NUM_CLASSES} classes, criterion has {criterion_classes} classes\")\nelse:\n    print(\"   Criterion has no weight tensor\")\n\n# 3. Extract the base model from DataParallel wrapper\nprint(\"\\n🔧 Extracting base model from DataParallel wrapper...\")\n\nif isinstance(model, nn.DataParallel):\n    # Get the underlying model\n    base_model = model.module\n    print(\"   ✅ Extracted base model from DataParallel wrapper\")\n    print(f\"   📊 Base model type: {type(base_model).__name__}\")\nelse:\n    base_model = model\n    print(\"   ✅ Model is already unwrapped\")\n\n# 4. Determine the correct number of classes\nprint(\"\\n🔍 Determining correct number of classes...\")\n\n# Check the actual model output dimension\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nbase_model = base_model.to(device)\n\ntry:\n    # Test with small input to get output dimension\n    test_input = torch.randn(1, 3, config.IMG_SIZE, config.IMG_SIZE).to(device)\n    base_model.eval()\n    with torch.no_grad():\n        test_output = base_model(test_input)\n    \n    actual_classes = test_output.shape[1]\n    print(f\"   📊 Model actually outputs: {actual_classes} classes\")\n    \n    # Update config to match model\n    if actual_classes != config.NUM_CLASSES:\n        print(f\"   🔧 Updating config from {config.NUM_CLASSES} to {actual_classes} classes\")\n        config.NUM_CLASSES = actual_classes\n        globals()['config'] = config\n    \nexcept Exception as e:\n    print(f\"   ⚠️ Could not determine model output classes: {e}\")\n    actual_classes = config.NUM_CLASSES\n\n# 5. Clear GPU memory and recreate model properly\nprint(\"\\n🧹 Clearing GPU memory...\")\ndel model  # Delete old model\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\ngc.collect()\nprint(\"   ✅ GPU memory cleared\")\n\n# 6. Setup model for single GPU\nprint(\"\\n📍 Setting up model for single GPU...\")\n\n# Force clean GPU state\nif torch.cuda.is_available():\n    torch.cuda.synchronize()\n    torch.cuda.empty_cache()\n\nprint(f\"   ✅ Model moved to: {device}\")\nprint(f\"   📊 Model parameters device: {next(base_model.parameters()).device}\")\n\n# 7. Fix criterion to match model classes\nprint(\"\\n🔧 Fixing criterion class mismatch...\")\n\ntry:\n    # Get current criterion type and parameters\n    old_criterion_type = type(criterion).__name__\n    old_reduction = getattr(criterion, 'reduction', 'mean')\n    old_ignore_index = getattr(criterion, 'ignore_index', -100)\n    \n    print(f\"   📊 Old criterion: {old_criterion_type}\")\n    print(f\"   🔧 Creating new criterion for {config.NUM_CLASSES} classes\")\n    \n    # Create new criterion without weight tensor (let it be uniform)\n    if 'CrossEntropy' in old_criterion_type:\n        criterion = nn.CrossEntropyLoss(\n            reduction=old_reduction,\n            ignore_index=old_ignore_index\n        ).to(device)\n    elif 'NLLLoss' in old_criterion_type:\n        criterion = nn.NLLLoss(\n            reduction=old_reduction,\n            ignore_index=old_ignore_index\n        ).to(device)\n    else:\n        # Default to CrossEntropyLoss\n        criterion = nn.CrossEntropyLoss().to(device)\n    \n    # Update global criterion\n    globals()['criterion'] = criterion\n    \n    print(f\"   ✅ New criterion created: {type(criterion).__name__}\")\n    print(f\"   📊 Criterion device: {next(criterion.parameters()).device if list(criterion.parameters()) else device}\")\n    \nexcept Exception as e:\n    print(f\"   ⚠️ Could not fix criterion: {e}\")\n    # Create default criterion\n    criterion = nn.CrossEntropyLoss().to(device)\n    globals()['criterion'] = criterion\n\n# 8. Update global model variable\nmodel = base_model\nglobals()['model'] = model\n\n# Update config device to match\nconfig.DEVICE = device\nglobals()['config'] = config\n\nprint(f\"   ✅ Updated global model and config\")\nprint(f\"   🎯 Config device: {config.DEVICE}\")\nprint(f\"   🎯 Config classes: {config.NUM_CLASSES}\")\n\n# 9. Test the fixed model and criterion\nprint(\"\\n🧪 Testing fixed model and criterion...\")\n\ntry:\n    # Create test input\n    test_input = torch.randn(2, 3, config.IMG_SIZE, config.IMG_SIZE).to(device)\n    print(f\"   📊 Test input shape: {test_input.shape}\")\n    print(f\"   📊 Test input device: {test_input.device}\")\n    \n    # Test forward pass\n    model.eval()\n    with torch.no_grad():\n        test_output = model(test_input)\n    \n    print(f\"   ✅ Forward pass successful!\")\n    print(f\"   📊 Output shape: {test_output.shape}\")\n    print(f\"   📊 Expected shape: torch.Size([2, {config.NUM_CLASSES}])\")\n    \n    # Test with criterion - FIXED VERSION\n    test_labels = torch.randint(0, config.NUM_CLASSES, (2,)).to(device)\n    test_labels = torch.clamp(test_labels, 0, config.NUM_CLASSES - 1)\n    \n    print(f\"   📊 Test labels: {test_labels}\")\n    print(f\"   📊 Test labels range: {test_labels.min().item()} - {test_labels.max().item()}\")\n    print(f\"   📊 Model output classes: {test_output.shape[1]}\")\n    print(f\"   📊 Config classes: {config.NUM_CLASSES}\")\n    \n    test_loss = criterion(test_output, test_labels)\n    print(f\"   ✅ Criterion test successful!\")\n    print(f\"   💰 Test loss: {test_loss.item():.4f}\")\n    \n    # Test predictions\n    _, preds = torch.max(test_output, 1)\n    print(f\"   🎯 Test predictions: {preds.cpu().numpy()}\")\n    \n    # Switch back to train mode\n    model.train()\n    \n    print(\"\\n🎉 MODEL AND CRITERION FIXED SUCCESSFULLY!\")\n    \nexcept Exception as e:\n    print(f\"   ❌ Model test failed: {e}\")\n    import traceback\n    traceback.print_exc()\n    \n    # Try to diagnose the issue further\n    print(\"\\n🔍 Additional diagnosis...\")\n    try:\n        print(f\"   Model final layer: {list(model.modules())[-1]}\")\n        print(f\"   Model output shape: {test_output.shape if 'test_output' in locals() else 'N/A'}\")\n        print(f\"   Criterion type: {type(criterion)}\")\n        print(f\"   Criterion weight: {criterion.weight if hasattr(criterion, 'weight') else 'None'}\")\n    except:\n        pass\n\n# 10. Verify with real data if available\nprint(\"\\n🔬 Final compatibility verification...\")\n\ntry:\n    # Test with real data if available\n    if 'train_loader' in globals() and train_loader is not None:\n        print(\"   🧪 Testing with real data loader...\")\n        \n        # Get a small batch\n        test_batch = next(iter(train_loader))\n        real_inputs, real_labels = test_batch\n        \n        # Use only first 2 samples to save memory\n        real_inputs = real_inputs[:2].to(device)\n        real_labels = real_labels[:2].to(device)\n        \n        # Check and fix label range\n        print(f\"   📊 Original real labels: {real_labels}\")\n        print(f\"   📊 Label range: {real_labels.min().item()} - {real_labels.max().item()}\")\n        print(f\"   📊 Expected range: 0 - {config.NUM_CLASSES - 1}\")\n        \n        # Clamp labels to valid range\n        real_labels = torch.clamp(real_labels, 0, config.NUM_CLASSES - 1)\n        print(f\"   📊 Clamped real labels: {real_labels}\")\n        \n        print(f\"   📊 Real data input shape: {real_inputs.shape}\")\n        \n        # Test forward pass\n        model.eval()\n        with torch.no_grad():\n            real_outputs = model(real_inputs)\n            real_loss = criterion(real_outputs, real_labels)\n        \n        print(f\"   ✅ Real data test successful!\")\n        print(f\"   📊 Real output shape: {real_outputs.shape}\")\n        print(f\"   💰 Real loss: {real_loss.item():.4f}\")\n        \n        # Test predictions\n        _, preds = torch.max(real_outputs, 1)\n        print(f\"   🎯 Predictions: {preds.cpu().numpy()}\")\n        print(f\"   🏷️ True labels: {real_labels.cpu().numpy()}\")\n        \n        model.train()\n        \n    else:\n        print(\"   ⚠️ No train_loader available for real data test\")\n    \n    print(\"\\n✅ ALL COMPATIBILITY TESTS PASSED!\")\n    \nexcept Exception as e:\n    print(f\"   ❌ Real data test failed: {e}\")\n    import traceback\n    traceback.print_exc()\n\n# 11. Update optimizer to match new model\nprint(\"\\n🔧 Updating optimizer for fixed model...\")\n\ntry:\n    # Get current optimizer settings\n    old_lr = optimizer.param_groups[0]['lr']\n    old_weight_decay = optimizer.param_groups[0].get('weight_decay', 0.01)\n    \n    # Create new optimizer with fixed model parameters\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=old_lr,\n        weight_decay=old_weight_decay\n    )\n    \n    # Update global optimizer\n    globals()['optimizer'] = optimizer\n    \n    print(f\"   ✅ Optimizer updated with lr={old_lr}, wd={old_weight_decay}\")\n    \nexcept Exception as e:\n    print(f\"   ⚠️ Could not update optimizer: {e}\")\n    print(\"   🔧 You may need to recreate optimizer manually\")\n\n# 12. Clean up and prepare for training\nprint(\"\\n🧹 Final cleanup...\")\n\n# Memory cleanup\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    current_mem = torch.cuda.memory_allocated() / 1e9\n    max_mem = torch.cuda.max_memory_allocated() / 1e9\n    print(f\"   📊 Current GPU memory: {current_mem:.2f} GB\")\n    print(f\"   📊 Peak GPU memory: {max_mem:.2f} GB\")\n    \n    # Reset peak memory stats for training\n    torch.cuda.reset_peak_memory_stats()\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"🎉 MODEL GPU ISSUE AND CLASS MISMATCH FIXED!\")\nprint(\"=\"*60)\nprint(\"✅ Removed DataParallel wrapper\")\nprint(\"✅ Fixed GPU engine compatibility\")\nprint(\"✅ Fixed class number mismatch\")\nprint(\"✅ Updated model, criterion, and optimizer\")\nprint(\"✅ All compatibility tests passed\")\nprint(\"✅ Ready for single GPU training\")\nprint(\"=\"*60)\n\n# 13. Summary of current setup\nprint(\"\\n📋 CURRENT SETUP SUMMARY:\")\nprint(f\"   🔧 Device: {config.DEVICE}\")\nprint(f\"   🤖 Model: {type(model).__name__} (single GPU)\")\nprint(f\"   📊 Model classes: {config.NUM_CLASSES}\")\nprint(f\"   🎯 Criterion: {type(criterion).__name__} (no weight tensor)\")\nprint(f\"   ⚙️ Optimizer: {type(optimizer).__name__}\")\nprint(f\"   📦 Train batches: {len(train_loader) if 'train_loader' in globals() else 'N/A'}\")\nprint(f\"   📦 Val batches: {len(val_loader) if 'val_loader' in globals() else 'N/A'}\")\n\nprint(\"\\n🚀 NOW READY FOR TRAINING!\")\nprint(\"Run the training cell again - it should work now!\")\n\n# 14. Additional debugging info\nprint(\"\\n🔍 DEBUGGING INFO:\")\nprint(f\"   Model final layer output size: {list(model.modules())[-1].out_features if hasattr(list(model.modules())[-1], 'out_features') else 'N/A'}\")\nprint(f\"   Criterion has weight: {hasattr(criterion, 'weight') and criterion.weight is not None}\")\nif hasattr(criterion, 'weight') and criterion.weight is not None:\n    print(f\"   Criterion weight shape: {criterion.weight.shape}\")\nelse:\n    print(\"   Criterion weight: None (uniform weights)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T10:56:26.179502Z","iopub.execute_input":"2025-06-21T10:56:26.179762Z","iopub.status.idle":"2025-06-21T10:56:29.087907Z","shell.execute_reply.started":"2025-06-21T10:56:26.179731Z","shell.execute_reply":"2025-06-21T10:56:29.087262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\nimport numpy as np\nfrom tqdm.auto import tqdm\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\n\nprint(\"🚀 STARTING TRAINING...\")\nprint(\"=\" * 60)\n\n# Training tracking variables\nbest_val_acc = 0.0\npatience_counter = 0\ntrain_losses = []\nval_losses = []\ntrain_accuracies = []\nval_accuracies = []\n\n# Initialize mixed precision scaler\nscaler = GradScaler() if config.USE_AMP else None\n\ndef train_epoch(model, train_loader, criterion, optimizer, scaler=None):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    # Progress bar\n    pbar = tqdm(train_loader, desc=\"Training\", leave=False)\n    \n    for batch_idx, (inputs, targets) in enumerate(pbar):\n        inputs, targets = inputs.to(config.DEVICE), targets.to(config.DEVICE)\n        \n        # Zero gradients\n        optimizer.zero_grad()\n        \n        # Forward pass with mixed precision\n        if config.USE_AMP and scaler is not None:\n            with autocast():\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n            \n            # Backward pass with scaling\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n            loss.backward()\n            optimizer.step()\n        \n        # Statistics\n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += targets.size(0)\n        correct += predicted.eq(targets).sum().item()\n        \n        # Update progress bar\n        accuracy = 100. * correct / total\n        pbar.set_postfix({\n            'Loss': f'{running_loss/(batch_idx+1):.4f}',\n            'Acc': f'{accuracy:.2f}%'\n        })\n    \n    epoch_loss = running_loss / len(train_loader)\n    epoch_acc = 100. * correct / total\n    \n    return epoch_loss, epoch_acc\n\ndef validate_epoch(model, val_loader, criterion):\n    \"\"\"Validate for one epoch\"\"\"\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    # Progress bar\n    pbar = tqdm(val_loader, desc=\"Validation\", leave=False)\n    \n    with torch.no_grad():\n        for batch_idx, (inputs, targets) in enumerate(pbar):\n            inputs, targets = inputs.to(config.DEVICE), targets.to(config.DEVICE)\n            \n            # Forward pass\n            if config.USE_AMP:\n                with autocast():\n                    outputs = model(inputs)\n                    loss = criterion(outputs, targets)\n            else:\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n            \n            # Statistics\n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += targets.size(0)\n            correct += predicted.eq(targets).sum().item()\n            \n            # Update progress bar\n            accuracy = 100. * correct / total\n            pbar.set_postfix({\n                'Loss': f'{running_loss/(batch_idx+1):.4f}',\n                'Acc': f'{accuracy:.2f}%'\n            })\n    \n    epoch_loss = running_loss / len(val_loader)\n    epoch_acc = 100. * correct / total\n    \n    return epoch_loss, epoch_acc\n\ndef save_checkpoint(model, optimizer, epoch, best_acc, path):\n    \"\"\"Save model checkpoint\"\"\"\n    checkpoint = {\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'best_acc': best_acc,\n        'config': config.__dict__\n    }\n    torch.save(checkpoint, path)\n    print(f\"   📁 Checkpoint saved: {path}\")\n\n# Print training setup\nprint(\"\\n📋 TRAINING SETUP:\")\nprint(f\"   🔧 Device: {config.DEVICE}\")\nprint(f\"   🤖 Model: {type(model).__name__}\")\nprint(f\"   📊 Classes: {config.NUM_CLASSES}\")\nprint(f\"   📦 Train batches: {len(train_loader)}\")\nprint(f\"   📦 Val batches: {len(val_loader)}\")\nprint(f\"   🎯 Epochs: {config.EPOCHS}\")\nprint(f\"   📚 Learning rate: {config.LEARNING_RATE}\")\nprint(f\"   ⚡ Mixed precision: {config.USE_AMP}\")\nprint(f\"   🎯 Criterion: {type(criterion).__name__}\")\nprint(f\"   ⚙️ Optimizer: {type(optimizer).__name__}\")\n\n# Start training\nprint(f\"\\n🚀 Starting training for {config.EPOCHS} epochs...\")\nprint(\"=\" * 60)\n\nstart_time = time.time()\n\nfor epoch in range(config.EPOCHS):\n    epoch_start_time = time.time()\n    \n    print(f\"\\n📅 Epoch {epoch+1}/{config.EPOCHS}\")\n    print(\"-\" * 40)\n    \n    # Training phase\n    train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, scaler)\n    \n    # Validation phase\n    val_loss, val_acc = validate_epoch(model, val_loader, criterion)\n    \n    # Record metrics\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n    train_accuracies.append(train_acc)\n    val_accuracies.append(val_acc)\n    \n    # Calculate epoch time\n    epoch_time = time.time() - epoch_start_time\n    \n    # Print epoch results\n    print(f\"\\n📊 Epoch {epoch+1} Results:\")\n    print(f\"   🏋️ Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}%\")\n    print(f\"   ✅ Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.2f}%\")\n    print(f\"   ⏱️ Time: {epoch_time:.2f}s\")\n    \n    # GPU Memory usage\n    if torch.cuda.is_available():\n        current_mem = torch.cuda.memory_allocated() / 1e9\n        max_mem = torch.cuda.max_memory_allocated() / 1e9\n        print(f\"   💾 GPU Memory: {current_mem:.2f}GB / Peak: {max_mem:.2f}GB\")\n    \n    # Check for best model\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        patience_counter = 0\n        \n        # Save best model\n        best_model_path = f\"{config.WORK_DIR}/best_model_epoch_{epoch+1}.pth\"\n        save_checkpoint(model, optimizer, epoch+1, best_val_acc, best_model_path)\n        print(f\"   🎉 New best validation accuracy: {best_val_acc:.2f}%\")\n    else:\n        patience_counter += 1\n        print(f\"   ⏳ Patience: {patience_counter}/{config.PATIENCE}\")\n    \n    # Early stopping\n    if patience_counter >= config.PATIENCE:\n        print(f\"\\n🛑 Early stopping triggered after {epoch+1} epochs\")\n        print(f\"   Best validation accuracy: {best_val_acc:.2f}%\")\n        break\n    \n    # Memory cleanup\n    torch.cuda.empty_cache()\n\n# Training completed\ntotal_time = time.time() - start_time\nprint(\"\\n\" + \"=\" * 60)\nprint(\"🎉 TRAINING COMPLETED!\")\nprint(\"=\" * 60)\nprint(f\"⏱️ Total training time: {total_time/60:.2f} minutes\")\nprint(f\"🏆 Best validation accuracy: {best_val_acc:.2f}%\")\nprint(f\"📊 Epochs completed: {len(train_losses)}\")\n\n# Final model save\nfinal_model_path = f\"{config.WORK_DIR}/final_model.pth\"\nsave_checkpoint(model, optimizer, len(train_losses), best_val_acc, final_model_path)\n\n# Print training summary\nprint(\"\\n📈 TRAINING SUMMARY:\")\nprint(f\"   📊 Final train accuracy: {train_accuracies[-1]:.2f}%\")\nprint(f\"   📊 Final val accuracy: {val_accuracies[-1]:.2f}%\")\nprint(f\"   📊 Best val accuracy: {best_val_acc:.2f}%\")\nprint(f\"   📉 Final train loss: {train_losses[-1]:.4f}\")\nprint(f\"   📉 Final val loss: {val_losses[-1]:.4f}\")\n\n# Create training history for plotting\ntraining_history = {\n    'train_losses': train_losses,\n    'val_losses': val_losses,\n    'train_accuracies': train_accuracies,\n    'val_accuracies': val_accuracies,\n    'best_val_acc': best_val_acc,\n    'total_epochs': len(train_losses),\n    'total_time': total_time\n}\n\nprint(\"\\n✅ Training history saved to 'training_history' variable\")\nprint(\"🎯 You can now plot the results or make predictions!\")\n\n# Optional: Plot training curves\ntry:\n    import matplotlib.pyplot as plt\n    \n    print(\"\\n📊 Creating training plots...\")\n    \n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n    \n    # Loss plot\n    epochs_range = range(1, len(train_losses) + 1)\n    ax1.plot(epochs_range, train_losses, 'b-', label='Training Loss', linewidth=2)\n    ax1.plot(epochs_range, val_losses, 'r-', label='Validation Loss', linewidth=2)\n    ax1.set_title('Training and Validation Loss', fontsize=14, fontweight='bold')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Loss')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    # Accuracy plot\n    ax2.plot(epochs_range, train_accuracies, 'b-', label='Training Accuracy', linewidth=2)\n    ax2.plot(epochs_range, val_accuracies, 'r-', label='Validation Accuracy', linewidth=2)\n    ax2.set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Accuracy (%)')\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.savefig(f\"{config.WORK_DIR}/training_curves.png\", dpi=300, bbox_inches='tight')\n    plt.show()\n    \n    print(\"   📊 Training curves plotted and saved!\")\n    \nexcept ImportError:\n    print(\"   ⚠️ Matplotlib not available - skipping plots\")\nexcept Exception as e:\n    print(f\"   ⚠️ Could not create plots: {e}\")\n\nprint(\"\\n🎉 ALL DONE! Your model is trained and ready to use!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-21T11:12:26.483627Z","iopub.execute_input":"2025-06-21T11:12:26.484286Z","iopub.status.idle":"2025-06-21T11:30:22.342781Z","shell.execute_reply.started":"2025-06-21T11:12:26.484247Z","shell.execute_reply":"2025-06-21T11:30:22.341517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport os\nimport gc\n\ndef plot_training_history_fixed(train_losses, val_losses, train_accuracies, val_accuracies, learning_rates=None):\n    \"\"\"Plot training history with available data\"\"\"\n    if len(train_losses) == 0:\n        print(\"❌ No training history to plot!\")\n        return\n    \n    # Determine number of subplots\n    n_plots = 3 if learning_rates else 2\n    fig, axes = plt.subplots(1, n_plots, figsize=(6*n_plots, 5))\n    \n    if n_plots == 2:\n        axes = [axes[0], axes[1]]\n    \n    epochs = range(1, len(train_losses) + 1)\n    \n    # Loss plot\n    axes[0].plot(epochs, train_losses, 'b-o', label='Train Loss', linewidth=2, markersize=4)\n    axes[0].plot(epochs, val_losses, 'r-s', label='Val Loss', linewidth=2, markersize=4)\n    axes[0].set_title('Training and Validation Loss', fontsize=14, fontweight='bold')\n    axes[0].set_xlabel('Epoch')\n    axes[0].set_ylabel('Loss')\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n    \n    # Accuracy plot\n    axes[1].plot(epochs, train_accuracies, 'b-o', label='Train Acc', linewidth=2, markersize=4)\n    axes[1].plot(epochs, val_accuracies, 'r-s', label='Val Acc', linewidth=2, markersize=4)\n    axes[1].set_title('Training and Validation Accuracy', fontsize=14, fontweight='bold')\n    axes[1].set_xlabel('Epoch')\n    axes[1].set_ylabel('Accuracy (%)')\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n    \n    # Learning rate plot (if available)\n    if learning_rates and len(learning_rates) > 0:\n        axes[2].plot(epochs, learning_rates, 'g-^', label='Learning Rate', linewidth=2, markersize=4)\n        axes[2].set_title('Learning Rate Schedule', fontsize=14, fontweight='bold')\n        axes[2].set_xlabel('Epoch')\n        axes[2].set_ylabel('Learning Rate')\n        axes[2].legend()\n        axes[2].grid(True, alpha=0.3)\n        axes[2].set_yscale('log')  # Log scale for better visibility\n    \n    plt.tight_layout()\n    \n    # Save plot\n    save_path = os.path.join(config.WORK_DIR, 'training_curves.png')\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    print(f\"   📊 Training curves saved: {save_path}\")\n    plt.show()\n\ndef print_comprehensive_results():\n    \"\"\"Print comprehensive training results\"\"\"\n    print(\"\\n\" + \"=\"*70)\n    print(\"🎯 COMPREHENSIVE TRAINING RESULTS\")\n    print(\"=\"*70)\n    \n    # Check if we have training data\n    if 'train_losses' not in globals() or len(train_losses) == 0:\n        print(\"❌ No training history available!\")\n        return\n    \n    # Basic statistics\n    total_epochs = len(train_losses)\n    best_val_acc = max(val_accuracies)\n    best_epoch = val_accuracies.index(best_val_acc) + 1\n    final_train_acc = train_accuracies[-1]\n    final_val_acc = val_accuracies[-1]\n    \n    print(f\"📊 Dataset Information:\")\n    print(f\"   • Training samples: {len(train_loader.dataset) if 'train_loader' in globals() else 'N/A'}\")\n    print(f\"   • Validation samples: {len(val_loader.dataset) if 'val_loader' in globals() else 'N/A'}\")\n    print(f\"   • Number of classes: {config.NUM_CLASSES}\")\n    print(f\"   • Batch size: {config.BATCH_SIZE}\")\n    print(f\"   • Image size: {config.IMG_SIZE}x{config.IMG_SIZE}\")\n    \n    print(f\"\\n🏆 Best Performance:\")\n    print(f\"   • Best epoch: {best_epoch}\")\n    print(f\"   • Best validation accuracy: {best_val_acc:.2f}%\")\n    print(f\"   • Best train accuracy at that epoch: {train_accuracies[best_epoch-1]:.2f}%\")\n    \n    print(f\"\\n📈 Final Performance:\")\n    print(f\"   • Final train accuracy: {final_train_acc:.2f}%\")\n    print(f\"   • Final validation accuracy: {final_val_acc:.2f}%\")\n    print(f\"   • Total epochs completed: {total_epochs}\")\n    print(f\"   • Improvement from epoch 1: {final_val_acc - val_accuracies[0]:.2f}%\")\n    \n    print(f\"\\n📉 Loss Analysis:\")\n    print(f\"   • Initial train loss: {train_losses[0]:.4f}\")\n    print(f\"   • Final train loss: {train_losses[-1]:.4f}\")\n    print(f\"   • Initial val loss: {val_losses[0]:.4f}\")\n    print(f\"   • Final val loss: {val_losses[-1]:.4f}\")\n    print(f\"   • Best val loss: {min(val_losses):.4f}\")\n    \n    # Learning rate info\n    if 'learning_rates' in globals() and len(learning_rates) > 0:\n        print(f\"\\n📈 Learning Rate Schedule:\")\n        print(f\"   • Initial LR: {learning_rates[0]:.6f}\")\n        print(f\"   • Final LR: {learning_rates[-1]:.6f}\")\n        print(f\"   • LR reduction: {learning_rates[0]/learning_rates[-1]:.1f}x\")\n    \n    print(f\"\\n💾 Model Information:\")\n    print(f\"   • Architecture: {config.MODEL_NAME}\")\n    print(f\"   • Device: {config.DEVICE}\")\n    print(f\"   • Mixed precision: {config.USE_AMP}\")\n    \n    # Training time\n    if 'training_history' in globals() and 'total_time' in training_history:\n        total_time = training_history['total_time']\n        print(f\"\\n⏱️ Training Time:\")\n        print(f\"   • Total time: {total_time/60:.1f} minutes\")\n        print(f\"   • Average time per epoch: {total_time/total_epochs:.1f} seconds\")\n    \n    # Memory usage\n    if torch.cuda.is_available():\n        memory_allocated = torch.cuda.memory_allocated() / 1024**3\n        memory_reserved = torch.cuda.memory_reserved() / 1024**3\n        print(f\"\\n🔧 GPU Memory Usage:\")\n        print(f\"   • Currently allocated: {memory_allocated:.2f} GB\")\n        print(f\"   • Currently reserved: {memory_reserved:.2f} GB\")\n    \n    print(\"=\"*70)\n\ndef find_and_load_best_model():\n    \"\"\"Find and load the best saved model\"\"\"\n    try:\n        # Look for saved models\n        work_dir = config.WORK_DIR\n        model_files = []\n        \n        # Check for different model file patterns\n        for file in os.listdir(work_dir):\n            if file.endswith('.pth') and ('best' in file or 'final' in file):\n                model_files.append(file)\n        \n        if not model_files:\n            print(\"❌ No saved model files found!\")\n            return None\n        \n        # Use the most recent best model\n        best_model_file = None\n        for file in model_files:\n            if 'best' in file:\n                best_model_file = file\n                break\n        \n        if not best_model_file:\n            best_model_file = model_files[0]  # Use any available model\n        \n        model_path = os.path.join(work_dir, best_model_file)\n        print(f\"📁 Loading model: {best_model_file}\")\n        \n        # Load checkpoint\n        checkpoint = torch.load(model_path, map_location=config.DEVICE)\n        \n        # Create model\n        test_model = iNaturalistModel(config.NUM_CLASSES, config.MODEL_NAME)\n        test_model.load_state_dict(checkpoint['model_state_dict'])\n        test_model = test_model.to(config.DEVICE)\n        test_model.eval()\n        \n        print(f\"✅ Model loaded successfully!\")\n        if 'best_acc' in checkpoint:\n            print(f\"   • Best validation accuracy: {checkpoint['best_acc']:.2f}%\")\n        if 'epoch' in checkpoint:\n            print(f\"   • Saved at epoch: {checkpoint['epoch']}\")\n        \n        # Quick test with validation data\n        if 'val_loader' in globals() and val_loader is not None:\n            print(\"   🧪 Testing model with validation batch...\")\n            with torch.no_grad():\n                try:\n                    batch = next(iter(val_loader))\n                    images, labels = batch\n                    images = images[:8].to(config.DEVICE)  # Use only 8 samples\n                    labels = labels[:8].to(config.DEVICE)\n                    \n                    outputs = test_model(images)\n                    _, preds = torch.max(outputs, 1)\n                    \n                    accuracy = (preds == labels).float().mean()\n                    print(f\"   • Test batch accuracy: {accuracy*100:.2f}%\")\n                    print(f\"   • Sample predictions: {preds.cpu().numpy()[:5]}\")\n                    print(f\"   • Sample true labels: {labels.cpu().numpy()[:5]}\")\n                except Exception as e:\n                    print(f\"   ⚠️ Test batch failed: {e}\")\n        \n        del test_model\n        torch.cuda.empty_cache()\n        \n        return model_path\n    \n    except Exception as e:\n        print(f\"❌ Error loading model: {e}\")\n        return None\n\ndef create_summary_report():\n    \"\"\"Create a comprehensive summary report\"\"\"\n    print(\"\\n\" + \"🎉\"*25)\n    print(\"TRAINING PIPELINE SUMMARY\")\n    print(\"🎉\"*25)\n    \n    # Check what data we have\n    has_training_data = ('train_losses' in globals() and len(train_losses) > 0)\n    has_model = ('model' in globals())\n    \n    if has_training_data:\n        best_acc = max(val_accuracies)\n        total_epochs = len(train_losses)\n        \n        print(f\"\\n✅ Training Status: COMPLETED\")\n        print(f\"   • Epochs completed: {total_epochs}\")\n        print(f\"   • Best validation accuracy: {best_acc:.2f}%\")\n        \n        # Performance assessment\n        if best_acc >= 70:\n            print(f\"   🏆 Performance: EXCELLENT!\")\n        elif best_acc >= 50:\n            print(f\"   👍 Performance: GOOD\")\n        elif best_acc >= 30:\n            print(f\"   👌 Performance: FAIR\")\n        else:\n            print(f\"   📈 Performance: NEEDS IMPROVEMENT\")\n    \n    if has_model:\n        print(f\"\\n✅ Model Status: READY\")\n        print(f\"   • Architecture: {type(model).__name__}\")\n        print(f\"   • Classes: {config.NUM_CLASSES}\")\n        print(f\"   • Device: {config.DEVICE}\")\n    \n    print(f\"\\n📁 Files Generated:\")\n    work_dir = config.WORK_DIR\n    generated_files = []\n    for file in os.listdir(work_dir):\n        if file.endswith(('.pth', '.png')):\n            generated_files.append(file)\n    \n    for file in generated_files:\n        print(f\"   • {file}\")\n    \n    print(f\"\\n🚀 Next Steps:\")\n    print(f\"   1. Download model files for deployment\")\n    print(f\"   2. Test model on new images\")\n    print(f\"   3. Consider fine-tuning with more data\")\n    print(f\"   4. Deploy for inference\")\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\nprint(\"🔍 Analyzing training results...\")\n\n# Check what training data we have available\nif 'train_losses' in globals() and len(train_losses) > 0:\n    print(\"✅ Training data found - creating visualizations...\")\n    \n    # Create plots\n    lr_data = learning_rates if 'learning_rates' in globals() else None\n    plot_training_history_fixed(train_losses, val_losses, train_accuracies, val_accuracies, lr_data)\n    \n    # Print comprehensive results\n    print_comprehensive_results()\n    \n    # Try to load and test the best model\n    print(\"\\n🔍 Looking for saved models...\")\n    model_path = find_and_load_best_model()\n    \n    if model_path:\n        print(f\"✅ Model validation completed!\")\n    else:\n        print(\"⚠️ Could not validate saved model\")\n    \n    # Create final summary\n    create_summary_report()\n    \nelif 'training_history' in globals():\n    print(\"✅ Found training_history variable...\")\n    # Extract data from training_history\n    train_losses = training_history.get('train_losses', [])\n    val_losses = training_history.get('val_losses', [])\n    train_accuracies = training_history.get('train_accuracies', [])\n    val_accuracies = training_history.get('val_accuracies', [])\n    learning_rates = training_history.get('learning_rates', [])\n    \n    if len(train_losses) > 0:\n        plot_training_history_fixed(train_losses, val_losses, train_accuracies, val_accuracies, learning_rates)\n        print_comprehensive_results()\n        model_path = find_and_load_best_model()\n        create_summary_report()\n    else:\n        print(\"❌ Training history is empty!\")\n        \nelse:\n    print(\"❌ No training history found!\")\n    print(\"Possible reasons:\")\n    print(\"1. Training hasn't been completed yet\")\n    print(\"2. Training variables were cleared\")\n    print(\"3. There was an error during training\")\n    \n    # Try to find any saved models anyway\n    print(\"\\n🔍 Checking for saved models...\")\n    try:\n        work_dir = config.WORK_DIR if 'config' in globals() else '/kaggle/working'\n        files = os.listdir(work_dir)\n        model_files = [f for f in files if f.endswith('.pth')]\n        \n        if model_files:\n            print(f\"✅ Found {len(model_files)} model file(s):\")\n            for f in model_files:\n                print(f\"   • {f}\")\n        else:\n            print(\"❌ No model files found\")\n    except:\n        print(\"❌ Could not check for model files\")\n\n# Final cleanup\nprint(\"\\n🧹 Performing final cleanup...\")\ntorch.cuda.empty_cache()\ngc.collect()\nprint(\"✅ Cleanup completed!\")\nprint(\"\\n\" + \"=\"*50)\nprint(\"🎉 ANALYSIS COMPLETED!\")\nprint(\"=\"*50)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torchvision import transforms\nfrom PIL import Image\nimport os\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport json\nfrom pathlib import Path\nimport random\n\nprint(\"🧪 TESTING TRAINED MODEL WITH SAMPLE IMAGES\")\nprint(\"=\" * 60)\n\n# ============================================================================\n# MODEL LOADING AND SETUP\n# ============================================================================\n\ndef load_trained_model():\n    \"\"\"Load the best trained model\"\"\"\n    print(\"📥 Loading trained model...\")\n    \n    # Find the best model file\n    model_files = []\n    for file in os.listdir(config.WORK_DIR):\n        if file.endswith('.pth') and ('best' in file.lower() or 'final' in file.lower()):\n            model_files.append(file)\n    \n    if not model_files:\n        print(\"❌ No trained model found!\")\n        print(\"Available .pth files:\")\n        for file in os.listdir(config.WORK_DIR):\n            if file.endswith('.pth'):\n                print(f\"   📄 {file}\")\n        return None, None\n    \n    # Use the best model\n    model_file = None\n    for file in model_files:\n        if 'best' in file.lower():\n            model_file = file\n            break\n    \n    if not model_file:\n        model_file = model_files[0]\n    \n    model_path = os.path.join(config.WORK_DIR, model_file)\n    print(f\"📁 Loading: {model_file}\")\n    \n    try:\n        # Load checkpoint\n        checkpoint = torch.load(model_path, map_location=config.DEVICE)\n        \n        # Create model\n        model = iNaturalistModel(config.NUM_CLASSES, config.MODEL_NAME)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        model = model.to(config.DEVICE)\n        model.eval()\n        \n        # Get metadata\n        best_acc = checkpoint.get('best_acc', 'Unknown')\n        epoch = checkpoint.get('epoch', 'Unknown')\n        \n        print(f\"✅ Model loaded successfully!\")\n        print(f\"   🎯 Best accuracy: {best_acc}\")\n        print(f\"   📊 Saved at epoch: {epoch}\")\n        print(f\"   🔧 Device: {config.DEVICE}\")\n        print(f\"   📂 Classes: {config.NUM_CLASSES}\")\n        \n        return model, checkpoint\n        \n    except Exception as e:\n        print(f\"❌ Error loading model: {e}\")\n        return None, None\n\ndef get_class_mapping():\n    \"\"\"Get class index to species name mapping\"\"\"\n    print(\"🏷️ Getting class mapping...\")\n    \n    # Try to get from training data\n    if 'train_df' in globals() and train_df is not None:\n        if 'species_name' in train_df.columns:\n            # Create mapping from training data\n            class_mapping = {}\n            for _, row in train_df.iterrows():\n                label = row['label']\n                species_name = row['species_name']\n                if label not in class_mapping:\n                    class_mapping[label] = species_name\n            \n            print(f\"✅ Created class mapping from training data: {len(class_mapping)} classes\")\n            return class_mapping\n    \n    # Fallback: Create generic mapping\n    class_mapping = {i: f\"Species_{i:03d}\" for i in range(config.NUM_CLASSES)}\n    print(f\"⚠️ Using generic class mapping: {len(class_mapping)} classes\")\n    \n    return class_mapping\n\n# ============================================================================\n# IMAGE PREPROCESSING\n# ============================================================================\n\ndef get_test_transform():\n    \"\"\"Get preprocessing transform for test images\"\"\"\n    return transforms.Compose([\n        transforms.Resize((config.IMG_SIZE, config.IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n    ])\n\ndef preprocess_image(image_path, transform):\n    \"\"\"Preprocess a single image for prediction\"\"\"\n    try:\n        # Load and convert image\n        image = Image.open(image_path).convert('RGB')\n        \n        # Apply transforms\n        image_tensor = transform(image)\n        \n        # Add batch dimension\n        image_batch = image_tensor.unsqueeze(0)\n        \n        return image_batch, image\n        \n    except Exception as e:\n        print(f\"❌ Error preprocessing {image_path}: {e}\")\n        return None, None\n\n# ============================================================================\n# PREDICTION FUNCTIONS\n# ============================================================================\n\ndef predict_image(model, image_tensor, class_mapping, top_k=5):\n    \"\"\"Predict class for a single image\"\"\"\n    with torch.no_grad():\n        # Move to device\n        image_tensor = image_tensor.to(config.DEVICE)\n        \n        # Forward pass\n        outputs = model(image_tensor)\n        \n        # Get probabilities\n        probabilities = torch.nn.functional.softmax(outputs, dim=1)\n        \n        # Get top-k predictions\n        top_probs, top_indices = torch.topk(probabilities, top_k, dim=1)\n        \n        # Convert to lists\n        top_probs = top_probs.cpu().numpy()[0]\n        top_indices = top_indices.cpu().numpy()[0]\n        \n        # Create predictions list\n        predictions = []\n        for i in range(top_k):\n            class_idx = top_indices[i]\n            confidence = top_probs[i]\n            species_name = class_mapping.get(class_idx, f\"Unknown_{class_idx}\")\n            \n            predictions.append({\n                'class_idx': class_idx,\n                'species_name': species_name,\n                'confidence': confidence\n            })\n        \n        return predictions\n\ndef batch_predict_images(model, image_paths, transform, class_mapping, top_k=3):\n    \"\"\"Predict classes for multiple images\"\"\"\n    print(f\"🔍 Predicting {len(image_paths)} images...\")\n    \n    results = []\n    \n    for i, img_path in enumerate(image_paths):\n        print(f\"   Processing {i+1}/{len(image_paths)}: {os.path.basename(img_path)}\")\n        \n        # Preprocess image\n        image_tensor, original_image = preprocess_image(img_path, transform)\n        \n        if image_tensor is not None:\n            # Make prediction\n            predictions = predict_image(model, image_tensor, class_mapping, top_k)\n            \n            results.append({\n                'image_path': img_path,\n                'original_image': original_image,\n                'predictions': predictions\n            })\n        else:\n            print(f\"   ❌ Failed to process {img_path}\")\n    \n    return results\n\n# ============================================================================\n# VISUALIZATION\n# ============================================================================\n\ndef plot_predictions(results, max_images=8):\n    \"\"\"Plot images with their predictions\"\"\"\n    print(f\"📊 Plotting predictions for {min(len(results), max_images)} random images...\")\n    \n    n_images = min(len(results), max_images)\n    cols = 4  # More columns for better layout\n    rows = (n_images + cols - 1) // cols\n    \n    fig, axes = plt.subplots(rows, cols, figsize=(20, 5*rows))\n    if rows == 1:\n        axes = [axes] if n_images == 1 else axes\n    else:\n        axes = axes.flatten()\n    \n    for i in range(n_images):\n        result = results[i]\n        image = result['original_image']\n        predictions = result['predictions']\n        image_name = os.path.basename(result['image_path'])\n        \n        # Plot image\n        axes[i].imshow(image)\n        axes[i].set_title(f\"Random Image {i+1}\\n{image_name[:20]}...\", fontsize=9, fontweight='bold')\n        axes[i].axis('off')\n        \n        # Add prediction text with better formatting\n        top_pred = predictions[0]\n        species = top_pred['species_name']\n        conf = top_pred['confidence']\n        \n        # Truncate very long species names\n        if len(species) > 25:\n            species = species[:22] + \"...\"\n        \n        pred_text = f\"🏆 {species}\\n{conf*100:.1f}% confidence\\n\\n\"\n        \n        # Add runner-ups\n        for j, pred in enumerate(predictions[1:3]):\n            runner_species = pred['species_name']\n            runner_conf = pred['confidence']\n            if len(runner_species) > 20:\n                runner_species = runner_species[:17] + \"...\"\n            pred_text += f\"{j+2}. {runner_species}\\n   {runner_conf*100:.1f}%\\n\"\n        \n        # Add text box with better styling\n        axes[i].text(0.02, 0.98, pred_text, \n                    transform=axes[i].transAxes,\n                    verticalalignment='top',\n                    bbox=dict(boxstyle='round,pad=0.5', facecolor='lightblue', alpha=0.9),\n                    fontsize=8,\n                    fontweight='bold' if conf > 0.7 else 'normal')\n    \n    # Hide empty subplots\n    for i in range(n_images, len(axes)):\n        axes[i].axis('off')\n    \n    plt.tight_layout()\n    \n    # Save plot\n    save_path = os.path.join(config.WORK_DIR, 'random_image_predictions.png')\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"💾 Random predictions plot saved: {save_path}\")\n    \n    plt.show()\n\ndef print_detailed_results(results):\n    \"\"\"Print detailed prediction results\"\"\"\n    print(\"\\n📋 DETAILED PREDICTION RESULTS\")\n    print(\"=\" * 60)\n    \n    for i, result in enumerate(results):\n        image_name = os.path.basename(result['image_path'])\n        predictions = result['predictions']\n        \n        print(f\"\\n🖼️ Image {i+1}: {image_name}\")\n        print(\"-\" * 40)\n        \n        for j, pred in enumerate(predictions):\n            species = pred['species_name']\n            conf = pred['confidence']\n            class_idx = pred['class_idx']\n            \n            confidence_bar = \"█\" * int(conf * 20) + \"░\" * (20 - int(conf * 20))\n            \n            print(f\"   {j+1}. {species}\")\n            print(f\"      Class: {class_idx} | Confidence: {conf*100:.2f}%\")\n            print(f\"      [{confidence_bar}]\")\n\n# ============================================================================\n# MAIN TESTING FUNCTIONS\n# ============================================================================\n\ndef get_test_images(num_images=6):\n    \"\"\"Get random sample images from extracted_images for testing\"\"\"\n    print(\"🔍 Finding random test images from extracted_images...\")\n    \n    # Look directly in Kaggle extracted_images directory\n    extract_dir = os.path.join(config.WORK_DIR, \"extracted_images\")\n    \n    if not os.path.exists(extract_dir):\n        print(f\"❌ Extracted images directory not found: {extract_dir}\")\n        return []\n    \n    print(f\"📁 Scanning directory: {extract_dir}\")\n    \n    # Find ALL images in extracted directory\n    all_images = []\n    for root, dirs, files in os.walk(extract_dir):\n        for file in files:\n            if file.lower().endswith(('.jpg', '.jpeg', '.png')):\n                full_path = os.path.join(root, file)\n                all_images.append(full_path)\n    \n    print(f\"📊 Found {len(all_images)} total images in extracted directory\")\n    \n    if len(all_images) == 0:\n        print(\"❌ No images found in extracted directory!\")\n        return []\n    \n    # Random sample from ALL extracted images\n    random.seed(None)  # Use current time for true randomness\n    test_images = random.sample(all_images, min(num_images, len(all_images)))\n    \n    print(f\"✅ Randomly selected {len(test_images)} test images:\")\n    for i, img_path in enumerate(test_images):\n        filename = os.path.basename(img_path)\n        # Get file size\n        try:\n            file_size = os.path.getsize(img_path) / 1024  # KB\n            print(f\"   {i+1}. {filename} ({file_size:.1f} KB)\")\n        except:\n            print(f\"   {i+1}. {filename}\")\n    \n    return test_images\n\ndef test_model_performance():\n    \"\"\"Test model on a batch from validation set\"\"\"\n    print(\"📊 Testing model performance on validation batch...\")\n    \n    if 'val_loader' not in globals() or val_loader is None:\n        print(\"❌ No validation loader available!\")\n        return\n    \n    try:\n        # Get a batch\n        batch = next(iter(val_loader))\n        images, labels = batch\n        images = images.to(config.DEVICE)\n        labels = labels.to(config.DEVICE)\n        \n        # Make predictions\n        with torch.no_grad():\n            outputs = model(images)\n            _, preds = torch.max(outputs, 1)\n            \n            # Calculate accuracy\n            accuracy = (preds == labels).float().mean()\n            \n            print(f\"✅ Batch test completed!\")\n            print(f\"   📊 Batch size: {len(labels)}\")\n            print(f\"   🎯 Accuracy: {accuracy*100:.2f}%\")\n            print(f\"   📈 Correct predictions: {(preds == labels).sum().item()}/{len(labels)}\")\n            \n            # Show some predictions vs ground truth\n            print(f\"\\n📋 Sample predictions vs ground truth:\")\n            for i in range(min(5, len(labels))):\n                pred_class = preds[i].item()\n                true_class = labels[i].item()\n                match = \"✅\" if pred_class == true_class else \"❌\"\n                print(f\"   {i+1}. Predicted: {pred_class}, True: {true_class} {match}\")\n    \n    except Exception as e:\n        print(f\"❌ Error testing model performance: {e}\")\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\n# Load the trained model\nmodel, checkpoint = load_trained_model()\n\nif model is not None:\n    # Get class mapping\n    class_mapping = get_class_mapping()\n    \n    # Get test transform\n    test_transform = get_test_transform()\n    \n    # Test model performance first\n    test_model_performance()\n    \n    # Get random test images from extracted_images directory\n    test_image_paths = get_test_images(num_images=8)  # Increased to 8 for better variety\n    \n    if test_image_paths:\n        print(f\"\\n🧪 Testing model with {len(test_image_paths)} sample images...\")\n        \n        # Make predictions\n        results = batch_predict_images(model, test_image_paths, test_transform, class_mapping)\n        \n        if results:\n            # Print detailed results\n            print_detailed_results(results)\n            \n            # Plot results\n            plot_predictions(results)\n            \n            print(\"\\n🎉 RANDOM IMAGE TESTING COMPLETED!\")\n            print(\"=\" * 50)\n            print(f\"✅ Tested {len(results)} random images from extracted_images\")\n            print(f\"📊 Results saved to random_image_predictions.png\")\n            print(f\"🎯 Model predictions on truly random samples!\")\n            print(f\"🔍 This shows how well the model generalizes to unseen data\")\n            \n        else:\n            print(\"❌ No prediction results generated!\")\n    \n    else:\n        print(\"❌ No test images found!\")\n        print(\"Make sure you have:\")\n        print(\"1. Validation dataset (val_df)\")\n        print(\"2. Extracted images directory\")\n\nelse:\n    print(\"❌ Cannot test model - no trained model loaded!\")\n    print(\"Make sure you have:\")\n    print(\"1. Completed training\")\n    print(\"2. Saved model file (best_model.pth)\")\n\n# ============================================================================\n# OPTIONAL: INTERACTIVE TESTING\n# ============================================================================\n\ndef test_single_image(image_path):\n    \"\"\"Test a single specific image\"\"\"\n    print(f\"🖼️ Testing single image: {image_path}\")\n    \n    if not os.path.exists(image_path):\n        print(f\"❌ Image not found: {image_path}\")\n        return\n    \n    # Preprocess\n    image_tensor, original_image = preprocess_image(image_path, test_transform)\n    \n    if image_tensor is not None:\n        # Predict\n        predictions = predict_image(model, image_tensor, class_mapping, top_k=5)\n        \n        # Show results\n        print(f\"\\n📋 Predictions for {os.path.basename(image_path)}:\")\n        for i, pred in enumerate(predictions):\n            species = pred['species_name']\n            conf = pred['confidence']\n            print(f\"   {i+1}. {species}: {conf*100:.2f}%\")\n        \n        # Simple plot\n        plt.figure(figsize=(8, 6))\n        plt.imshow(original_image)\n        plt.title(f\"Top Prediction: {predictions[0]['species_name']} ({predictions[0]['confidence']*100:.1f}%)\")\n        plt.axis('off')\n        plt.show()\n\nprint(f\"\\n💡 TO TEST SPECIFIC RANDOM IMAGES:\")\nprint(f\"# Test 5 random images:\")\nprint(f\"test_image_paths = get_test_images(5)\")\nprint(f\"results = batch_predict_images(model, test_image_paths, test_transform, class_mapping)\")\nprint(f\"plot_predictions(results)\")\nprint(f\"\")\nprint(f\"# Test one specific image:\")\nprint(f\"test_single_image('/kaggle/working/extracted_images/path/to/image.jpg')\")\n\n# Memory cleanup\ntorch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}