{"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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-08T01:06:31.65241Z","iopub.execute_input":"2026-05-08T01:06:31.653254Z","iopub.status.idle":"2026-05-08T01:06:42.633231Z","shell.execute_reply.started":"2026-05-08T01:06:31.653222Z","shell.execute_reply":"2026-05-08T01:06:42.63224Z"}},"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-08T04:11:50.424114Z","iopub.execute_input":"2026-05-08T04:11:50.42479Z","iopub.status.idle":"2026-05-08T04:11:50.491688Z","shell.execute_reply.started":"2026-05-08T04:11:50.424759Z","shell.execute_reply":"2026-05-08T04:11:50.490812Z"}},"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-08T04:11:57.996656Z","iopub.execute_input":"2026-05-08T04:11:57.997438Z","iopub.status.idle":"2026-05-08T04:11:58.741036Z","shell.execute_reply.started":"2026-05-08T04:11:57.997402Z","shell.execute_reply":"2026-05-08T04:11:58.739827Z"}},"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,"jupyter":{"source_hidden":true}},"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-08T01:06:53.396308Z","iopub.execute_input":"2026-05-08T01:06:53.397115Z","iopub.status.idle":"2026-05-08T01:06:53.404615Z","shell.execute_reply.started":"2026-05-08T01:06:53.39708Z","shell.execute_reply":"2026-05-08T01:06:53.403715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T05:05:59.998396Z","iopub.execute_input":"2026-05-08T05:05:59.999322Z","iopub.status.idle":"2026-05-08T05:06:00.034544Z","shell.execute_reply.started":"2026-05-08T05:05:59.999292Z","shell.execute_reply":"2026-05-08T05:06:00.033764Z"}},"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-08T01:06:56.32034Z","iopub.execute_input":"2026-05-08T01:06:56.321064Z","iopub.status.idle":"2026-05-08T01:06:56.326394Z","shell.execute_reply.started":"2026-05-08T01:06:56.321032Z","shell.execute_reply":"2026-05-08T01:06:56.32543Z"}},"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-08T01:07:01.274401Z","iopub.execute_input":"2026-05-08T01:07:01.274802Z","iopub.status.idle":"2026-05-08T01:07:09.76539Z","shell.execute_reply.started":"2026-05-08T01:07:01.274762Z","shell.execute_reply":"2026-05-08T01:07:09.764264Z"}},"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-08T01:07:09.776749Z","iopub.execute_input":"2026-05-08T01:07:09.777161Z","iopub.status.idle":"2026-05-08T01:07:09.805543Z","shell.execute_reply.started":"2026-05-08T01:07:09.777123Z","shell.execute_reply":"2026-05-08T01:07:09.804683Z"}},"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-08T04:07:45.268021Z","iopub.execute_input":"2026-05-08T04:07:45.268957Z","iopub.status.idle":"2026-05-08T04:07:48.80497Z","shell.execute_reply.started":"2026-05-08T04:07:45.2689Z","shell.execute_reply":"2026-05-08T04:07:48.802402Z"}},"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=3,#7,\n            stride=2,\n            padding=1,#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-08T01:07:13.932345Z","iopub.execute_input":"2026-05-08T01:07:13.932754Z","iopub.status.idle":"2026-05-08T01:07:13.955946Z","shell.execute_reply.started":"2026-05-08T01:07:13.932703Z","shell.execute_reply":"2026-05-08T01:07:13.954938Z"}},"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-08T01:07:13.95675Z","iopub.execute_input":"2026-05-08T01:07:13.95704Z","iopub.status.idle":"2026-05-08T01:07:14.01161Z","shell.execute_reply.started":"2026-05-08T01:07:13.957009Z","shell.execute_reply":"2026-05-08T01:07:14.00793Z"}},"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\")\nprint(f\"Using device: {DEVICE}\")\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    {\"params\": model.layer0.parameters(), \"lr\": 1e-4},\n    {\"params\": model.layer1.parameters(), \"lr\": 1e-4},\n\n    {\"params\": model.layer2.parameters(), \"lr\": 5e-4},\n    {\"params\": model.layer3.parameters(), \"lr\": 5e-4},\n\n    {\"params\": model.layer4.parameters(), \"lr\": 1e-3},\n    {\"params\": model.fc.parameters(), \"lr\": 1e-3}\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T01:15:26.103541Z","iopub.execute_input":"2026-05-08T01:15:26.104012Z","iopub.status.idle":"2026-05-08T01:15:26.456496Z","shell.execute_reply.started":"2026-05-08T01:15:26.10398Z","shell.execute_reply":"2026-05-08T01:15:26.455908Z"}},"outputs":[],"execution_count":null},{"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-08T01:15:28.47501Z","iopub.execute_input":"2026-05-08T01:15:28.475651Z","iopub.status.idle":"2026-05-08T01:15:28.480956Z","shell.execute_reply.started":"2026-05-08T01:15:28.47562Z","shell.execute_reply":"2026-05-08T01:15:28.480041Z"}},"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T01:15:31.703045Z","iopub.execute_input":"2026-05-08T01:15:31.703921Z","iopub.status.idle":"2026-05-08T01:15:31.709123Z","shell.execute_reply.started":"2026-05-08T01:15:31.703879Z","shell.execute_reply":"2026-05-08T01:15:31.708379Z"}},"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T01:15:33.738239Z","iopub.execute_input":"2026-05-08T01:15:33.739011Z","iopub.status.idle":"2026-05-08T01:15:33.745425Z","shell.execute_reply.started":"2026-05-08T01:15:33.738959Z","shell.execute_reply":"2026-05-08T01:15:33.744518Z"}},"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T01:15:36.216832Z","iopub.execute_input":"2026-05-08T01:15:36.217537Z","iopub.status.idle":"2026-05-08T01:15:36.223464Z","shell.execute_reply.started":"2026-05-08T01:15:36.217506Z","shell.execute_reply":"2026-05-08T01:15:36.222528Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Full Training loop**","metadata":{}},{"cell_type":"code","source":"NUM_EPOCHS = 25\nTHRESHOLD  = 0.3\nSAVE_PATH  = \"s3_pd1_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-08T01:15:39.219834Z","iopub.execute_input":"2026-05-08T01:15:39.22028Z","iopub.status.idle":"2026-05-08T03:28:34.176419Z","shell.execute_reply.started":"2026-05-08T01:15:39.220251Z","shell.execute_reply":"2026-05-08T03:28:34.175591Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"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\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-08T03:33:47.116728Z","iopub.execute_input":"2026-05-08T03:33:47.117488Z","iopub.status.idle":"2026-05-08T03:33:48.334676Z","shell.execute_reply.started":"2026-05-08T03:33:47.117432Z","shell.execute_reply":"2026-05-08T03:33:48.333454Z"}},"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    thresholds = np.arange(0.05, 0.65, 0.05)\n    f1_scores = []\n\n    best_t, best_f1 = 0.3, 0.0\n\n    print(f\"\\n{'Threshold':>10}  |  {'Macro F1':>10}\")\n    print(\"-\" * 26)\n\n    for t in thresholds:\n        f1 = compute_macro_f1(all_preds, all_targets, t)\n        f1_scores.append(f1)\n\n        marker = \"  ← BEST\" if f1 > best_f1 else \"\"\n        print(f\"  {t:.2f}        |    {f1:.5f}{marker}\")\n\n        if f1 > best_f1:\n            best_f1, best_t = f1, t\n\n    print(f\"\\nBest threshold : {best_t:.2f}  |  Best F1 : {best_f1:.5f}\")\n\n    return float(best_t), float(best_f1), thresholds, f1_scores, all_preds, all_targets\n\n\nBEST_THRESHOLD, BEST_F1, thresholds, f1_scores, val_probs_all, val_targets_all = find_best_threshold(\n    model, val_loader\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T03:34:00.13713Z","iopub.execute_input":"2026-05-08T03:34:00.138331Z","iopub.status.idle":"2026-05-08T03:34:25.438408Z","shell.execute_reply.started":"2026-05-08T03:34:00.138296Z","shell.execute_reply":"2026-05-08T03:34:25.437404Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Threshold vs Macro F1 Visualization**","metadata":{}},{"cell_type":"code","source":"best_idx = int(np.argmax(f1_scores))\n\nplt.figure(figsize=(10, 4))\n\nplt.bar(\n    thresholds,\n    f1_scores,\n    width=0.04,\n    color=[\"#E74C3C\" if i == best_idx else \"#3498DB\" for i in range(len(thresholds))]\n)\n\nplt.axvline(\n    x=thresholds[best_idx],\n    color=\"red\",\n    linestyle=\"--\",\n    label=f\"Best = {thresholds[best_idx]:.2f} (F1={f1_scores[best_idx]:.4f})\"\n)\n\nplt.xlabel(\"Threshold\")\nplt.ylabel(\"Macro F1 Score\")\nplt.title(\"Threshold Search — Finding Optimal Decision Boundary\")\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()\n\ndefault_idx = np.argmin(np.abs(thresholds - 0.50))\n\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[default_idx]:.5f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T03:34:25.440412Z","iopub.execute_input":"2026-05-08T03:34:25.440926Z","iopub.status.idle":"2026-05-08T03:34:25.811472Z","shell.execute_reply.started":"2026-05-08T03:34:25.440838Z","shell.execute_reply":"2026-05-08T03:34:25.81065Z"}},"outputs":[],"execution_count":null},{"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-08T03:34:25.812388Z","iopub.execute_input":"2026-05-08T03:34:25.81276Z","iopub.status.idle":"2026-05-08T03:36:55.514313Z","shell.execute_reply.started":"2026-05-08T03:34:25.812734Z","shell.execute_reply":"2026-05-08T03:36:55.513486Z"}},"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-08T03:36:55.516263Z","iopub.execute_input":"2026-05-08T03:36:55.516553Z","iopub.status.idle":"2026-05-08T03:36:55.595723Z","shell.execute_reply.started":"2026-05-08T03:36:55.516523Z","shell.execute_reply":"2026-05-08T03:36:55.594776Z"}},"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-08T03:36:55.596657Z","iopub.execute_input":"2026-05-08T03:36:55.597049Z","iopub.status.idle":"2026-05-08T03:36:55.651711Z","shell.execute_reply.started":"2026-05-08T03:36:55.597024Z","shell.execute_reply":"2026-05-08T03:36:55.651112Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Plotting Per Class Macro F1 Score**","metadata":{}},{"cell_type":"code","source":"val_preds_all = (val_probs_all >= BEST_THRESHOLD).astype(int)\n\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)\n\nfor i, (name, f1_val) in enumerate(zip(CLASS_NAMES, per_class_f1)):\n    bar = \"▓\" * int(f1_val * 15)\n    print(f\"  {i:3d}  {name:<26}  {f1_val:.3f}  {bar}\")\n\nprint(\"  \" + \"─\" * 42)\nprint(f\"  Macro Average F1 : {per_class_f1.mean():.4f}\")\n\nplt.figure(figsize=(15, 5))\n\ncolors = [\n    \"#2ECC71\" if f >= 0.5 else \"#E74C3C\" if f < 0.2 else \"#F39C12\"\n    for f in per_class_f1\n]\n\nplt.bar(range(28), per_class_f1, color=colors)\n\npatches = [\n    mpatches.Patch(color=\"#2ECC71\", label=\"F1 ≥ 0.5\"),\n    mpatches.Patch(color=\"#F39C12\", label=\"0.2 ≤ F1 < 0.5\"),\n    mpatches.Patch(color=\"#E74C3C\", label=\"F1 < 0.2\"),\n]\n\nplt.axhline(\n    y=per_class_f1.mean(),\n    color=\"navy\",\n    linestyle=\"--\",\n    label=f\"Macro avg = {per_class_f1.mean():.4f}\"\n)\n\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(handles=patches)\nplt.tight_layout()\nplt.savefig(\"per_class_F1.png\", dpi=100)\nplt.show()\n\nprint(\"Saved per_class_F1.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T03:36:55.652632Z","iopub.execute_input":"2026-05-08T03:36:55.653242Z","iopub.status.idle":"2026-05-08T03:36:56.488277Z","shell.execute_reply.started":"2026-05-08T03:36:55.653216Z","shell.execute_reply":"2026-05-08T03:36:56.48752Z"}},"outputs":[],"execution_count":null}]}