{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10418,"databundleVersionId":862236}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nimport numpy as np\nimport pandas as pd\n\nimport einops\n\nimport cv2\nfrom PIL import Image\nimport albumentations as A\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torch.nn.functional as F\nfrom torchvision.models import resnet34, ResNet34_Weights\n\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport time\nimport tempfile\nimport torch.optim as optim\nfrom sklearn.metrics import f1_score\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:19.121079Z","iopub.execute_input":"2026-05-07T03:08:19.121434Z","iopub.status.idle":"2026-05-07T03:08:31.971471Z","shell.execute_reply.started":"2026-05-07T03:08:19.121407Z","shell.execute_reply":"2026-05-07T03:08:31.970513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"p = \"/kaggle/input/competitions/human-protein-atlas-image-classification/train/\"\nimages = []\n\nfor n in [\n    \"00070df0-bbc3-11e8-b2bc-ac1f6b6435d0\", \n    \"000a6c98-bb9b-11e8-b2b9-ac1f6b6435d0\", \n    \"000a9596-bbc4-11e8-b2bc-ac1f6b6435d0\", \n    \"000c99ba-bba4-11e8-b2b9-ac1f6b6435d0\"\n]:\n    image = [\n        cv2.imread(f\"{p}{n}_{c}.png\", cv2.IMREAD_GRAYSCALE).astype(np.float32)/255 for c in ['red','green','blue','yellow']\n    ]\n    images += [np.stack(image, axis=-1)]\n\nimages[0].shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:37.921444Z","iopub.execute_input":"2026-05-07T03:08:37.922062Z","iopub.status.idle":"2026-05-07T03:08:38.004464Z","shell.execute_reply.started":"2026-05-07T03:08:37.922028Z","shell.execute_reply":"2026-05-07T03:08:38.003751Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualization(images):\n    \n    fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n    \n    for i, image in enumerate(images):\n        \n        composite = np.zeros_like(image[:,:,:3])\n        \n        composite[:,:,0] = image[:,:,0] + image[:,:,3] * .5\n        composite[:,:,1] = image[:,:,1] + image[:,:,3] * .5\n        composite[:,:,2] = image[:,:,2]\n        \n        composite = np.clip(composite, 0, 1)\n        \n        axes[i].imshow(composite)\n        axes[i].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\nvisualization(images)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:41.205281Z","iopub.execute_input":"2026-05-07T03:08:41.206019Z","iopub.status.idle":"2026-05-07T03:08:42.055556Z","shell.execute_reply.started":"2026-05-07T03:08:41.205986Z","shell.execute_reply":"2026-05-07T03:08:42.054362Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = A.Compose([\n    A.HorizontalFlip(p=.5),\n    A.VerticalFlip(p=.3),\n    A.ShiftScaleRotate(\n        shift_limit=.1, \n        scale_limit=.2, \n        rotate_limit=30, \n        p=.5\n    )\n])\n\nvisualization([transform(image=image)['image'] for image in images])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:45.525016Z","iopub.execute_input":"2026-05-07T03:08:45.525338Z","iopub.status.idle":"2026-05-07T03:08:46.267990Z","shell.execute_reply.started":"2026-05-07T03:08:45.525312Z","shell.execute_reply":"2026-05-07T03:08:46.267151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ProteinAtlasDataset(Dataset):\n    def __init__(self, csv_file, root_dir, transform=None):\n        # works for BOTH file path and dataframe\n        self.df = pd.read_csv(csv_file) if isinstance(csv_file, str) else csv_file.copy()\n\n        self.root_dir = root_dir\n        self.transform = transform\n        self.num_classes = 28\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        img_id = self.df.iloc[idx, 0]\n        label_str = str(self.df.iloc[idx, 1]).split()\n\n        channels = ['red', 'green', 'blue', 'yellow']\n        img_arrays = []\n\n        for ch in channels:\n            img_path = os.path.join(self.root_dir, f\"{img_id}_{ch}.png\")\n            img = Image.open(img_path).convert('L')\n            img_arrays.append(np.array(img))\n\n        image = np.stack(img_arrays, axis=0).astype(np.float32) / 255.0\n        image = torch.tensor(image, dtype=torch.float32)\n\n        target = torch.zeros(self.num_classes, dtype=torch.float32)\n        for label in label_str:\n            if label != 'nan':\n                target[int(label)] = 1.0\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:51.251730Z","iopub.execute_input":"2026-05-07T03:08:51.252663Z","iopub.status.idle":"2026-05-07T03:08:51.260273Z","shell.execute_reply.started":"2026-05-07T03:08:51.252616Z","shell.execute_reply":"2026-05-07T03:08:51.259300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transforms = T.Compose([\n    T.Resize((256, 256), antialias=True),\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomVerticalFlip(p=0.5),\n    T.RandomRotation(degrees=30),\n])\n\nval_transforms = T.Compose([\n    T.Resize((256, 256), antialias=True)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:53.741465Z","iopub.execute_input":"2026-05-07T03:08:53.742437Z","iopub.status.idle":"2026-05-07T03:08:53.748714Z","shell.execute_reply.started":"2026-05-07T03:08:53.742389Z","shell.execute_reply":"2026-05-07T03:08:53.747642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# DATA LOADING + SPLIT + DATALOADERS\n# ================================\n\nimport pandas as pd\nfrom torch.utils.data import DataLoader\nimport torchvision.transforms as T\n\n# Paths\nCSV_PATH = \"/kaggle/input/competitions/human-protein-atlas-image-classification/train.csv\"\nIMG_DIR  = \"/kaggle/input/competitions/human-protein-atlas-image-classification/train/\"\nTEST_DIR = \"/kaggle/input/competitions/human-protein-atlas-image-classification/test/\"\nTEST_SUBMISSION = \"/kaggle/input/competitions/human-protein-atlas-image-classification/sample_submission.csv\"\n\n# ----------------\n# Load + split data\n# ----------------\ndf_full = pd.read_csv(CSV_PATH).sample(frac=1, random_state=42).reset_index(drop=True)\n\nsplit_idx = int(len(df_full) * 0.9)\ntrain_df = df_full.iloc[:split_idx].reset_index(drop=True)\nval_df   = df_full.iloc[split_idx:].reset_index(drop=True)\n\nprint(f\"Train: {len(train_df)} | Val: {len(val_df)}\")\n\n# ----------------\n# Transforms (augmentation)\n# ----------------\ntrain_transforms = T.Compose([\n    T.Resize((256, 256), antialias=True),\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomVerticalFlip(p=0.5),\n    T.RandomRotation(degrees=30),\n])\n\nval_transforms = T.Compose([\n    T.Resize((256, 256), antialias=True)\n])\n\n# ----------------\n# Datasets\n# ----------------\ntrain_dataset = ProteinAtlasDataset(\n    csv_file=train_df,   # IMPORTANT: direct dataframe use\n    root_dir=IMG_DIR,\n    transform=train_transforms\n)\n\nval_dataset = ProteinAtlasDataset(\n    csv_file=val_df,\n    root_dir=IMG_DIR,\n    transform=val_transforms\n)\n\n# ----------------\n# DataLoaders\n# ----------------\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=32,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=32,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n\n# ----------------\n# Test loader (NO CSV LABELS, ONLY FOR INFERENCE)\n# ----------------\ntest_df = pd.read_csv(TEST_SUBMISSION)\n\ntest_dataset = ProteinAtlasDataset(\n    csv_file=test_df,\n    root_dir=TEST_DIR,\n    transform=val_transforms\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=32,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=True\n)\n\n# ----------------\n# Sanity check\n# ----------------\nb_img, b_lbl = next(iter(train_loader))\nv_img, v_lbl = next(iter(val_loader))\n\nprint(f\"Train batch: {b_img.shape}, {b_lbl.shape}\")\nprint(f\"Val batch  : {v_img.shape}, {v_lbl.shape}\")\nprint(\"All loaders ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:08:55.552951Z","iopub.execute_input":"2026-05-07T03:08:55.553725Z","iopub.status.idle":"2026-05-07T03:09:04.175470Z","shell.execute_reply.started":"2026-05-07T03:08:55.553692Z","shell.execute_reply":"2026-05-07T03:09:04.174240Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Batch Visualization Function**","metadata":{}},{"cell_type":"markdown","source":"***Display a Batch-1 of Images from PyTorch DataLoader***","metadata":{}},{"cell_type":"code","source":"def display_imgs_pytorch(images):\n   \n    images = images.detach().cpu()\n\n    columns = 4\n    bs = images.shape[0]\n    rows = min((bs + columns - 1) // columns, 4)\n\n    fig = plt.figure(figsize=(columns * 4, rows * 4))\n\n    for idx in range(rows * columns):\n\n        if idx >= bs:\n            break\n\n        img = images[idx]                 # shape: (4, H, W)\n        img = img.permute(1, 2, 0)         # shape: (H, W, 4)\n        img = img.numpy()\n\n        # Create RGB composite from 4-channel image\n        composite = np.zeros_like(img[:, :, :3])\n\n        composite[:, :, 0] = img[:, :, 0] + img[:, :, 3] * 0.5  # red + yellow\n        composite[:, :, 1] = img[:, :, 1] + img[:, :, 3] * 0.5  # green + yellow\n        composite[:, :, 2] = img[:, :, 2]                       # blue\n\n        composite = np.clip(composite, 0, 1)\n\n        ax = fig.add_subplot(rows, columns, idx + 1)\n        ax.imshow(composite)\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:09:10.928622Z","iopub.execute_input":"2026-05-07T03:09:10.929528Z","iopub.status.idle":"2026-05-07T03:09:10.936934Z","shell.execute_reply.started":"2026-05-07T03:09:10.929491Z","shell.execute_reply":"2026-05-07T03:09:10.936162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_iter = iter(train_loader)\nimages, labels = next(data_iter) #Batch 1 Images\n\n\nprint(f\"Image batch shape: {images.shape}\")\nprint(f\"Label batch shape: {labels.shape}\")\nprint(\"Displaying Batch 1\")\ndisplay_imgs_pytorch(images)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:09:13.385237Z","iopub.execute_input":"2026-05-07T03:09:13.385534Z","iopub.status.idle":"2026-05-07T03:09:17.568062Z","shell.execute_reply.started":"2026-05-07T03:09:13.385493Z","shell.execute_reply":"2026-05-07T03:09:17.567051Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Pretrained ResNet34 Model 4 Channel image input**","metadata":{}},{"cell_type":"markdown","source":"***Model Architecture***","metadata":{}},{"cell_type":"code","source":"class Resnet34_4(nn.Module):\n    def __init__(self, num_classes=28, pretrained=True):\n        super().__init__()\n\n        # Load pretrained ResNet34\n        if pretrained:\n            encoder = resnet34(weights=ResNet34_Weights.IMAGENET1K_V1)\n        else:\n            encoder = resnet34(weights=None)\n\n        # conv1 that accepts 4 channels\n        self.conv1 = nn.Conv2d(\n            in_channels=4,\n            out_channels=64,\n            kernel_size=7,\n            stride=2,\n            padding=3,\n            bias=False\n        )\n\n        # Initialize 4-channel conv weights\n        if pretrained:\n            w = encoder.conv1.weight.data   # shape: (64, 3, 7, 7)\n\n            self.conv1.weight = nn.Parameter(\n                torch.cat(\n                    (\n                        w,\n                        0.5 * (w[:, 0:1, :, :] + w[:, 2:3, :, :])\n                    ),\n                    dim=1\n                )\n            )\n\n        # Copy ResNet34 layers\n        self.bn1 = encoder.bn1\n        self.relu = encoder.relu\n        self.maxpool = encoder.maxpool\n\n        # Correct ResNet order: conv → bn → relu → maxpool\n        self.layer0 = nn.Sequential(\n            self.conv1,\n            self.bn1,\n            self.relu,\n            self.maxpool\n        )\n\n        self.layer1 = encoder.layer1\n        self.layer2 = encoder.layer2\n        self.layer3 = encoder.layer3\n        self.layer4 = encoder.layer4\n\n        # Classification head for 28 protein classes\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        self.fc = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        x = self.layer0(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\n        x = self.fc(x)\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:09:36.829453Z","iopub.execute_input":"2026-05-07T03:09:36.830703Z","iopub.status.idle":"2026-05-07T03:09:36.840037Z","shell.execute_reply.started":"2026-05-07T03:09:36.830635Z","shell.execute_reply":"2026-05-07T03:09:36.839165Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"***Focal Loss for Multi-label Classification***","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2):\n        super().__init__()\n        self.gamma = gamma\n\n    def forward(self, input, target):\n        if target.size() != input.size():\n            raise ValueError(\n                f\"Target size {target.size()} must be same as input size {input.size()}\"\n            )\n\n        max_val = (-input).clamp(min=0)\n\n        loss = input - input * target + max_val + (\n            (-max_val).exp() + (-input - max_val).exp()\n        ).log()\n\n        invprobs = F.logsigmoid(-input * (target * 2.0 - 1.0))\n\n        loss = torch.exp(invprobs * self.gamma) * loss\n\n        return loss.sum(dim=1).mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:09:42.630269Z","iopub.execute_input":"2026-05-07T03:09:42.630540Z","iopub.status.idle":"2026-05-07T03:09:42.636650Z","shell.execute_reply.started":"2026-05-07T03:09:42.630517Z","shell.execute_reply":"2026-05-07T03:09:42.635739Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"***Accuracy Function***","metadata":{}},{"cell_type":"code","source":"def acc(preds, targs, th=0.0):\n    preds = (preds > th).int()\n    targs = targs.int()\n\n    return (preds == targs).float().mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:09:47.916185Z","iopub.execute_input":"2026-05-07T03:09:47.916836Z","iopub.status.idle":"2026-05-07T03:09:47.920798Z","shell.execute_reply.started":"2026-05-07T03:09:47.916805Z","shell.execute_reply":"2026-05-07T03:09:47.920017Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"***Create model, Loss, and Optimizer (with 3 learning rates)***","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = Resnet34_4(\n    num_classes=28,\n    pretrained=True\n).to(device)\n\ncriterion = FocalLoss(gamma=2)\n\noptimizer = torch.optim.Adam([\n\n    {\"params\": model.layer0.parameters(), \"lr\": 1e-5},\n    {\"params\": model.layer1.parameters(), \"lr\": 1e-5},\n\n    {\"params\": model.layer2.parameters(), \"lr\": 5e-5},\n    {\"params\": model.layer3.parameters(), \"lr\": 5e-5},\n\n    {\"params\": model.layer4.parameters(), \"lr\": 1e-4},\n    {\"params\": model.fc.parameters(), \"lr\": 1e-4}\n\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T03:11:25.972893Z","iopub.execute_input":"2026-05-07T03:11:25.973443Z","iopub.status.idle":"2026-05-07T03:11:26.945647Z","shell.execute_reply.started":"2026-05-07T03:11:25.973412Z","shell.execute_reply":"2026-05-07T03:11:26.944946Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Val/Test Split and Loaders**","metadata":{}},{"cell_type":"code","source":"# Member 2 built the optimizer with differential LRs — reuse it directly.\n# Member 3 adds a scheduler on top.\n\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,   # Member 2's optimizer\n    mode     = \"max\",     # watching F1 — higher is better\n    patience = 3,         # wait 3 epochs before reducing LR\n    factor   = 0.3,       # new_lr = old_lr * 0.3\n    min_lr   = 1e-8,\n)\n\nprint(\"Scheduler ready\")\nprint(f\"Current LRs: {[pg['lr'] for pg in optimizer.param_groups]}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:10:41.018015Z","iopub.execute_input":"2026-05-06T19:10:41.018591Z","iopub.status.idle":"2026-05-06T19:10:41.023690Z","shell.execute_reply.started":"2026-05-06T19:10:41.018561Z","shell.execute_reply":"2026-05-06T19:10:41.022847Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Macro F1 Function**","metadata":{}},{"cell_type":"code","source":"def compute_macro_f1(preds, targets, threshold=0.3):\n    preds_bin = (preds >= threshold).astype(int)\n    return f1_score(targets, preds_bin, average=\"macro\", zero_division=0)\n\nprint(\"✅ compute_macro_f1() ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:10:44.363148Z","iopub.execute_input":"2026-05-06T19:10:44.363606Z","iopub.status.idle":"2026-05-06T19:10:44.368766Z","shell.execute_reply.started":"2026-05-06T19:10:44.363574Z","shell.execute_reply":"2026-05-06T19:10:44.367908Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Train 1 epoch**","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, epoch_num):\n    model.train()\n    total_loss = 0.0\n    n_batches  = len(loader)\n    t0         = time.time()\n\n    for i, (images, labels) in enumerate(loader):\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss    = criterion(outputs, labels)\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n\n        optimizer.step()\n        total_loss += loss.item()\n\n        if (i + 1) % 100 == 0:\n            print(f\"  [Epoch {epoch_num}] Batch {i+1}/{n_batches} | Loss: {loss.item():.4f} | Time: {time.time()-t0:.0f}s\")\n\n    return total_loss / n_batches\n\nprint(\"✅ train_one_epoch() ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:10:47.848743Z","iopub.execute_input":"2026-05-06T19:10:47.849008Z","iopub.status.idle":"2026-05-06T19:10:47.855295Z","shell.execute_reply.started":"2026-05-06T19:10:47.848985Z","shell.execute_reply":"2026-05-06T19:10:47.854442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, loader, criterion, threshold=0.3):\n    model.eval()\n    total_loss  = 0.0\n    all_preds   = []\n    all_targets = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True)\n\n            outputs = model(images)\n            loss    = criterion(outputs, labels)\n            probs   = torch.sigmoid(outputs)\n\n            total_loss += loss.item()\n            all_preds.append(probs.detach().cpu().numpy())\n            all_targets.append(labels.detach().cpu().numpy())\n\n    all_preds   = np.vstack(all_preds)\n    all_targets = np.vstack(all_targets)\n    macro_f1    = compute_macro_f1(all_preds, all_targets, threshold)\n\n    return total_loss / len(loader), macro_f1\n\nprint(\"✅ validate() ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:10:51.173167Z","iopub.execute_input":"2026-05-06T19:10:51.174259Z","iopub.status.idle":"2026-05-06T19:10:51.180651Z","shell.execute_reply.started":"2026-05-06T19:10:51.174225Z","shell.execute_reply":"2026-05-06T19:10:51.179860Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**FulL Training lOOP**","metadata":{}},{"cell_type":"code","source":"DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {DEVICE}\")\n\nmodel = model.to(DEVICE)\nNUM_EPOCHS = 30\nTHRESHOLD  = 0.3\nSAVE_PATH  = \"best_model.pth\"\n\ndef train_model(model, train_loader, val_loader,\n                optimizer, criterion, scheduler,\n                num_epochs, threshold, save_path):\n\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_f1\": [], \"lr\": []}\n    best_f1 = 0.0\n\n    print(f\"\\n{'='*55}\")\n    print(f\"  TRAINING STARTED — {num_epochs} epochs on {DEVICE}\")\n    print(f\"{'='*55}\")\n\n    for epoch in range(1, num_epochs + 1):\n        t0         = time.time()\n        lrs = [pg[\"lr\"] for pg in optimizer.param_groups]\n        print(f\"LRS: {[f'{lr:.1e}' for lr in lrs]}\")\n\n\n        train_loss          = train_one_epoch(model, train_loader, optimizer, criterion, epoch)\n        val_loss, val_f1    = validate(model, val_loader, criterion, threshold)\n\n        scheduler.step(val_f1)\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_f1\"].append(val_f1)\n        history[\"lr\"].append(lrs)\n\n        is_best = val_f1 > best_f1\n        print(f\"  Train Loss : {train_loss:.5f}\")\n        print(f\"  Val   Loss : {val_loss:.5f}\")\n        print(f\"  Val   F1   : {val_f1:.5f}  {'← NEW BEST' if is_best else ''}\")\n        print(f\"  Time       : {time.time()-t0:.1f}s\")\n\n        if is_best:\n            best_f1 = val_f1\n            torch.save({\n                \"epoch\"               : epoch,\n                \"model_state_dict\"    : model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"val_f1\"              : best_f1,\n                \"threshold\"           : threshold,\n            }, save_path)\n            print(f\"Saved → {save_path}\")\n\n    print(f\"\\n{'='*55}\")\n    print(f\"  DONE  |  Best Val F1 : {best_f1:.5f}\")\n    print(f\"{'='*55}\")\n    return history\n\n\nhistory = train_model(\n    model        = model,\n    train_loader = train_loader,\n    val_loader   = val_loader,\n    optimizer    = optimizer,\n    criterion    = criterion,\n    scheduler    = scheduler,\n    num_epochs   = NUM_EPOCHS,\n    threshold    = THRESHOLD,\n    save_path    = SAVE_PATH,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:19:46.814501Z","iopub.execute_input":"2026-05-06T19:19:46.814722Z","iopub.status.idle":"2026-05-06T19:19:46.835295Z","shell.execute_reply.started":"2026-05-06T19:19:46.814700Z","shell.execute_reply":"2026-05-06T19:19:46.834171Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Plot training curves**","metadata":{}},{"cell_type":"code","source":"# ── Training Curves (FIXED) ───────────────────────────────────────────────────\n\nepochs = range(1, len(history[\"train_loss\"]) + 1)\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 4))\nfig.suptitle(\"Training Curves — Member 3\", fontsize=13, fontweight=\"bold\")\n\n# ── Plot 1: Loss ──\naxes[0].plot(epochs, history[\"train_loss\"], \"b-o\", markersize=3, label=\"Train Loss\")\naxes[0].plot(epochs, history[\"val_loss\"],   \"r-o\", markersize=3, label=\"Val Loss\")\naxes[0].set_title(\"Loss per Epoch\")\naxes[0].set_xlabel(\"Epoch\")\naxes[0].set_ylabel(\"Focal Loss\")\naxes[0].legend()\naxes[0].grid(True, alpha=0.3)\n# Annotation: explain the gap\naxes[0].annotate(\"Overfitting gap\", xy=(15, 0.6), fontsize=8, color=\"gray\")\n\n# ── Plot 2: Macro F1 ──\nbest_ep = int(np.argmax(history[\"val_f1\"])) + 1\nbest_f1_val = max(history[\"val_f1\"])\naxes[1].plot(epochs, history[\"val_f1\"], \"g-o\", markersize=3, label=\"Val Macro F1\")\naxes[1].axvline(x=best_ep, color=\"red\", linestyle=\"--\",\n                alpha=0.5, label=f\"Best ep={best_ep} ({best_f1_val:.3f})\")\naxes[1].set_title(\"Validation Macro F1\")\naxes[1].set_xlabel(\"Epoch\")\naxes[1].set_ylabel(\"Macro F1\")\naxes[1].legend()\naxes[1].grid(True, alpha=0.3)\n\n# ── Plot 3: Learning Rate (FIXED — only FC layer, clean legend) ──\n# history[\"lr\"] is a list of lists — extract FC layer (last group) only\nfc_lrs = [lr_group[-1] for lr_group in history[\"lr\"]]\naxes[2].plot(epochs, fc_lrs, \"m-o\", markersize=3, label=\"FC Layer LR (highest)\")\naxes[2].set_title(\"Learning Rate (FC Layer)\")\naxes[2].set_xlabel(\"Epoch\")\naxes[2].set_ylabel(\"LR (log scale)\")\naxes[2].set_yscale(\"log\")\naxes[2].legend()\naxes[2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(\"training_curves_fixed.png\", dpi=120)\nplt.show()\nprint(f\"Best epoch: {best_ep}  |  Best Val F1: {best_f1_val:.5f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:04:12.707183Z","iopub.status.idle":"2026-05-06T19:04:12.707679Z","shell.execute_reply.started":"2026-05-06T19:04:12.707431Z","shell.execute_reply":"2026-05-06T19:04:12.707464Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Load best model+find best threshold**","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(SAVE_PATH, map_location=DEVICE)\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\nprint(f\"Loaded epoch {checkpoint['epoch']}  |  Val F1: {checkpoint['val_f1']:.5f}\")\n\ndef find_best_threshold(model, loader):\n    model.eval()\n    all_preds, all_targets = [], []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            probs = torch.sigmoid(model(images.to(DEVICE)))\n            all_preds.append(probs.cpu().numpy())\n            all_targets.append(labels.numpy())\n\n    all_preds   = np.vstack(all_preds)\n    all_targets = np.vstack(all_targets)\n\n    best_t, best_f1 = 0.3, 0.0\n\n    print(f\"\\n{'Threshold':>10}  |  {'Macro F1':>10}\")\n    print(\"-\" * 26)\n    for t in np.arange(0.05, 0.65, 0.05):\n        f1     = compute_macro_f1(all_preds, all_targets, t)\n        marker = \"  ← BEST\" if f1 > best_f1 else \"\"\n        print(f\"  {t:.2f}        |    {f1:.5f}{marker}\")\n        if f1 > best_f1:\n            best_f1, best_t = f1, t\n\n    print(f\"\\n✅ Best threshold : {best_t:.2f}  |  Best F1 : {best_f1:.5f}\")\n    return float(best_t), float(best_f1)\n\n\nBEST_THRESHOLD, BEST_F1 = find_best_threshold(model, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:04:12.710041Z","iopub.status.idle":"2026-05-06T19:04:12.710488Z","shell.execute_reply.started":"2026-05-06T19:04:12.710266Z","shell.execute_reply":"2026-05-06T19:04:12.710293Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Threshold Search Visualization**","metadata":{}},{"cell_type":"code","source":"# ── Threshold Search Visualization ───────────────────────────────────────────\n# This is Member 3's key contribution — visualize it for the presentation\n\nthresholds = np.arange(0.05, 0.65, 0.05)\nf1_scores  = []\n\n# Recompute (model already loaded with best weights)\nmodel.eval()\nall_preds, all_targets = [], []\nwith torch.no_grad():\n    for images, labels in val_loader:\n        probs = torch.sigmoid(model(images.to(DEVICE)))\n        all_preds.append(probs.cpu().numpy())\n        all_targets.append(labels.numpy())\nall_preds   = np.vstack(all_preds)\nall_targets = np.vstack(all_targets)\n\nfor t in thresholds:\n    f1_scores.append(compute_macro_f1(all_preds, all_targets, t))\n\nbest_idx = int(np.argmax(f1_scores))\n\nplt.figure(figsize=(10, 4))\nbars = plt.bar(thresholds, f1_scores, width=0.04,\n               color=[\"#E74C3C\" if i == best_idx else \"#3498DB\"\n                      for i in range(len(thresholds))])\nplt.axvline(x=thresholds[best_idx], color=\"red\", linestyle=\"--\",\n            label=f\"Best = {thresholds[best_idx]:.2f} (F1={f1_scores[best_idx]:.4f})\")\nplt.xlabel(\"Threshold\")\nplt.ylabel(\"Macro F1 Score\")\nplt.title(\"Threshold Search — Finding Optimal Decision Boundary\\n\"\n          \"Default 0.5 is NOT always optimal for multi-label classification\")\nplt.legend()\nplt.grid(True, alpha=0.3, axis=\"y\")\nplt.xticks(thresholds, [f\"{t:.2f}\" for t in thresholds], rotation=45)\nplt.tight_layout()\nplt.savefig(\"threshold_search.png\", dpi=120)\nplt.show()\nprint(f\"Best threshold : {thresholds[best_idx]:.2f}\")\nprint(f\"Best Macro F1  : {f1_scores[best_idx]:.5f}\")\nprint(f\"vs default 0.5 : {f1_scores[list(thresholds).index(0.50) if 0.50 in thresholds else -1]:.5f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T19:04:12.714472Z","iopub.status.idle":"2026-05-06T19:04:12.714929Z","shell.execute_reply.started":"2026-05-06T19:04:12.714708Z","shell.execute_reply":"2026-05-06T19:04:12.714734Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**FINAL TEST PREDICTION + SUBMISSION FILE**","metadata":{}},{"cell_type":"code","source":"\nmodel.eval()\n\nall_test_probs = []\n\nwith torch.no_grad():\n    for images, _ in test_loader:\n        images = images.to(DEVICE, non_blocking=True)\n\n        outputs = model(images)\n        probs = torch.sigmoid(outputs)\n\n        all_test_probs.append(probs.cpu().numpy())\n\nall_test_probs = np.vstack(all_test_probs)\n\nprint(\"Test probability shape:\", all_test_probs.shape)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**CONVERT PROBABILITIES TO LABEL STRINGS**","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CONVERT PROBABILITIES TO LABEL STRINGS\n# ============================================================\n\nthreshold = BEST_THRESHOLD\n\npredicted_labels = []\n\nfor probs in all_test_probs:\n    label_indices = np.where(probs >= threshold)[0]\n\n    # If no class passes threshold, choose the class with highest probability\n    if len(label_indices) == 0:\n        label_indices = [int(np.argmax(probs))]\n\n    label_string = \" \".join(map(str, label_indices))\n    predicted_labels.append(label_string)\n\nprint(\"Example predictions:\")\nprint(predicted_labels[:10])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**CREATE KAGGLE SUBMISSION FILE**","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CREATE KAGGLE SUBMISSION FILE\n# ============================================================\n\nsubmission = pd.DataFrame({\n    \"Id\": test_df[\"Id\"].values,\n    \"Predicted\": predicted_labels\n})\n\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(\"submission.csv created successfully!\")\nprint(submission.head(10))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify the submission file is in the correct Kaggle format\n# before uploading — catches any formatting issues early\n\nsub = pd.read_csv(\"submission.csv\")\n\nprint(\"── Submission Validation ────────────────────────────\")\nprint(f\"Columns       : {sub.columns.tolist()}\")\n# MUST be: ['Id', 'Predicted']\n\nprint(f\"Total rows    : {len(sub)}\")\n# MUST be: 11702\n\nprint(f\"Null values   : {sub.isnull().sum().sum()}\")\n# MUST be: 0\n\nprint(f\"Empty strings : {(sub['Predicted'] == '').sum()}\")\n# MUST be: 0\n\n# Check label range — all predicted labels must be 0 to 27\nall_labels = []\nfor pred_str in sub[\"Predicted\"]:\n    for label in str(pred_str).split():\n        all_labels.append(int(label))\n\nall_labels = np.array(all_labels)\nprint(f\"Label range   : {all_labels.min()} to {all_labels.max()}\")\n# MUST be: 0 to 27\n\nprint(f\"Avg labels/img: {len(all_labels) / len(sub):.2f}\")\n\nprint(\"\\n✅ Submission format valid — ready to upload to Kaggle\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analyze which of the 28 classes the model handles well\n# This is useful for the presentation and the report\n\nfrom sklearn.metrics import f1_score as sk_f1\n\n# Use the val set (we have ground truth labels there)\nmodel.eval()\nval_probs_all   = []\nval_targets_all = []\n\nwith torch.no_grad():\n    for images, labels in val_loader:\n        probs = torch.sigmoid(model(images.to(DEVICE)))\n        val_probs_all.append(probs.cpu().numpy())\n        val_targets_all.append(labels.numpy())\n\nval_probs_all   = np.vstack(val_probs_all)\nval_targets_all = np.vstack(val_targets_all)\nval_preds_all   = (val_probs_all >= BEST_THRESHOLD).astype(int)\n\n# Per-class F1\nper_class_f1 = sk_f1(\n    val_targets_all,\n    val_preds_all,\n    average       = None,\n    zero_division = 0\n)\n\nCLASS_NAMES = [\n    \"Nucleoplasm\", \"Nuclear Membrane\", \"Nucleoli\", \"Nucleoli FC\",\n    \"Nuclear Speckles\", \"Nuclear Bodies\", \"Endoplasmic Reticulum\",\n    \"Golgi Apparatus\", \"Intermediate Filaments\", \"Actin Filaments\",\n    \"Focal Adhesion\", \"Microtubules\", \"Mitotic Spindle\", \"Centrosome\",\n    \"Lipid Droplets\", \"Plasma Membrane\", \"Cell Junctions\", \"Mitochondria\",\n    \"Aggresome\", \"Cytosol\", \"Vesicles\", \"Negative\", \"Unspecified\",\n    \"Class 23\", \"Class 24\", \"Class 25\", \"Class 26\", \"Class 27\"\n]\n\nprint(f\"── Per-Class F1 (threshold={BEST_THRESHOLD}) ──────────────────\")\nprint(f\"  {'Class':>3}  {'Name':<26}  {'F1':>6}\")\nprint(\"  \" + \"─\" * 42)\nfor i, (name, f1) in enumerate(zip(CLASS_NAMES, per_class_f1)):\n    bar = \"▓\" * int(f1 * 15)\n    print(f\"  {i:3d}  {name:<26}  {f1:.3f}  {bar}\")\nprint(\"  \" + \"─\" * 42)\nprint(f\"  Macro Average F1 : {per_class_f1.mean():.4f}\")\n\n# Plot\nplt.figure(figsize=(15, 5))\ncolors = [\"#2ECC71\" if f >= 0.5 else \"#E74C3C\" if f < 0.2 else \"#F39C12\"\n          for f in per_class_f1]\nplt.bar(range(28), per_class_f1, color=colors)\nplt.axhline(y=per_class_f1.mean(), color=\"navy\", linestyle=\"--\",\n            label=f\"Macro avg = {per_class_f1.mean():.4f}\")\nplt.xlabel(\"Class Label (0–27)\")\nplt.ylabel(\"F1 Score\")\nplt.title(f\"Per-Class F1 Score at threshold = {BEST_THRESHOLD}\")\nplt.xticks(range(28), range(28))\nplt.legend()\n\nimport matplotlib.patches as mpatches\npatches = [\n    mpatches.Patch(color=\"#2ECC71\", label=\"F1 ≥ 0.5  (good)\"),\n    mpatches.Patch(color=\"#F39C12\", label=\"0.2 ≤ F1 < 0.5  (ok)\"),\n    mpatches.Patch(color=\"#E74C3C\", label=\"F1 < 0.2  (poor)\"),\n]\nplt.legend(handles=patches)\nplt.tight_layout()\nplt.savefig(\"per_class_f1.png\", dpi=100)\nplt.show()\nprint(\"Saved per_class_f1.png\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"\"\"\n{'='*58}\n  PIPELINE COMPLETE — FINAL SUMMARY\n{'='*58}\n\n  Model         : Resnet34_4 (4-channel pretrained ResNet34)\n  Loss          : FocalLoss (gamma=2)\n  Optimizer     : Adam (differential LR)\n  Epochs        : {checkpoint['epoch']}\n  Threshold     : {BEST_THRESHOLD}\n\n  Val Macro F1  : {BEST_F1:.5f}\n  Submission    : submission.csv ({len(submission_df)} predictions)\n\n  Files saved:\n    best_model.pth         — best model weights\n    training_curves.png    — loss / F1 / LR plots\n    threshold_search.png   — threshold optimization chart\n    per_class_f1.png       — per-class breakdown\n    submission.csv         — upload this to Kaggle","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}