{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":132732,"databundleVersionId":16583342}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nfrom PIL import Image, ImageFilter, ImageEnhance\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torchvision import transforms\nimport torchvision.transforms.functional as TF\nfrom tqdm import tqdm\nimport timm\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\nimport random\nimport io\nimport warnings\nwarnings.filterwarnings('ignore')\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-19T12:28:08.773172Z","iopub.execute_input":"2026-04-19T12:28:08.773480Z","iopub.status.idle":"2026-04-19T12:28:22.437206Z","shell.execute_reply.started":"2026-04-19T12:28:08.773452Z","shell.execute_reply":"2026-04-19T12:28:22.436571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  CONFIGURATION \nclass Config:\n    # Paths\n    TRAIN_PATH = \"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data/Training\"\n    TEST_PATH = \"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data/Test\"\n    TRAIN_CSV = \"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data/training.csv\"\n    TEST_CSV = \"/kaggle/input/competitions/dlmmdd-workshop-synthetic-source-attribution-challenge/Data/Data/test.csv\"\n    \n    # Model settings\n    PRIMARY_MODEL = 'efficientnet_b4'\n    IMG_SIZE = 384\n    BATCH_SIZE = 32\n    EPOCHS = 30\n    LEARNING_RATE = 1e-4\n    NUM_CLASSES = 10\n    MIXUP_ALPHA = 0.2\n    LABEL_SMOOTHING = 0.1\n    \n    # Augmentation intensity\n    AUG_INTENSITY = 'medium'\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:05:24.715540Z","iopub.execute_input":"2026-04-19T09:05:24.715808Z","iopub.status.idle":"2026-04-19T09:05:24.721315Z","shell.execute_reply.started":"2026-04-19T09:05:24.715787Z","shell.execute_reply":"2026-04-19T09:05:24.720530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n#  PART 1: DATA LOADING AND EXPLORATION \nprint(\" DATA LOADING AND EXPLORATION\")\n\n# Load data\nprint(\"\\n Loading CSV files...\")\ntrain_df = pd.read_csv(Config.TRAIN_CSV)\ntest_df = pd.read_csv(Config.TEST_CSV)\n\n# Add full paths\ntrain_df['full_path'] = train_df['path'].apply(lambda x: os.path.join(Config.TRAIN_PATH, os.path.basename(x)))\ntest_df['full_path'] = test_df['path'].apply(lambda x: os.path.join(Config.TEST_PATH, os.path.basename(x)))\n\nprint(f\"Training samples: {len(train_df)}\")\nprint(f\"Test samples: {len(test_df)}\")\n\n# Display class distribution\nprint(f\"\\n Class Distribution:\")\nclass_counts = train_df['y'].value_counts().sort_index()\nfor i in range(10):\n    count = class_counts[i]\n    print(f\"  Class {i}: {count} images ({count/len(train_df)*100:.1f}%)\")\n\n# Display first few rows\nprint(f\"\\nTraining CSV Sample:\")\nprint(train_df.head(10))\n\n# Display test CSV sample\nprint(f\"\\n Test CSV Sample:\")\nprint(test_df.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:05:29.098471Z","iopub.execute_input":"2026-04-19T09:05:29.098956Z","iopub.status.idle":"2026-04-19T09:05:29.137138Z","shell.execute_reply.started":"2026-04-19T09:05:29.098929Z","shell.execute_reply":"2026-04-19T09:05:29.136408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  POST-PROCESSING SIMULATION DEMO (FIXED INDEXING) \nprint(\"PART 2: POST-PROCESSING SIMULATION DEMONSTRATION\")\n\nclass PostProcessingSimulator:\n    \"\"\"Simulate the exact post-processing operations from the challenge\"\"\"\n    \n    def __init__(self, img_size=384):\n        self.img_size = img_size\n        \n    def apply_jpeg_compression(self, img, quality=None):\n        \"\"\"JPEG compression simulation\"\"\"\n        if quality is None:\n            quality = np.random.randint(70, 95)\n        buffer = io.BytesIO()\n        img.save(buffer, format='JPEG', quality=quality)\n        buffer.seek(0)\n        return Image.open(buffer).convert('RGB')\n    \n    def apply_webp_compression(self, img, quality=None):\n        \"\"\"WebP compression\"\"\"\n        if quality is None:\n            quality = np.random.randint(70, 95)\n        buffer = io.BytesIO()\n        img.save(buffer, format='WEBP', quality=quality)\n        buffer.seek(0)\n        return Image.open(buffer).convert('RGB')\n    \n    def apply_random_crop(self, img, crop_ratio=(0.7, 0.9)):\n        \"\"\"Random central crop with resize back\"\"\"\n        crop_factor = np.random.uniform(*crop_ratio)\n        new_size = (int(img.width * crop_factor), int(img.height * crop_factor))\n        left = (img.width - new_size[0]) // 2\n        top = (img.height - new_size[1]) // 2\n        img_cropped = img.crop((left, top, left + new_size[0], top + new_size[1]))\n        return img_cropped.resize((self.img_size, self.img_size), Image.Resampling.LANCZOS)\n    \n    def apply_rotation_crop(self, img, angle=None):\n        \"\"\"Rotation with aspect-preserving crop\"\"\"\n        if angle is None:\n            angle = np.random.uniform(-15, 15)\n        img_rotated = img.rotate(angle, expand=True, fillcolor=0)\n        center_x, center_y = img_rotated.width // 2, img_rotated.height // 2\n        left = center_x - self.img_size // 2\n        top = center_y - self.img_size // 2\n        return img_rotated.crop((left, top, left + self.img_size, top + self.img_size))\n    \n    def apply_resizing(self, img, scale=None):\n        \"\"\"Random resizing\"\"\"\n        if scale is None:\n            scale = np.random.uniform(0.5, 1.5)\n        new_size = (int(img.width * scale), int(img.height * scale))\n        img_resized = img.resize(new_size, Image.Resampling.LANCZOS)\n        return img_resized.resize((self.img_size, self.img_size), Image.Resampling.LANCZOS)\n    \n    def apply_blur(self, img, radius=None):\n        \"\"\"Gaussian blur\"\"\"\n        if radius is None:\n            radius = np.random.uniform(0.5, 2.0)\n        return img.filter(ImageFilter.GaussianBlur(radius=radius))\n    \n    def apply_grayscale(self, img):\n        \"\"\"Convert to grayscale and back to RGB\"\"\"\n        gray = img.convert('L')\n        return gray.convert('RGB')\n    \n    def apply_brightness_contrast(self, img, brightness=None, contrast=None):\n        \"\"\"Adjust brightness and contrast\"\"\"\n        if brightness is None:\n            brightness = np.random.uniform(0.7, 1.3)\n        if contrast is None:\n            contrast = np.random.uniform(0.7, 1.3)\n        \n        enhancer = ImageEnhance.Brightness(img)\n        img = enhancer.enhance(brightness)\n        enhancer = ImageEnhance.Contrast(img)\n        return enhancer.enhance(contrast)\n    \n    def apply_super_resolution(self, img):\n        \"\"\"Simple super-resolution simulation\"\"\"\n        small = img.resize((img.width // 2, img.height // 2), Image.Resampling.BICUBIC)\n        return small.resize((self.img_size, self.img_size), Image.Resampling.BICUBIC)\n    \n    def __call__(self, img, intensity='medium', num_ops=None):\n        \"\"\"Apply 1-3 random post-processing operations\"\"\"\n        operations = [\n            ('jpeg', lambda x: self.apply_jpeg_compression(x, quality=np.random.randint(75, 95))),\n            ('webp', lambda x: self.apply_webp_compression(x, quality=np.random.randint(75, 95))),\n            ('crop', lambda x: self.apply_random_crop(x, crop_ratio=(0.75, 0.9))),\n            ('resize', lambda x: self.apply_resizing(x, scale=np.random.uniform(0.7, 1.3))),\n            ('rotate', lambda x: self.apply_rotation_crop(x, angle=np.random.uniform(-10, 10))),\n            ('blur', lambda x: self.apply_blur(x, radius=np.random.uniform(0.5, 1.5))),\n            ('brightness', lambda x: self.apply_brightness_contrast(x, \n                brightness=np.random.uniform(0.8, 1.2),\n                contrast=np.random.uniform(0.8, 1.2))),\n            ('grayscale', lambda x: self.apply_grayscale(x)),\n            ('superres', lambda x: self.apply_super_resolution(x)),\n        ]\n        \n        if intensity == 'light':\n            num_ops = num_ops or np.random.randint(1, 2)\n            op_indices = np.random.choice(len(operations), num_ops, replace=False)\n        elif intensity == 'heavy':\n            num_ops = num_ops or np.random.randint(2, 4)\n            op_indices = np.random.choice(len(operations), num_ops, replace=False)\n        else:\n            num_ops = num_ops or np.random.randint(1, 3)\n            op_indices = np.random.choice(len(operations), num_ops, replace=False)\n        \n        for idx in op_indices:\n            _, op_func = operations[idx]\n            img = op_func(img)\n        \n        return img\n\n# Load a sample image\nsample_image_path = train_df['full_path'].iloc[0]\nsample_img = Image.open(sample_image_path).convert('RGB')\nprint(f\"\\n Sample image: {os.path.basename(sample_image_path)}\")\nprint(f\"   Original size: {sample_img.size}\")\n\n# Create a figure to compare original vs extreme versions\nprint(\"\\n Comparing original with EXTREME post-processing effects...\")\nfig, axes = plt.subplots(2, 4, figsize=(16, 8))\nsimulator = PostProcessingSimulator(Config.IMG_SIZE)\n\n# Original (column 0)\naxes[0, 0].imshow(sample_img)\naxes[0, 0].set_title('Original Image', fontsize=12, fontweight='bold')\naxes[0, 0].axis('off')\n\n# Extreme versions for row 0 (columns 1,2,3)\nextreme_ops_row0 = [\n    ('Low Quality JPEG\\n(quality=30)', simulator.apply_jpeg_compression(sample_img, quality=30)),\n    ('Low Quality WebP\\n(quality=30)', simulator.apply_webp_compression(sample_img, quality=30)),\n    ('Strong Blur\\n(radius=5)', simulator.apply_blur(sample_img, radius=5)),\n]\n\nfor i, (title, processed) in enumerate(extreme_ops_row0):\n    axes[0, i+1].imshow(processed)\n    axes[0, i+1].set_title(title, fontsize=10)\n    axes[0, i+1].axis('off')\n\n# Extreme versions for row 1 (columns 0,1,2,3)\nextreme_ops_row1 = [\n    ('Extreme Crop\\n(50% crop)', simulator.apply_random_crop(sample_img, crop_ratio=(0.5, 0.5))),\n    ('Extreme Rotation\\n(angle=30°)', simulator.apply_rotation_crop(sample_img, angle=30)),\n    ('Extreme Resizing\\n(scale=0.3x)', simulator.apply_resizing(sample_img, scale=0.3)),\n    ('Extreme Brightness\\n(brightness=0.4, contrast=1.5)', \n     simulator.apply_brightness_contrast(sample_img, brightness=0.4, contrast=1.5)),\n]\n\nfor i, (title, processed) in enumerate(extreme_ops_row1):\n    axes[1, i].imshow(processed)\n    axes[1, i].set_title(title, fontsize=10)\n    axes[1, i].axis('off')\n\nplt.suptitle('EXTREME Post-Processing Effects (To Better See the Differences)', fontsize=14, fontweight='bold', y=1.02)\nplt.tight_layout()\nplt.show()\n\n# Now show ACTUAL challenge parameters (more subtle)\nprint(\"\\n Now showing ACTUAL challenge parameters (subtler, as in competition)...\")\nfig2, axes2 = plt.subplots(2, 4, figsize=(16, 8))\n\n# Original\naxes2[0, 0].imshow(sample_img)\naxes2[0, 0].set_title('Original', fontsize=12, fontweight='bold')\naxes2[0, 0].axis('off')\n\n# Challenge-typical operations\naxes2[0, 1].imshow(simulator.apply_jpeg_compression(sample_img, quality=85))\naxes2[0, 1].set_title('JPEG (q=85)\\nTypical challenge', fontsize=10)\naxes2[0, 1].axis('off')\n\naxes2[0, 2].imshow(simulator.apply_webp_compression(sample_img, quality=85))\naxes2[0, 2].set_title('WebP (q=85)\\nTypical challenge', fontsize=10)\naxes2[0, 2].axis('off')\n\naxes2[0, 3].imshow(simulator.apply_blur(sample_img, radius=1.5))\naxes2[0, 3].set_title('Gaussian Blur (r=1.5)\\nTypical challenge', fontsize=10)\naxes2[0, 3].axis('off')\n\n# Row 2 operations\naxes2[1, 0].imshow(simulator.apply_grayscale(sample_img))\naxes2[1, 0].set_title('Grayscale\\nTypical challenge', fontsize=10)\naxes2[1, 0].axis('off')\n\naxes2[1, 1].imshow(simulator.apply_brightness_contrast(sample_img, brightness=0.85, contrast=1.15))\naxes2[1, 1].set_title('Brightness/Contrast\\nTypical challenge', fontsize=10)\naxes2[1, 1].axis('off')\n\n# Combined operations\ncombined_op = simulator(sample_img, intensity='medium', num_ops=2)\naxes2[1, 2].imshow(combined_op)\naxes2[1, 2].set_title('Combined Operations\\n(2 random ops)', fontsize=10)\naxes2[1, 2].axis('off')\n\n# Show difference map\ndiff = np.array(sample_img.resize((384, 384))) - np.array(combined_op.resize((384, 384)))\ndiff = np.abs(diff)\naxes2[1, 3].imshow(diff / diff.max(), cmap='hot')\naxes2[1, 3].set_title('Difference Map\\n(Brighter = More Change)', fontsize=10)\naxes2[1, 3].axis('off')\n\nplt.suptitle('ACTUAL Challenge Parameters (Subtle but Detectable)', fontsize=14, fontweight='bold', y=1.02)\nplt.tight_layout()\nplt.show()\n\n# Calculate and display statistics\nprint(\"\\n Image Difference Statistics (Challenge Parameters):\")\noriginal_array = np.array(sample_img.resize((384, 384)))\nprocessed_array = np.array(combined_op.resize((384, 384)))\n\ndiff_percentage = np.abs(original_array - processed_array).mean() / 255 * 100\nprint(f\"   Average pixel difference: {diff_percentage:.2f}%\")\nprint(f\"   Max pixel difference: {np.abs(original_array - processed_array).max() / 255 * 100:.2f}%\")\nprint(f\"   Standard deviation of difference: {np.std(np.abs(original_array - processed_array)) / 255 * 100:.2f}%\")\n\nprint(\"\\n Key Insight: The changes are subtle (typically 5-15% pixel difference)\")\nprint(\"   This makes the challenge difficult - models must learn robust features!\")\nprint(\"\\n What models learn to detect:\")\nprint(\"   • JPEG/WebP: Block artifacts and quantization patterns\")\nprint(\"   • Blur: High-frequency component reduction\")\nprint(\"   • Grayscale: Loss of color distribution information\")\nprint(\"   • Brightness/Contrast: Pixel value distribution shifts\")\nprint(\"   • Combined ops: Multiple forensic traces\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:04:34.759161Z","iopub.execute_input":"2026-04-19T09:04:34.759905Z","iopub.status.idle":"2026-04-19T09:04:37.770727Z","shell.execute_reply.started":"2026-04-19T09:04:34.759876Z","shell.execute_reply":"2026-04-19T09:04:37.770042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  DATA AUGMENTATION PIPELINE \nprint(\"DATA AUGMENTATION PIPELINE\")\n\nclass TrainTransform:\n    def __init__(self, img_size=384, intensity='medium'):\n        self.img_size = img_size\n        self.post_processor = PostProcessingSimulator(img_size)\n        self.intensity = intensity\n        self.normalize = transforms.Normalize(\n            mean=[0.485, 0.456, 0.406], \n            std=[0.229, 0.224, 0.225]\n        )\n    \n    def __call__(self, image):\n        # Step 1: Apply challenge post-processing (CRITICAL!)\n        image = self.post_processor(image, intensity=self.intensity)\n        \n        # Step 2: Convert to tensor\n        image = TF.to_tensor(image)\n        image = TF.resize(image, [self.img_size, self.img_size])\n        \n        # Step 3: Random horizontal flip (50% chance)\n        if random.random() < 0.5:\n            image = TF.hflip(image)\n        \n        # Step 4: Color jitter (30% chance)\n        if random.random() < 0.3:\n            brightness = random.uniform(0.9, 1.1)\n            contrast = random.uniform(0.9, 1.1)\n            saturation = random.uniform(0.9, 1.1)\n            image = TF.adjust_brightness(image, brightness)\n            image = TF.adjust_contrast(image, contrast)\n            image = TF.adjust_saturation(image, saturation)\n        \n        # Step 5: Normalize (ImageNet stats)\n        image = self.normalize(image)\n        \n        return image\n\nclass TestTransform:\n    def __init__(self, img_size=384):\n        self.img_size = img_size\n        self.normalize = transforms.Normalize(\n            mean=[0.485, 0.456, 0.406], \n            std=[0.229, 0.224, 0.225]\n        )\n    \n    def __call__(self, image):\n        # Just resize, convert to tensor, and normalize (no augmentation)\n        image = TF.to_tensor(image)\n        image = TF.resize(image, [self.img_size, self.img_size])\n        image = self.normalize(image)\n        return image\n\nprint(\"TrainTransform and TestTransform defined\")\nprint(\"   - Training: Challenge post-processing + augmentations\")\nprint(\"   - Testing: Only resize + normalize\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:06:16.718678Z","iopub.execute_input":"2026-04-19T09:06:16.718983Z","iopub.status.idle":"2026-04-19T09:06:16.728074Z","shell.execute_reply.started":"2026-04-19T09:06:16.718957Z","shell.execute_reply":"2026-04-19T09:06:16.727521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  DATASET AND DATALOADER SETUP\nprint(\" DATASET AND DATALOADER SETUP\")\n\nclass ImageDataset(Dataset):\n    def __init__(self, df, transform=None, is_train=True):\n        self.df = df\n        self.transform = transform\n        self.is_train = is_train\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        img_path = self.df.iloc[idx]['full_path']\n        try:\n            image = Image.open(img_path).convert('RGB')\n        except Exception as e:\n            print(f\"Error loading {img_path}: {e}\")\n            # Fallback to black image\n            image = Image.new('RGB', (384, 384), (0, 0, 0))\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        if self.is_train:\n            label = self.df.iloc[idx]['y']\n            return image, label\n        else:\n            return image, self.df.iloc[idx]['ID']\n\n# Create train/validation split\nprint(\"\\n Creating train/validation split...\")\ntrain_data, val_data = train_test_split(\n    train_df, \n    test_size=0.2, \n    stratify=train_df['y'], \n    random_state=42\n)\n\nprint(f\" Training samples: {len(train_data)}\")\nprint(f\" Validation samples: {len(val_data)}\")\n\n# Check class distribution in split\nprint(\"\\nTraining set class distribution:\")\nfor i in range(10):\n    count = (train_data['y'] == i).sum()\n    print(f\"   Class {i}: {count} ({count/len(train_data)*100:.1f}%)\")\n\nprint(\"\\n Validation set class distribution:\")\nfor i in range(10):\n    count = (val_data['y'] == i).sum()\n    print(f\"   Class {i}: {count} ({count/len(val_data)*100:.1f}%)\")\n\n# Create datasets\nprint(\"\\n Creating datasets...\")\ntrain_transform = TrainTransform(Config.IMG_SIZE, intensity=Config.AUG_INTENSITY)\nval_transform = TestTransform(Config.IMG_SIZE)\ntest_transform = TestTransform(Config.IMG_SIZE)\n\ntrain_dataset = ImageDataset(train_data, transform=train_transform, is_train=True)\nval_dataset = ImageDataset(val_data, transform=val_transform, is_train=True)\ntest_dataset = ImageDataset(test_df, transform=test_transform, is_train=False)\n\nprint(f\" Train dataset: {len(train_dataset)} samples\")\nprint(f\" Validation dataset: {len(val_dataset)} samples\")\nprint(f\" Test dataset: {len(test_dataset)} samples\")\n\n# Create dataloaders\nprint(\"\\n Creating dataloaders...\")\ntrain_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True)\ntest_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True)\n\nprint(f\" Train batches: {len(train_loader)}\")\nprint(f\" Validation batches: {len(val_loader)}\")\nprint(f\" Test batches: {len(test_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:06:26.679909Z","iopub.execute_input":"2026-04-19T09:06:26.680220Z","iopub.status.idle":"2026-04-19T09:06:26.707632Z","shell.execute_reply.started":"2026-04-19T09:06:26.680182Z","shell.execute_reply":"2026-04-19T09:06:26.706854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#   MODEL ARCHITECTURE\nprint(\" MODEL ARCHITECTURE\")\n\nclass EnhancedModel(nn.Module):\n    def __init__(self, model_name='efficientnet_b4', num_classes=10, dropout_rate=0.3):\n        super().__init__()\n        # Load pretrained backbone\n        self.backbone = timm.create_model(model_name, pretrained=True, num_classes=0)\n        \n        # Get feature dimension\n        if hasattr(self.backbone, 'num_features'):\n            num_features = self.backbone.num_features\n        else:\n            num_features = self.backbone.classifier.in_features\n        \n        print(f\"   Backbone: {model_name}\")\n        print(f\"   Feature dimension: {num_features}\")\n        \n        # Custom classifier head (removed BatchNorm for stability)\n        self.classifier = nn.Sequential(\n            nn.Dropout(dropout_rate),\n            nn.Linear(num_features, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout_rate),\n            nn.Linear(512, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\n\n# Initialize model\nprint(f\"\\n Creating model: {Config.PRIMARY_MODEL}\")\nmodel = EnhancedModel(\n    model_name=Config.PRIMARY_MODEL, \n    num_classes=Config.NUM_CLASSES,\n    dropout_rate=0.3\n).to(device)\n\n# Count parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"\\n Model Statistics:\")\nprint(f\"   Total parameters: {total_params:,}\")\nprint(f\"   Trainable parameters: {trainable_params:,}\")\nprint(f\"   Non-trainable parameters: {total_params - trainable_params:,}\")\n\n# Test forward pass with a proper batch (not single sample)\nprint(\"\\n Testing forward pass...\")\ntest_batch = next(iter(train_loader))[0][:4].to(device)  # Use batch of 4\nwith torch.no_grad():\n    output = model(test_batch)\nprint(f\"   Input shape: {test_batch.shape}\")\nprint(f\"   Output shape: {output.shape}\")\nprint(f\"    Model working correctly!\")\n\n# Test with single sample (for inference)\nprint(\"\\n Testing inference mode (single sample)...\")\nmodel.eval()\ntest_single = next(iter(train_loader))[0][:1].to(device)\nwith torch.no_grad():\n    output_single = model(test_single)\nprint(f\"   Input shape: {test_single.shape}\")\nprint(f\"   Output shape: {output_single.shape}\")\nprint(f\"    Model works for both batch and single samples!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:06:32.393597Z","iopub.execute_input":"2026-04-19T09:06:32.393884Z","iopub.status.idle":"2026-04-19T09:06:45.092441Z","shell.execute_reply.started":"2026-04-19T09:06:32.393859Z","shell.execute_reply":"2026-04-19T09:06:45.091571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#   TRAINING FUNCTIONS \nprint(\"TRAINING SETUP\")\n\ndef mixup_data(x, y, alpha=0.2):\n    \"\"\"Mixup augmentation - creates convex combinations of pairs of examples\"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n    \n    batch_size = x.size()[0]\n    index = torch.randperm(batch_size).to(x.device)\n    \n    mixed_x = lam * x + (1 - lam) * x[index, :]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    \"\"\"Mixup loss function\"\"\"\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\ndef train_model(model, train_loader, val_loader, epochs, device, use_mixup=True):\n    \"\"\"Main training loop with mixed precision and mixup\"\"\"\n    \n    # Loss function with label smoothing\n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTHING)\n    \n    # Optimizer with weight decay\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LEARNING_RATE, weight_decay=0.05)\n    \n    # Learning rate scheduler - Cosine annealing with warm restarts\n    scheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=5, T_mult=2)\n    \n    # Mixed precision training\n    scaler = GradScaler()\n    \n    best_accuracy = 0\n    train_losses = []\n    val_accs = []\n    learning_rates = []\n    \n    print(f\"\\n Training Configuration:\")\n    print(f\"   Epochs: {epochs}\")\n    print(f\"   Batch size: {Config.BATCH_SIZE}\")\n    print(f\"   Learning rate: {Config.LEARNING_RATE}\")\n    print(f\"   Weight decay: 0.05\")\n    print(f\"   Mixup: {'Enabled' if use_mixup else 'Disabled'} (α={Config.MIXUP_ALPHA})\")\n    print(f\"   Label smoothing: {Config.LABEL_SMOOTHING}\")\n    print(f\"   Mixed precision: Enabled\")\n    print(f\"   Scheduler: CosineAnnealingWarmRestarts\")\n    \n    for epoch in range(epochs):\n        # ========== TRAINING PHASE ==========\n        model.train()\n        train_loss = 0\n        train_correct = 0\n        \n        pbar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{epochs}')\n        for batch_idx, (images, labels) in enumerate(pbar):\n            images, labels = images.to(device), labels.to(device)\n            \n            # Apply mixup after epoch 2 (allows model to learn basic patterns first)\n            use_mixup_this_batch = use_mixup and epoch > 2\n            if use_mixup_this_batch:\n                images, labels_a, labels_b, lam = mixup_data(images, labels, Config.MIXUP_ALPHA)\n            \n            optimizer.zero_grad()\n            \n            # Mixed precision forward pass\n            with autocast():\n                outputs = model(images)\n                if use_mixup_this_batch:\n                    loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n                else:\n                    loss = criterion(outputs, labels)\n            \n            # Backward pass with mixed precision\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            \n            train_loss += loss.item()\n            \n            # Calculate accuracy (only when not using mixup)\n            if not use_mixup_this_batch:\n                preds = outputs.argmax(dim=1)\n                train_correct += (preds == labels).sum().item()\n                current_acc = train_correct / ((batch_idx+1) * Config.BATCH_SIZE)\n            else:\n                current_acc = 0\n            \n            # Update progress bar\n            pbar.set_postfix({\n                'loss': f'{loss.item():.4f}', \n                'acc': f'{current_acc:.4f}',\n                'lr': f'{optimizer.param_groups[0][\"lr\"]:.2e}'\n            })\n        \n        avg_train_loss = train_loss / len(train_loader)\n        train_acc = train_correct / len(train_loader.dataset) if not (use_mixup and epoch > 2) else 0\n        \n        #  VALIDATION PHASE \n        model.eval()\n        val_correct = 0\n        val_loss = 0\n        \n        with torch.no_grad():\n            for images, labels in val_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item()\n                val_correct += (outputs.argmax(dim=1) == labels).sum().item()\n        \n        val_acc = val_correct / len(val_loader.dataset)\n        avg_val_loss = val_loss / len(val_loader)\n        \n        # Update scheduler\n        scheduler.step()\n        current_lr = optimizer.param_groups[0]['lr']\n        learning_rates.append(current_lr)\n        \n        # Store metrics\n        train_losses.append(avg_train_loss)\n        val_accs.append(val_acc)\n        \n        # Print epoch summary\n        print(f'\\n Epoch {epoch+1}/{epochs}:')\n        print(f'   Train Loss: {avg_train_loss:.4f} | Train Acc: {train_acc:.4f}')\n        print(f'   Val Loss: {avg_val_loss:.4f} | Val Acc: {val_acc:.4f}')\n        print(f'   Learning Rate: {current_lr:.2e}')\n        \n        # Save best model\n        if val_acc > best_accuracy:\n            best_accuracy = val_acc\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_acc': val_acc,\n                'val_loss': avg_val_loss\n            }, 'best_model.pth')\n            print(f'   Saved best model with val_acc: {val_acc:.4f}')\n    \n    # Plot training curves\n    print(\"\\n Plotting training curves...\")\n    fig, axes = plt.subplots(1, 3, figsize=(15, 4))\n    \n    # Loss plot\n    axes[0].plot(train_losses)\n    axes[0].set_title('Training Loss Over Epochs')\n    axes[0].set_xlabel('Epoch')\n    axes[0].set_ylabel('Loss')\n    axes[0].grid(True)\n    \n    # Accuracy plot\n    axes[1].plot(val_accs)\n    axes[1].set_title('Validation Accuracy Over Epochs')\n    axes[1].set_xlabel('Epoch')\n    axes[1].set_ylabel('Accuracy')\n    axes[1].grid(True)\n    \n    # Learning rate plot\n    axes[2].plot(learning_rates)\n    axes[2].set_title('Learning Rate Schedule')\n    axes[2].set_xlabel('Epoch')\n    axes[2].set_ylabel('Learning Rate')\n    axes[2].set_yscale('log')\n    axes[2].grid(True)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    return best_accuracy, train_losses, val_accs\n\n# Verify everything is ready\nprint(\"\\n All training functions defined and ready!\")\nprint(\"   - Mixup augmentation: Ready\")\nprint(\"   - Mixed precision training: Ready\")\nprint(\"   - Cosine annealing scheduler: Ready\")\nprint(\"   - Model checkpointing: Ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:06:07.332659Z","iopub.execute_input":"2026-04-19T09:06:07.333466Z","iopub.status.idle":"2026-04-19T09:06:07.353386Z","shell.execute_reply.started":"2026-04-19T09:06:07.333432Z","shell.execute_reply":"2026-04-19T09:06:07.352763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  TRAINING EXECUTION WITH LOSS SAVING \nprint(\"TRAINING EXECUTION\")\n\ndef train_model_fresh(model, train_loader, val_loader, epochs, device, use_mixup=True):\n    \"\"\"Main training loop with proper accuracy tracking and loss saving\"\"\"\n    \n    # Loss function with label smoothing\n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTHING)\n    \n    # Optimizer with weight decay\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LEARNING_RATE, weight_decay=0.05)\n    \n    # Learning rate scheduler\n    scheduler = optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=5, T_mult=2)\n    \n    # Mixed precision training\n    scaler = GradScaler()\n    \n    best_accuracy = 0\n    train_losses = []\n    val_losses = []  # ADDED: Store validation losses\n    val_accs = []\n    learning_rates = []\n    train_accs = []  # ADDED: Store training accuracies\n    \n    print(f\"\\nTraining Configuration:\")\n    print(f\"   Epochs: {epochs}\")\n    print(f\"   Batch size: {Config.BATCH_SIZE}\")\n    print(f\"   Learning rate: {Config.LEARNING_RATE}\")\n    print(f\"   Mixup: {'Enabled' if use_mixup else 'Disabled'} (α={Config.MIXUP_ALPHA})\")\n    print(f\"   Label smoothing: {Config.LABEL_SMOOTHING}\")\n    print(f\"   Start: FRESH FROM EPOCH 1\")\n    \n    for epoch in range(epochs):\n        # ========== TRAINING PHASE ==========\n        model.train()\n        train_loss = 0\n        train_correct = 0\n        train_total = 0\n        \n        pbar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{epochs}')\n        for batch_idx, (images, labels) in enumerate(pbar):\n            images, labels = images.to(device), labels.to(device)\n            \n            # Apply mixup after epoch 2\n            use_mixup_this_batch = use_mixup and epoch > 2\n            if use_mixup_this_batch:\n                images, labels_a, labels_b, lam = mixup_data(images, labels, Config.MIXUP_ALPHA)\n            \n            optimizer.zero_grad()\n            \n            # Mixed precision forward pass\n            with autocast():\n                outputs = model(images)\n                if use_mixup_this_batch:\n                    loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n                else:\n                    loss = criterion(outputs, labels)\n            \n            # Backward pass\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            \n            train_loss += loss.item()\n            \n            # Track accuracy\n            if not use_mixup_this_batch:\n                preds = outputs.argmax(dim=1)\n                train_correct += (preds == labels).sum().item()\n                train_total += labels.size(0)\n                current_acc = train_correct / train_total if train_total > 0 else 0\n            else:\n                # For mixup batches, track using original labels\n                preds = outputs.argmax(dim=1)\n                batch_correct = (preds == labels_a).sum().item()\n                train_correct += batch_correct\n                train_total += labels_a.size(0)\n                current_acc = train_correct / train_total if train_total > 0 else 0\n            \n            # Update progress bar\n            pbar.set_postfix({\n                'loss': f'{loss.item():.4f}', \n                'acc': f'{current_acc:.4f}',\n                'lr': f'{optimizer.param_groups[0][\"lr\"]:.2e}'\n            })\n        \n        avg_train_loss = train_loss / len(train_loader)\n        train_acc = train_correct / train_total if train_total > 0 else 0\n        \n        # ========== VALIDATION PHASE ==========\n        model.eval()\n        val_correct = 0\n        val_total = 0\n        val_loss = 0\n        \n        with torch.no_grad():\n            for images, labels in val_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                val_loss += loss.item()\n                val_correct += (outputs.argmax(dim=1) == labels).sum().item()\n                val_total += labels.size(0)\n        \n        val_acc = val_correct / val_total if val_total > 0 else 0\n        avg_val_loss = val_loss / len(val_loader)\n        \n        # Update scheduler\n        scheduler.step()\n        current_lr = optimizer.param_groups[0]['lr']\n        learning_rates.append(current_lr)\n        \n        # Store metrics\n        train_losses.append(avg_train_loss)\n        val_losses.append(avg_val_loss)  # ADDED: Store validation loss\n        val_accs.append(val_acc)\n        train_accs.append(train_acc)  # ADDED: Store training accuracy\n        \n        # Print epoch summary\n        print(f'\\nEpoch {epoch+1}/{epochs}:')\n        print(f'   Train Loss: {avg_train_loss:.4f} | Train Acc: {train_acc:.4f} ({train_acc*100:.2f}%)')\n        print(f'   Val Loss: {avg_val_loss:.4f} | Val Acc: {val_acc:.4f} ({val_acc*100:.2f}%)')\n        print(f'   Learning Rate: {current_lr:.2e}')\n        \n        # Save best model with full history\n        if val_acc > best_accuracy:\n            best_accuracy = val_acc\n            torch.save({\n                'epoch': epoch + 1,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_acc': val_acc,\n                'val_loss': avg_val_loss,\n                'train_acc': train_acc,\n                'train_loss': avg_train_loss,\n                # ADDED: Save full history\n                'train_losses': train_losses,\n                'val_losses': val_losses,\n                'val_accs': val_accs,\n                'train_accs': train_accs,\n                'learning_rates': learning_rates\n            }, 'best_model.pth')\n            print(f'    Saved best model with val_acc: {val_acc:.4f} ({val_acc*100:.2f}%)')\n        \n        # Early stopping if no improvement for 5 epochs \n        if epoch > 8 and val_accs[-1] <= max(val_accs[-6:-1]) and val_accs[-1] < 0.97:\n            print(f\"\\n Early stopping - no improvement for 5 epochs\")\n            break\n    \n    # Save final checkpoint with all history\n    torch.save({\n        'train_losses': train_losses,\n        'val_losses': val_losses,\n        'val_accs': val_accs,\n        'train_accs': train_accs,\n        'learning_rates': learning_rates,\n        'best_accuracy': best_accuracy\n    }, 'training_history.pth')\n    print(\"\\n Training history saved to 'training_history.pth'\")\n    \n    # Plot training curves with both losses\n    print(\"\\nPlotting training curves...\")\n    fig, axes = plt.subplots(2, 2, figsize=(14, 10))\n    \n    # Loss plot\n    axes[0, 0].plot(train_losses, 'b-', label='Train Loss', linewidth=2)\n    axes[0, 0].plot(val_losses, 'r-', label='Val Loss', linewidth=2)\n    axes[0, 0].set_title('Training and Validation Loss', fontsize=12, fontweight='bold')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    # Accuracy plot\n    axes[0, 1].plot(val_accs, 'g-', label='Val Acc', linewidth=2)\n    axes[0, 1].plot(train_accs, 'orange', label='Train Acc', linewidth=2, alpha=0.7)\n    axes[0, 1].set_title('Training and Validation Accuracy', fontsize=12, fontweight='bold')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Accuracy')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    # Loss gap plot\n    loss_gap = [abs(t - v) for t, v in zip(train_losses, val_losses)]\n    axes[1, 0].plot(loss_gap, 'purple', linewidth=2)\n    axes[1, 0].set_title('Train-Val Loss Gap (Overfitting Indicator)', fontsize=12, fontweight='bold')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('Loss Gap')\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    # Learning rate plot\n    axes[1, 1].plot(learning_rates, 'brown', linewidth=2)\n    axes[1, 1].set_title('Learning Rate Schedule', fontsize=12, fontweight='bold')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('Learning Rate')\n    axes[1, 1].set_yscale('log')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Print summary\n    print(\"TRAINING SUMMARY\")\n    print(f\"Best Validation Accuracy: {best_accuracy:.4f} ({best_accuracy*100:.2f}%)\")\n    print(f\"Best Validation Loss: {min(val_losses):.4f} (Epoch {val_losses.index(min(val_losses))+1})\")\n    print(f\"Best Training Loss: {min(train_losses):.4f} (Epoch {train_losses.index(min(train_losses))+1})\")\n    print(f\"Final Train Loss: {train_losses[-1]:.4f}\")\n    print(f\"Final Val Loss: {val_losses[-1]:.4f}\")\n    print(f\"Loss Gap: {abs(train_losses[-1] - val_losses[-1]):.4f}\")\n    \n    return best_accuracy, train_losses, val_losses, val_accs\n\n#  RE-INITIALIZE MODEL FOR FRESH START \nprint(\"\\n Re-initializing model for training...\")\nmodel = EnhancedModel(\n    model_name=Config.PRIMARY_MODEL, \n    num_classes=Config.NUM_CLASSES,\n    dropout_rate=0.3\n).to(device)\n\n# Count parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\" Model re-initialized: {total_params:,} total parameters\")\nprint(f\" Trainable parameters: {trainable_params:,}\")\n\n#  START TRAINING \nprint(\"STARTING TRAINING\")\n\nbest_acc, train_losses, val_losses, val_accs = train_model_fresh(\n    model, \n    train_loader, \n    val_loader, \n    Config.EPOCHS, \n    device, \n    use_mixup=True\n)\n\nprint(f\"TRAINING COMPLETE!\")\nprint(f\"Best Validation Accuracy: {best_acc:.4f} ({best_acc*100:.2f}%)\")\n\n#  LOAD AND DISPLAY SAVED HISTORY \nprint(\"LOADING SAVED TRAINING HISTORY\")\n\nhistory = torch.load('training_history.pth')\nprint(\"\\nSaved data:\")\nprint(f\"   Train losses: {len(history['train_losses'])} values\")\nprint(f\"   Val losses: {len(history['val_losses'])} values\")\nprint(f\"   Val accuracies: {len(history['val_accs'])} values\")\nprint(f\"   Best accuracy: {history['best_accuracy']:.4f}\")\n\n# Create DataFrame for easy viewing\nhistory_df = pd.DataFrame({\n    'Epoch': range(1, len(history['train_losses']) + 1),\n    'Train Loss': history['train_losses'],\n    'Val Loss': history['val_losses'],\n    'Val Acc': [acc*100 for acc in history['val_accs']],\n    'Train Acc': [acc*100 for acc in history['train_accs']]\n})\n\nprint(\"\\nTraining History Table:\")\nprint(history_df.to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T09:11:06.342463Z","iopub.execute_input":"2026-04-19T09:11:06.342851Z","iopub.status.idle":"2026-04-19T10:29:03.179149Z","shell.execute_reply.started":"2026-04-19T09:11:06.342813Z","shell.execute_reply":"2026-04-19T10:29:03.177782Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  DISPLAY TRAINING RESULTS FROM SAVED DATA \nprint(\"LOADING AND DISPLAYING  ACTUAL TRAINING RESULTS\")\n\n# Load the saved checkpoint to get training history\ncheckpoint = torch.load('best_model.pth')\n\n# Display what's in the checkpoint\nprint(\"\\n Checkpoint contents:\")\nprint(f\"   Keys in checkpoint: {checkpoint.keys()}\")\n\n# If you saved training history in checkpoint\nif 'train_losses' in checkpoint:\n    train_losses = checkpoint['train_losses']\n    val_losses = checkpoint['val_losses']\n    val_accs = checkpoint['val_accs']\n    epochs = list(range(1, len(train_losses) + 1))\n    \n    print(\"\\n TRAINING HISTORY TABLE:\")\n    print(\"-\" * 85)\n    print(f\"{'Epoch':^8} {'Train Loss':^12} {'Val Loss':^12} {'Val Acc':^12}\")\n    print(\"-\" * 85)\n    \n    for i in range(len(epochs)):\n        print(f\"{epochs[i]:^8} {train_losses[i]:^12.4f} {val_losses[i]:^12.4f} {val_accs[i]:^11.2f}%\")\n    print(\"-\" * 85)\n    \n    # Plot the data\n    import matplotlib.pyplot as plt\n    \n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    \n    # Loss plot\n    axes[0].plot(epochs, train_losses, 'b-o', label='Train Loss', linewidth=2)\n    axes[0].plot(epochs, val_losses, 'r-s', label='Val Loss', linewidth=2)\n    axes[0].set_xlabel('Epoch')\n    axes[0].set_ylabel('Loss')\n    axes[0].set_title('Training History - Loss')\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n    \n    # Accuracy plot\n    axes[1].plot(epochs, val_accs, 'g-o', label='Val Acc', linewidth=2)\n    axes[1].set_xlabel('Epoch')\n    axes[1].set_ylabel('Accuracy (%)')\n    axes[1].set_title('Training History - Validation Accuracy')\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    print(f\"\\n Best Validation Accuracy: {max(val_accs):.4f}%\")\n    print(f\"   Best Validation Loss: {min(val_losses):.4f}\")\n    \nelse:\n    print(\"\\n Training history not saved in checkpoint\")\n    print(\"   Only these values were saved:\")\n    for key in checkpoint.keys():\n        if not isinstance(checkpoint[key], dict):\n            print(f\"   - {key}: {checkpoint[key]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:30:11.220583Z","iopub.execute_input":"2026-04-19T10:30:11.220958Z","iopub.status.idle":"2026-04-19T10:30:12.052669Z","shell.execute_reply.started":"2026-04-19T10:30:11.220923Z","shell.execute_reply":"2026-04-19T10:30:12.052004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  DISPLAY ACTUAL TRAINING RESULTS FROM SAVED DATA \nprint(\"\\n\" + \"=\"*60)\nprint(\"ACTUAL TRAINING RESULTS\")\nprint(\"=\"*60)\n\n# Load the checkpoint\ncheckpoint = torch.load('best_model.pth')\n\nprint(\"\\n BEST MODEL INFORMATION:\")\nprint(f\"   Best Epoch: {checkpoint['epoch']}\")\nprint(f\"   Best Validation Accuracy: {checkpoint['val_acc']:.4f} ({checkpoint['val_acc']*100:.2f}%)\")\nprint(f\"   Best Validation Loss: {checkpoint['val_loss']:.4f}\")\nprint(f\"   Training Accuracy at best epoch: {checkpoint['train_acc']:.4f} ({checkpoint['train_acc']*100:.2f}%)\")\n\n# Data from  new training run (epochs 1-16)\nprint(\"\\n COMPLETE TRAINING HISTORY (from console output):\")\nprint(\"-\" * 85)\nprint(f\"{'Epoch':^8} {'Train Loss':^12} {'Val Loss':^12} {'Train Acc':^12} {'Val Acc':^12}\")\nprint(\"-\" * 85)\n\n#  actual data from the latest training run\ntraining_data = [\n    (1, 1.8735, 1.0613, 41.46, 80.93),\n    (2, 1.0114, 0.7701, 78.95, 91.93),\n    (3, 0.8225, 0.6994, 88.95, 94.43),\n    (4, 1.0180, 0.7025, 47.80, 94.21),\n    (5, 1.0168, 0.6851, 51.48, 94.57),\n    (6, 0.9926, 0.6554, 52.79, 95.79),\n    (7, 0.9991, 0.6380, 52.66, 96.43),\n    (8, 0.9222, 0.6340, 57.54, 96.43),\n    (9, 0.9299, 0.6004, 53.27, 97.36),\n    (10, 0.9269, 0.6006, 58.34, 97.14),\n    (11, 0.8491, 0.5809, 55.88, 97.50),\n    (12, 0.8578, 0.5811, 54.02, 97.36),\n    (13, 0.8785, 0.5792, 53.75, 97.36),\n    (14, 0.8597, 0.5948, 55.79, 97.07),\n    (15, 0.8896, 0.5850, 51.73, 97.64),\n    (16, 0.8368, 0.5918, 59.64, 97.64),\n]\n\nfor epoch, train_loss, val_loss, train_acc, val_acc in training_data:\n    print(f\"{epoch:^8} {train_loss:^12.4f} {val_loss:^12.4f} {train_acc:^11.2f}% {val_acc:^11.2f}%\")\nprint(\"-\" * 85)\n\n# Summary statistics\nprint(\"\\n SUMMARY STATISTICS:\")\ntrain_losses = [t[1] for t in training_data]\nval_losses = [t[2] for t in training_data]\ntrain_accs = [t[3] for t in training_data]\nval_accs = [t[4] for t in training_data]\n\nprint(f\"   Best Train Loss: {min(train_losses):.4f} (Epoch {train_losses.index(min(train_losses))+1})\")\nprint(f\"   Best Val Loss:   {min(val_losses):.4f} (Epoch {val_losses.index(min(val_losses))+1})\")\nprint(f\"   Best Val Acc:    {max(val_accs):.2f}% (Epoch {val_accs.index(max(val_accs))+1})\")\nprint(f\"   Final Val Acc:   {val_accs[-1]:.2f}%\")\n\n# Mixup effect\nprint(\"\\nMIXUP EFFECT ANALYSIS:\")\nprint(f\"   Before Mixup (Epoch 3): Train Acc={train_accs[2]:.2f}%, Val Acc={val_accs[2]:.2f}%\")\nprint(f\"   After Mixup (Epoch 8):  Train Acc={train_accs[7]:.2f}%, Val Acc={val_accs[7]:.2f}%\")\nprint(f\"   Improvement: Val Acc +{val_accs[7] - val_accs[2]:.2f}%\")\n\n# Create visualization\nimport matplotlib.pyplot as plt\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\n# Loss plot\nepochs = [t[0] for t in training_data]\naxes[0].plot(epochs, train_losses, 'b-o', label='Train Loss', linewidth=2, markersize=8)\naxes[0].plot(epochs, val_losses, 'r-s', label='Val Loss', linewidth=2, markersize=8)\naxes[0].axvline(x=3.5, color='g', linestyle='--', linewidth=2, label='Mixup Started', alpha=0.7)\naxes[0].set_xlabel('Epoch', fontsize=12)\naxes[0].set_ylabel('Loss', fontsize=12)\naxes[0].set_title('Training vs Validation Loss', fontsize=14, fontweight='bold')\naxes[0].legend(fontsize=10)\naxes[0].grid(True, alpha=0.3)\n\n# Accuracy plot\naxes[1].plot(epochs, train_accs, 'b-o', label='Train Acc', linewidth=2, markersize=8)\naxes[1].plot(epochs, val_accs, 'r-s', label='Val Acc', linewidth=2, markersize=8)\naxes[1].axvline(x=3.5, color='g', linestyle='--', linewidth=2, label='Mixup Started', alpha=0.7)\naxes[1].set_xlabel('Epoch', fontsize=12)\naxes[1].set_ylabel('Accuracy (%)', fontsize=12)\naxes[1].set_title('Training vs Validation Accuracy', fontsize=14, fontweight='bold')\naxes[1].legend(fontsize=10)\naxes[1].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.show()\n\nprint(\"\\n\" + \"=\"*60)\nprint(f\"BEST VALIDATION ACCURACY: {checkpoint['val_acc']*100:.2f}%\")\nprint(f\"   Achieved at Epoch: {checkpoint['epoch']}\")\nprint(\"=\"*60)\n\n# Additional analysis\nprint(\"\\nPERFORMANCE SUMMARY:\")\nprint(f\"   Starting Val Acc: {val_accs[0]:.2f}%\")\nprint(f\"   Peak Val Acc: {max(val_accs):.2f}%\")\nprint(f\"   Improvement: +{max(val_accs) - val_accs[0]:.2f}%\")\nprint(f\"   Starting Val Loss: {val_losses[0]:.4f}\")\nprint(f\"   Best Val Loss: {min(val_losses):.4f}\")\nprint(f\"   Loss Reduction: {val_losses[0] - min(val_losses):.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:34:46.407821Z","iopub.execute_input":"2026-04-19T10:34:46.408624Z","iopub.status.idle":"2026-04-19T10:34:47.245676Z","shell.execute_reply.started":"2026-04-19T10:34:46.408590Z","shell.execute_reply":"2026-04-19T10:34:47.245066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  GENERATE PREDICTIONS USING BEST MODEL \nprint(\"GENERATING TEST PREDICTIONS\")\n\n# Load the best model\nprint(\"\\nLoading best model...\")\ncheckpoint = torch.load('best_model.pth')\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\n\nprint(f\"Loaded model from Epoch {checkpoint['epoch']}\")\nprint(f\"Validation Accuracy: {checkpoint['val_acc']*100:.2f}%\")\nprint(f\"Validation Loss: {checkpoint['val_loss']:.4f}\")\n\n# Test-Time Augmentation function\ndef predict_with_tta(model, images, device, use_tta=True):\n    \"\"\"Enhanced prediction with Test-Time Augmentation\"\"\"\n    with torch.no_grad():\n        # Original prediction\n        outputs = model(images)\n        \n        if use_tta:\n            # Horizontal flip augmentation\n            outputs_flip = model(torch.flip(images, dims=[3]))\n            outputs = (outputs + outputs_flip) / 2\n    \n    return outputs\n\n# Make predictions\nprint(\"\\nMaking predictions on test set...\")\npredictions = []\nimage_ids = []\nuse_tta = True\n\nprint(f\"Test-Time Augmentation: {'Enabled' if use_tta else 'Disabled'}\")\nprint(f\"Test samples: {len(test_dataset)}\")\nprint(f\"Test batches: {len(test_loader)}\")\n\nwith torch.no_grad():\n    for images, ids in tqdm(test_loader, desc=\"Predicting\"):\n        images = images.to(device)\n        \n        # Get predictions with TTA\n        outputs = predict_with_tta(model, images, device, use_tta=use_tta)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n        \n        predictions.extend(preds)\n        image_ids.extend(ids.numpy() if torch.is_tensor(ids) else ids)\n\nprint(f\"\\nGenerated {len(predictions)} predictions\")\n\n# Create submission DataFrame\nprint(\"\\nCreating submission file...\")\nsubmission = pd.DataFrame({\n    'ID': image_ids,\n    'TARGET': predictions\n})\n\n# Sort by ID for consistency\nsubmission = submission.sort_values('ID').reset_index(drop=True)\n\n# Verify submission format\nprint(f\"\\nSubmission validation:\")\nprint(f\"Total predictions: {len(submission)}\")\nprint(f\"Unique IDs: {submission['ID'].nunique()}\")\nprint(f\"Missing values: {submission.isnull().sum().sum()}\")\nprint(f\"Target range: {submission['TARGET'].min()} - {submission['TARGET'].max()}\")\n\n# Display prediction distribution\nprint(f\"\\nPrediction distribution:\")\npred_dist = submission['TARGET'].value_counts().sort_index()\nfor i in range(10):\n    count = pred_dist.get(i, 0)\n    percentage = count/len(submission)*100\n    bar = '*' * int(percentage/2)\n    print(f\"   Class {i}: {count:4d} ({percentage:5.2f}%) {bar}\")\n\n# Save submission\nsubmission.to_csv('submission.csv', index=False)\nprint(f\"\\nSubmission saved to 'submission.csv'\")\n\n# Display sample\nprint(f\"\\nSample predictions (first 20 rows):\")\nprint(submission.head(20))\n\n# Quick confidence check on validation set\nprint(\"\\nQuick validation check (on validation set):\")\nmodel.eval()\nval_correct = 0\nval_total = 0\n\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Validating\", leave=False):\n        images, labels = images.to(device), labels.to(device)\n        outputs = predict_with_tta(model, images, device, use_tta=use_tta)\n        preds = outputs.argmax(dim=1)\n        val_correct += (preds == labels).sum().item()\n        val_total += labels.size(0)\n\nval_accuracy = val_correct / val_total\nprint(f\"Validation accuracy with TTA: {val_accuracy:.4f} ({val_accuracy*100:.2f}%)\")\nprint(f\"Saved model accuracy: {checkpoint['val_acc']*100:.2f}%\")\n\nprint(\"PREDICTIONS COMPLETE\")\nprint(\"\\nFiles created:\")\nprint(\"   - submission.csv (ready for Kaggle upload)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:31:34.908732Z","iopub.execute_input":"2026-04-19T10:31:34.909281Z","iopub.status.idle":"2026-04-19T10:34:08.760732Z","shell.execute_reply.started":"2026-04-19T10:31:34.909251Z","shell.execute_reply":"2026-04-19T10:34:08.759878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#  MODEL PERFORMANCE METRICS \nprint(\"MODEL PERFORMANCE METRICS\")\n\n# Load best model\ncheckpoint = torch.load('best_model.pth')\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\n\nprint(f\"\\nModel loaded from Epoch {checkpoint['epoch']}\")\nprint(f\"Best Validation Accuracy: {checkpoint['val_acc']*100:.2f}%\")\n\n#  CONFUSION MATRIX \nprint(\"1. CONFUSION MATRIX\")\n\nfrom sklearn.metrics import confusion_matrix, classification_report, accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\nfrom sklearn.preprocessing import label_binarize\nimport seaborn as sns\n\n# Get predictions on validation set\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Computing predictions\"):\n        images, labels = images.to(device), labels.to(device)\n        outputs = predict_with_tta(model, images, device, use_tta=True)\n        preds = outputs.argmax(dim=1)\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n# Calculate confusion matrix\ncm = confusion_matrix(all_labels, all_preds)\n\n# Display confusion matrix\nprint(\"\\nConfusion Matrix:\")\nprint(\"Rows: True Labels, Columns: Predicted Labels\")\nprint(\"-\" * 60)\nprint(\"     \", end=\"\")\nfor i in range(10):\n    print(f\"  C{i} \", end=\"\")\nprint()\nfor i in range(10):\n    print(f\"C{i}  \", end=\"\")\n    for j in range(10):\n        print(f\"{cm[i,j]:4d} \", end=\"\")\n    print()\nprint(\"-\" * 60)\n\n# Plot confusion matrix\nplt.figure(figsize=(12, 10))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=range(10), yticklabels=range(10))\nplt.title('Confusion Matrix - Synthetic Image Attribution', fontsize=14, fontweight='bold')\nplt.xlabel('Predicted Class', fontsize=12)\nplt.ylabel('True Class', fontsize=12)\nplt.tight_layout()\nplt.show()\n\n# CLASSIFICATION REPORT \nprint(\"2. CLASSIFICATION REPORT\")\n\n# Calculate per-class metrics\nreport = classification_report(all_labels, all_preds, target_names=[f'Class_{i}' for i in range(10)], output_dict=True)\n\nprint(\"\\nPer-Class Metrics:\")\nprint(\"-\" * 80)\nprint(f\"{'Class':<10} {'Precision':<12} {'Recall':<12} {'F1-Score':<12} {'Support':<10}\")\nprint(\"-\" * 80)\nfor i in range(10):\n    class_name = f'Class_{i}'\n    print(f\"{class_name:<10} {report[class_name]['precision']:<12.4f} {report[class_name]['recall']:<12.4f} {report[class_name]['f1-score']:<12.4f} {report[class_name]['support']:<10.0f}\")\nprint(\"-\" * 80)\nprint(f\"{'Macro Avg':<10} {report['macro avg']['precision']:<12.4f} {report['macro avg']['recall']:<12.4f} {report['macro avg']['f1-score']:<12.4f}\")\nprint(f\"{'Weighted Avg':<10} {report['weighted avg']['precision']:<12.4f} {report['weighted avg']['recall']:<12.4f} {report['weighted avg']['f1-score']:<12.4f}\")\n\n# ACCURACY METRICS \nprint(\"3. ACCURACY METRICS\")\n\naccuracy = accuracy_score(all_labels, all_preds)\nprint(f\"\\nOverall Accuracy: {accuracy:.4f} ({accuracy*100:.2f}%)\")\n\n# Per-class accuracy\nclass_accuracy = []\nfor i in range(10):\n    class_correct = cm[i,i]\n    class_total = sum(cm[i,:])\n    class_acc = class_correct / class_total if class_total > 0 else 0\n    class_accuracy.append(class_acc)\n    print(f\"Class {i} Accuracy: {class_acc:.4f} ({class_acc*100:.2f}%) - {class_correct}/{class_total}\")\n\nprint(f\"\\nAverage Class Accuracy: {np.mean(class_accuracy):.4f} ({np.mean(class_accuracy)*100:.2f}%)\")\n\n#  PRECISION, RECALL, F1 SCORES \nprint(\"4. PRECISION, RECALL, F1 SCORES\")\n\nprecision_macro = precision_score(all_labels, all_preds, average='macro')\nrecall_macro = recall_score(all_labels, all_preds, average='macro')\nf1_macro = f1_score(all_labels, all_preds, average='macro')\n\nprecision_weighted = precision_score(all_labels, all_preds, average='weighted')\nrecall_weighted = recall_score(all_labels, all_preds, average='weighted')\nf1_weighted = f1_score(all_labels, all_preds, average='weighted')\n\nprint(f\"\\nMacro Average:\")\nprint(f\"   Precision: {precision_macro:.4f} ({precision_macro*100:.2f}%)\")\nprint(f\"   Recall:    {recall_macro:.4f} ({recall_macro*100:.2f}%)\")\nprint(f\"   F1-Score:  {f1_macro:.4f} ({f1_macro*100:.2f}%)\")\n\nprint(f\"\\nWeighted Average:\")\nprint(f\"   Precision: {precision_weighted:.4f} ({precision_weighted*100:.2f}%)\")\nprint(f\"   Recall:    {recall_weighted:.4f} ({recall_weighted*100:.2f}%)\")\nprint(f\"   F1-Score:  {f1_weighted:.4f} ({f1_weighted*100:.2f}%)\")\n\n#  MISCLASSIFICATION ANALYSIS \nprint(\"5. MISCLASSIFICATION ANALYSIS\")\n\n# Find most confused pairs\nmisclassified = []\nfor i in range(10):\n    for j in range(10):\n        if i != j and cm[i,j] > 0:\n            misclassified.append((i, j, cm[i,j]))\n\n# Sort by number of misclassifications\nmisclassified.sort(key=lambda x: x[2], reverse=True)\n\nprint(\"\\nTop 5 Most Confused Class Pairs:\")\nprint(\"-\" * 50)\nfor i, (true_class, pred_class, count) in enumerate(misclassified[:5]):\n    print(f\"{i+1}. Class {true_class} -> Class {pred_class}: {count} misclassifications\")\n\n#  SUMMARY REPORT \nprint(\"6. SUMMARY REPORT\")\n\nprint(f\"\"\"\nMODEL PERFORMANCE SUMMARY\n-------------------------\nBest Validation Accuracy: {checkpoint['val_acc']*100:.2f}%\nOverall Test Accuracy:     {accuracy*100:.2f}%\n\nPer-Class Metrics (Average):\n   Precision: {precision_macro*100:.2f}%\n   Recall:    {recall_macro*100:.2f}%\n   F1-Score:  {f1_macro*100:.2f}%\n\nConfusion Matrix Statistics:\n   Total Predictions: {len(all_preds)}\n   Correct Predictions: {sum(cm[i,i] for i in range(10))}\n   Misclassifications: {len(all_preds) - sum(cm[i,i] for i in range(10))}\n   Accuracy per class range: {min(class_accuracy)*100:.2f}% - {max(class_accuracy)*100:.2f}%\n\nModel Configuration:\n   Architecture: {Config.PRIMARY_MODEL}\n   Image Size: {Config.IMG_SIZE}x{Config.IMG_SIZE}\n   Batch Size: {Config.BATCH_SIZE}\n   Best Epoch: {checkpoint['epoch']}\n\"\"\")\n\n# VISUALIZATION DASHBOARD \nprint(\"\\nCreating performance dashboard...\")\nfig, axes = plt.subplots(2, 2, figsize=(14, 12))\n\n# Per-class accuracy bar chart\naxes[0,0].bar(range(10), class_accuracy, color='steelblue', edgecolor='black')\naxes[0,0].set_xlabel('Class', fontsize=12)\naxes[0,0].set_ylabel('Accuracy', fontsize=12)\naxes[0,0].set_title('Per-Class Accuracy', fontsize=14, fontweight='bold')\naxes[0,0].set_ylim([0, 1])\naxes[0,0].set_xticks(range(10))\naxes[0,0].grid(True, alpha=0.3)\n\n# Metrics comparison\nmetrics = ['Precision', 'Recall', 'F1-Score']\nmacro_scores = [precision_macro, recall_macro, f1_macro]\nweighted_scores = [precision_weighted, recall_weighted, f1_weighted]\n\nx = np.arange(len(metrics))\nwidth = 0.35\n\naxes[0,1].bar(x - width/2, macro_scores, width, label='Macro Avg', color='steelblue', edgecolor='black')\naxes[0,1].bar(x + width/2, weighted_scores, width, label='Weighted Avg', color='lightcoral', edgecolor='black')\naxes[0,1].set_xlabel('Metrics', fontsize=12)\naxes[0,1].set_ylabel('Score', fontsize=12)\naxes[0,1].set_title('Precision, Recall, F1-Score Comparison', fontsize=14, fontweight='bold')\naxes[0,1].set_xticks(x)\naxes[0,1].set_xticklabels(metrics)\naxes[0,1].legend()\naxes[0,1].set_ylim([0, 1])\naxes[0,1].grid(True, alpha=0.3)\n\n# Misclassification heatmap\nmisclass_matrix = cm.copy()\nfor i in range(10):\n    misclass_matrix[i,i] = 0\nsns.heatmap(misclass_matrix, annot=True, fmt='d', cmap='Reds', ax=axes[1,0])\naxes[1,0].set_title('Misclassification Heatmap (Diagonal Removed)', fontsize=14, fontweight='bold')\naxes[1,0].set_xlabel('Predicted Class', fontsize=12)\naxes[1,0].set_ylabel('True Class', fontsize=12)\n\n# Summary text\naxes[1,1].axis('off')\nsummary_text = f\"\"\"\nPERFORMANCE SUMMARY\n\nOverall Accuracy: {accuracy*100:.2f}%\nBest Val Accuracy: {checkpoint['val_acc']*100:.2f}%\n\nMacro Averages:\n- Precision: {precision_macro*100:.2f}%\n- Recall:    {recall_macro*100:.2f}%\n- F1-Score:  {f1_macro*100:.2f}%\n\nBest Class: Class {class_accuracy.index(max(class_accuracy))}\nAccuracy: {max(class_accuracy)*100:.2f}%\n\nWorst Class: Class {class_accuracy.index(min(class_accuracy))}\nAccuracy: {min(class_accuracy)*100:.2f}%\n\nTotal Misclassifications: {len(all_preds) - sum(cm[i,i] for i in range(10))}\n\"\"\"\naxes[1,1].text(0.1, 0.5, summary_text, fontsize=10, verticalalignment='center',\n               fontfamily='monospace', bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n\nplt.suptitle('Model Performance Dashboard', fontsize=16, fontweight='bold')\nplt.tight_layout()\nplt.show()\n\nprint(\"METRICS COMPUTATION COMPLETE\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:35:07.973908Z","iopub.execute_input":"2026-04-19T10:35:07.974523Z","iopub.status.idle":"2026-04-19T10:35:51.288279Z","shell.execute_reply.started":"2026-04-19T10:35:07.974492Z","shell.execute_reply":"2026-04-19T10:35:51.287471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============ MISCLASSIFICATION ANALYSIS FOR IMPROVEMENT ============\nprint(\"\\n\" + \"=\"*60)\nprint(\"DETAILED MISCLASSIFICATION ANALYSIS\")\nprint(\"=\"*60)\n\n# Load best model\ncheckpoint = torch.load('best_model.pth')\nmodel.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()\n\n# Get all predictions and logits\nall_preds = []\nall_labels = []\nall_logits = []\n\nwith torch.no_grad():\n    for images, labels in tqdm(val_loader, desc=\"Computing predictions\"):\n        images, labels = images.to(device), labels.to(device)\n        outputs = model(images)\n        probs = torch.softmax(outputs, dim=1)\n        preds = outputs.argmax(dim=1)\n        \n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n        all_logits.extend(probs.cpu().numpy())\n\nall_preds = np.array(all_preds)\nall_labels = np.array(all_labels)\nall_logits = np.array(all_logits)\n\n# Find misclassified samples\nmisclassified_idx = np.where(all_preds != all_labels)[0]\nprint(f\"\\n📊 Total misclassifications: {len(misclassified_idx)} out of {len(all_labels)} ({len(misclassified_idx)/len(all_labels)*100:.2f}%)\")\n\n# Analyze by class\nprint(\"\\n\" + \"=\"*60)\nprint(\"MISCLASSIFICATION BY CLASS\")\nprint(\"=\"*60)\n\nclass_misclass = {}\nfor class_id in range(10):\n    class_indices = np.where(all_labels == class_id)[0]\n    class_mis = np.where(all_preds[class_indices] != all_labels[class_indices])[0]\n    mis_rate = len(class_mis) / len(class_indices) * 100\n    class_misclass[class_id] = {\n        'total': len(class_indices),\n        'misclassified': len(class_mis),\n        'rate': mis_rate\n    }\n    print(f\"Class {class_id}: {len(class_mis)}/{len(class_indices)} misclassified ({mis_rate:.2f}%)\")\n\n# Find low confidence misclassifications\nprint(\"\\n\" + \"=\"*60)\nprint(\"CONFIDENCE ANALYSIS ON MISCLASSIFICATIONS\")\nprint(\"=\"*60)\n\ncorrect_confidences = []\nmis_confidences = []\n\nfor i in range(len(all_labels)):\n    confidence = all_logits[i][all_labels[i]]\n    if all_preds[i] == all_labels[i]:\n        correct_confidences.append(confidence)\n    else:\n        mis_confidences.append(confidence)\n\nprint(f\"\\nCorrect predictions - Avg confidence: {np.mean(correct_confidences):.4f}\")\nprint(f\"Misclassified - Avg confidence: {np.mean(mis_confidences):.4f}\")\nprint(f\"Confidence gap: {np.mean(correct_confidences) - np.mean(mis_confidences):.4f}\")\n\n# Find specific problematic samples\nprint(\"\\n\" + \"=\"*60)\nprint(\"PROBLEMATIC CLASS PAIRS ANALYSIS\")\nprint(\"=\"*60)\n\nconfusion_pairs = {}\nfor i in misclassified_idx:\n    true_class = all_labels[i]\n    pred_class = all_preds[i]\n    pair = (true_class, pred_class)\n    if pair not in confusion_pairs:\n        confusion_pairs[pair] = 0\n    confusion_pairs[pair] += 1\n\nprint(\"\\nTop misclassification pairs:\")\nsorted_pairs = sorted(confusion_pairs.items(), key=lambda x: x[1], reverse=True)\nfor (true_class, pred_class), count in sorted_pairs[:10]:\n    print(f\"   Class {true_class} -> Class {pred_class}: {count} times\")\n\n# Analyze confidence for confused pairs\nprint(\"\\n\" + \"=\"*60)\nprint(\"CONFIDENCE FOR CONFUSED PAIRS\")\nprint(\"=\"*60)\n\nfor (true_class, pred_class), count in sorted_pairs[:5]:\n    pair_confidences = []\n    for i in misclassified_idx:\n        if all_labels[i] == true_class and all_preds[i] == pred_class:\n            pair_confidences.append(all_logits[i][pred_class])\n    \n    if pair_confidences:\n        print(f\"Class {true_class} -> Class {pred_class}:\")\n        print(f\"   Count: {count}\")\n        print(f\"   Avg confidence in wrong class: {np.mean(pair_confidences):.4f}\")\n        print(f\"   True class confidence: {np.mean([all_logits[i][true_class] for i in misclassified_idx if all_labels[i]==true_class and all_preds[i]==pred_class]):.4f}\")\n\n# Find high confidence misclassifications (most problematic)\nprint(\"\\n\" + \"=\"*60)\nprint(\"HIGHEST CONFIDENCE MISCLASSIFICATIONS (Most Problematic)\")\nprint(\"=\"*60)\n\nmis_confidences_list = []\nfor i in misclassified_idx:\n    wrong_confidence = all_logits[i][all_preds[i]]\n    mis_confidences_list.append((i, all_labels[i], all_preds[i], wrong_confidence))\n\nmis_confidences_list.sort(key=lambda x: x[3], reverse=True)\n\nprint(\"\\nTop 10 most confident wrong predictions:\")\nfor i, (idx, true_class, pred_class, conf) in enumerate(mis_confidences_list[:10]):\n    print(f\"   {i+1}. Sample {idx}: Class {true_class} -> Class {pred_class} (confidence: {conf:.4f})\")\n\n# Improvement suggestions\nprint(\"\\n\" + \"=\"*60)\nprint(\"IMPROVEMENT SUGGESTIONS\")\nprint(\"=\"*60)\n\nprint(\"\\nBased on analysis:\")\n\n# Find which classes need most improvement\nworst_classes = sorted([(c, d['rate']) for c, d in class_misclass.items()], key=lambda x: x[1], reverse=True)\nprint(f\"\\n1. Focus on worst-performing classes:\")\nfor class_id, rate in worst_classes[:3]:\n    print(f\"   - Class {class_id}: {rate:.2f}% error rate\")\n\nprint(f\"\\n2. Address specific confusion pairs:\")\nfor (true_class, pred_class), count in sorted_pairs[:3]:\n    print(f\"   - Class {true_class} is often confused with Class {pred_class}\")\n\nprint(f\"\\n3. Low confidence area:\")\nif np.mean(mis_confidences) < 0.7:\n    print(f\"   - Model is unsure on {len(misclassified_idx)} samples (avg conf: {np.mean(mis_confidences):.3f})\")\n    print(f\"   - Consider: More data augmentation or ensemble methods\")\n\n# Visualize confusion matrix with percentages\nprint(\"\\n\" + \"=\"*60)\nprint(\"CONFUSION MATRIX WITH PERCENTAGES\")\nprint(\"=\"*60)\n\nfrom sklearn.metrics import confusion_matrix\nimport seaborn as sns\n\ncm = confusion_matrix(all_labels, all_preds)\ncm_percent = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis] * 100\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 6))\n\n# Absolute confusion matrix\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=axes[0])\naxes[0].set_title('Confusion Matrix (Absolute)', fontsize=12, fontweight='bold')\naxes[0].set_xlabel('Predicted Class')\naxes[0].set_ylabel('True Class')\n\n# Percentage confusion matrix\nsns.heatmap(cm_percent, annot=True, fmt='.1f', cmap='YlOrRd', ax=axes[1])\naxes[1].set_title('Confusion Matrix (% of True Class)', fontsize=12, fontweight='bold')\naxes[1].set_xlabel('Predicted Class')\naxes[1].set_ylabel('True Class')\n\nplt.tight_layout()\nplt.show()\n\n# Summary\nprint(\"\\n\" + \"=\"*60)\nprint(\"SUMMARY AND RECOMMENDATIONS\")\nprint(\"=\"*60)\n\nprint(f\"\"\"\nOverall Performance: {len(misclassified_idx)}/{len(all_labels)} errors ({len(misclassified_idx)/len(all_labels)*100:.2f}%)\n\nMain Issues:\n1. Class 6 and Class 7 confusion (most significant)\n2. Class 0 has higher error rate than average\n3. Some misclassifications have high confidence\n\nRecommended Actions:\n1. Collect more training samples for Class 6 and 7\n2. Add class-specific augmentation for confused pairs\n3. Consider ensemble with different architecture for these classes\n4. Analyze original source images to understand visual similarities\n\nExpected improvement potential: +0.5-1.0% with targeted fixes\n\"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:38:31.137820Z","iopub.execute_input":"2026-04-19T10:38:31.138576Z","iopub.status.idle":"2026-04-19T10:39:10.825782Z","shell.execute_reply.started":"2026-04-19T10:38:31.138530Z","shell.execute_reply":"2026-04-19T10:39:10.824953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============ TARGETED AUGMENTATION FOR CLASS 6 & 7 ============\nprint(\"\\nAdding targeted augmentation for confused classes\")\n\nclass TargetedAugmentation:\n    def __init__(self, confused_pairs=[(6,7), (7,6)]):\n        self.confused_pairs = confused_pairs\n    \n    def __call__(self, image, class_id):\n        # If image belongs to confused classes, apply extra augmentation\n        if class_id in [6, 7]:\n            # Add stronger augmentation for these classes\n            if random.random() < 0.5:\n                # Random rotation between classes\n                angle = random.uniform(-25, 25)\n                image = image.rotate(angle, expand=True, fillcolor=0)\n            \n            if random.random() < 0.5:\n                # Add slight blur to confuse features\n                image = image.filter(ImageFilter.GaussianBlur(radius=random.uniform(0.5, 1.5)))\n            \n            if random.random() < 0.3:\n                # Adjust contrast to make features less distinct\n                enhancer = ImageEnhance.Contrast(image)\n                image = enhancer.enhance(random.uniform(0.7, 1.3))\n        \n        return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:40:48.325466Z","iopub.execute_input":"2026-04-19T10:40:48.325811Z","iopub.status.idle":"2026-04-19T10:40:48.332972Z","shell.execute_reply.started":"2026-04-19T10:40:48.325777Z","shell.execute_reply":"2026-04-19T10:40:48.332317Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============ ANALYZE WHY CLASS 6 AND 7 ARE CONFUSED ============\nprint(\"\\nAnalyzing source similarity between Class 6 and 7\")\n\n# Find specific misclassified samples\nmisclassified_6_as_7 = []\nmisclassified_7_as_6 = []\n\nfor i in range(len(all_labels)):\n    if all_labels[i] == 6 and all_preds[i] == 7:\n        misclassified_6_as_7.append(i)\n    elif all_labels[i] == 7 and all_preds[i] == 6:\n        misclassified_7_as_6.append(i)\n\nprint(f\"\\nSamples where Class 6 misclassified as Class 7: {len(misclassified_6_as_7)}\")\nprint(f\"Samples where Class 7 misclassified as Class 6: {len(misclassified_7_as_6)}\")\n\n# Display sample indices for inspection\nprint(f\"\\nSample indices to investigate:\")\nprint(f\"   Class 6 -> 7: {misclassified_6_as_7[:5]}\")\nprint(f\"   Class 7 -> 6: {misclassified_7_as_6[:5]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-19T10:41:10.079028Z","iopub.execute_input":"2026-04-19T10:41:10.079336Z","iopub.status.idle":"2026-04-19T10:41:10.086009Z","shell.execute_reply.started":"2026-04-19T10:41:10.079310Z","shell.execute_reply":"2026-04-19T10:41:10.085169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}