{"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":20270,"databundleVersionId":1222630},{"sourceType":"datasetVersion","sourceId":1353811,"datasetId":762203,"databundleVersionId":1386220}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Confirming paths","metadata":{}},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    print(dirname)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-25T12:24:00.355486Z","iopub.execute_input":"2026-05-25T12:24:00.355816Z","iopub.status.idle":"2026-05-25T12:27:34.662696Z","shell.execute_reply.started":"2026-05-25T12:24:00.355774Z","shell.execute_reply":"2026-05-25T12:27:34.661672Z"}},"outputs":[],"execution_count":null},{"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    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-05-26T13:33:40.212324Z","iopub.execute_input":"2026-05-26T13:33:40.212627Z","iopub.status.idle":"2026-05-26T13:34:00.939269Z","shell.execute_reply.started":"2026-05-26T13:33:40.212590Z","shell.execute_reply":"2026-05-26T13:34:00.938682Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Reproducibility","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-05-26T13:47:32.268035Z","iopub.execute_input":"2026-05-26T13:47:32.268556Z","iopub.status.idle":"2026-05-26T13:47:32.297690Z","shell.execute_reply.started":"2026-05-26T13:47:32.268509Z","shell.execute_reply":"2026-05-26T13:47:32.297116Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load 2020 dataset","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(\n    \"/kaggle/input/competitions/siim-isic-melanoma-classification/train.csv\"\n)\n\nimg_dir = \"/kaggle/input/competitions/siim-isic-melanoma-classification/jpeg/train\"\n\ndf[\"image_path\"] = df[\"image_name\"].apply(\n    lambda x: os.path.join(img_dir, f\"{x}.jpg\")\n)\n\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:34.539626Z","iopub.execute_input":"2026-05-26T13:47:34.540337Z","iopub.status.idle":"2026-05-26T13:47:34.717991Z","shell.execute_reply.started":"2026-05-26T13:47:34.540306Z","shell.execute_reply":"2026-05-26T13:47:34.717280Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Missing Values Handling ","metadata":{}},{"cell_type":"code","source":"# MISSINGNESS INDICATORS\n\ndf[\"age_missing\"] = (\n    df[\"age_approx\"].isna().astype(int)\n)\n\ndf[\"sex_missing\"] = (\n    df[\"sex\"].isna().astype(int)\n)\n\ndf[\"site_missing\"] = (\n    df[\"anatom_site_general_challenge\"]\n    .isna()\n    .astype(int)\n)\n\n# IMPUTATION\n\ndf[\"age_approx\"] = df[\"age_approx\"].fillna(\n    df[\"age_approx\"].median()\n)\n\ndf[\"sex\"] = df[\"sex\"].map({\n    \"male\": 1,\n    \"female\": 0\n})\n\ndf[\"sex\"] = df[\"sex\"].fillna(-1)\n\ndf[\"anatom_site_general_challenge\"] = (\n    df[\"anatom_site_general_challenge\"]\n    .fillna(\"unknown\")\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:37.464731Z","iopub.execute_input":"2026-05-26T13:47:37.465321Z","iopub.status.idle":"2026-05-26T13:47:37.497503Z","shell.execute_reply.started":"2026-05-26T13:47:37.465284Z","shell.execute_reply":"2026-05-26T13:47:37.496421Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Encode Anatomical Site","metadata":{}},{"cell_type":"code","source":"le = LabelEncoder()\n\ndf[\"site_encoded\"] = le.fit_transform(\n    df[\"anatom_site_general_challenge\"]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:39.937141Z","iopub.execute_input":"2026-05-26T13:47:39.937841Z","iopub.status.idle":"2026-05-26T13:47:39.948259Z","shell.execute_reply.started":"2026-05-26T13:47:39.937808Z","shell.execute_reply":"2026-05-26T13:47:39.947612Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\n    df[\n        [\n            'image_path',\n            'age_approx',\n            'sex',\n            'site_encoded',\n            'age_missing',\n            'sex_missing',\n            'site_missing',\n            'target'\n        ]\n    ].head()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:41.073881Z","iopub.execute_input":"2026-05-26T13:47:41.074471Z","iopub.status.idle":"2026-05-26T13:47:41.085871Z","shell.execute_reply.started":"2026-05-26T13:47:41.074441Z","shell.execute_reply":"2026-05-26T13:47:41.085057Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## PATIENT-WISE SPLIT\nThis avoids leakage ","metadata":{}},{"cell_type":"code","source":"df[\"fold\"] = -1\n\nsgkf = StratifiedGroupKFold(\n    n_splits=5,\n    shuffle=True,\n    random_state=42\n)\n\nfor fold, (train_idx, val_idx) in enumerate(\n    sgkf.split(\n        df,\n        y=df[\"target\"],\n        groups=df[\"patient_id\"]\n    )\n):\n\n    df.loc[val_idx, \"fold\"] = fold\n\ntrain_df = df[df.fold != 0].reset_index(drop=True)\nval_df   = df[df.fold == 0].reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:43.756175Z","iopub.execute_input":"2026-05-26T13:47:43.756434Z","iopub.status.idle":"2026-05-26T13:47:44.498861Z","shell.execute_reply.started":"2026-05-26T13:47:43.756414Z","shell.execute_reply":"2026-05-26T13:47:44.497886Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Normalize Metadata","metadata":{}},{"cell_type":"code","source":"mean = train_df[\"age_approx\"].mean()\nstd  = train_df[\"age_approx\"].std()\n\ntrain_df[\"age_approx\"] = (\n    train_df[\"age_approx\"] - mean\n) / std\n\nval_df[\"age_approx\"] = (\n    val_df[\"age_approx\"] - mean\n) / std","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:46.210396Z","iopub.execute_input":"2026-05-26T13:47:46.211230Z","iopub.status.idle":"2026-05-26T13:47:46.217601Z","shell.execute_reply.started":"2026-05-26T13:47:46.211198Z","shell.execute_reply":"2026-05-26T13:47:46.216859Z"}},"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-05-26T13:47:48.903899Z","iopub.execute_input":"2026-05-26T13:47:48.904709Z","iopub.status.idle":"2026-05-26T13:47:48.910711Z","shell.execute_reply.started":"2026-05-26T13:47:48.904675Z","shell.execute_reply":"2026-05-26T13:47:48.909821Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class ISICDataset(Dataset):\n\n    def __init__(self, df, transforms=None):\n\n        self.df = df.reset_index(drop=True)\n        self.transforms = transforms\n\n        self.meta_features = [\n            \"age_approx\",\n            \"sex\",\n            \"site_encoded\",\n            \"age_missing\",\n            \"sex_missing\",\n            \"site_missing\"\n        ]\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.loc[idx]\n\n        image = cv2.imread(row.image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        if self.transforms:\n            image = self.transforms(image=image)[\"image\"]\n\n        meta = torch.tensor(\n            row[self.meta_features].values.astype(np.float32)\n        )\n\n        target = torch.tensor(row.target).float()\n\n        return image, meta, target","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:51.607885Z","iopub.execute_input":"2026-05-26T13:47:51.608692Z","iopub.status.idle":"2026-05-26T13:47:51.614211Z","shell.execute_reply.started":"2026-05-26T13:47:51.608659Z","shell.execute_reply":"2026-05-26T13:47:51.613430Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Hierarchical SpatialGated 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(meta_features, channels)\n        self.beta  = nn.Linear(meta_features, channels)\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-05-26T13:47:54.196979Z","iopub.execute_input":"2026-05-26T13:47:54.197274Z","iopub.status.idle":"2026-05-26T13:47:54.202174Z","shell.execute_reply.started":"2026-05-26T13:47:54.197251Z","shell.execute_reply":"2026-05-26T13:47:54.201493Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Hierarchical FiLM Model","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# PHASE 3B\n# BIDIRECTIONAL ADAPTIVE GATED HIERARCHICAL FiLM\n# ============================================================\n\nclass HierarchicalFiLMAdaptiveGatedModel(nn.Module):\n\n    def __init__(self, meta_features=6):\n\n        super().__init__()\n\n        # ====================================================\n        # EfficientNet Backbone\n        # ====================================================\n        self.backbone = timm.create_model(\n            \"efficientnet_b3\",\n            pretrained=True,\n            features_only=True\n        )\n\n        channels = [24, 32, 48, 136, 384]\n\n        # ====================================================\n        # Hierarchical Spatial FiLM Layers\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        # ====================================================\n        # Metadata Encoder 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        # ADAPTIVE BIDIRECTIONAL GATE\n        # Uses BOTH image + metadata\n        # ====================================================\n        self.gate = nn.Sequential(\n\n            nn.Linear(384 + 128, 128),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.2),\n\n            nn.Linear(128, 1),\n\n            nn.Sigmoid()\n        )\n\n        # ====================================================\n        # Final Classifier\n        # ====================================================\n        self.classifier = nn.Sequential(\n\n            nn.Linear(384 + 128, 256),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.4),\n\n            nn.Linear(256, 1)\n        )\n\n    def forward(self, image, meta):\n\n        # ====================================================\n        # Backbone Features\n        # ====================================================\n        features = self.backbone(image)\n\n        # ====================================================\n        # Hierarchical FiLM Conditioning\n        # ====================================================\n        f1 = self.film1(features[0], meta)\n\n        f2 = self.film2(features[1], meta)\n\n        f3 = self.film3(features[2], meta)\n\n        f4 = self.film4(features[3], meta)\n\n        f5 = self.film5(features[4], meta)\n\n        # ====================================================\n        # Global Image Feature\n        # ====================================================\n        img_feat = F.adaptive_avg_pool2d(f5, 1)\n\n        img_feat = img_feat.view(\n            img_feat.size(0),\n            -1\n        )\n\n        # ====================================================\n        # Metadata Feature\n        # ====================================================\n        meta_feat = self.meta_branch(meta)\n\n        # ====================================================\n        # Joint Representation\n        # ====================================================\n        joint_feat = torch.cat([\n            img_feat,\n            meta_feat\n        ], dim=1)\n\n        # ====================================================\n        # Adaptive Gate\n        # alpha -> image importance\n        # (1-alpha) -> metadata importance\n        # ====================================================\n        alpha = self.gate(joint_feat)\n\n        # ====================================================\n        # Gated Features\n        # ====================================================\n        gated_img_feat = alpha * img_feat\n\n        gated_meta_feat = (\n            1 - alpha\n        ) * meta_feat\n\n        # ====================================================\n        # Final Fusion\n        # ====================================================\n        fused = torch.cat([\n            gated_img_feat,\n            gated_meta_feat\n        ], dim=1)\n\n        # ====================================================\n        # Classification\n        # ====================================================\n        out = self.classifier(fused)\n\n        return out.squeeze(1), alpha","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:47:56.225569Z","iopub.execute_input":"2026-05-26T13:47:56.225854Z","iopub.status.idle":"2026-05-26T13:47:56.236466Z","shell.execute_reply.started":"2026-05-26T13:47:56.225829Z","shell.execute_reply":"2026-05-26T13:47:56.235571Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Focal Loss Cell","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n\n    def __init__(self, alpha=1, gamma=2):\n\n        super().__init__()\n\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(self, logits, targets):\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 = self.alpha * (1 - pt) ** self.gamma * bce\n\n        return loss.mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:48:01.584309Z","iopub.execute_input":"2026-05-26T13:48:01.584770Z","iopub.status.idle":"2026-05-26T13:48:01.590226Z","shell.execute_reply.started":"2026-05-26T13:48:01.584737Z","shell.execute_reply":"2026-05-26T13:48:01.589481Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DataLoaders\n\nUse num_workers=0","metadata":{}},{"cell_type":"code","source":"train_dataset = ISICDataset(\n    train_df,\n    transforms=get_train_transforms()\n)\n\nval_dataset = ISICDataset(\n    val_df,\n    transforms=get_valid_transforms()\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=16,\n    shuffle=True,\n    num_workers=0\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=16,\n    shuffle=False,\n    num_workers=0\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:48:03.927743Z","iopub.execute_input":"2026-05-26T13:48:03.928494Z","iopub.status.idle":"2026-05-26T13:48:03.942357Z","shell.execute_reply.started":"2026-05-26T13:48:03.928464Z","shell.execute_reply":"2026-05-26T13:48:03.941495Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation Function","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import average_precision_score\n\n\ndef validate(model, loader, threshold=0.5):\n\n    model.eval()\n\n    preds = []\n    targets_list = []\n    gate_values = []\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, alpha = model(images, meta)\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            # -----------------------------------------\n            # Store gate values\n            # -----------------------------------------\n            gate_values.extend(\n                alpha.cpu().numpy().ravel()\n            )\n\n    preds = np.array(preds)\n    targets_list = np.array(targets_list)\n    gate_values = np.array(gate_values)\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 = confusion_matrix(\n        targets_list,\n        binary_preds\n    ).ravel()\n\n    specificity = tn / (tn + fp)\n\n    fpr, tpr, _ = roc_curve(\n        targets_list,\n        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    return (\n        metrics,\n        fpr,\n        tpr,\n        preds,\n        targets_list,\n        gate_values\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:48:06.904336Z","iopub.execute_input":"2026-05-26T13:48:06.904960Z","iopub.status.idle":"2026-05-26T13:48:06.913339Z","shell.execute_reply.started":"2026-05-26T13:48:06.904931Z","shell.execute_reply":"2026-05-26T13:48:06.912742Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Setup","metadata":{}},{"cell_type":"code","source":"device = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nmodel = HierarchicalFiLMAdaptiveGatedModel(\n    meta_features=6\n).to(device)\n\ncriterion = FocalLoss(\n    alpha=1,\n    gamma=2\n)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:48:14.193041Z","iopub.execute_input":"2026-05-26T13:48:14.193717Z","iopub.status.idle":"2026-05-26T13:48:17.026693Z","shell.execute_reply.started":"2026-05-26T13:48:14.193686Z","shell.execute_reply":"2026-05-26T13:48:17.025963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=5\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:48:17.028093Z","iopub.execute_input":"2026-05-26T13:48:17.028450Z","iopub.status.idle":"2026-05-26T13:48:17.032616Z","shell.execute_reply.started":"2026-05-26T13:48:17.028422Z","shell.execute_reply":"2026-05-26T13:48:17.031843Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Function","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader):\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        outputs, alpha = model(images, meta)\n\n        loss = criterion(outputs, targets)\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            max_norm=1.0\n        )\n\n        optimizer.step()\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-05-26T13:48:19.509509Z","iopub.execute_input":"2026-05-26T13:48:19.510252Z","iopub.status.idle":"2026-05-26T13:48:19.516126Z","shell.execute_reply.started":"2026-05-26T13:48:19.510223Z","shell.execute_reply":"2026-05-26T13:48:19.515510Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Training Loop","metadata":{}},{"cell_type":"code","source":"history = {\n    \"train_loss\": [],\n    \"auc\": [],\n    \"pr_auc\": [],\n    \"f1\": [],\n    \"recall\": [],\n    \"specificity\": []\n}\n\nbest_auc = 0\n\nEPOCHS = 5\n\nfor epoch in range(EPOCHS):\n\n    print(f\"\\n===== Epoch {epoch+1}/{EPOCHS} =====\")\n\n    train_loss = train_one_epoch(\n        model,\n        train_loader\n    )\n\n    metrics, fpr, tpr, preds, targets_list, gate_values = validate(\n        model,\n        val_loader,\n        threshold=0.5\n    )\n\n    scheduler.step()\n\n    if metrics[\"AUC\"] > best_auc:\n\n        best_auc = metrics[\"AUC\"]\n\n        torch.save(\n            model.state_dict(),\n            \"best_hierarchical_film.pth\"\n        )\n\n        print(\"✅ Best model saved!\")\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(\n        metrics[\"Specificity\"]\n    )\n\n    print(f\"Loss        : {train_loss:.4f}\")\n    print(f\"AUC         : {metrics['AUC']:.4f}\")\n    print(f\"PR-AUC      : {metrics['PR_AUC']:.4f}\")\n    print(f\"F1          : {metrics['F1']:.4f}\")\n    print(f\"Recall      : {metrics['Recall']:.4f}\")\n    print(f\"Specificity : {metrics['Specificity']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-26T13:48:22.291508Z","iopub.execute_input":"2026-05-26T13:48:22.292246Z","iopub.status.idle":"2026-05-26T13:49:53.147830Z","shell.execute_reply.started":"2026-05-26T13:48:22.292215Z","shell.execute_reply":"2026-05-26T13:49:53.146689Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Adaptive Gate Analysis","metadata":{}},{"cell_type":"code","source":"print(\"===================================\")\nprint(\"ADAPTIVE GATE ANALYSIS\")\nprint(\"===================================\")\n\nprint(f\"\\nMean alpha: {gate_values.mean():.4f}\")\n\nprint(\n    f\"Image dominant (>0.5): \"\n    f\"{(gate_values > 0.5).mean():.4f}\"\n)\n\nprint(\n    f\"Metadata dominant (<0.5): \"\n    f\"{(gate_values < 0.5).mean():.4f}\"\n)\n\nmelanoma_alpha = gate_values[\n    targets_list == 1\n]\n\nbenign_alpha = gate_values[\n    targets_list == 0\n]\n\nprint(\n    f\"\\nMean melanoma alpha: \"\n    f\"{melanoma_alpha.mean():.4f}\"\n)\n\nprint(\n    f\"Mean benign alpha: \"\n    f\"{benign_alpha.mean():.4f}\"\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Missingness aware Gate Anlysis","metadata":{}},{"cell_type":"code","source":"print(\"===================================\")\nprint(\"MISSINGNESS-AWARE GATE ANALYSIS\")\nprint(\"===================================\")\n\n# ============================================================\n# AGE MISSING\n# ============================================================\n\nage_missing_alpha = gate_values[\n    val_df[\"age_missing\"].values == 1\n]\n\nage_present_alpha = gate_values[\n    val_df[\"age_missing\"].values == 0\n]\n\nprint(\"\\nAGE\")\nprint(f\"Missing count : {len(age_missing_alpha)}\")\nprint(f\"Present count : {len(age_present_alpha)}\")\n\nif len(age_missing_alpha) > 0:\n    print(\n        f\"Mean alpha (missing): \"\n        f\"{age_missing_alpha.mean():.4f}\"\n    )\n\nprint(\n    f\"Mean alpha (present): \"\n    f\"{age_present_alpha.mean():.4f}\"\n)\n\n# ============================================================\n# SEX MISSING\n# ============================================================\n\nsex_missing_alpha = gate_values[\n    val_df[\"sex_missing\"].values == 1\n]\n\nsex_present_alpha = gate_values[\n    val_df[\"sex_missing\"].values == 0\n]\n\nprint(\"\\nSEX\")\nprint(f\"Missing count : {len(sex_missing_alpha)}\")\nprint(f\"Present count : {len(sex_present_alpha)}\")\n\nif len(sex_missing_alpha) > 0:\n    print(\n        f\"Mean alpha (missing): \"\n        f\"{sex_missing_alpha.mean():.4f}\"\n    )\n\nprint(\n    f\"Mean alpha (present): \"\n    f\"{sex_present_alpha.mean():.4f}\"\n)\n\n# ============================================================\n# SITE MISSING\n# ============================================================\n\nsite_missing_alpha = gate_values[\n    val_df[\"site_missing\"].values == 1\n]\n\nsite_present_alpha = gate_values[\n    val_df[\"site_missing\"].values == 0\n]\n\nprint(\"\\nSITE\")\nprint(f\"Missing count : {len(site_missing_alpha)}\")\nprint(f\"Present count : {len(site_present_alpha)}\")\n\nif len(site_missing_alpha) > 0:\n    print(\n        f\"Mean alpha (missing): \"\n        f\"{site_missing_alpha.mean():.4f}\"\n    )\n\nprint(\n    f\"Mean alpha (present): \"\n    f\"{site_present_alpha.mean():.4f}\"\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Statistical Validation","metadata":{}},{"cell_type":"code","source":"from scipy.stats import mannwhitneyu\n\nprint(\"===================================\")\nprint(\"STATISTICAL TESTS\")\nprint(\"===================================\")\n\n# ============================================================\n# AGE\n# ============================================================\n\nif len(age_missing_alpha) > 0:\n\n    stat, p = mannwhitneyu(\n        age_missing_alpha,\n        age_present_alpha\n    )\n\n    print(f\"\\nAGE missing vs present p-value: {p:.6f}\")\n\n# ============================================================\n# SEX\n# ============================================================\n\nif len(sex_missing_alpha) > 0:\n\n    stat, p = mannwhitneyu(\n        sex_missing_alpha,\n        sex_present_alpha\n    )\n\n    print(f\"SEX missing vs present p-value: {p:.6f}\")\n\n# ============================================================\n# SITE\n# ============================================================\n\nif len(site_missing_alpha) > 0:\n\n    stat, p = mannwhitneyu(\n        site_missing_alpha,\n        site_present_alpha\n    )\n\n    print(f\"SITE missing vs present p-value: {p:.6f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Alpha DIstrinution Visualization(Histogram)","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nplt.hist(\n    gate_values,\n    bins=30,\n    edgecolor='black'\n)\n\nplt.xlabel(\"Alpha Value\")\n\nplt.ylabel(\"Frequency\")\n\nplt.title(\n    \"Distribution of Adaptive Gate Values\"\n)\n\nplt.grid()\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Class-Wise Alpha(Box Plot)","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(6,5))\n\nplt.boxplot(\n    [benign_alpha, melanoma_alpha],\n    labels=[\"Benign\", \"Melanoma\"]\n)\n\nplt.ylabel(\"Alpha\")\n\nplt.title(\"Alpha Distribution by Class\")\n\nplt.grid()\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Statistical Test","metadata":{}},{"cell_type":"code","source":"from scipy.stats import mannwhitneyu\n\nstat, p = mannwhitneyu(\n    benign_alpha,\n    melanoma_alpha\n)\n\nprint(f\"Mann-Whitney U p-value: {p:.6f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ROC Curve","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(7,7))\n\nplt.plot(\n    fpr,\n    tpr,\n    label=f\"AUC={metrics['AUC']:.4f}\"\n)\n\nplt.plot([0,1], [0,1], linestyle=\"--\")\n\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\n\nplt.title(\"ROC Curve\")\n\nplt.legend()\nplt.grid()\n\nplt.savefig(\"roc_curve.png\", dpi=300)\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Threshold Search ","metadata":{}},{"cell_type":"code","source":"thresholds = np.arange(0.1, 1.0, 0.05)\n\nbest_f1 = 0\nbest_threshold = 0.5\nbest_preds = None\nbest_targets = None\n\nfor t in thresholds:\n\n    metrics, _, _, preds, targets_list, gate_values = validate(\n        model,\n        val_loader,\n        threshold=t\n    )\n\n    if metrics[\"F1\"] > best_f1:\n\n        best_f1 = metrics[\"F1\"]\n        best_threshold = t\n        best_preds = preds\n        best_targets = targets_list\n\nprint(f\"Best Threshold: {best_threshold:.2f}\")\nprint(f\"Best F1: {best_f1:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Confusion Matrix","metadata":{}},{"cell_type":"code","source":"binary_preds = (\n    best_preds > best_threshold\n).astype(int)\n\ncm = confusion_matrix(\n    best_targets,\n    binary_preds\n)\n\ndisp = ConfusionMatrixDisplay(\n    confusion_matrix=cm,\n    display_labels=[\n        \"Benign\",\n        \"Melanoma\"\n    ]\n)\n\ndisp.plot(cmap=\"Blues\")\n\nplt.title(\n    f\"Confusion Matrix (t={best_threshold:.2f})\"\n)\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Metric Curves Across Epochs","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(10,5))\n\nplt.plot(history[\"auc\"], label=\"AUC\")\nplt.plot(history[\"pr_auc\"], label=\"PR-AUC\")\nplt.plot(history[\"f1\"], label=\"F1\")\nplt.plot(history[\"recall\"], label=\"Recall\")\nplt.plot(history[\"specificity\"], label=\"Specificity\")\n\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Score\")\n\nplt.title(\"Validation Metrics Across Epochs\")\n\nplt.legend()\nplt.grid()\n\nplt.savefig(\"metrics.png\", dpi=300)\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = {\n    \"best_auc\": best_auc,\n    \"history\": history\n}\n\ntorch.save(results, \"training_results.pth\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.savefig(\"roc_curve.png\", dpi=300)\n\nplt.savefig(\"metrics.png\", dpi=300)\n\nplt.savefig(\"confusion_matrix.png\", dpi=300)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This Phase 3B architecture tests:\n\n> α=g(v,m)\n\ninstead of:\n\n> α=g(m)\n\nSo now:\n\n* image ambiguity,\n* visual confidence,\n* metadata priors,\n* lesion characteristics\n\njointly determine modality dominance.","metadata":{}}]}