{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    roc_curve,\n    confusion_matrix,\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score,\n    average_precision_score,\n    ConfusionMatrixDisplay\n)\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom tqdm.auto import tqdm\nimport copy","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:51.769380Z","iopub.execute_input":"2026-06-23T10:39:51.769582Z","iopub.status.idle":"2026-06-23T10:39:58.322685Z","shell.execute_reply.started":"2026-06-23T10:39:51.769561Z","shell.execute_reply":"2026-06-23T10:39:58.321761Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## seeds\n","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed=42):\n\n    random.seed(seed)\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.324360Z","iopub.execute_input":"2026-06-23T10:39:58.324824Z","iopub.status.idle":"2026-06-23T10:39:58.332419Z","shell.execute_reply.started":"2026-06-23T10:39:58.324798Z","shell.execute_reply":"2026-06-23T10:39:58.331504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(\n    \"/kaggle/input/notebooks/seekersgame/isic-md-heirarchicalfilm-create-folds/isic_fold_assignments.csv\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.333431Z","iopub.execute_input":"2026-06-23T10:39:58.333731Z","iopub.status.idle":"2026-06-23T10:39:58.547309Z","shell.execute_reply.started":"2026-06-23T10:39:58.333709Z","shell.execute_reply":"2026-06-23T10:39:58.546264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.548533Z","iopub.execute_input":"2026-06-23T10:39:58.549439Z","iopub.status.idle":"2026-06-23T10:39:58.567004Z","shell.execute_reply.started":"2026-06-23T10:39:58.549408Z","shell.execute_reply":"2026-06-23T10:39:58.566177Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Augmentation","metadata":{}},{"cell_type":"code","source":"def get_train_transforms():\n\n    return A.Compose([\n\n        A.Resize(512,512),\n\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n\n        A.RandomRotate90(p=0.5),\n\n        A.Affine(\n            scale=(0.95, 1.05),\n            translate_percent=(0.05, 0.05),\n            rotate=(-20, 20),\n            p=0.5\n        ),\n\n        A.ColorJitter(\n            brightness=0.2,\n            contrast=0.2,\n            saturation=0.2,\n            hue=0.1,\n            p=0.5\n        ),\n\n        A.Normalize(),\n\n        ToTensorV2()\n    ])\n\n\ndef get_valid_transforms():\n\n    return A.Compose([\n\n        A.Resize(512,512),\n\n        A.Normalize(),\n\n        ToTensorV2()\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.569158Z","iopub.execute_input":"2026-06-23T10:39:58.569504Z","iopub.status.idle":"2026-06-23T10:39:58.575882Z","shell.execute_reply.started":"2026-06-23T10:39:58.569478Z","shell.execute_reply":"2026-06-23T10:39:58.575179Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class ISICDataset(Dataset):\n\n    def __init__(\n        self,\n        df,\n        transforms=None\n    ):\n\n        self.df = (\n            df.reset_index(drop=True)\n        )\n\n        self.transforms = transforms\n\n        self.meta_features = [\n            \"age_approx\",\n            \"sex\",\n            \"site_encoded\"\n        ]\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.loc[idx]\n\n        image = cv2.imread(\n            row.image_path\n        )\n\n        image = cv2.cvtColor(\n            image,\n            cv2.COLOR_BGR2RGB\n        )\n\n        if self.transforms:\n\n            image = self.transforms(\n                image=image\n            )[\"image\"]\n\n        meta = torch.tensor(\n            row[self.meta_features]\n            .values\n            .astype(np.float32)\n        )\n\n        target = torch.tensor(\n            row.target\n        ).float()\n\n        return image, meta, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.576981Z","iopub.execute_input":"2026-06-23T10:39:58.577799Z","iopub.status.idle":"2026-06-23T10:39:58.589123Z","shell.execute_reply.started":"2026-06-23T10:39:58.577769Z","shell.execute_reply":"2026-06-23T10:39:58.588348Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Single FiLM","metadata":{}},{"cell_type":"code","source":"class SpatialFiLM(nn.Module):\n\n    def __init__(self, channels, meta_features):\n\n        super().__init__()\n\n        self.gamma = nn.Linear(\n            meta_features,\n            channels\n        )\n\n        self.beta = nn.Linear(\n            meta_features,\n            channels\n        )\n\n    def forward(self, x, meta):\n\n        gamma = self.gamma(meta)\n        beta = self.beta(meta)\n\n        gamma = gamma.unsqueeze(-1).unsqueeze(-1)\n        beta  = beta.unsqueeze(-1).unsqueeze(-1)\n\n        return gamma * x + beta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.590095Z","iopub.execute_input":"2026-06-23T10:39:58.590383Z","iopub.status.idle":"2026-06-23T10:39:58.605635Z","shell.execute_reply.started":"2026-06-23T10:39:58.590350Z","shell.execute_reply":"2026-06-23T10:39:58.604977Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## HierarchicalFiLMModel","metadata":{}},{"cell_type":"code","source":"class HierarchicalFiLMModel(nn.Module):\n\n    def __init__(self, meta_features=3):\n\n        super().__init__()\n\n        self.backbone = timm.create_model(\n            \"efficientnet_b3\",\n            pretrained=True,\n            features_only=True\n        )\n\n        channels = [\n            24,\n            32,\n            48,\n            136,\n            384\n        ]\n\n        self.film1 = SpatialFiLM(\n            channels[0],\n            meta_features\n        )\n\n        self.film2 = SpatialFiLM(\n            channels[1],\n            meta_features\n        )\n\n        self.film3 = SpatialFiLM(\n            channels[2],\n            meta_features\n        )\n\n        self.film4 = SpatialFiLM(\n            channels[3],\n            meta_features\n        )\n\n        self.film5 = SpatialFiLM(\n            channels[4],\n            meta_features\n        )\n\n        self.meta_branch = nn.Sequential(\n\n            nn.Linear(\n                meta_features,\n                128\n            ),\n\n            nn.ReLU(),\n            nn.BatchNorm1d(128),\n            nn.Dropout(0.3)\n        )\n\n        self.classifier = nn.Sequential(\n\n            nn.Linear(\n                752,\n                256\n            ),\n\n            nn.ReLU(),\n            nn.Dropout(0.4),\n            nn.Linear(\n                256,\n                1\n            )\n        )\n\n    def forward(\n        self,\n        image,\n        meta\n    ):\n\n        features = self.backbone(image)\n        f1 = self.film1(\n            features[0],\n            meta\n        )\n        f2 = self.film2(\n            features[1],\n            meta\n        )\n        f3 = self.film3(\n            features[2],\n            meta\n        )\n        f4 = self.film4(\n            features[3],\n            meta\n        )\n        f5 = self.film5(\n            features[4],\n            meta\n        )\n        p1 = F.adaptive_avg_pool2d(\n            f1,\n            1\n        ).flatten(1)\n        p2 = F.adaptive_avg_pool2d(\n            f2,\n            1\n        ).flatten(1)\n        p3 = F.adaptive_avg_pool2d(\n            f3,\n            1\n        ).flatten(1)\n        p4 = F.adaptive_avg_pool2d(\n            f4,\n            1\n        ).flatten(1)\n        p5 = F.adaptive_avg_pool2d(\n            f5,\n            1\n        ).flatten(1)\n        img_feat = torch.cat(\n            [p1, p2, p3, p4, p5],\n            dim=1\n        )\n        meta_feat = self.meta_branch(\n            meta\n        )\n        x = torch.cat(\n            [img_feat, meta_feat],\n            dim=1\n        )\n\n        out = self.classifier(x)\n\n        return out.squeeze(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.606486Z","iopub.execute_input":"2026-06-23T10:39:58.606780Z","iopub.status.idle":"2026-06-23T10:39:58.619947Z","shell.execute_reply.started":"2026-06-23T10:39:58.606751Z","shell.execute_reply":"2026-06-23T10:39:58.619144Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Focal Loss Cell","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n\n    def __init__(\n        self,\n        alpha=1,\n        gamma=2\n    ):\n\n        super().__init__()\n\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(\n        self,\n        logits,\n        targets\n    ):\n\n        bce = F.binary_cross_entropy_with_logits(\n            logits,\n            targets,\n            reduction=\"none\"\n        )\n\n        probs = torch.sigmoid(\n            logits\n        )\n\n        pt = torch.where(\n            targets == 1,\n            probs,\n            1 - probs\n        )\n\n        loss = (\n            self.alpha *\n            (1 - pt) ** self.gamma *\n            bce\n        )\n\n        return loss.mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.620931Z","iopub.execute_input":"2026-06-23T10:39:58.621305Z","iopub.status.idle":"2026-06-23T10:39:58.635675Z","shell.execute_reply.started":"2026-06-23T10:39:58.621281Z","shell.execute_reply":"2026-06-23T10:39:58.634849Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation Helper","metadata":{}},{"cell_type":"code","source":"def safe_confusion_matrix(\n    y_true,\n    y_pred\n):\n\n    cm = confusion_matrix(\n        y_true,\n        y_pred,\n        labels=[0, 1]\n    )\n\n    return cm.ravel()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.636689Z","iopub.execute_input":"2026-06-23T10:39:58.637032Z","iopub.status.idle":"2026-06-23T10:39:58.649848Z","shell.execute_reply.started":"2026-06-23T10:39:58.636974Z","shell.execute_reply":"2026-06-23T10:39:58.649154Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation Function","metadata":{}},{"cell_type":"code","source":"def validate(\n    model,\n    loader,\n    threshold=0.5\n):\n\n    model.eval()\n\n    preds = []\n    targets_list = []\n\n    with torch.no_grad():\n\n        for images, meta, targets in tqdm(\n            loader,\n            desc=\"Validation\",\n            leave=False\n        ):\n\n            images = images.to(device)\n            meta = meta.to(device)\n\n            outputs = model(\n                images,\n                meta\n            )\n\n            probs = torch.sigmoid(\n                outputs\n            )\n\n            preds.extend(\n                probs.cpu()\n                .numpy()\n                .ravel()\n            )\n\n            targets_list.extend(\n                targets.numpy()\n                .ravel()\n            )\n\n    preds = np.array(preds)\n    targets_list = np.array(targets_list)\n\n    auc = roc_auc_score(\n        targets_list,\n        preds\n    )\n\n    pr_auc = average_precision_score(\n        targets_list,\n        preds\n    )\n\n    binary_preds = (\n        preds > threshold\n    ).astype(int)\n\n    accuracy = accuracy_score(\n        targets_list,\n        binary_preds\n    )\n\n    precision = precision_score(\n        targets_list,\n        binary_preds,\n        zero_division=0\n    )\n\n    recall = recall_score(\n        targets_list,\n        binary_preds,\n        zero_division=0\n    )\n\n    f1 = f1_score(\n        targets_list,\n        binary_preds,\n        zero_division=0\n    )\n\n    tn, fp, fn, tp = safe_confusion_matrix(\n        targets_list,\n        binary_preds\n    )\n\n    specificity = (\n        tn / (tn + fp)\n    )\n\n    fpr, tpr, _ = roc_curve(\n        targets_list,\n        preds\n    )\n\n    metrics = {\n\n        \"AUC\": auc,\n\n        \"PR_AUC\": pr_auc,\n\n        \"Accuracy\": accuracy,\n\n        \"Precision\": precision,\n\n        \"Recall\": recall,\n\n        \"F1\": f1,\n\n        \"Specificity\": specificity\n    }\n\n    return (\n        metrics,\n        fpr,\n        tpr,\n        preds,\n        targets_list\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.650838Z","iopub.execute_input":"2026-06-23T10:39:58.651156Z","iopub.status.idle":"2026-06-23T10:39:58.664406Z","shell.execute_reply.started":"2026-06-23T10:39:58.651103Z","shell.execute_reply":"2026-06-23T10:39:58.663697Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Function","metadata":{}},{"cell_type":"code","source":"scaler = torch.amp.GradScaler(\"cuda\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.665429Z","iopub.execute_input":"2026-06-23T10:39:58.665776Z","iopub.status.idle":"2026-06-23T10:39:58.952685Z","shell.execute_reply.started":"2026-06-23T10:39:58.665753Z","shell.execute_reply":"2026-06-23T10:39:58.951863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(\n    model,\n    loader,\n    optimizer,\n    criterion\n):\n\n    model.train()\n\n    running_loss = 0\n\n    progress_bar = tqdm(\n        loader,\n        desc=\"Training\",\n        leave=False\n    )\n\n    for images, meta, targets in progress_bar:\n\n        images = images.to(device)\n        meta = meta.to(device)\n        targets = targets.to(device)\n\n        optimizer.zero_grad()\n\n        with torch.amp.autocast(\"cuda\"):\n\n            outputs = model(\n                images,\n                meta\n            )\n\n            loss = criterion(\n                outputs,\n                targets\n            )\n\n        scaler.scale(loss).backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            max_norm=1.0\n        )\n\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item()\n\n        progress_bar.set_postfix(\n            loss=f\"{loss.item():.4f}\"\n        )\n\n    return running_loss / len(loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.953800Z","iopub.execute_input":"2026-06-23T10:39:58.954612Z","iopub.status.idle":"2026-06-23T10:39:58.966129Z","shell.execute_reply.started":"2026-06-23T10:39:58.954583Z","shell.execute_reply":"2026-06-23T10:39:58.965543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\n    \"cuda\"\n    if torch.cuda.is_available()\n    else \"cpu\"\n)\n\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.968967Z","iopub.execute_input":"2026-06-23T10:39:58.969308Z","iopub.status.idle":"2026-06-23T10:39:58.983339Z","shell.execute_reply.started":"2026-06-23T10:39:58.969285Z","shell.execute_reply":"2026-06-23T10:39:58.982692Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Fold History Template","metadata":{}},{"cell_type":"code","source":"def create_history():\n\n    return {\n\n        \"train_loss\": [],\n\n        \"auc\": [],\n\n        \"pr_auc\": [],\n\n        \"f1\": [],\n\n        \"recall\": [],\n\n        \"specificity\": []\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.984354Z","iopub.execute_input":"2026-06-23T10:39:58.984636Z","iopub.status.idle":"2026-06-23T10:39:58.995679Z","shell.execute_reply.started":"2026-06-23T10:39:58.984614Z","shell.execute_reply":"2026-06-23T10:39:58.994962Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## CV Training Loop","metadata":{}},{"cell_type":"code","source":"EPOCHS = 5\n\nBATCH_SIZE = 32\n\nFOLDS_TO_RUN = [0,1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:58.996435Z","iopub.execute_input":"2026-06-23T10:39:58.996730Z","iopub.status.idle":"2026-06-23T10:39:59.007918Z","shell.execute_reply.started":"2026-06-23T10:39:58.996699Z","shell.execute_reply":"2026-06-23T10:39:59.006983Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FOLD_RESULTS = {}\n\nfor fold in FOLDS_TO_RUN:\n\n    print(\n        f\"\\n{'='*20} \"\n        f\"FOLD {fold}\"\n        f\" {'='*20}\\n\"\n    )\n\n    train_df = (\n        df[df.fold != fold]\n        .copy()\n        .reset_index(drop=True)\n    )\n\n    val_df = (\n        df[df.fold == fold]\n        .copy()\n        .reset_index(drop=True)\n    )\n\n    val_idx = (\n        df[df.fold == fold]\n        .index\n        .values\n    )\n\n    mean = train_df[\"age_approx\"].mean()\n    std = train_df[\"age_approx\"].std()\n\n    train_df[\"age_approx\"] = (\n        train_df[\"age_approx\"] - mean\n    ) / std\n\n    val_df[\"age_approx\"] = (\n        val_df[\"age_approx\"] - mean\n    ) / std\n\n    train_dataset = ISICDataset(\n        train_df,\n        transforms=get_train_transforms()\n    )\n\n    val_dataset = ISICDataset(\n        val_df,\n        transforms=get_valid_transforms()\n    )\n\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True,\n        persistent_workers=True\n    )\n\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=2,\n        pin_memory=True,\n        persistent_workers=True\n    )\n\n    model = HierarchicalFiLMModel(\n        meta_features=3\n    ).to(device)\n\n    criterion = FocalLoss(\n        alpha=1,\n        gamma=2\n    )\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=1e-4,\n        weight_decay=1e-4\n    )\n\n    scheduler = (\n        torch.optim.lr_scheduler\n        .CosineAnnealingLR(\n            optimizer,\n            T_max=EPOCHS\n        )\n    )\n\n    history = create_history()\n\n    roc_history = []\n\n    best_auc = 0\n\n\n    for epoch in range(EPOCHS):\n\n        print(\n            f\"Fold {fold} | \"\n            f\"Epoch {epoch+1}/{EPOCHS}\"\n        )\n\n        train_loss = train_one_epoch(\n            model,\n            train_loader,\n            optimizer,\n            criterion\n        )\n\n        metrics, fpr, tpr, preds, targets = validate(\n            model,\n            val_loader,\n            threshold=0.5\n        )\n\n        scheduler.step()\n\n        history[\"train_loss\"].append(\n            train_loss\n        )\n\n        history[\"auc\"].append(\n            metrics[\"AUC\"]\n        )\n\n        history[\"pr_auc\"].append(\n            metrics[\"PR_AUC\"]\n        )\n\n        history[\"f1\"].append(\n            metrics[\"F1\"]\n        )\n\n        history[\"recall\"].append(\n            metrics[\"Recall\"]\n        )\n\n        history[\"specificity\"].append(\n            metrics[\"Specificity\"]\n        )\n\n        roc_history.append(\n            (\n                fpr,\n                tpr,\n                metrics[\"AUC\"]\n            )\n        )\n\n        print(metrics)\n\n        if metrics[\"AUC\"] > best_auc:\n\n            best_auc = metrics[\"AUC\"]\n\n            best_state = copy.deepcopy(\n                model.state_dict()\n            )\n\n            print(\n                \"✅ Best model saved\"\n            )\n\n\n    torch.save(\n        best_state,\n        f\"best_hierarchical_film_fold_{fold}.pth\"\n    )\n\n    model.load_state_dict(\n        best_state\n    )\n\n    _, _, _, preds, targets = validate(\n        model,\n        val_loader,\n        threshold=0.5\n    )\n\n    torch.save(\n        {\n            \"fold\": fold,\n\n            \"best_auc\": best_auc,\n\n            \"history\": history,\n\n            \"roc_history\": roc_history,\n\n            \"preds\": preds,\n\n            \"targets\": targets,\n\n            \"val_idx\": val_idx\n        },\n\n        f\"fold_{fold}_results.pth\"\n    )\n\n    print(\n        f\"Fold {fold} complete.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-23T10:39:59.008850Z","iopub.execute_input":"2026-06-23T10:39:59.009059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}