{"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 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\nimport torch.optim as optim\nfrom torchvision.models import resnet34, ResNet34_Weights\n\nimport time\n\nfrom sklearn.metrics import f1_score\n\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-07T15:01:16.198055Z","iopub.execute_input":"2026-05-07T15:01:16.198767Z","iopub.status.idle":"2026-05-07T15:01:27.470621Z","shell.execute_reply.started":"2026-05-07T15:01:16.198736Z","shell.execute_reply":"2026-05-07T15:01:27.470005Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ***Section--Member 1***","metadata":{}},{"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-07T15:01:36.908238Z","iopub.execute_input":"2026-05-07T15:01:36.909030Z","iopub.status.idle":"2026-05-07T15:01:37.336408Z","shell.execute_reply.started":"2026-05-07T15:01:36.908979Z","shell.execute_reply":"2026-05-07T15:01:37.335754Z"}},"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-07T15:01:42.929543Z","iopub.execute_input":"2026-05-07T15:01:42.929806Z","iopub.status.idle":"2026-05-07T15:01:43.640315Z","shell.execute_reply.started":"2026-05-07T15:01:42.929783Z","shell.execute_reply":"2026-05-07T15:01:43.639243Z"}},"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-07T15:01:49.835343Z","iopub.execute_input":"2026-05-07T15:01:49.835647Z","iopub.status.idle":"2026-05-07T15:01:50.511800Z","shell.execute_reply.started":"2026-05-07T15:01:49.835623Z","shell.execute_reply":"2026-05-07T15:01:50.510774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ProteinAtlasDataset(Dataset):\n    def __init__(self, csv_file, root_dir, transform=None):\n\n        if isinstance(csv_file, pd.DataFrame):\n            self.df = csv_file.reset_index(drop=True)\n\n        else:\n            self.df = pd.read_csv(csv_file).reset_index(drop=True)\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            \n            img = Image.open(img_path).convert('L')\n            img_arrays.append(np.array(img))\n            \n        image_tensor = torch.tensor(np.stack(img_arrays), dtype=torch.float32) / 255.0\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_tensor = self.transform(image_tensor)\n            \n        return image_tensor, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T15:01:58.284998Z","iopub.execute_input":"2026-05-07T15:01:58.285554Z","iopub.status.idle":"2026-05-07T15:01:58.293461Z","shell.execute_reply.started":"2026-05-07T15:01:58.285527Z","shell.execute_reply":"2026-05-07T15:01:58.292748Z"}},"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-07T15:02:03.589372Z","iopub.execute_input":"2026-05-07T15:02:03.590169Z","iopub.status.idle":"2026-05-07T15:02:03.595523Z","shell.execute_reply.started":"2026-05-07T15:02:03.590123Z","shell.execute_reply":"2026-05-07T15:02:03.594641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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# ----------------\n# Datasets\n# ----------------\ntrain_dataset = ProteinAtlasDataset(\n    csv_file=train_df,\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\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-07T15:02:23.244878Z","iopub.execute_input":"2026-05-07T15:02:23.245154Z","iopub.status.idle":"2026-05-07T15:02:35.835070Z","shell.execute_reply.started":"2026-05-07T15:02:23.245131Z","shell.execute_reply":"2026-05-07T15:02:35.834204Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ***Section-- Member 2***","metadata":{}},{"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-07T15:02:51.583488Z","iopub.execute_input":"2026-05-07T15:02:51.583783Z","iopub.status.idle":"2026-05-07T15:02:51.590614Z","shell.execute_reply.started":"2026-05-07T15:02:51.583752Z","shell.execute_reply":"2026-05-07T15:02:51.589874Z"}},"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-07T15:02:59.332473Z","iopub.execute_input":"2026-05-07T15:02:59.332969Z","iopub.status.idle":"2026-05-07T15:03:04.045284Z","shell.execute_reply.started":"2026-05-07T15:02:59.332918Z","shell.execute_reply":"2026-05-07T15:03:04.044269Z"}},"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        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-07T15:03:33.356435Z","iopub.execute_input":"2026-05-07T15:03:33.357047Z","iopub.status.idle":"2026-05-07T15:03:33.363866Z","shell.execute_reply.started":"2026-05-07T15:03:33.357022Z","shell.execute_reply":"2026-05-07T15:03:33.363084Z"}},"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-07T15:03:48.328791Z","iopub.execute_input":"2026-05-07T15:03:48.329463Z","iopub.status.idle":"2026-05-07T15:03:48.334782Z","shell.execute_reply.started":"2026-05-07T15:03:48.329436Z","shell.execute_reply":"2026-05-07T15:03:48.333819Z"}},"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}, #Very Low LR (Conservative Fine-Tuning with Minimal Changes)\n    {\"params\": model.layer1.parameters(), \"lr\": 1e-5},\n\n    {\"params\": model.layer2.parameters(), \"lr\": 3e-5}, #Low to Moderate LR (Gradual Adaptation)\n    {\"params\": model.layer3.parameters(), \"lr\": 1e-4}, #Moderate LR (Moderate Learning)\n\n    {\"params\": model.layer4.parameters(), \"lr\": 3e-4}, #High LR (Rapid Adaptation)\n    {\"params\": model.fc.parameters(), \"lr\": 3e-4}\n\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T15:19:31.631123Z","iopub.execute_input":"2026-05-07T15:19:31.631758Z","iopub.status.idle":"2026-05-07T15:19:32.583566Z","shell.execute_reply.started":"2026-05-07T15:19:31.631726Z","shell.execute_reply":"2026-05-07T15:19:32.582971Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ***Section=--Member 3***","metadata":{}},{"cell_type":"markdown","source":"# **Define Scheduler**","metadata":{}},{"cell_type":"code","source":"scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode     = \"max\",    \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-07T15:23:17.011486Z","iopub.execute_input":"2026-05-07T15:23:17.011757Z","iopub.status.idle":"2026-05-07T15:23:17.017271Z","shell.execute_reply.started":"2026-05-07T15:23:17.011734Z","shell.execute_reply":"2026-05-07T15:23:17.016118Z"}},"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-07T15:23:20.787135Z","iopub.execute_input":"2026-05-07T15:23:20.787576Z","iopub.status.idle":"2026-05-07T15:23:20.792593Z","shell.execute_reply.started":"2026-05-07T15:23:20.787546Z","shell.execute_reply":"2026-05-07T15:23:20.791786Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Training 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-07T15:23:24.666262Z","iopub.execute_input":"2026-05-07T15:23:24.667158Z","iopub.status.idle":"2026-05-07T15:23:24.674531Z","shell.execute_reply.started":"2026-05-07T15:23:24.667113Z","shell.execute_reply":"2026-05-07T15:23:24.673645Z"}},"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-07T15:23:28.408960Z","iopub.execute_input":"2026-05-07T15:23:28.409583Z","iopub.status.idle":"2026-05-07T15:23:28.416264Z","shell.execute_reply.started":"2026-05-07T15:23:28.409553Z","shell.execute_reply":"2026-05-07T15:23:28.415353Z"}},"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\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-07T15:25:56.099040Z","iopub.execute_input":"2026-05-07T15:25:56.099662Z","iopub.status.idle":"2026-05-07T18:24:54.050375Z","shell.execute_reply.started":"2026-05-07T15:25:56.099630Z","shell.execute_reply":"2026-05-07T18:24:54.049477Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Plotting Curves**","metadata":{}},{"cell_type":"code","source":"epochs = range(1, len(history[\"train_loss\"]) + 1)\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 4))\n\n# Plot 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\naxes[0].annotate(\"Overfitting gap\", xy=(15, 0.6), fontsize=8, color=\"gray\")\n\n# Plot Macro F1 Score\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 Learning Rate\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.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-07T18:26:46.499385Z","iopub.execute_input":"2026-05-07T18:26:46.500028Z","iopub.status.idle":"2026-05-07T18:26:47.556792Z","shell.execute_reply.started":"2026-05-07T18:26:46.499988Z","shell.execute_reply":"2026-05-07T18:26:47.555988Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Loading Best Model and Finding 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-07T18:28:07.277250Z","iopub.execute_input":"2026-05-07T18:28:07.278024Z","iopub.status.idle":"2026-05-07T18:28:34.507069Z","shell.execute_reply.started":"2026-05-07T18:28:07.277993Z","shell.execute_reply":"2026-05-07T18:28:34.506078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Threshold vs Macro F1 Visualization**","metadata":{}},{"cell_type":"code","source":"thresholds = 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\")\n#plt.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-07T18:30:00.279712Z","iopub.execute_input":"2026-05-07T18:30:00.280440Z","iopub.status.idle":"2026-05-07T18:30:24.582571Z","shell.execute_reply.started":"2026-05-07T18:30:00.280402Z","shell.execute_reply":"2026-05-07T18:30:24.581878Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ***Section--Member 4***","metadata":{}},{"cell_type":"markdown","source":"# **Inference (prediction) on the test dataset**","metadata":{}},{"cell_type":"code","source":"model.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,"execution":{"iopub.status.busy":"2026-05-07T18:31:10.749320Z","iopub.execute_input":"2026-05-07T18:31:10.749771Z","iopub.status.idle":"2026-05-07T18:34:02.501796Z","shell.execute_reply.started":"2026-05-07T18:31:10.749740Z","shell.execute_reply":"2026-05-07T18:34:02.501028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Converting Predicted Probabilities into Protein Class Labels (Kaggle Requirements)**","metadata":{}},{"cell_type":"code","source":"threshold = 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,"execution":{"iopub.status.busy":"2026-05-07T18:36:09.021455Z","iopub.execute_input":"2026-05-07T18:36:09.021737Z","iopub.status.idle":"2026-05-07T18:36:09.096936Z","shell.execute_reply.started":"2026-05-07T18:36:09.021713Z","shell.execute_reply":"2026-05-07T18:36:09.096155Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Creating Final Kaggle Submission File**","metadata":{}},{"cell_type":"code","source":"submission = 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(5))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-07T18:36:15.808246Z","iopub.execute_input":"2026-05-07T18:36:15.808956Z","iopub.status.idle":"2026-05-07T18:36:15.859517Z","shell.execute_reply.started":"2026-05-07T18:36:15.808925Z","shell.execute_reply":"2026-05-07T18:36:15.858581Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Plotting Per Class Macro F1 Score**","metadata":{}},{"cell_type":"code","source":"# Use the val set\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 = f1_score(\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 (at 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,"execution":{"iopub.status.busy":"2026-05-07T18:36:52.807295Z","iopub.execute_input":"2026-05-07T18:36:52.808086Z","iopub.status.idle":"2026-05-07T18:37:19.535415Z","shell.execute_reply.started":"2026-05-07T18:36:52.808052Z","shell.execute_reply":"2026-05-07T18:37:19.534469Z"}},"outputs":[],"execution_count":null}]}