{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:36.617126Z","iopub.execute_input":"2025-11-27T23:24:36.617860Z","iopub.status.idle":"2025-11-27T23:24:55.931865Z","shell.execute_reply.started":"2025-11-27T23:24:36.617833Z","shell.execute_reply":"2025-11-27T23:24:55.930851Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-28T11:33:09.518165Z","iopub.execute_input":"2025-11-28T11:33:09.518657Z","iopub.status.idle":"2025-11-28T11:33:09.569882Z","shell.execute_reply.started":"2025-11-28T11:33:09.518629Z","shell.execute_reply":"2025-11-28T11:33:09.568979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom PIL import Image\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport os\nfrom tqdm import tqdm\nfrom sklearn.model_selection import StratifiedKFold\nimport warnings\nwarnings.filterwarnings('ignore')\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-28T13:59:38.916Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"CONFIG","metadata":{}},{"cell_type":"code","source":"\n# =============================================================================\n# CONFIGURATION\n# =============================================================================\nclass Config:\n    # Paths\n    DATA_PATH = \"../input/cassava-leaf-disease-classification\"\n    TRAIN_PATH = \"../input/cassava-leaf-disease-classification/train_images/\"\n    TEST_PATH = \"../input/cassava-leaf-disease-classification/test_images/\"\n    \n    # Model\n    IMG_SIZE = 384\n    N_CLASSES = 5\n    \n    # Training\n    MODE = \"train\"  # \"train\" or \"inference\"\n    N_FOLDS = 5\n    TRAIN_FOLDS = [0, 1, 2, 3, 4]  # Which folds to train\n    N_EPOCHS = 15\n    BATCH_SIZE = 16\n    ACCUMULATION_STEPS = 2  # Gradient accumulation\n    \n    # Optimizer\n    LR = 1e-3\n    WEIGHT_DECAY = 1e-4\n    SCHEDULER = \"cosine\"  # \"cosine\" or \"step\"\n    \n    # Augmentation\n    MIXUP_ALPHA = 0.2\n    LABEL_SMOOTHING = 0.1\n    \n    # Other\n    SEED = 42\n    NUM_WORKERS = 4\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    USE_AMP = True  # Mixed precision training\n    \n    # Inference\n    USE_TTA = True\n    \ncfg = Config()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-28T13:59:38.914Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# SEED\n# =============================================================================\ndef seed_everything(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(cfg.SEED)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-11-28T13:59:38.916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# CUSTOM CNN MODEL - BUILT FROM SCRATCH\n# =============================================================================\nclass ConvBlock(nn.Module):\n    \"\"\"Convolutional block with BatchNorm and activation\"\"\"\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1):\n        super().__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, bias=False)\n        self.bn = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n    \n    def forward(self, x):\n        return self.relu(self.bn(self.conv(x)))\n\nclass ResidualBlock(nn.Module):\n    \"\"\"Residual block with skip connection\"\"\"\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        self.conv1 = ConvBlock(in_channels, out_channels, stride=stride)\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(out_channels, out_channels, 3, 1, 1, bias=False),\n            nn.BatchNorm2d(out_channels)\n        )\n        \n        # Skip connection\n        self.skip = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.skip = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, 1, stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n        \n        self.relu = nn.ReLU(inplace=True)\n    \n    def forward(self, x):\n        residual = x\n        out = self.conv1(x)\n        out = self.conv2(out)\n        out += self.skip(residual)\n        out = self.relu(out)\n        return out\n\nclass CassavaNet(nn.Module):\n    \"\"\"Custom CNN for Cassava Classification - Built from scratch\"\"\"\n    def __init__(self, num_classes=5):\n        super().__init__()\n        \n        # Initial convolution\n        self.conv1 = ConvBlock(3, 64, kernel_size=7, stride=2, padding=3)\n        self.pool1 = nn.MaxPool2d(3, stride=2, padding=1)\n        \n        # Residual blocks\n        self.layer1 = self._make_layer(64, 64, 2, stride=1)\n        self.layer2 = self._make_layer(64, 128, 2, stride=2)\n        self.layer3 = self._make_layer(128, 256, 2, stride=2)\n        self.layer4 = self._make_layer(256, 512, 2, stride=2)\n        \n        # Global pooling and classifier\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        self.dropout = nn.Dropout(0.5)\n        self.fc = nn.Linear(512, num_classes)\n        \n        # Initialize weights\n        self._initialize_weights()\n    \n    def _make_layer(self, in_channels, out_channels, num_blocks, stride):\n        layers = []\n        layers.append(ResidualBlock(in_channels, out_channels, stride))\n        for _ in range(1, num_blocks):\n            layers.append(ResidualBlock(out_channels, out_channels, stride=1))\n        return nn.Sequential(*layers)\n    \n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.normal_(m.weight, 0, 0.01)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.pool1(x)\n        \n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        \n        x = self.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.dropout(x)\n        x = self.fc(x)\n        \n        return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:55.973248Z","iopub.execute_input":"2025-11-27T23:24:55.973645Z","iopub.status.idle":"2025-11-27T23:24:55.982770Z","shell.execute_reply.started":"2025-11-27T23:24:55.973629Z","shell.execute_reply":"2025-11-27T23:24:55.982084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# DATA AUGMENTATION\n# =============================================================================\ndef get_train_transforms():\n    return A.Compose([\n        A.RandomResizedCrop(size=(cfg.IMG_SIZE, cfg.IMG_SIZE), scale=(0.8, 1.0)),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.15, rotate_limit=45, p=0.5),\n        A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=20, p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5),\n        A.CoarseDropout(max_holes=8, max_height=cfg.IMG_SIZE//8, max_width=cfg.IMG_SIZE//8, \n                        fill_value=0, p=0.5),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2(),\n    ])\n\ndef get_valid_transforms():\n    return A.Compose([\n        A.Resize(height=cfg.IMG_SIZE, width=cfg.IMG_SIZE),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2(),\n    ])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:55.983551Z","iopub.execute_input":"2025-11-27T23:24:55.983792Z","iopub.status.idle":"2025-11-27T23:24:55.990216Z","shell.execute_reply.started":"2025-11-27T23:24:55.983777Z","shell.execute_reply":"2025-11-27T23:24:55.989458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# DATASET\n# =============================================================================\nclass CassavaDataset(Dataset):\n    def __init__(self, df, img_path, transforms=None):\n        self.df = df.reset_index(drop=True)\n        self.img_path = img_path\n        self.transforms = transforms\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        img_name = self.df.loc[idx, 'image_id']\n        label = self.df.loc[idx, 'label']\n        img_path = os.path.join(self.img_path, img_name)\n        \n        image = Image.open(img_path).convert('RGB')\n        image = np.array(image)\n        \n        if self.transforms:\n            image = self.transforms(image=image)['image']\n        \n        return image, label\n\nclass TestDataset(Dataset):\n    def __init__(self, image_ids, img_path, transforms=None):\n        self.image_ids = image_ids\n        self.img_path = img_path\n        self.transforms = transforms\n    \n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self, idx):\n        img_name = self.image_ids[idx]\n        img_path = os.path.join(self.img_path, img_name)\n        \n        image = Image.open(img_path).convert('RGB')\n        image = np.array(image)\n        \n        if self.transforms:\n            image = self.transforms(image=image)['image']\n        \n        return image\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:55.991035Z","iopub.execute_input":"2025-11-27T23:24:55.991268Z","iopub.status.idle":"2025-11-27T23:24:56.004561Z","shell.execute_reply.started":"2025-11-27T23:24:55.991249Z","shell.execute_reply":"2025-11-27T23:24:56.003995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# MIXUP\n# =============================================================================\ndef mixup_data(x, y, alpha=0.2):\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    \n    return mixed_x, y_a, y_b, lam\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:56.005299Z","iopub.execute_input":"2025-11-27T23:24:56.005589Z","iopub.status.idle":"2025-11-27T23:24:56.015641Z","shell.execute_reply.started":"2025-11-27T23:24:56.005569Z","shell.execute_reply":"2025-11-27T23:24:56.014924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# TRAINING\n# =============================================================================\ndef train_one_epoch(model, loader, criterion, optimizer, scheduler, scaler, epoch):\n    model.train()\n    total_loss = 0\n    total_acc = 0\n    \n    pbar = tqdm(loader, desc=f'Epoch {epoch+1} [TRAIN]')\n    optimizer.zero_grad()\n    \n    for i, (images, labels) in enumerate(pbar):\n        images = images.to(cfg.DEVICE)\n        labels = labels.to(cfg.DEVICE)\n        \n        # Mixup\n        if cfg.MIXUP_ALPHA > 0 and np.random.rand() < 0.5:\n            images, labels_a, labels_b, lam = mixup_data(images, labels, cfg.MIXUP_ALPHA)\n            use_mixup = True\n        else:\n            use_mixup = False\n        \n        # Forward with AMP\n        if cfg.USE_AMP:\n            with autocast():\n                outputs = model(images)\n                if use_mixup:\n                    loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n                else:\n                    loss = criterion(outputs, labels)\n                loss = loss / cfg.ACCUMULATION_STEPS\n            \n            scaler.scale(loss).backward()\n            \n            if (i + 1) % cfg.ACCUMULATION_STEPS == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n                if scheduler and cfg.SCHEDULER == \"step\":\n                    scheduler.step()\n        else:\n            outputs = model(images)\n            if use_mixup:\n                loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n            else:\n                loss = criterion(outputs, labels)\n            loss = loss / cfg.ACCUMULATION_STEPS\n            loss.backward()\n            \n            if (i + 1) % cfg.ACCUMULATION_STEPS == 0:\n                optimizer.step()\n                optimizer.zero_grad()\n                if scheduler and cfg.SCHEDULER == \"step\":\n                    scheduler.step()\n        \n        acc = (outputs.argmax(dim=1) == labels).float().mean()\n        total_loss += loss.item() * cfg.ACCUMULATION_STEPS\n        total_acc += acc.item()\n        \n        pbar.set_postfix({'loss': loss.item() * cfg.ACCUMULATION_STEPS, 'acc': acc.item()})\n    \n    return total_loss / len(loader), total_acc / len(loader)\n\ndef validate(model, loader, criterion, epoch):\n    model.eval()\n    total_loss = 0\n    total_acc = 0\n    \n    pbar = tqdm(loader, desc=f'Epoch {epoch+1} [VALID]')\n    \n    with torch.no_grad():\n        for images, labels in pbar:\n            images = images.to(cfg.DEVICE)\n            labels = labels.to(cfg.DEVICE)\n            \n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            acc = (outputs.argmax(dim=1) == labels).float().mean()\n            \n            total_loss += loss.item()\n            total_acc += acc.item()\n            \n            pbar.set_postfix({'loss': loss.item(), 'acc': acc.item()})\n    \n    return total_loss / len(loader), total_acc / len(loader)\n\ndef train_fold(fold, train_df, valid_df):\n    print(f'\\n{\"=\"*70}')\n    print(f'FOLD {fold + 1}')\n    print(f'{\"=\"*70}')\n    \n    # Datasets\n    train_dataset = CassavaDataset(train_df, cfg.TRAIN_PATH, get_train_transforms())\n    valid_dataset = CassavaDataset(valid_df, cfg.TRAIN_PATH, get_valid_transforms())\n    \n    # Loaders\n    train_loader = DataLoader(\n        train_dataset, batch_size=cfg.BATCH_SIZE, shuffle=True,\n        num_workers=cfg.NUM_WORKERS, pin_memory=True, drop_last=True\n    )\n    valid_loader = DataLoader(\n        valid_dataset, batch_size=cfg.BATCH_SIZE * 2, shuffle=False,\n        num_workers=cfg.NUM_WORKERS, pin_memory=True\n    )\n    \n    # Model\n    model = CassavaNet(num_classes=cfg.N_CLASSES)\n    model.to(cfg.DEVICE)\n    \n    # Loss and optimizer\n    criterion = nn.CrossEntropyLoss(label_smoothing=cfg.LABEL_SMOOTHING)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.LR, weight_decay=cfg.WEIGHT_DECAY)\n    \n    # Scheduler\n    if cfg.SCHEDULER == \"cosine\":\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg.N_EPOCHS)\n    else:\n        scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n    \n    scaler = GradScaler() if cfg.USE_AMP else None\n    \n    # Training loop\n    best_acc = 0\n    for epoch in range(cfg.N_EPOCHS):\n        train_loss, train_acc = train_one_epoch(\n            model, train_loader, criterion, optimizer, scheduler, scaler, epoch\n        )\n        valid_loss, valid_acc = validate(model, valid_loader, criterion, epoch)\n        \n        if cfg.SCHEDULER == \"cosine\":\n            scheduler.step()\n        \n        print(f'\\nEpoch {epoch+1}/{cfg.N_EPOCHS}')\n        print(f'Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}')\n        print(f'Valid Loss: {valid_loss:.4f}, Valid Acc: {valid_acc:.4f}')\n        print(f'LR: {optimizer.param_groups[0][\"lr\"]:.6f}')\n        \n        if valid_acc > best_acc:\n            best_acc = valid_acc\n            print(f'✅ New best! Saving model...')\n            torch.save({\n                'model_state_dict': model.state_dict(),\n                'fold': fold,\n                'epoch': epoch,\n                'accuracy': best_acc\n            }, f'cassava_fold{fold}_best.pth')\n    \n    print(f'\\nFold {fold+1} Best Accuracy: {best_acc:.4f}')\n    return best_acc\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:56.017554Z","iopub.execute_input":"2025-11-27T23:24:56.017787Z","iopub.status.idle":"2025-11-27T23:24:56.028932Z","shell.execute_reply.started":"2025-11-27T23:24:56.017772Z","shell.execute_reply":"2025-11-27T23:24:56.028281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# INFERENCE\n# =============================================================================\ndef predict_tta(model, image_path):\n    \"\"\"Predict with Test Time Augmentation\"\"\"\n    model.eval()\n    \n    image = Image.open(image_path).convert('RGB')\n    image = np.array(image)\n    \n    # TTA transforms\n    tta_transforms = [\n        get_valid_transforms(),\n        A.Compose([\n            A.Resize(height=cfg.IMG_SIZE, width=cfg.IMG_SIZE),\n            A.HorizontalFlip(p=1.0),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ]),\n        A.Compose([\n            A.Resize(height=cfg.IMG_SIZE, width=cfg.IMG_SIZE),\n            A.VerticalFlip(p=1.0),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ]),\n        A.Compose([\n            A.Resize(height=cfg.IMG_SIZE, width=cfg.IMG_SIZE),\n            A.Rotate(limit=15, p=1.0),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ]),\n    ]\n    \n    predictions = []\n    with torch.no_grad():\n        for transform in tta_transforms:\n            aug_img = transform(image=image)['image']\n            aug_img = aug_img.unsqueeze(0).to(cfg.DEVICE)\n            output = model(aug_img)\n            probs = F.softmax(output, dim=1)\n            predictions.append(probs.cpu().numpy())\n    \n    return np.mean(predictions, axis=0)\n\ndef inference():\n    print(f'\\n{\"=\"*70}')\n    print('INFERENCE MODE')\n    print(f'{\"=\"*70}')\n    \n    # Find models\n    import glob\n    model_files = sorted(glob.glob('cassava_fold*_best.pth'))\n    \n    if len(model_files) == 0:\n        print(\"❌ ERROR: No trained models found!\")\n        print(\"Please run training first with cfg.MODE = 'train'\")\n        return None\n    \n    print(f'Found {len(model_files)} trained model(s)')\n    \n    # Load test data\n    test_df = pd.read_csv(os.path.join(cfg.DATA_PATH, 'sample_submission.csv'))\n    \n    all_predictions = []\n    \n    for model_file in model_files:\n        print(f'\\nLoading: {model_file}')\n        \n        model = CassavaNet(num_classes=cfg.N_CLASSES)\n        checkpoint = torch.load(model_file)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        model.to(cfg.DEVICE)\n        model.eval()\n        \n        predictions = []\n        \n        for img_id in tqdm(test_df['image_id'], desc='Predicting'):\n            img_path = os.path.join(cfg.TEST_PATH, img_id)\n            \n            if cfg.USE_TTA:\n                pred = predict_tta(model, img_path)\n            else:\n                image = Image.open(img_path).convert('RGB')\n                image = np.array(image)\n                image = get_valid_transforms()(image=image)['image']\n                image = image.unsqueeze(0).to(cfg.DEVICE)\n                \n                with torch.no_grad():\n                    output = model(image)\n                    pred = F.softmax(output, dim=1).cpu().numpy()\n            \n            predictions.append(pred[0])\n        \n        all_predictions.append(np.array(predictions))\n    \n    # Ensemble\n    final_predictions = np.mean(all_predictions, axis=0)\n    final_labels = np.argmax(final_predictions, axis=1)\n    \n    # Create submission\n    submission = pd.DataFrame({\n        'image_id': test_df['image_id'],\n        'label': final_labels\n    })\n    \n    submission.to_csv('submission.csv', index=False)\n    \n    print('\\n✅ Submission saved!')\n    print(f'\\nPrediction distribution:')\n    print(submission.label.value_counts().sort_index())\n    \n    return submission\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:56.029547Z","iopub.execute_input":"2025-11-27T23:24:56.029700Z","iopub.status.idle":"2025-11-27T23:24:56.039784Z","shell.execute_reply.started":"2025-11-27T23:24:56.029688Z","shell.execute_reply":"2025-11-27T23:24:56.039128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# =============================================================================\n# MAIN\n# =============================================================================\ndef main():\n    print(f'Device: {cfg.DEVICE}')\n    print(f'Mode: {cfg.MODE}')\n    print(f'Image Size: {cfg.IMG_SIZE}')\n    print(f'Custom CNN Architecture: CassavaNet')\n    \n    if cfg.MODE == 'train':\n        # Load data\n        df = pd.read_csv(os.path.join(cfg.DATA_PATH, 'train.csv'))\n        \n        # K-Fold CV\n        skf = StratifiedKFold(n_splits=cfg.N_FOLDS, shuffle=True, random_state=cfg.SEED)\n        \n        fold_scores = []\n        \n        for fold, (train_idx, valid_idx) in enumerate(skf.split(df, df['label'])):\n            if fold not in cfg.TRAIN_FOLDS:\n                continue\n            \n            train_df = df.iloc[train_idx]\n            valid_df = df.iloc[valid_idx]\n            \n            best_acc = train_fold(fold, train_df, valid_df)\n            fold_scores.append(best_acc)\n        \n        print(f'\\n{\"=\"*70}')\n        print('TRAINING COMPLETE')\n        print(f'Average CV: {np.mean(fold_scores):.4f} ± {np.std(fold_scores):.4f}')\n        print(f'{\"=\"*70}')\n        \n        print('\\n✅ To run inference, set: cfg.MODE = \"inference\"')\n    \n    elif cfg.MODE == 'inference':\n        submission = inference()\n        if submission is not None:\n            print('\\n✅ Ready to submit!')\n    \n    else:\n        print(f'Invalid MODE: {cfg.MODE}. Use \"train\" or \"inference\"')\n\nif __name__ == '__main__':\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-28T11:44:33.762642Z","iopub.status.idle":"2025-11-28T11:44:33.763014Z","shell.execute_reply.started":"2025-11-28T11:44:33.762884Z","shell.execute_reply":"2025-11-28T11:44:33.762901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# RUN INFERENCE AFTER TRAINING\n# ADD THIS AS A NEW CELL AFTER THE MAIN CELL\n# =============================================================================\n\nimport gc\nimport time\n\n# Clear memory after training\ngc.collect()\ntorch.cuda.empty_cache()\ntime.sleep(2)\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"SWITCHING TO INFERENCE MODE\")\nprint(\"=\"*70)\n\n# Change mode to inference\ncfg.MODE = \"inference\"\n\n# Run main again for inference\nmain()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:56.054876Z","iopub.execute_input":"2025-11-27T23:24:56.055424Z","iopub.status.idle":"2025-11-27T23:24:56.071503Z","shell.execute_reply.started":"2025-11-27T23:24:56.055404Z","shell.execute_reply":"2025-11-27T23:24:56.070801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:56.072295Z","iopub.execute_input":"2025-11-27T23:24:56.072508Z","iopub.status.idle":"2025-11-27T23:24:56.085776Z","shell.execute_reply.started":"2025-11-27T23:24:56.072493Z","shell.execute_reply":"2025-11-27T23:24:56.085108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-27T23:24:56.086567Z","iopub.execute_input":"2025-11-27T23:24:56.086820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}