{"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":"markdown","source":"## ISIC MD-MissingnessAwareGating+FiLMAsymmetricalLearningRateImageDominant(TrainFolds0,1)","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport copy\nimport random\n\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.metrics import (\n    roc_auc_score,\n    roc_curve,\n    precision_recall_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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:12.133926Z","iopub.execute_input":"2026-07-28T04:16:12.134898Z","iopub.status.idle":"2026-07-28T04:16:26.543183Z","shell.execute_reply.started":"2026-07-28T04:16:12.134867Z","shell.execute_reply":"2026-07-28T04:16:26.542239Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n\n    random.seed(seed)\n\n    np.random.seed(seed)\n\n    torch.manual_seed(seed)\n\n    torch.cuda.manual_seed(seed)\n\n    torch.backends.cudnn.deterministic = True\n\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:26.544567Z","iopub.execute_input":"2026-07-28T04:16:26.545189Z","iopub.status.idle":"2026-07-28T04:16:26.558811Z","shell.execute_reply.started":"2026-07-28T04:16:26.545144Z","shell.execute_reply":"2026-07-28T04:16:26.557341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(\n\n    \"/kaggle/input/notebooks/farhanaz08/isic-md-missingnessawaregating-film-create-folds/isic_missingness_fold_assignments.csv\"\n\n)\n\nprint(df.shape)\n\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:26.560136Z","iopub.execute_input":"2026-07-28T04:16:26.560554Z","iopub.status.idle":"2026-07-28T04:16:26.961420Z","shell.execute_reply.started":"2026-07-28T04:16:26.560516Z","shell.execute_reply":"2026-07-28T04:16:26.960627Z"}},"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\n        A.VerticalFlip(p=0.5),\n\n        A.RandomRotate90(p=0.5),\n\n        A.Affine(\n\n            scale=(0.95,1.05),\n\n            translate_percent=(0.05,0.05),\n\n            rotate=(-20,20),\n\n            p=0.5\n\n        ),\n\n        A.ColorJitter(\n\n            brightness=0.2,\n\n            contrast=0.2,\n\n            saturation=0.2,\n\n            hue=0.1,\n\n            p=0.5\n\n        ),\n\n        A.Normalize(),\n\n        ToTensorV2()\n\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\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:26.963102Z","iopub.execute_input":"2026-07-28T04:16:26.963395Z","iopub.status.idle":"2026-07-28T04:16:26.969111Z","shell.execute_reply.started":"2026-07-28T04:16:26.963372Z","shell.execute_reply":"2026-07-28T04:16:26.968414Z"}},"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 = df.reset_index(drop=True)\n        self.transforms = transforms\n        self.meta_features = [\n            \"age_approx\",\n            \"sex\",\n            \"site_encoded\"\n\n        ]\n\n\n    def __len__(self):\n        return len(self.df)\n\n\n    def __getitem__(self, idx):\n        row = self.df.loc[idx]\n        image = cv2.imread(row.image_path)\n        image = cv2.cvtColor(\n            image,\n            cv2.COLOR_BGR2RGB\n        )\n\n        if self.transforms:\n            image = self.transforms(\n                image=image\n            )[\"image\"]\n\n\n        meta = torch.tensor(\n            row[self.meta_features]\n            .values\n            .astype(np.float32)\n        )\n\n\n        missing = torch.tensor(\n\n            [\n                row.age_missing,\n                row.sex_missing,\n                row.site_missing\n            ],\n            dtype=torch.float32\n        )\n\n\n        target = torch.tensor(\n            row.target,\n            dtype=torch.float32\n        )\n\n\n        return (\n            image,\n            meta,\n            missing,\n            target\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:26.969903Z","iopub.execute_input":"2026-07-28T04:16:26.970183Z","iopub.status.idle":"2026-07-28T04:16:26.983843Z","shell.execute_reply.started":"2026-07-28T04:16:26.970161Z","shell.execute_reply":"2026-07-28T04:16:26.983109Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Spatial FiLM","metadata":{}},{"cell_type":"code","source":"class SpatialFiLM(nn.Module):\n\n    def __init__(\n        self,\n        channels,\n        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\n    def forward(\n        self,\n        x,\n        meta\n    ):\n        gamma = self.gamma(meta)\n        beta = self.beta(meta)\n        gamma = gamma.unsqueeze(-1).unsqueeze(-1)\n        beta = beta.unsqueeze(-1).unsqueeze(-1)\n\n        return gamma * x + beta\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:26.984747Z","iopub.execute_input":"2026-07-28T04:16:26.985412Z","iopub.status.idle":"2026-07-28T04:16:27.000216Z","shell.execute_reply.started":"2026-07-28T04:16:26.985378Z","shell.execute_reply":"2026-07-28T04:16:26.999585Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Missingness Encoder","metadata":{}},{"cell_type":"code","source":"class MissingnessEncoder(nn.Module):\n\n    def __init__(\n\n        self,\n\n        missing_features=3,\n\n        hidden_dim=32\n\n    ):\n\n        super().__init__()\n\n        self.encoder = nn.Sequential(\n\n            nn.Linear(\n\n                missing_features,\n\n                hidden_dim\n\n            ),\n\n            nn.ReLU(),\n\n            nn.Linear(\n\n                hidden_dim,\n\n                hidden_dim\n\n            ),\n\n            nn.ReLU()\n\n        )\n\n\n    def forward(\n\n        self,\n\n        missing\n\n    ):\n\n        return self.encoder(\n\n            missing\n\n        )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:27.001687Z","iopub.execute_input":"2026-07-28T04:16:27.001919Z","iopub.status.idle":"2026-07-28T04:16:27.014202Z","shell.execute_reply.started":"2026-07-28T04:16:27.001900Z","shell.execute_reply":"2026-07-28T04:16:27.013512Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##  Missingness Gate","metadata":{}},{"cell_type":"code","source":"class MissingnessGate(nn.Module):\n\n    def __init__(\n        self,\n        feature_dim,\n        gate_dim=32\n    ):\n\n        super().__init__()\n\n        self.gate = nn.Sequential(\n            nn.Linear(\n                gate_dim,\n                feature_dim\n            ),\n            nn.Sigmoid()\n        )\n        \n        self.last_weights = None\n\n    def forward(\n        self,\n        features,\n        gate_embedding\n    ):\n\n        weights = self.gate(\n\n            gate_embedding\n\n        )\n\n        self.last_weights = weights.detach()\n\n        weights = weights.unsqueeze(-1).unsqueeze(-1)\n\n        return features * weights\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:27.015092Z","iopub.execute_input":"2026-07-28T04:16:27.015438Z","iopub.status.idle":"2026-07-28T04:16:27.030131Z","shell.execute_reply.started":"2026-07-28T04:16:27.015408Z","shell.execute_reply":"2026-07-28T04:16:27.029298Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## MissingnessAwareHierarchicalFiLMModel","metadata":{}},{"cell_type":"code","source":"class MissingnessAwareHierarchicalFiLMModel(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        # ----------------------------\n        # Hierarchical FiLM\n        # ----------------------------\n        self.film1 = SpatialFiLM(channels[0], meta_features)\n        self.film2 = SpatialFiLM(channels[1], meta_features)\n        self.film3 = SpatialFiLM(channels[2], meta_features)\n        self.film4 = SpatialFiLM(channels[3], meta_features)\n        self.film5 = SpatialFiLM(channels[4], meta_features)\n\n        # ----------------------------\n        # Missingness Encoder\n        # ----------------------------\n        self.missing_encoder = MissingnessEncoder(\n            missing_features=3,\n            hidden_dim=32\n        )\n\n        # ----------------------------\n        # Missingness Gates\n        # ----------------------------\n        self.gate1 = MissingnessGate(channels[0], gate_dim=32)\n        self.gate2 = MissingnessGate(channels[1], gate_dim=32)\n        self.gate3 = MissingnessGate(channels[2], gate_dim=32)\n        self.gate4 = MissingnessGate(channels[3], gate_dim=32)\n        self.gate5 = MissingnessGate(channels[4], gate_dim=32)\n\n        # ----------------------------\n        # Metadata branch\n        # ----------------------------\n        self.meta_branch = nn.Sequential(\n\n            nn.Linear(meta_features, 128),\n\n            nn.ReLU(),\n\n            nn.BatchNorm1d(128),\n\n            nn.Dropout(0.3)\n\n        )\n\n        # ----------------------------\n        # Classifier\n        # ----------------------------\n        self.classifier = nn.Sequential(\n\n            nn.Linear(752, 256),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.4),\n\n            nn.Linear(256, 1)\n\n        )\n\n    def forward(\n        self,\n        image,\n        meta,\n        missing\n    ):\n\n        # Backbone features\n        features = self.backbone(image)\n\n        # Missingness embedding\n        gate_embedding = self.missing_encoder(missing)\n\n        # Stage 1\n        f1 = self.film1(features[0], meta)\n        f1 = self.gate1(f1, gate_embedding)\n\n        # Stage 2\n        f2 = self.film2(features[1], meta)\n        f2 = self.gate2(f2, gate_embedding)\n\n        # Stage 3\n        f3 = self.film3(features[2], meta)\n        f3 = self.gate3(f3, gate_embedding)\n\n        # Stage 4\n        f4 = self.film4(features[3], meta)\n        f4 = self.gate4(f4, gate_embedding)\n\n        # Stage 5\n        f5 = self.film5(features[4], meta)\n        f5 = self.gate5(f5, gate_embedding)\n\n        # Global Average Pooling\n        p1 = F.adaptive_avg_pool2d(f1, 1).flatten(1)\n        p2 = F.adaptive_avg_pool2d(f2, 1).flatten(1)\n        p3 = F.adaptive_avg_pool2d(f3, 1).flatten(1)\n        p4 = F.adaptive_avg_pool2d(f4, 1).flatten(1)\n        p5 = F.adaptive_avg_pool2d(f5, 1).flatten(1)\n\n        img_feat = torch.cat(\n            [p1, p2, p3, p4, p5],\n            dim=1\n        )\n\n        meta_feat = self.meta_branch(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), self.gate5.last_weights","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:27.031035Z","iopub.execute_input":"2026-07-28T04:16:27.031288Z","iopub.status.idle":"2026-07-28T04:16:27.046168Z","shell.execute_reply.started":"2026-07-28T04:16:27.031268Z","shell.execute_reply":"2026-07-28T04:16:27.045479Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Focal Loss","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(logits)\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-07-28T04:16:27.048321Z","iopub.execute_input":"2026-07-28T04:16:27.048594Z","iopub.status.idle":"2026-07-28T04:16:27.062360Z","shell.execute_reply.started":"2026-07-28T04:16:27.048573Z","shell.execute_reply":"2026-07-28T04:16:27.061687Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Safe Confusion Matrix","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-07-28T04:16:27.063189Z","iopub.execute_input":"2026-07-28T04:16:27.063646Z","iopub.status.idle":"2026-07-28T04:16:27.075280Z","shell.execute_reply.started":"2026-07-28T04:16:27.063620Z","shell.execute_reply":"2026-07-28T04:16:27.074360Z"}},"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    gate_values = []\n    missing_flags = []\n\n    with torch.no_grad():\n\n        for images, meta, missing, targets in tqdm(\n            loader,\n            desc=\"Validation\",\n            leave=False\n        ):\n\n            images = images.to(device)\n            meta = meta.to(device)\n            missing = missing.to(device)\n\n            outputs, alpha = model(\n                images,\n                meta,\n                missing\n            )\n\n            probs = torch.sigmoid(outputs)\n\n            preds.extend(\n                probs.cpu().numpy().ravel()\n            )\n\n            targets_list.extend(\n                targets.numpy().ravel()\n            )\n\n            gate_values.extend(\n                alpha.cpu().numpy().mean(axis=1)\n            )\n\n            missing_flags.extend(\n                (missing.sum(dim=1) > 0)\n                .cpu()\n                .numpy()\n            )\n\n    preds = np.array(preds)\n    targets_list = np.array(targets_list)\n    gate_values = np.array(gate_values)\n    missing_flags = np.array(missing_flags)\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 = tn / (tn + fp)\n\n    fpr, tpr, _ = roc_curve(\n        targets_list,\n        preds\n    )\n\n    precision_curve, recall_curve, _ = precision_recall_curve(\n        targets_list,\n        preds\n    )\n    \n    cm = confusion_matrix(\n        targets_list,\n        binary_preds\n    )\n\n    metrics = {\n\n        \"AUC\": auc,\n        \"PR_AUC\": pr_auc,\n        \"Accuracy\": accuracy,\n        \"Precision\": precision,\n        \"Recall\": recall,\n        \"F1\": f1,\n        \"Specificity\": specificity\n\n    }\n\n    return (\n        metrics,\n        fpr,\n        tpr,\n        precision_curve,\n        recall_curve,\n        cm,\n        preds,\n        targets_list,\n        gate_values,\n        missing_flags\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:27.076325Z","iopub.execute_input":"2026-07-28T04:16:27.076868Z","iopub.status.idle":"2026-07-28T04:16:27.091929Z","shell.execute_reply.started":"2026-07-28T04:16:27.076845Z","shell.execute_reply":"2026-07-28T04:16:27.091384Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Mixed Precision Scaler","metadata":{}},{"cell_type":"code","source":"scaler = torch.amp.GradScaler(\"cuda\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:27.092718Z","iopub.execute_input":"2026-07-28T04:16:27.093344Z","iopub.status.idle":"2026-07-28T04:16:27.378899Z","shell.execute_reply.started":"2026-07-28T04:16:27.093321Z","shell.execute_reply":"2026-07-28T04:16:27.377782Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Function","metadata":{}},{"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, missing, targets in progress_bar:\n\n        images = images.to(device)\n        meta = meta.to(device)\n        missing = missing.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                missing\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-07-28T04:16:27.379925Z","iopub.execute_input":"2026-07-28T04:16:27.380230Z","iopub.status.idle":"2026-07-28T04:16:27.392822Z","shell.execute_reply.started":"2026-07-28T04:16:27.380196Z","shell.execute_reply":"2026-07-28T04:16:27.391846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Device","metadata":{}},{"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-07-28T04:16:27.393816Z","iopub.execute_input":"2026-07-28T04:16:27.394143Z","iopub.status.idle":"2026-07-28T04:16:27.406186Z","shell.execute_reply.started":"2026-07-28T04:16:27.394109Z","shell.execute_reply":"2026-07-28T04:16:27.405445Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## History Template","metadata":{}},{"cell_type":"code","source":"def create_history():\n\n    return {\n\n        \"train_loss\": [],\n        \"auc\": [],\n        \"pr_auc\": [],\n        \"f1\": [],\n        \"recall\": [],\n        \"specificity\": [],\n\n        # global gate behavior\n        \"gate_mean\": [],\n        \"gate_std\": [],\n\n        # thesis-level diagnostics (optional but powerful)\n        \"gate_melanoma\": [],\n        \"gate_benign\": [],\n        \"gate_missing\": [],\n        \"gate_present\": []\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-28T04:16:27.407135Z","iopub.execute_input":"2026-07-28T04:16:27.407425Z","iopub.status.idle":"2026-07-28T04:16:27.418999Z","shell.execute_reply.started":"2026-07-28T04:16:27.407404Z","shell.execute_reply":"2026-07-28T04:16:27.418404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 5\n\nBATCH_SIZE = 32\n\n# Change this between runs\n# Run 1: [0,1]\n# Run 2: [2,3]\n# Run 3: [4]\n\nFOLDS_TO_RUN = [2,3]\n\nFOLD_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    # --------------------------\n    # Standardize Age\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    # --------------------------\n    # Datasets\n    # --------------------------\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    # --------------------------\n    # Model\n    # --------------------------\n\n    model = MissingnessAwareHierarchicalFiLMModel(\n        meta_features=3\n    ).to(device)\n\n    criterion = FocalLoss(\n        alpha=1,\n        gamma=2\n    )\n    # --------------------------\n    # Image Dominant\n    # --------------------------\n    optimizer = torch.optim.AdamW(\n\n        [\n    \n            {\n    \n                \"params\":model.backbone.parameters(),\n                \"lr\":5e-4\n    \n            },\n    \n            {\n    \n                \"params\":\n                    list(model.film1.parameters()) +\n                    list(model.film2.parameters()) +\n                    list(model.film3.parameters()) +\n                    list(model.film4.parameters()) +\n                    list(model.film5.parameters()),\n    \n                \"lr\":5e-4\n    \n            },\n    \n            {\n    \n                \"params\":\n                    list(model.missing_encoder.parameters()) +\n    \n                    list(model.gate1.parameters()) +\n                    list(model.gate2.parameters()) +\n                    list(model.gate3.parameters()) +\n                    list(model.gate4.parameters()) +\n                    list(model.gate5.parameters()),\n    \n                \"lr\":1e-5\n    \n            },\n    \n            {\n    \n                \"params\":\n                    list(model.meta_branch.parameters()) +\n                    list(model.classifier.parameters()),\n    \n                \"lr\":1e-5\n    \n            }\n    \n        ],\n    \n        weight_decay=1e-4\n    \n    )\n\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer,\n        T_max=EPOCHS\n    )\n\n    history = create_history()\n\n    roc_history = []\n\n    best_auc = 0\n\n    best_state = None\n\n    # --------------------------\n    # Epoch Loop\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, pr_precision, pr_recall, cm, preds, targets, gate_values, missing_flag = validate(\n            model,\n            val_loader,\n            threshold=0.5\n        )\n\n        scheduler.step()\n\n        history[\"train_loss\"].append(train_loss)\n        history[\"auc\"].append(metrics[\"AUC\"])\n        history[\"pr_auc\"].append(metrics[\"PR_AUC\"])\n        history[\"f1\"].append(metrics[\"F1\"])\n        history[\"recall\"].append(metrics[\"Recall\"])\n        history[\"specificity\"].append(metrics[\"Specificity\"])\n        history[\"gate_mean\"].append(np.mean(gate_values))\n        history[\"gate_std\"].append(np.std(gate_values))\n        history[\"gate_melanoma\"].append(np.mean(gate_values[targets == 1]))\n        history[\"gate_benign\"].append(np.mean(gate_values[targets == 0]))\n        history[\"gate_missing\"].append(np.mean(gate_values[missing_flag == 1]))\n        history[\"gate_present\"].append(np.mean(gate_values[missing_flag == 0]))\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(\"✅ Best model saved\")\n\n    # --------------------------\n    # Save Best Model\n    # --------------------------\n\n    torch.save(\n        best_state,\n        f\"best_missingness_film_fold_{fold}.pth\"\n    )\n\n    model.load_state_dict(best_state)\n\n    _, fpr, tpr, pr_precision, pr_recall, cm, preds, targets, gate_values, missing_flag = validate(\n        model,\n        val_loader,\n        threshold=0.5\n    )\n\n    # ----------------------------\n    # Gate statistics\n    # ----------------------------\n    \n    gate_stats = {\n    \n        \"mean\": float(np.mean(gate_values)),\n        \"std\": float(np.std(gate_values)),\n        \"median\": float(np.median(gate_values)),\n        \"min\": float(np.min(gate_values)),\n        \"max\": float(np.max(gate_values)),\n    \n        \"melanoma_mean\": float(np.mean(gate_values[targets == 1])) if np.any(targets == 1) else None,\n        \"benign_mean\": float(np.mean(gate_values[targets == 0])) if np.any(targets == 0) else None,\n        \n        \"missing_mean\": float(np.mean(gate_values[missing_flag])) if np.any(missing_flag) else None,\n        \"present_mean\": float(np.mean(gate_values[~missing_flag])) if np.any(~missing_flag) else None,\n    }\n    \n    # ----------------------------\n    # Best Threshold\n    # ----------------------------\n    \n    thresholds = np.arange(0.05, 0.96, 0.01)\n    \n    best_threshold = 0.5\n    best_f1 = -1\n    \n    for t in thresholds:\n    \n        pred = (preds > t).astype(int)\n    \n        score = f1_score(targets, pred)\n    \n        if score > best_f1:\n    \n            best_f1 = score\n            best_threshold = t\n\n    torch.save(\n\n        {\n\n            \"fold\": fold,\n            \"best_auc\": best_auc,\n            \"history\": history,\n            \"roc_history\": roc_history,       \n            \"preds\": preds,\n            \"targets\": targets,\n            \"val_idx\": val_idx,\n            \"gate_values\": gate_values,\n            \"gate_stats\": gate_stats,\n            \"best_threshold\": best_threshold,\n            \"best_f1\": best_f1,\n            \"best_fpr\": fpr,\n            \"best_tpr\": tpr,\n            \"pr_precision\": pr_precision,\n            \"pr_recall\": pr_recall,\n            \"confusion_matrix\": cm,\n            \"age_missing\": val_df[\"age_missing\"].values,\n            \"sex_missing\": val_df[\"sex_missing\"].values,\n            \"site_missing\": val_df[\"site_missing\"].values,\n            \"age\": val_df[\"age_approx\"].values,\n            \"sex\": val_df[\"sex\"].values,\n            \"site\": val_df[\"site_encoded\"].values,\n            \"image_name\": val_df[\"image_name\"].values\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-07-28T04:16:27.420387Z","iopub.execute_input":"2026-07-28T04:16:27.420605Z","execution_failed":"2026-07-28T04:17:47.709Z"}},"outputs":[],"execution_count":null}]}