{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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\n\n# Use the kagglehub client library to attach Kaggle resources like competitions, datasets, and models to your session\n# Learn more about kagglehub: https://github.com/Kaggle/kagglehub/blob/main/README.md\n\nimport kagglehub\n# kagglehub.dataset_download('<owner>/<dataset-slug>')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ResNet50 Binary Skin Lesion Classification\n# Kaggle-safe version with:\n# - num_workers=0\n# - automatic checkpointing\n# - automatic resume\n# - EMA\n# - MixUp\n# - WeightedRandomSampler\n# - 2-phase training\n# - validation AUC\n# - test-time augmentation\n# ============================================================\n\nimport os\nimport copy\nimport random\nimport warnings\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\n\nfrom torch.utils.data import DataLoader, WeightedRandomSampler\nfrom torchvision import datasets, transforms, models\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    confusion_matrix,\n    precision_score,\n    recall_score\n)\n\nfrom tqdm.auto import tqdm\n\n\n# ============================================================\n# 1. SETTINGS\n# ============================================================\n\nwarnings.filterwarnings(\n    \"ignore\",\n    message=\".*does not have many workers.*\"\n)\n\nSEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\n\ntorch.manual_seed(SEED)\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\n\nCUDA = torch.cuda.is_available()\nDEVICE = torch.device(\"cuda\" if CUDA else \"cpu\")\n\nprint(\"Device:\", DEVICE)\nprint(\"CUDA available:\", CUDA)\n\nif CUDA:\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n\n\n# ------------------------------------------------------------\n# Training configuration\n# ------------------------------------------------------------\n\nIMG_SIZE = 224\n\nBATCH_SIZE = 64\n\nHEAD_EPOCHS = 10\nFINE_TUNE_EPOCHS = 40\n\nHEAD_LR = 1e-3\nFINE_TUNE_LR = 3e-5\n\nWEIGHT_DECAY = 1e-4\n\nWARMUP_EPOCHS = 3\n\nOVERSAMPLE_MALIG = 5\n\nMIXUP_ALPHA = 0.3\n\nEMA_DECAY = 0.999\n\nCLASSES = [\"benign\", \"malignant\"]\n\n\n# ============================================================\n# 2. KAGGLE WORKING DIRECTORY\n# ============================================================\n\nWORK = \"/kaggle/working\"\n\nos.makedirs(WORK, exist_ok=True)\n\nBEST_MODEL_PATH = os.path.join(\n    WORK,\n    \"ResNet50_best.pt\"\n)\n\nHEAD_CHECKPOINT_PATH = os.path.join(\n    WORK,\n    \"ResNet50_head_checkpoint.pt\"\n)\n\nFINETUNE_CHECKPOINT_PATH = os.path.join(\n    WORK,\n    \"ResNet50_finetune_checkpoint.pt\"\n)\n\nprint(\"\\nWorking directory:\", WORK)\n\n\n# ============================================================\n# 3. FIND DATASET\n# ============================================================\n\ncandidates = []\n\nfor root, dirs, _files in os.walk(\"/kaggle/input\"):\n\n    dirset = set(d.lower() for d in dirs)\n\n    if (\n        \"train\" in dirset\n        and (\"val\" in dirset or \"validation\" in dirset)\n        and \"test\" in dirset\n    ):\n        candidates.append(root)\n\n\nif not candidates:\n\n    raise FileNotFoundError(\n        \"Could not find a folder under /kaggle/input containing \"\n        \"train/, val(idation)/, and test/ subfolders.\"\n    )\n\n\nDATA_ROOT = candidates[0]\n\nprint(\"\\nUsing dataset root:\")\nprint(DATA_ROOT)\n\n\n# ============================================================\n# 4. RESOLVE TRAIN / VAL / TEST DIRECTORIES\n# ============================================================\n\ndef resolve_split_dir(root, wanted):\n\n    for name in os.listdir(root):\n\n        if (\n            name.lower() == wanted\n            or (\n                wanted == \"val\"\n                and name.lower() == \"validation\"\n            )\n        ):\n            return os.path.join(root, name)\n\n    raise FileNotFoundError(\n        f\"Could not find a '{wanted}' folder under {root}\"\n    )\n\n\nSPLIT_DIRS = {\n    s: resolve_split_dir(DATA_ROOT, s)\n    for s in [\"train\", \"val\", \"test\"]\n}\n\n\nprint(\"\\nDataset splits:\")\n\nfor s, d in SPLIT_DIRS.items():\n    print(f\"{s}: {d}\")\n\n\n# ============================================================\n# 5. TRANSFORMS\n# ============================================================\n\ntrain_tf = transforms.Compose([\n\n    transforms.RandomResizedCrop(\n        IMG_SIZE,\n        scale=(0.75, 1.0),\n        ratio=(0.9, 1.1)\n    ),\n\n    transforms.RandomHorizontalFlip(),\n\n    transforms.RandomVerticalFlip(),\n\n    transforms.RandomRotation(25),\n\n    transforms.ColorJitter(\n        brightness=0.3,\n        contrast=0.2,\n        saturation=0.4,\n        hue=0.02\n    ),\n\n    transforms.RandomAffine(\n        degrees=0,\n        translate=(0.05, 0.05),\n        scale=(0.95, 1.05)\n    ),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        [0.485, 0.456, 0.406],\n        [0.229, 0.224, 0.225]\n    ),\n\n    transforms.RandomErasing(\n        p=0.2,\n        scale=(0.02, 0.08)\n    ),\n])\n\n\neval_tf = transforms.Compose([\n\n    transforms.Resize(\n        (IMG_SIZE, IMG_SIZE)\n    ),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        [0.485, 0.456, 0.406],\n        [0.229, 0.224, 0.225]\n    ),\n])\n\n\n# Test-time horizontal flip\neval_tf_flip = transforms.Compose([\n\n    transforms.Resize(\n        (IMG_SIZE, IMG_SIZE)\n    ),\n\n    transforms.RandomHorizontalFlip(\n        p=1.0\n    ),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        [0.485, 0.456, 0.406],\n        [0.229, 0.224, 0.225]\n    ),\n])\n\n\n# ============================================================\n# 6. DATASETS\n# ============================================================\n\ntrain_ds = datasets.ImageFolder(\n    SPLIT_DIRS[\"train\"],\n    transform=train_tf\n)\n\n\nassert train_ds.class_to_idx == {\n    \"benign\": 0,\n    \"malignant\": 1\n}, train_ds.class_to_idx\n\n\ntargets = np.array(train_ds.targets)\n\n\nprint(\"\\nTraining dataset:\")\nprint(\"Total:\", len(train_ds))\nprint(\"Benign:\", (targets == 0).sum())\nprint(\"Malignant:\", (targets == 1).sum())\n\n\n# ============================================================\n# 7. WEIGHTED OVERSAMPLING\n# ============================================================\n\nweights = np.where(\n    targets == 1,\n    float(OVERSAMPLE_MALIG),\n    1.0\n)\n\n\nsampler = WeightedRandomSampler(\n\n    torch.as_tensor(\n        weights,\n        dtype=torch.double\n    ),\n\n    num_samples=int(weights.sum()),\n\n    replacement=True\n)\n\n\n# ============================================================\n# 8. DATALOADERS\n#\n# IMPORTANT:\n# num_workers=0\n#\n# This is the major Kaggle stability fix.\n# ============================================================\n\nloaders = {}\n\n\nloaders[\"train\"] = DataLoader(\n\n    train_ds,\n\n    batch_size=BATCH_SIZE,\n\n    sampler=sampler,\n\n    # IMPORTANT:\n    # Prevents Kaggle multiprocessing worker crashes\n    num_workers=0,\n\n    pin_memory=CUDA,\n\n    drop_last=True\n)\n\n\neval_ds = {}\n\n\nfor s in [\"val\", \"test\"]:\n\n    ds = datasets.ImageFolder(\n        SPLIT_DIRS[s],\n        transform=eval_tf\n    )\n\n    ds_flip = datasets.ImageFolder(\n        SPLIT_DIRS[s],\n        transform=eval_tf_flip\n    )\n\n\n    assert ds.class_to_idx == {\n        \"benign\": 0,\n        \"malignant\": 1\n    }, ds.class_to_idx\n\n\n    eval_ds[s] = ds\n\n\n    loaders[s] = DataLoader(\n\n        ds,\n\n        batch_size=BATCH_SIZE,\n\n        shuffle=False,\n\n        # IMPORTANT\n        num_workers=0,\n\n        pin_memory=CUDA\n    )\n\n\n    loaders[s + \"_flip\"] = DataLoader(\n\n        ds_flip,\n\n        batch_size=BATCH_SIZE,\n\n        shuffle=False,\n\n        # IMPORTANT\n        num_workers=0,\n\n        pin_memory=CUDA\n    )\n\n\n    print(\n        f\"{s} images: {len(ds)}\"\n    )\n\n\n# ============================================================\n# 9. MODEL\n# ============================================================\n\ndef build_model():\n\n    model = models.resnet50(\n        weights=models.ResNet50_Weights.IMAGENET1K_V2\n    )\n\n    in_features = model.fc.in_features\n\n    model.fc = nn.Sequential(\n\n        nn.Dropout(\n            p=0.4,\n            inplace=True\n        ),\n\n        nn.Linear(\n            in_features,\n            2\n        )\n    )\n\n    return model\n\n\nmodel = build_model().to(DEVICE)\n\n\nprint(\"\\nResNet50 loaded.\")\n\n\n# ============================================================\n# 10. EMA MODEL\n# ============================================================\n\nema_model = copy.deepcopy(model).to(DEVICE)\n\n\nfor p in ema_model.parameters():\n    p.requires_grad_(False)\n\n\n@torch.no_grad()\ndef update_ema(\n    ema_model,\n    model,\n    decay\n):\n\n    for ema_p, p in zip(\n        ema_model.state_dict().values(),\n        model.state_dict().values()\n    ):\n\n        if ema_p.dtype.is_floating_point:\n\n            ema_p.mul_(decay).add_(\n                p.detach(),\n                alpha=1 - decay\n            )\n\n        else:\n\n            ema_p.copy_(p)\n\n\n# ============================================================\n# 11. FREEZE / UNFREEZE BACKBONE\n# ============================================================\n\ndef set_backbone_trainable(\n    model,\n    trainable\n):\n\n    for name, p in model.named_parameters():\n\n        if not name.startswith(\"fc\"):\n\n            p.requires_grad_(trainable)\n\n\n# ============================================================\n# 12. MIXUP\n# ============================================================\n\ndef mixup(\n    x,\n    y,\n    alpha\n):\n\n    if alpha <= 0:\n\n        return (\n            x,\n            y,\n            y,\n            1.0\n        )\n\n\n    lam = float(\n        np.random.beta(\n            alpha,\n            alpha\n        )\n    )\n\n\n    idx = torch.randperm(\n        x.size(0),\n        device=x.device\n    )\n\n\n    x_mixed = (\n        lam * x\n        + (1 - lam) * x[idx]\n    )\n\n\n    return (\n        x_mixed,\n        y,\n        y[idx],\n        lam\n    )\n\n\n# ============================================================\n# 13. EVALUATION\n# ============================================================\n\n@torch.no_grad()\ndef eval_split(\n    model,\n    split,\n    tta=False\n):\n\n    model.eval()\n\n    all_probs = []\n\n    final_trues = None\n\n\n    loader_names = (\n\n        [split, split + \"_flip\"]\n        if tta\n        else\n        [split]\n    )\n\n\n    for loader_name in loader_names:\n\n        loader_probs = []\n\n        trues = []\n\n\n        for x, y in loaders[loader_name]:\n\n            x = x.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n\n            with torch.amp.autocast(\n                \"cuda\",\n                enabled=CUDA\n            ):\n\n                output = model(x)\n\n                probability = torch.softmax(\n                    output.float(),\n                    dim=1\n                )[:, 1]\n\n\n            loader_probs.extend(\n                probability.cpu().tolist()\n            )\n\n            trues.extend(\n                y.tolist()\n            )\n\n\n        all_probs.append(\n            np.array(loader_probs)\n        )\n\n\n        if final_trues is None:\n\n            final_trues = np.array(\n                trues\n            )\n\n\n    probs = np.mean(\n        all_probs,\n        axis=0\n    )\n\n\n    return (\n        final_trues,\n        probs\n    )\n\n\n# ============================================================\n# 14. LOSS + AMP\n# ============================================================\n\ncriterion = nn.CrossEntropyLoss(\n    label_smoothing=0.05\n)\n\n\nscaler = torch.amp.GradScaler(\n    \"cuda\",\n    enabled=CUDA\n)\n\n\n# ============================================================\n# 15. LR SCHEDULER\n# ============================================================\n\ndef make_scheduler(\n    optimizer,\n    warmup_epochs,\n    total_epochs,\n    steps_per_epoch\n):\n\n    warmup_steps = (\n        warmup_epochs\n        * steps_per_epoch\n    )\n\n\n    total_steps = (\n        total_epochs\n        * steps_per_epoch\n    )\n\n\n    def lr_lambda(step):\n\n        if step < warmup_steps:\n\n            return (\n                step\n                / max(\n                    1,\n                    warmup_steps\n                )\n            )\n\n\n        progress = (\n            step - warmup_steps\n        ) / max(\n            1,\n            total_steps - warmup_steps\n        )\n\n\n        return 0.5 * (\n            1\n            + np.cos(\n                np.pi * progress\n            )\n        )\n\n\n    return torch.optim.lr_scheduler.LambdaLR(\n        optimizer,\n        lr_lambda\n    )\n\n\n# ============================================================\n# 16. SAVE CHECKPOINT\n# ============================================================\n\ndef save_checkpoint(\n    path,\n    model,\n    ema_model,\n    optimizer,\n    scheduler,\n    scaler,\n    epoch,\n    best_auc,\n    phase_name\n):\n\n    checkpoint = {\n\n        \"epoch\": epoch,\n\n        \"phase_name\": phase_name,\n\n        \"model_state_dict\":\n            model.state_dict(),\n\n        \"ema_state_dict\":\n            ema_model.state_dict(),\n\n        \"optimizer_state_dict\":\n            optimizer.state_dict(),\n\n        \"scheduler_state_dict\":\n            scheduler.state_dict(),\n\n        \"scaler_state_dict\":\n            scaler.state_dict(),\n\n        \"best_auc\":\n            best_auc,\n\n        \"seed\":\n            SEED\n    }\n\n\n    torch.save(\n        checkpoint,\n        path\n    )\n\n\n# ============================================================\n# 17. TRAINING PHASE\n# ============================================================\n\ndef train_phase(\n    model,\n    epochs,\n    lr,\n    warmup_epochs,\n    phase_name,\n    checkpoint_path,\n    resume=True\n):\n\n    global ema_model\n    global scaler\n\n\n    # --------------------------------------------------------\n    # Optimizer\n    # --------------------------------------------------------\n\n    optimizer = torch.optim.AdamW(\n\n        filter(\n            lambda p: p.requires_grad,\n            model.parameters()\n        ),\n\n        lr=lr,\n\n        weight_decay=WEIGHT_DECAY\n    )\n\n\n    # --------------------------------------------------------\n    # Scheduler\n    # --------------------------------------------------------\n\n    scheduler = make_scheduler(\n\n        optimizer,\n\n        warmup_epochs,\n\n        epochs,\n\n        len(loaders[\"train\"])\n    )\n\n\n    # --------------------------------------------------------\n    # Resume variables\n    # --------------------------------------------------------\n\n    start_epoch = 0\n\n    best_auc = -1.0\n\n    best_state = None\n\n\n    # --------------------------------------------------------\n    # Try to resume\n    # --------------------------------------------------------\n\n    if resume and os.path.exists(\n        checkpoint_path\n    ):\n\n        print(\n            f\"\\nFound checkpoint:\"\n        )\n\n        print(checkpoint_path)\n\n        print(\"Attempting to resume...\")\n\n\n        checkpoint = torch.load(\n            checkpoint_path,\n            map_location=DEVICE\n        )\n\n\n        saved_phase = checkpoint.get(\n            \"phase_name\",\n            None\n        )\n\n\n        if saved_phase == phase_name:\n\n            model.load_state_dict(\n                checkpoint[\"model_state_dict\"]\n            )\n\n\n            ema_model.load_state_dict(\n                checkpoint[\"ema_state_dict\"]\n            )\n\n\n            optimizer.load_state_dict(\n                checkpoint[\"optimizer_state_dict\"]\n            )\n\n\n            scheduler.load_state_dict(\n                checkpoint[\"scheduler_state_dict\"]\n            )\n\n\n            if (\n                \"scaler_state_dict\"\n                in checkpoint\n            ):\n\n                scaler.load_state_dict(\n                    checkpoint[\n                        \"scaler_state_dict\"\n                    ]\n                )\n\n\n            start_epoch = (\n                checkpoint[\"epoch\"] + 1\n            )\n\n\n            best_auc = checkpoint.get(\n                \"best_auc\",\n                -1.0\n            )\n\n\n            print(\n                f\"Resuming {phase_name} \"\n                f\"from epoch \"\n                f\"{start_epoch + 1}/{epochs}\"\n            )\n\n        else:\n\n            print(\n                \"Checkpoint belongs to \"\n                f\"phase '{saved_phase}', \"\n                f\"not '{phase_name}'.\"\n            )\n\n            print(\n                \"Starting this phase \"\n                \"from the beginning.\"\n            )\n\n\n    # --------------------------------------------------------\n    # Training loop\n    # --------------------------------------------------------\n\n    for ep in range(\n        start_epoch,\n        epochs\n    ):\n\n        model.train()\n\n\n        running_loss = 0.0\n\n        batch_count = 0\n\n\n        progress_bar = tqdm(\n\n            loaders[\"train\"],\n\n            desc=(\n                f\"[{phase_name}] \"\n                f\"epoch {ep + 1}/{epochs}\"\n            ),\n\n            leave=True\n        )\n\n\n        for x, y in progress_bar:\n\n            x = x.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            y = y.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n\n            # ------------------------------------------------\n            # MixUp\n            # ------------------------------------------------\n\n            (\n                x_mixed,\n                y_a,\n                y_b,\n                lam\n            ) = mixup(\n                x,\n                y,\n                MIXUP_ALPHA\n            )\n\n\n            # ------------------------------------------------\n            # Zero gradients\n            # ------------------------------------------------\n\n            optimizer.zero_grad(\n                set_to_none=True\n            )\n\n\n            # ------------------------------------------------\n            # Forward\n            # ------------------------------------------------\n\n            with torch.amp.autocast(\n                \"cuda\",\n                enabled=CUDA\n            ):\n\n                output = model(\n                    x_mixed\n                )\n\n\n                loss = (\n                    lam\n                    * criterion(\n                        output,\n                        y_a\n                    )\n                    +\n                    (1 - lam)\n                    * criterion(\n                        output,\n                        y_b\n                    )\n                )\n\n\n            # ------------------------------------------------\n            # Backward\n            # ------------------------------------------------\n\n            scaler.scale(\n                loss\n            ).backward()\n\n\n            scaler.step(\n                optimizer\n            )\n\n\n            scaler.update()\n\n\n            scheduler.step()\n\n\n            # ------------------------------------------------\n            # EMA\n            # ------------------------------------------------\n\n            update_ema(\n                ema_model,\n                model,\n                EMA_DECAY\n            )\n\n\n            # ------------------------------------------------\n            # Loss tracking\n            # ------------------------------------------------\n\n            running_loss += loss.item()\n\n            batch_count += 1\n\n\n            progress_bar.set_postfix(\n                loss=f\"{loss.item():.4f}\"\n            )\n\n\n        # ====================================================\n        # VALIDATION\n        # ====================================================\n\n        val_true, val_prob = eval_split(\n            ema_model,\n            \"val\",\n            tta=False\n        )\n\n\n        auc = roc_auc_score(\n            val_true,\n            val_prob\n        )\n\n\n        avg_loss = (\n            running_loss\n            / max(\n                1,\n                batch_count\n            )\n        )\n\n\n        current_lr = optimizer.param_groups[0][\"lr\"]\n\n\n        print(\n            f\"\\n[{phase_name}] \"\n            f\"epoch {ep + 1:>2}/{epochs}\"\n        )\n\n        print(\n            f\"Loss: {avg_loss:.5f}\"\n        )\n\n        print(\n            f\"LR: {current_lr:.8f}\"\n        )\n\n        print(\n            f\"Val AUC (EMA): {auc:.4f}\"\n        )\n\n\n        # ====================================================\n        # BEST MODEL\n        # ====================================================\n\n        if auc > best_auc:\n\n            best_auc = auc\n\n\n            best_state = {\n\n                k: v.detach()\n                .cpu()\n                .clone()\n\n                for k, v\n                in ema_model.state_dict().items()\n            }\n\n\n            torch.save(\n                best_state,\n                BEST_MODEL_PATH\n            )\n\n\n            print(\n                \"★ New best model saved!\"\n            )\n\n            print(\n                f\"Best AUC: {best_auc:.4f}\"\n            )\n\n\n        # ====================================================\n        # SAVE CHECKPOINT AFTER EVERY EPOCH\n        # ====================================================\n\n        save_checkpoint(\n\n            checkpoint_path,\n\n            model,\n\n            ema_model,\n\n            optimizer,\n\n            scheduler,\n\n            scaler,\n\n            ep,\n\n            best_auc,\n\n            phase_name\n        )\n\n\n        print(\n            f\"Checkpoint saved:\"\n            f\" {checkpoint_path}\"\n        )\n\n\n        # ====================================================\n        # MEMORY CLEANUP\n        # ====================================================\n\n        if CUDA:\n\n            torch.cuda.empty_cache()\n\n\n    return (\n        best_auc,\n        best_state\n    )\n\n\n# ============================================================\n# 18. PHASE 1\n# ============================================================\n\nprint(\n    \"\\n\"\n    \"=\" * 60\n)\n\nprint(\n    \"PHASE 1:\"\n)\n\nprint(\n    \"Training classifier head\"\n)\n\nprint(\n    \"Backbone frozen\"\n)\n\nprint(\n    \"=\" * 60\n)\n\n\nset_backbone_trainable(\n    model,\n    False\n)\n\n\n# Make EMA match current model\n\nema_model.load_state_dict(\n    model.state_dict()\n)\n\n\nbest_auc1, best_state1 = train_phase(\n\n    model=model,\n\n    epochs=HEAD_EPOCHS,\n\n    lr=HEAD_LR,\n\n    warmup_epochs=1,\n\n    phase_name=\"head\",\n\n    checkpoint_path=HEAD_CHECKPOINT_PATH,\n\n    resume=True\n)\n\n\n# ============================================================\n# 19. LOAD BEST PHASE 1 MODEL\n# ============================================================\n\nif best_state1 is None:\n\n    # If the phase was already completed and\n    # best_state1 is unavailable in this Python session,\n    # load the best saved model.\n\n    if os.path.exists(\n        BEST_MODEL_PATH\n    ):\n\n        print(\n            \"\\nLoading best saved model \"\n            \"from Phase 1.\"\n        )\n\n        best_state1 = torch.load(\n            BEST_MODEL_PATH,\n            map_location=\"cpu\"\n        )\n\n    else:\n\n        raise RuntimeError(\n            \"Could not find Phase 1 best model.\"\n        )\n\n\nmodel.load_state_dict(\n    best_state1\n)\n\nema_model.load_state_dict(\n    best_state1\n)\n\n\n# ============================================================\n# 20. PHASE 2\n# ============================================================\n\nprint(\n    \"\\n\"\n    \"=\" * 60\n)\n\nprint(\n    \"PHASE 2:\"\n)\n\nprint(\n    \"Fine-tuning full ResNet50\"\n)\n\nprint(\n    \"Backbone unfrozen\"\n)\n\nprint(\n    \"=\" * 60\n)\n\n\nset_backbone_trainable(\n    model,\n    True\n)\n\n\nbest_auc2, best_state2 = train_phase(\n\n    model=model,\n\n    epochs=FINE_TUNE_EPOCHS,\n\n    lr=FINE_TUNE_LR,\n\n    warmup_epochs=WARMUP_EPOCHS,\n\n    phase_name=\"finetune\",\n\n    checkpoint_path=FINETUNE_CHECKPOINT_PATH,\n\n    resume=True\n)\n\n\n# ============================================================\n# 21. SELECT BEST MODEL\n# ============================================================\n\nbest_auc = max(\n    best_auc1,\n    best_auc2\n)\n\n\nif (\n    best_auc2 >= best_auc1\n    and best_state2 is not None\n):\n\n    best_state = best_state2\n\nelse:\n\n    best_state = best_state1\n\n\n# If best_state is unavailable, load saved best model\n\nif best_state is None:\n\n    print(\n        \"\\nLoading best model from disk...\"\n    )\n\n    best_state = torch.load(\n        BEST_MODEL_PATH,\n        map_location=\"cpu\"\n    )\n\n\nema_model.load_state_dict(\n    best_state\n)\n\n\nprint(\n    \"\\n\"\n    + \"=\" * 60\n)\n\nprint(\n    f\"BEST VALIDATION AUC: {best_auc:.4f}\"\n)\n\nprint(\n    \"=\" * 60\n)\n\n\n# ============================================================\n# 22. TEST EVALUATION WITH TTA\n# ============================================================\n\nprint(\n    \"\\nRunning test evaluation...\"\n)\n\n\ntest_true, test_prob = eval_split(\n\n    ema_model,\n\n    \"test\",\n\n    tta=True\n)\n\n\ntest_auc = roc_auc_score(\n    test_true,\n    test_prob\n)\n\n\n# ============================================================\n# 23. CLASSIFICATION METRICS\n# ============================================================\n\ny_pred = (\n    test_prob >= 0.5\n).astype(int)\n\n\ntn, fp, fn, tp = confusion_matrix(\n\n    test_true,\n\n    y_pred,\n\n    labels=[0, 1]\n\n).ravel()\n\n\nrecall = recall_score(\n\n    test_true,\n\n    y_pred,\n\n    zero_division=0\n)\n\n\nprecision = precision_score(\n\n    test_true,\n\n    y_pred,\n\n    zero_division=0\n)\n\n\n# ============================================================\n# 24. PRINT RESULTS\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 60\n)\n\nprint(\n    \"FINAL TEST RESULTS\"\n)\n\nprint(\n    \"=\" * 60\n)\n\nprint(\n    f\"Test AUC (EMA + TTA): \"\n    f\"{test_auc:.4f}\"\n)\n\nprint(\n    f\"Threshold: 0.5\"\n)\n\nprint(\n    f\"TP = {tp}\"\n)\n\nprint(\n    f\"TN = {tn}\"\n)\n\nprint(\n    f\"FP = {fp}\"\n)\n\nprint(\n    f\"FN = {fn}\"\n)\n\nprint(\n    f\"Recall = {recall:.4f}\"\n)\n\nprint(\n    f\"Precision = {precision:.4f}\"\n)\n\n\n# ============================================================\n# 25. SAVE PREDICTIONS\n# ============================================================\n\npredictions_path = os.path.join(\n    WORK,\n    \"preds_ResNet50.npz\"\n)\n\n\nnp.savez(\n\n    predictions_path,\n\n    y_true=test_true,\n\n    y_prob=test_prob\n)\n\n\n# ============================================================\n# 26. FINAL FILE INFORMATION\n# ============================================================\n\nprint(\n    \"\\n\"\n    + \"=\" * 60\n)\n\nprint(\n    \"FILES SAVED\"\n)\n\nprint(\n    \"=\" * 60\n)\n\nprint(\n    f\"Best model:\"\n)\n\nprint(\n    BEST_MODEL_PATH\n)\n\nprint(\n    f\"\\nHead checkpoint:\"\n)\n\nprint(\n    HEAD_CHECKPOINT_PATH\n)\n\nprint(\n    f\"\\nFine-tune checkpoint:\"\n)\n\nprint(\n    FINETUNE_CHECKPOINT_PATH\n)\n\nprint(\n    f\"\\nPredictions:\"\n)\n\nprint(\n    predictions_path\n)\n\nprint(\n    \"\\nTraining completed successfully.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:46:49.813787Z","iopub.execute_input":"2026-09-07T11:46:49.814112Z","iopub.status.idle":"2026-09-07T12:32:15.973641Z","shell.execute_reply.started":"2026-09-07T11:46:49.814086Z","shell.execute_reply":"2026-09-07T12:32:15.972697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# ResNet50 Model Evaluation Suite\n# Binary Classification:\n#   0 = benign\n#   1 = malignant\n#\n# Evaluations:\n#   1. Confusion Matrix\n#   2. Class-wise Precision / Recall / F1\n#   3. ROC Curve + AUC\n#   4. Precision-Recall Curve\n#   5. Normalized Confusion Matrix\n#   6. F1 Score vs Confidence Threshold\n# ============================================================\n\n\n# ============================================================\n# 0. IMPORTS\n# ============================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfrom sklearn.metrics import (\n    confusion_matrix,\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_curve,\n    roc_auc_score,\n    precision_recall_curve,\n    average_precision_score,\n    classification_report\n)\n\n\nsns.set_style(\"whitegrid\")\n\n\n# ============================================================\n# 1. LOAD RESNET50 PREDICTIONS\n# ============================================================\n\nPRED_PATH = \"/kaggle/working/preds_ResNet50.npz\"\n\n\nif not os.path.exists(PRED_PATH):\n\n    raise FileNotFoundError(\n        f\"Could not find:\\n{PRED_PATH}\\n\\n\"\n        \"Make sure you have run the ResNet50 training/evaluation \"\n        \"cell first.\"\n    )\n\n\ndata = np.load(PRED_PATH)\n\n\ny_true = data[\"y_true\"]\ny_prob = data[\"y_prob\"]\n\n\n# ------------------------------------------------------------\n# Your training code saves malignant probability in y_prob.\n# y_prob is therefore a 1D array:\n#\n# y_prob[i] = probability that image i is malignant\n# ------------------------------------------------------------\n\ny_true = np.asarray(y_true).astype(int)\ny_prob = np.asarray(y_prob).astype(float)\n\n\n# Convert probability to prediction using threshold 0.5\n\ny_pred = (\n    y_prob >= 0.5\n).astype(int)\n\n\nclass_names = [\n    \"benign\",\n    \"malignant\"\n]\n\n\nprint(\"=\" * 60)\nprint(\"RESNET50 EVALUATION\")\nprint(\"=\" * 60)\n\nprint(f\"Number of test samples: {len(y_true)}\")\n\nprint(\n    f\"Benign samples: {(y_true == 0).sum()}\"\n)\n\nprint(\n    f\"Malignant samples: {(y_true == 1).sum()}\"\n)\n\nprint(\n    f\"Probability range: \"\n    f\"{y_prob.min():.4f} - {y_prob.max():.4f}\"\n)\n\nprint(\n    f\"Prediction threshold: 0.50\"\n)\n\n\n# ============================================================\n# 2. CONFUSION MATRIX\n# ============================================================\n\ncm = confusion_matrix(\n    y_true,\n    y_pred,\n    labels=[0, 1]\n)\n\n\nplt.figure(figsize=(7, 6))\n\n\nsns.heatmap(\n    cm,\n    annot=True,\n    fmt=\"d\",\n    cmap=\"Blues\",\n    xticklabels=class_names,\n    yticklabels=class_names,\n    cbar=True\n)\n\n\nplt.xlabel(\n    \"Predicted Label\",\n    fontsize=12\n)\n\nplt.ylabel(\n    \"True Label\",\n    fontsize=12\n)\n\nplt.title(\n    \"ResNet50 Confusion Matrix\",\n    fontsize=15,\n    fontweight=\"bold\"\n)\n\nplt.tight_layout()\nplt.show()\n\n\n# ------------------------------------------------------------\n# Print TP/TN/FP/FN\n# ------------------------------------------------------------\n\ntn, fp, fn, tp = cm.ravel()\n\n\nprint(\"\\nConfusion Matrix Values\")\nprint(\"-\" * 40)\n\nprint(f\"True Negative  (TN): {tn}\")\nprint(f\"False Positive (FP): {fp}\")\nprint(f\"False Negative (FN): {fn}\")\nprint(f\"True Positive  (TP): {tp}\")\n\n\n# ============================================================\n# 3. CLASS-WISE PRECISION / RECALL / F1\n# ============================================================\n\nprecision = precision_score(\n    y_true,\n    y_pred,\n    labels=[0, 1],\n    average=None,\n    zero_division=0\n)\n\n\nrecall = recall_score(\n    y_true,\n    y_pred,\n    labels=[0, 1],\n    average=None,\n    zero_division=0\n)\n\n\nf1 = f1_score(\n    y_true,\n    y_pred,\n    labels=[0, 1],\n    average=None,\n    zero_division=0\n)\n\n\nmetrics_df = pd.DataFrame(\n\n    {\n        \"Precision\": precision,\n        \"Recall\": recall,\n        \"F1-score\": f1\n    },\n\n    index=class_names\n)\n\n\nprint(\"\\n\")\nprint(\"=\" * 60)\nprint(\"CLASS-WISE PERFORMANCE\")\nprint(\"=\" * 60)\n\ndisplay(\n    metrics_df.round(4)\n)\n\n\n# ------------------------------------------------------------\n# Heatmap\n# ------------------------------------------------------------\n\nplt.figure(figsize=(8, 4))\n\n\nsns.heatmap(\n    metrics_df,\n    annot=True,\n    fmt=\".3f\",\n    cmap=\"YlGnBu\",\n    vmin=0,\n    vmax=1\n)\n\n\nplt.title(\n    \"Class-wise Precision / Recall / F1\",\n    fontsize=15,\n    fontweight=\"bold\"\n)\n\nplt.xlabel(\"Metrics\")\nplt.ylabel(\"Class\")\n\nplt.tight_layout()\nplt.show()\n\n\n# ------------------------------------------------------------\n# Classification report\n# ------------------------------------------------------------\n\nprint(\"\\nClassification Report\")\nprint(\"=\" * 60)\n\nprint(\n    classification_report(\n        y_true,\n        y_pred,\n        labels=[0, 1],\n        target_names=class_names,\n        digits=4,\n        zero_division=0\n    )\n)\n\n\n# ============================================================\n# 4. ROC CURVE + AUC\n# ============================================================\n\nfpr, tpr, thresholds_roc = roc_curve(\n    y_true,\n    y_prob\n)\n\n\nroc_auc = roc_auc_score(\n    y_true,\n    y_prob\n)\n\n\nplt.figure(figsize=(8, 7))\n\n\nplt.plot(\n    fpr,\n    tpr,\n    lw=2.5,\n    label=f\"ResNet50 (AUC = {roc_auc:.4f})\"\n)\n\n\nplt.plot(\n    [0, 1],\n    [0, 1],\n    \"k--\",\n    lw=1.5,\n    label=\"Random classifier\"\n)\n\n\nplt.xlim(\n    [0.0, 1.0]\n)\n\nplt.ylim(\n    [0.0, 1.05]\n)\n\n\nplt.xlabel(\n    \"False Positive Rate\",\n    fontsize=12\n)\n\nplt.ylabel(\n    \"True Positive Rate\",\n    fontsize=12\n)\n\n\nplt.title(\n    \"ROC Curve — ResNet50\",\n    fontsize=15,\n    fontweight=\"bold\"\n)\n\n\nplt.legend(\n    loc=\"lower right\"\n)\n\n\nplt.tight_layout()\nplt.show()\n\n\nprint(\n    f\"\\nROC AUC = {roc_auc:.4f}\"\n)\n\n\n# ============================================================\n# 5. PRECISION-RECALL CURVE\n# ============================================================\n\npr_precision, pr_recall, pr_thresholds = (\n    precision_recall_curve(\n        y_true,\n        y_prob\n    )\n)\n\n\naverage_precision = average_precision_score(\n    y_true,\n    y_prob\n)\n\n\nplt.figure(figsize=(8, 7))\n\n\nplt.plot(\n    pr_recall,\n    pr_precision,\n    lw=2.5,\n    label=(\n        f\"ResNet50 \"\n        f\"(AP = {average_precision:.4f})\"\n    )\n)\n\n\nplt.xlabel(\n    \"Recall\",\n    fontsize=12\n)\n\nplt.ylabel(\n    \"Precision\",\n    fontsize=12\n)\n\n\nplt.title(\n    \"Precision-Recall Curve — ResNet50\",\n    fontsize=15,\n    fontweight=\"bold\"\n)\n\n\nplt.legend(\n    loc=\"lower left\"\n)\n\n\nplt.xlim(\n    [0.0, 1.0]\n)\n\nplt.ylim(\n    [0.0, 1.05]\n)\n\n\nplt.tight_layout()\nplt.show()\n\n\nprint(\n    f\"Average Precision (AP) = \"\n    f\"{average_precision:.4f}\"\n)\n\n\n# ============================================================\n# 6. NORMALIZED CONFUSION MATRIX\n# ============================================================\n\ncm_normalized = confusion_matrix(\n\n    y_true,\n\n    y_pred,\n\n    labels=[0, 1],\n\n    normalize=\"true\"\n)\n\n\nplt.figure(figsize=(7, 6))\n\n\nsns.heatmap(\n\n    cm_normalized,\n\n    annot=True,\n\n    fmt=\".3f\",\n\n    cmap=\"Purples\",\n\n    xticklabels=class_names,\n\n    yticklabels=class_names,\n\n    vmin=0,\n\n    vmax=1\n)\n\n\nplt.xlabel(\n    \"Predicted Label\",\n    fontsize=12\n)\n\nplt.ylabel(\n    \"True Label\",\n    fontsize=12\n)\n\n\nplt.title(\n    \"Normalized Confusion Matrix — ResNet50\",\n    fontsize=15,\n    fontweight=\"bold\"\n)\n\n\nplt.tight_layout()\nplt.show()\n\n\n# ============================================================\n# 7. F1 SCORE VS CONFIDENCE THRESHOLD\n# ============================================================\n\nthresholds = np.linspace(\n    0.01,\n    0.99,\n    99\n)\n\n\nf1_scores = []\nprecision_scores = []\nrecall_scores = []\n\n\nfor threshold in thresholds:\n\n    predictions = (\n        y_prob >= threshold\n    ).astype(int)\n\n\n    f1_scores.append(\n        f1_score(\n            y_true,\n            predictions,\n            zero_division=0\n        )\n    )\n\n\n    precision_scores.append(\n        precision_score(\n            y_true,\n            predictions,\n            zero_division=0\n        )\n    )\n\n\n    recall_scores.append(\n        recall_score(\n            y_true,\n            predictions,\n            zero_division=0\n        )\n    )\n\n\nf1_scores = np.array(\n    f1_scores\n)\n\nprecision_scores = np.array(\n    precision_scores\n)\n\nrecall_scores = np.array(\n    recall_scores\n)\n\n\n# ------------------------------------------------------------\n# Best threshold\n# ------------------------------------------------------------\n\nbest_idx = np.argmax(\n    f1_scores\n)\n\n\nbest_threshold = thresholds[\n    best_idx\n]\n\nbest_f1 = f1_scores[\n    best_idx\n]\n\nbest_precision = precision_scores[\n    best_idx\n]\n\nbest_recall = recall_scores[\n    best_idx\n]\n\n\n# ------------------------------------------------------------\n# Plot\n# ------------------------------------------------------------\n\nplt.figure(figsize=(9, 7))\n\n\nplt.plot(\n    thresholds,\n    f1_scores,\n    lw=2.5,\n    label=\"F1 Score\"\n)\n\n\nplt.plot(\n    thresholds,\n    precision_scores,\n    lw=2,\n    linestyle=\"--\",\n    label=\"Precision\"\n)\n\n\nplt.plot(\n    thresholds,\n    recall_scores,\n    lw=2,\n    linestyle=\":\",\n    label=\"Recall\"\n)\n\n\nplt.axvline(\n    best_threshold,\n    linestyle=\"--\",\n    lw=1.5,\n    label=(\n        f\"Best F1 threshold = \"\n        f\"{best_threshold:.2f}\"\n    )\n)\n\n\nplt.scatter(\n    [best_threshold],\n    [best_f1],\n    s=80,\n    zorder=5\n)\n\n\nplt.xlabel(\n    \"Confidence Threshold\",\n    fontsize=12\n)\n\nplt.ylabel(\n    \"Score\",\n    fontsize=12\n)\n\n\nplt.title(\n    \"F1 / Precision / Recall vs Confidence Threshold\",\n    fontsize=15,\n    fontweight=\"bold\"\n)\n\n\nplt.xlim(\n    0,\n    1\n)\n\nplt.ylim(\n    0,\n    1.05\n)\n\n\nplt.legend(\n    loc=\"best\"\n)\n\n\nplt.tight_layout()\nplt.show()\n\n\n# ============================================================\n# 8. BEST THRESHOLD RESULTS\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 60)\nprint(\"OPTIMAL CONFIDENCE THRESHOLD\")\nprint(\"=\" * 60)\n\nprint(\n    f\"Best threshold : {best_threshold:.2f}\"\n)\n\nprint(\n    f\"F1-score       : {best_f1:.4f}\"\n)\n\nprint(\n    f\"Precision      : {best_precision:.4f}\"\n)\n\nprint(\n    f\"Recall         : {best_recall:.4f}\"\n)\n\n\n# ============================================================\n# 9. COMPARE 0.50 VS OPTIMAL THRESHOLD\n# ============================================================\n\nbest_predictions = (\n    y_prob >= best_threshold\n).astype(int)\n\n\ndefault_f1 = f1_score(\n    y_true,\n    y_pred,\n    zero_division=0\n)\n\n\ndefault_precision = precision_score(\n    y_true,\n    y_pred,\n    zero_division=0\n)\n\n\ndefault_recall = recall_score(\n    y_true,\n    y_pred,\n    zero_division=0\n)\n\n\ncomparison_df = pd.DataFrame(\n\n    {\n        \"Threshold\": [\n            0.50,\n            best_threshold\n        ],\n\n        \"Precision\": [\n            default_precision,\n            best_precision\n        ],\n\n        \"Recall\": [\n            default_recall,\n            best_recall\n        ],\n\n        \"F1-score\": [\n            default_f1,\n            best_f1\n        ]\n    },\n\n    index=[\n        \"Default threshold\",\n        \"Optimal F1 threshold\"\n    ]\n)\n\n\nprint(\"\\n\")\nprint(\"=\" * 60)\nprint(\"THRESHOLD COMPARISON\")\nprint(\"=\" * 60)\n\ndisplay(\n    comparison_df.round(4)\n)\n\n\n# ============================================================\n# 10. FINAL SUMMARY\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 60)\nprint(\"FINAL RESNET50 EVALUATION SUMMARY\")\nprint(\"=\" * 60)\n\nprint(\n    f\"ROC AUC              : {roc_auc:.4f}\"\n)\n\nprint(\n    f\"Average Precision    : {average_precision:.4f}\"\n)\n\nprint(\n    f\"F1 @ 0.50            : {default_f1:.4f}\"\n)\n\nprint(\n    f\"Precision @ 0.50     : {default_precision:.4f}\"\n)\n\nprint(\n    f\"Recall @ 0.50        : {default_recall:.4f}\"\n)\n\nprint(\n    f\"Best F1 threshold    : {best_threshold:.2f}\"\n)\n\nprint(\n    f\"Best F1              : {best_f1:.4f}\"\n)\n\nprint(\n    f\"Precision @ best     : {best_precision:.4f}\"\n)\n\nprint(\n    f\"Recall @ best        : {best_recall:.4f}\"\n)\n\nprint(\n    f\"\\nTP = {tp} | TN = {tn} | \"\n    f\"FP = {fp} | FN = {fn}\"\n)\n\nprint(\"=\" * 60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T12:38:28.869637Z","iopub.execute_input":"2026-09-07T12:38:28.870142Z","iopub.status.idle":"2026-09-07T12:38:30.316082Z","shell.execute_reply.started":"2026-09-07T12:38:28.870114Z","shell.execute_reply":"2026-09-07T12:38:30.315354Z"}},"outputs":[],"execution_count":null}]}