{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"databundleVersionId":17216535,"datasetId":10409512,"mountSlug":"datasets/jemimambabacha/explainability","sourceId":16234422,"sourceType":"datasetVersion"},{"databundleVersionId":8017664,"datasetId":4646801,"mountSlug":"datasets/ctvmnn/ham10000","sourceId":7910038,"sourceType":"datasetVersion"},{"databundleVersionId":1333500,"datasetId":752995,"mountSlug":"datasets/tschandl/ham10000-lesion-segmentations","sourceId":1301322,"sourceType":"datasetVersion"}],"dockerImageVersionId":31401,"isGpuEnabled":true,"isInternetEnabled":true,"language":"python","sourceType":"notebook"},"papermill":{"default_parameters":{},"duration":15622.095748,"end_time":"2026-07-17T05:21:22.500680+00:00","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2026-07-17T01:01:00.404932+00:00","version":"2.7.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","collapsed":true,"execution":{"iopub.execute_input":"2026-07-17T01:01:02.876718Z","iopub.status.busy":"2026-07-17T01:01:02.875968Z","iopub.status.idle":"2026-07-17T01:01:15.953057Z","shell.execute_reply":"2026-07-17T01:01:15.952232Z"},"jupyter":{"outputs_hidden":true},"papermill":{"duration":13.134702,"end_time":"2026-07-17T01:01:16.001853+00:00","exception":false,"start_time":"2026-07-17T01:01:02.867151+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# IMPORTS\nos.environ[\"PYTORCH_ALLOC_CONF\"] = \"expandable_segments:True\"  # must be set before CUDA initialises\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nimport torchvision.transforms as transforms\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport logging\nimport random\nimport gc\nimport torch.optim as optim\n\nfrom pathlib import Path\nfrom skimage.color import rgb2lab\nfrom PIL import Image\nfrom tqdm import tqdm\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder, label_binarize\nfrom sklearn.metrics import (\n    confusion_matrix, accuracy_score, f1_score,\n    roc_auc_score, average_precision_score,\n    balanced_accuracy_score, recall_score\n)\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\n# ── Device setup ──────────────────────────────────────────────\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU:   {torch.cuda.get_device_name(0)}\")\n    print(f\"CUDA:  {torch.version.cuda}\")\n\n# ── Logging ───────────────────────────────────────────────────\nlogging.basicConfig(level=logging.INFO)\nlogger = logging.getLogger(__name__)\n\nsns.set_style('whitegrid')\nplt.rcParams['figure.figsize'] = (12, 6)\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:16.108003Z","iopub.status.busy":"2026-07-17T01:01:16.107176Z","iopub.status.idle":"2026-07-17T01:01:29.345911Z","shell.execute_reply":"2026-07-17T01:01:29.344980Z"},"papermill":{"duration":13.286549,"end_time":"2026-07-17T01:01:29.347539+00:00","exception":false,"start_time":"2026-07-17T01:01:16.060990+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\n# Forces CUDA to use deterministic algorithms — may slow training slightly\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:29.416374Z","iopub.status.busy":"2026-07-17T01:01:29.415836Z","iopub.status.idle":"2026-07-17T01:01:29.426659Z","shell.execute_reply":"2026-07-17T01:01:29.426117Z"},"papermill":{"duration":0.046091,"end_time":"2026-07-17T01:01:29.428107+00:00","exception":false,"start_time":"2026-07-17T01:01:29.382016+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# main dataset path\nham10000_path = Path('/kaggle/input/datasets/ctvmnn/ham10000')\n# Metadata CSV (contains labels, age, sex, etc.)\nmetadata_path = Path('/kaggle/input/datasets/ctvmnn/ham10000/HAM10000_metadata.csv')\n# segmentation masks dataset\nsegmentation_path = Path('/kaggle/input/datasets/tschandl/ham10000-lesion-segmentations')\n\n# Check if data exists and has actually loaded\nif ham10000_path.exists():\n    logger.info(f\"Found HAM10000 data: {ham10000_path}\")\n    image_files = list(ham10000_path.rglob('*.jpg'))\n    logger.info(f\"Total images: {len(image_files)}\")\nelse:\n    logger.warning(f\"HAM10000 data not found at {ham10000_path}\")\n    logger.info(\"Please ensure Kaggle dataset is added to notebook\")","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:29.496890Z","iopub.status.busy":"2026-07-17T01:01:29.496459Z","iopub.status.idle":"2026-07-17T01:01:32.350050Z","shell.execute_reply":"2026-07-17T01:01:32.349083Z"},"papermill":{"duration":2.889647,"end_time":"2026-07-17T01:01:32.351899+00:00","exception":false,"start_time":"2026-07-17T01:01:29.462252+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load metadata\nif metadata_path.exists():\n    df = pd.read_csv(metadata_path)# read csv into dataframe\n    logger.info(f\"Metadata shape: {df.shape}\")\n    logger.info(f\"Diagnosis distribution:\\n{df['dx'].value_counts()}\")\n    logger.info(f\"\\nAge distribution:\\n{df['age'].describe()}\")\n    display(df.head())","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:32.423309Z","iopub.status.busy":"2026-07-17T01:01:32.422530Z","iopub.status.idle":"2026-07-17T01:01:32.490080Z","shell.execute_reply":"2026-07-17T01:01:32.489158Z"},"papermill":{"duration":0.104383,"end_time":"2026-07-17T01:01:32.491560+00:00","exception":false,"start_time":"2026-07-17T01:01:32.387177+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Find image directory automatically\nimage_dir = ham10000_path / 'HAM10000_images' / 'HAM10000_images'\n\nprint(image_dir)\n# Create full image paths from image_id\ndf['image_path'] = df['image_id'].apply(lambda x: image_dir / f\"{x}.jpg\")\n\n# Check if images exist (sanity check)\nmissing_images = df['image_path'].apply(lambda x: not x.exists()).sum()\nprint(f\"Missing images: {missing_images}\")\n\n# Encode diagnosis labels (dx → numeric)\nle = LabelEncoder()\ndf['label'] = le.fit_transform(df['dx'])\n\n# Show label mapping\nlabel_mapping = dict(zip(le.classes_, le.transform(le.classes_)))\nprint(\"Label mapping:\")\nprint(label_mapping)\n\n# Preview dataframe\ndf[['image_id', 'dx', 'label', 'image_path']].head()","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:32.564323Z","iopub.status.busy":"2026-07-17T01:01:32.563564Z","iopub.status.idle":"2026-07-17T01:01:32.746936Z","shell.execute_reply":"2026-07-17T01:01:32.746114Z"},"papermill":{"duration":0.221512,"end_time":"2026-07-17T01:01:32.748419+00:00","exception":false,"start_time":"2026-07-17T01:01:32.526907+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Collect all mask files\nmask_files = list(segmentation_path.rglob('*.png'))\n\nprint(f\"Total masks found: {len(mask_files)}\")\n\n# Key: image_id → Value: mask path\nmask_dict = {}\n\nfor path in mask_files:\n    image_id = path.name.replace('_segmentation.png', '')\n    mask_dict[image_id] = path\n\n# Check a few entries\nsample_keys = list(mask_dict.keys())[:5]\nprint(\"\\nSample mask mappings:\")\nfor k in sample_keys:\n    print(k, \"->\", mask_dict[k])\n\n# keep masks aligned\ndf['has_mask'] = df['image_id'].apply(lambda x: x in mask_dict)\n\nprint(\"\\nMask availability:\")\nprint(df['has_mask'].value_counts())","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:32.817526Z","iopub.status.busy":"2026-07-17T01:01:32.817148Z","iopub.status.idle":"2026-07-17T01:01:35.479361Z","shell.execute_reply":"2026-07-17T01:01:35.478394Z"},"papermill":{"duration":2.697629,"end_time":"2026-07-17T01:01:35.480738+00:00","exception":false,"start_time":"2026-07-17T01:01:32.783109+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Skin Tone Estimation (ITA Scoring)\nWe estimate each patient's skin tone using the **Individual Typology Angle (ITA)** — a measure computed in LAB colour space from the *surrounding skin pixels only* (lesion is masked out). The ITA score is then mapped to one of five **Fitzpatrick groups**, which is used throughout training and evaluation as our fairness axis.","metadata":{"papermill":{"duration":0.034053,"end_time":"2026-07-17T01:01:35.549941+00:00","exception":false,"start_time":"2026-07-17T01:01:35.515888+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# SKIN TONE DISENTANGLEMENT: ITA Score Computation\n# ITA (Individual Typology Angle) estimates skin tone from the surrounding skin,\n# NOT the lesion itself. We flip the segmentation mask so the lesion is blocked\n# and only the healthy skin pixels contribute to the ITA calculation.\n# Formula: ITA = arctan((L - 50) / b) * 180 / π  (LAB color space)\n# Fitzpatrick mapping: >55 very light, 41-55 light, 28-41 medium, 10-28 dark, <10 very dark\n\nimport cv2\nfrom skimage.color import rgb2lab\n\ndef compute_ita_score(image_path, mask_path=None):\n    \"\"\"\n    Compute ITA score from skin pixels only (lesion masked out).\n    If no mask is available, falls back to whole-image ITA.\n    \"\"\"\n    try:\n        img = cv2.imread(str(image_path))\n        if img is None:\n            return None\n        img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        h, w = img_rgb.shape[:2]\n\n        if mask_path is not None:\n            # Load segmentation mask (white=lesion, black=skin)\n            mask = cv2.imread(str(mask_path), cv2.IMREAD_GRAYSCALE)\n            if mask is not None:\n                mask = cv2.resize(mask, (w, h))\n                # FLIP the mask: we want skin pixels (where mask=0), not lesion pixels\n                skin_mask = (mask < 128)  # True where skin is (lesion is excluded)\n                \n                # Need at least 100 skin pixels for a reliable ITA estimate\n                if skin_mask.sum() < 100:\n                    skin_mask = np.ones((h, w), dtype=bool)  # fallback: use whole image\n            else:\n                skin_mask = np.ones((h, w), dtype=bool)\n        else:\n            skin_mask = np.ones((h, w), dtype=bool)\n\n        # Convert to LAB\n        img_lab = rgb2lab(img_rgb / 255.0)\n        L = img_lab[:, :, 0]\n        b = img_lab[:, :, 2]\n\n        # Compute mean L and b over SKIN pixels only\n        L_mean = np.mean(L[skin_mask])\n        b_mean = np.mean(b[skin_mask])\n\n        if b_mean != 0:\n            ita = np.arctan((L_mean - 50) / b_mean) * 180 / np.pi\n        else:\n            ita = 0\n\n        return ita\n\n    except Exception as e:\n        logger.warning(f\"Error computing ITA for {image_path}: {e}\")\n        return None\n\n# Compute ITA scores — pass mask path for each image where available\nprint(\"Computing ITA scores from skin-only pixels (lesion masked out)...\")\n\ndef get_ita_with_mask(row):\n    mask_path = mask_dict.get(row['image_id'], None)\n    return compute_ita_score(row['image_path'], mask_path)\n\ndf['ita_score'] = df.apply(get_ita_with_mask, axis=1)\n\nfailed_ita = df['ita_score'].isna().sum()\nprint(f\"Failed ITA computations: {failed_ita}\")\n\nif failed_ita > 0:\n    median_ita = df['ita_score'].median()\n    df['ita_score'].fillna(median_ita, inplace=True)\n\nprint(f\"\\nITA Score Statistics:\")\nprint(f\"  Min: {df['ita_score'].min():.2f}\")\nprint(f\"  Max: {df['ita_score'].max():.2f}\")\nprint(f\"  Mean: {df['ita_score'].mean():.2f}\")\nprint(f\"  Median: {df['ita_score'].median():.2f}\")\n\ndef ita_to_fitzpatrick(ita_score):\n    if ita_score > 55:\n        return 0   # Very Light\n    elif ita_score > 41:\n        return 1   # Light\n    elif ita_score > 28:\n        return 2   # Medium\n    elif ita_score > 10:\n        return 3   # Dark\n    else:\n        return 4   # Very Dark\n\ndf['fitzpatrick_group'] = df['ita_score'].apply(ita_to_fitzpatrick)\nprint(f\"\\nFitzpatrick Group Distribution:\")\nprint(df['fitzpatrick_group'].value_counts().sort_index())\nprint(\"✅ ITA scores computed from skin pixels only — Fitzpatrick groups assigned\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:01:35.619327Z","iopub.status.busy":"2026-07-17T01:01:35.618556Z","iopub.status.idle":"2026-07-17T01:11:23.032920Z","shell.execute_reply":"2026-07-17T01:11:23.032213Z"},"papermill":{"duration":587.485543,"end_time":"2026-07-17T01:11:23.069270+00:00","exception":false,"start_time":"2026-07-17T01:01:35.583727+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FAIRNESS-AWARE TRAIN/VAL/TEST SPLIT\ndf['stratify_col'] = df['label'].astype(str) + '_' + df['fitzpatrick_group'].astype(str)\n\ndef safe_stratified_split(data, test_size, stratify_col, fallback_col, random_state=42):\n    \"\"\"\n    Safely split data with stratification, falling back when groups are too small.\n    \"\"\"\n    strat_counts = data[stratify_col].value_counts()\n    if strat_counts.min() >= 2:\n        return train_test_split(\n            data,\n            test_size=test_size,\n            stratify=data[stratify_col],\n            random_state=random_state\n        )\n    else:\n        print(f\"WARNING: Stratification by {stratify_col} is not possible because some groups have only one sample.\")\n        print(f\"Falling back to stratification by {fallback_col}.\")\n        return train_test_split(\n            data,\n            test_size=test_size,\n            stratify=data[fallback_col],\n            random_state=random_state\n        )\n\n# Split into train (70%) and temp (30%)\ntrain_df, temp_df = safe_stratified_split(\n    df,\n    test_size=0.3,\n    stratify_col='stratify_col',\n    fallback_col='label',\n    random_state=42\n)\n\n# Split temp into validation (15%) and test (15%)\nval_df, test_df = safe_stratified_split(\n    temp_df,\n    test_size=0.5,\n    stratify_col='stratify_col',\n    fallback_col='label',\n    random_state=42\n)\n\n# Print sizes\nprint(f\"Train size: {len(train_df)}\")\nprint(f\"Validation size: {len(val_df)}\")\nprint(f\"Test size: {len(test_df)}\")\n\n# Check class distribution\nprint(\"\\nTrain diagnosis distribution:\")\nprint(train_df['dx'].value_counts(normalize=True))\n\nprint(\"\\nValidation diagnosis distribution:\")\nprint(val_df['dx'].value_counts(normalize=True))\n\nprint(\"\\nTest diagnosis distribution:\")\nprint(test_df['dx'].value_counts(normalize=True))\n\n# FAIRNESS CHECK: Skin Tone Distribution Across Splits\nprint(\"\\n\" + \"=\"*60)\nprint(\"FAIRNESS CHECK: Skin Tone (ITA/Fitzpatrick) Distribution\")\nprint(\"=\"*60)\n\nfor split_name, split_df in [(\"Train\", train_df), (\"Val\", val_df), (\"Test\", test_df)]:\n    print(f\"\\n{split_name} Fitzpatrick Distribution:\")\n    fitz_counts = split_df['fitzpatrick_group'].value_counts(normalize=True).sort_index()\n    for fitz_group, pct in fitz_counts.items():\n        print(f\"  Group {fitz_group}: {pct*100:.1f}%\")\n    \n    # Also show ITA statistics\n    print(f\"  ITA Mean: {split_df['ita_score'].mean():.2f}\")\n    print(f\"  ITA Std: {split_df['ita_score'].std():.2f}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:23.137278Z","iopub.status.busy":"2026-07-17T01:11:23.136826Z","iopub.status.idle":"2026-07-17T01:11:23.173342Z","shell.execute_reply":"2026-07-17T01:11:23.172433Z"},"papermill":{"duration":0.072739,"end_time":"2026-07-17T01:11:23.175097+00:00","exception":false,"start_time":"2026-07-17T01:11:23.102358+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Fairness-Aware Data Split\nThe train/val/test split is **stratified jointly by lesion class and Fitzpatrick group**, ensuring every skin tone is represented in every split — not just in the overall dataset. This prevents evaluation leakage where a skin tone group only appears at test time.","metadata":{"papermill":{"duration":0.035568,"end_time":"2026-07-17T01:11:23.244980+00:00","exception":false,"start_time":"2026-07-17T01:11:23.209412+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# NOTE: Age and sex are deliberately excluded from this model.\n# Kept for backward compatibility only. NOT used as a model input.\ndef age_to_label(age):\n    if age < 30:\n        return 0\n    elif age < 60:\n        return 1\n    else:\n        return 2\n\ntrain_df['age_label'] = train_df['age'].apply(age_to_label)\nval_df['age_label']   = val_df['age'].apply(age_to_label)\ntest_df['age_label']  = test_df['age'].apply(age_to_label)\n\nprint(\"age_label added (not used as input — skin-tone fairness only) ✅\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:23.314870Z","iopub.status.busy":"2026-07-17T01:11:23.314140Z","iopub.status.idle":"2026-07-17T01:11:23.325197Z","shell.execute_reply":"2026-07-17T01:11:23.324436Z"},"papermill":{"duration":0.047449,"end_time":"2026-07-17T01:11:23.326565+00:00","exception":false,"start_time":"2026-07-17T01:11:23.279116+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Dataset & Preprocessing — Gentle Contrast Normalisation\nEach image is preprocessed with **gentle contrast normalisation**: we subtract the mean colour of the surrounding skin region (from the segmentation mask) from the entire image. This removes the *absolute skin tone baseline* while preserving all internal lesion texture and contrast.\n\nAugmentation during training includes heavy colour jitter, random grayscale, Gaussian blur, and random flips/rotations — all applied consistently to the lesion mask to keep them aligned.","metadata":{"papermill":{"duration":0.033328,"end_time":"2026-07-17T01:11:23.393160+00:00","exception":false,"start_time":"2026-07-17T01:11:23.359832+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# CUSTOM DATASET WITH SKIN-TONE-AWARE AUGMENTATION\n# Key insight: To disentangle lesion morphology from skin tone,\n# we need augmentation that forces the model to NOT rely on color/lighting.\n# Strategy: Heavy color jitter + hue shifts simulate the same lesion on different skin tones.\n\nclass SkinDataset(Dataset):\n    def __init__(self, df, mask_dict, train=True, input_size=224):\n        self.df = df.reset_index(drop=True)\n        self.mask_dict = mask_dict\n        self.train = train\n        self.input_size = input_size\n\n        self.resize    = transforms.Resize((input_size, input_size))\n        self.to_tensor = transforms.ToTensor()\n\n        # Heavy color jitter to break skin tone correlation\n        self.color_jitter = transforms.ColorJitter(\n            brightness=0.6,\n            contrast=0.6,\n            saturation=0.6,\n            hue=0.3\n        )\n        # RandomGrayscale properly assigned (was a floating expression before)\n        self.random_grayscale = transforms.RandomGrayscale(p=0.1)\n\n        self.gaussian_blur = transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 2.0))\n\n        self.normalize = transforms.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n        )\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_id         = row['image_id']\n        label            = row['label']\n        img_path         = row['image_path']\n        fitzpatrick_group = row['fitzpatrick_group']\n\n        image = Image.open(img_path).convert(\"RGB\")\n\n        mask_path = self.mask_dict.get(image_id)\n        if mask_path is not None:\n            mask = Image.open(mask_path).convert(\"L\")\n        else:\n            mask = Image.new(\"L\", image.size, color=255)  # default: full mask\n\n        image = self.resize(image)\n        mask  = self.resize(mask)\n\n        if self.train:\n            # Spatial transforms applied consistently to image & mask\n            if random.random() > 0.5:\n                image = transforms.functional.hflip(image)\n                mask  = transforms.functional.hflip(mask)\n\n            angle = random.uniform(-15, 15)\n            image = transforms.functional.rotate(image, angle)\n            mask  = transforms.functional.rotate(mask, angle)\n\n            # Color augmentation (image only)\n            if random.random() > 0.3:\n                image = self.color_jitter(image)\n\n            # RandomGrayscale now actually runs\n            image = self.random_grayscale(image)\n\n            if random.random() > 0.5:\n                image = transforms.functional.adjust_hue(image, random.uniform(-0.1, 0.1))\n\n            if random.random() > 0.7:\n                image = self.gaussian_blur(image)\n\n        image = self.to_tensor(image)\n        mask  = self.to_tensor(mask)\n        skin_region = (mask < 0.5).expand_as(image)\n        skin_pixels = image[skin_region]\n        if skin_pixels.numel() > 0:\n            skin_mean = skin_pixels.mean()\n            image = image - skin_mean  # shift only, no std division\n\n        image = self.normalize(image)\n\n        return image, mask, label, fitzpatrick_group\n\n# CREATE DATASETS — input_size=224 matches EfficientNet-B3 setup\ntrain_dataset = SkinDataset(train_df, mask_dict, train=True,  input_size=224)\nval_dataset   = SkinDataset(val_df,   mask_dict, train=False, input_size=224)\ntest_dataset  = SkinDataset(test_df,  mask_dict, train=False, input_size=224)\n\n# WEIGHTED SAMPLING — combines class imbalance AND Fitzpatrick group imbalance\n# This ensures dark-skinned samples (groups 3 and 4) are seen more often,\n# directly addressing the accuracy gap for those groups.\n\nclass_counts  = df['label'].value_counts().sort_index().values\nclass_weights = 1.0 / np.sqrt(class_counts)\nclass_weights = class_weights / class_weights.sum() * len(class_weights)\n\nfitz_counts   = df['fitzpatrick_group'].value_counts().sort_index().values\nfitz_weights  = 1.0 / np.sqrt(fitz_counts)\nfitz_weights  = fitz_weights / fitz_weights.sum() * len(fitz_weights)\n\nclass_weight_map = dict(enumerate(class_weights))\nfitz_weight_map  = dict(enumerate(fitz_weights))\n\n# Combined weight = class weight * fitzpatrick weight\nsample_weights = (\n    train_df['label'].map(class_weight_map) *\n    train_df['fitzpatrick_group'].map(fitz_weight_map)\n).values\n\n# Extra boost: DF and BKL samples within Group 3 or Group 4\nboost_mask_df  = (train_df['label'].values == 5) & (train_df['fitzpatrick_group'].values >= 3)\nboost_mask_bkl = (train_df['label'].values == 4) & (train_df['fitzpatrick_group'].values >= 3)\nsample_weights[boost_mask_df]  *= 1.5\nsample_weights[boost_mask_bkl] *= 1.3\n\n# Aggressive boost for ALL classes in Group 3 (Dark, Type IV)\n# Group 3 has only ~44 test samples — the model has barely seen this group.\n# 3x upsampling ensures it appears as often as larger groups during training.\nboost_mask_group3 = (train_df['fitzpatrick_group'].values == 3)\nsample_weights[boost_mask_group3] *= 3.0\n\nweighted_sampler = WeightedRandomSampler(\n    weights=sample_weights,\n    num_samples=len(train_df),\n    replacement=True\n)\n\ntrain_loader = DataLoader(train_dataset, batch_size=16, sampler=weighted_sampler, num_workers=0, generator=torch.Generator().manual_seed(SEED))\nval_loader   = DataLoader(val_dataset,   batch_size=16, shuffle=False, num_workers=0)\ntest_loader  = DataLoader(test_dataset,  batch_size=16, shuffle=False, num_workers=0)\ntest_loader  = DataLoader(test_dataset,  batch_size=16, shuffle=False, num_workers=0)\n\nprint(\"✅ DataLoaders ready with class + Fitzpatrick weighted sampling\")\nprint(f\"  Train batches : {len(train_loader)}\")\nprint(f\"  Val batches   : {len(val_loader)}\")\nprint(f\"  Test batches  : {len(test_loader)}\")","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:23.461401Z","iopub.status.busy":"2026-07-17T01:11:23.461080Z","iopub.status.idle":"2026-07-17T01:11:23.510062Z","shell.execute_reply":"2026-07-17T01:11:23.509324Z"},"papermill":{"duration":0.085272,"end_time":"2026-07-17T01:11:23.511431+00:00","exception":false,"start_time":"2026-07-17T01:11:23.426159+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get one batch\nimages, masks, labels, age_labels = next(iter(train_loader))\n\n# Pick first sample\nimage = images[0]\nmask = masks[0]\n\n# Apply mask (segmentation)\nsegmented = image * mask\n\n# Convert tensors to numpy for plotting\nimage_np = image.permute(1, 2, 0).numpy()\nmask_np = mask.squeeze().numpy()\nsegmented_np = segmented.permute(1, 2, 0).numpy()\n\n# Plot\nplt.figure(figsize=(15,5))\n\nplt.subplot(1,3,1)\nplt.title(\"Original Image\")\nplt.imshow(image_np)\nplt.axis('off')\n\nplt.subplot(1,3,2)\nplt.title(\"Mask\")\nplt.imshow(mask_np, cmap='gray')\nplt.axis('off')\n\nplt.subplot(1,3,3)\nplt.title(\"Segmented Image\")\nplt.imshow(segmented_np)\nplt.axis('off')\n\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:23.583017Z","iopub.status.busy":"2026-07-17T01:11:23.582626Z","iopub.status.idle":"2026-07-17T01:11:24.388592Z","shell.execute_reply":"2026-07-17T01:11:24.387614Z"},"papermill":{"duration":0.843633,"end_time":"2026-07-17T01:11:24.390324+00:00","exception":false,"start_time":"2026-07-17T01:11:23.546691+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Select only useful columns for inspection\nview_df = df[['image_id', 'dx', 'label', 'age', 'sex', 'localization', 'has_mask']]\n\ndisplay(view_df.head())\n\nprint(\"\\nDataset shape:\", view_df.shape)","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:24.463418Z","iopub.status.busy":"2026-07-17T01:11:24.463121Z","iopub.status.idle":"2026-07-17T01:11:24.475073Z","shell.execute_reply":"2026-07-17T01:11:24.474196Z"},"papermill":{"duration":0.049138,"end_time":"2026-07-17T01:11:24.476636+00:00","exception":false,"start_time":"2026-07-17T01:11:24.427498+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Phase 1 — SimCLR Contrastive Pretraining\nBefore supervised training, we pretrain the EfficientNet-B3 backbone using **SimCLR** (Simple Contrastive Learning of Representations). The model learns to produce similar embeddings for two differently-augmented views of the same image — with no labels. This gives the backbone a strong, generalised visual foundation before it sees any diagnosis labels.","metadata":{"papermill":{"duration":0.034554,"end_time":"2026-07-17T01:11:24.551109+00:00","exception":false,"start_time":"2026-07-17T01:11:24.516555+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# GRADIENT REVERSAL LAYER\n#   Forward pass : passes input through unchanged\n#   Backward pass: NEGATES the gradient — forces encoder away from skin-tone-predictive features\n\nclass GradReverse(torch.autograd.Function):\n    @staticmethod\n    def forward(ctx, x):\n        return x.clone()\n\n    @staticmethod\n    def backward(ctx, grad_output):\n        return -grad_output\n\n\n# DERM-AID MODEL: EfficientNet-B3 + Conv2D heads + Adversarial Skin Tone Head\n#\n# Pipeline:\n#   Image (masked) → EfficientNet-B3 backbone (spatial feature map [B, 1536, 7, 7])\n#                  → Conv2D projection head    (preserves spatial structure → [B, 256, 1, 1])\n#                  → Flatten                   ([B, 256])\n#                      ├─→ Lesion classifier   (Conv2D → flatten → FC → 7 classes)\n#                      └─→ Adversary (via GRL) (FC → 5 Fitzpatrick groups)\n#\n# Why Conv2D instead of Linear for the projection/classifier heads?\n#   Linear layers operate on a globally-pooled vector — spatial information about\n#   WHERE in the image a feature appears is lost before the classifier sees it.\n#   Conv2D layers operate on the spatial feature map directly, letting the model\n#   attend to specific regions (e.g. lesion border vs centre) before pooling.\n#   This is especially useful here because the lesion segmentation mask already\n#   focuses the image — Conv2D can exploit that spatial structure.\n\nclass DermAidModel(nn.Module):\n\n    def __init__(self, num_classes=7, num_fitzpatrick=5):\n        super().__init__()\n\n        # ── 1. BACKBONE: EfficientNet-B3 (features only, no pooling) ──\n        base = models.efficientnet_b3(weights=\"DEFAULT\")\n        # Use only the feature extractor — stops before global avg pool\n        # Output: [B, 1536, 7, 7] for 224px input\n        self.backbone = base.features\n        backbone_channels = 1536\n\n        # ── 2. CONV2D PROJECTION HEAD ─────────────────────────────────\n        self.projection = nn.Sequential(\n            nn.Conv2d(backbone_channels, 512, kernel_size=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Dropout2d(0.2),\n            nn.Conv2d(512, 256, kernel_size=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.AdaptiveAvgPool2d(1),   # [B, 256, 7, 7] → [B, 256, 1, 1]\n        )\n\n        # ── 3. LESION CLASSIFIER ──────────────────────────────────────\n        self.lesion_classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(256, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(128, 64),\n            nn.ReLU(),\n            nn.Linear(64, num_classes),\n        )\n\n        # ── 4. SKIN-TONE ADVERSARY ────────────────────────────────────\n        self.skin_tone_adversary = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(256, 128),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(128, num_fitzpatrick),\n        )\n    def forward(self, x):\n        # Step 1: spatial feature map from EfficientNet-B3\n        features = self.backbone(x)          # [B, 1536, 7, 7]\n\n        # Step 2: Conv2D projection → compact spatial embedding\n        embedding_spatial = self.projection(features)  # [B, 256, 1, 1]\n\n        # Step 3: lesion prediction\n        lesion_logits = self.lesion_classifier(embedding_spatial)   # [B, num_classes]\n\n        # Step 4: adversarial skin-tone prediction via gradient reversal\n        reversed_embedding = GradReverse.apply(embedding_spatial)\n        skin_tone_logits   = self.skin_tone_adversary(reversed_embedding)  # [B, num_fitzpatrick]\n\n        # Return flattened embedding for SimCLR contrastive loss\n        embedding_flat = embedding_spatial.view(embedding_spatial.size(0), -1)  # [B, 256]\n\n        return lesion_logits, skin_tone_logits, embedding_flat\n\n\n# ── Instantiate ───────────────────────────────────────────────\nmodel = DermAidModel(\n    num_classes=len(df['label'].unique()),\n    num_fitzpatrick=5\n).to(device)\n\ntotal_params     = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(\"✅ DermAidModel initialised — EfficientNet-B3 + Conv2D projection heads\")\nprint(f\"   Total parameters     : {total_params:,}\")\nprint(f\"   Trainable parameters : {trainable_params:,}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:24.625605Z","iopub.status.busy":"2026-07-17T01:11:24.625180Z","iopub.status.idle":"2026-07-17T01:11:25.540907Z","shell.execute_reply":"2026-07-17T01:11:25.540004Z"},"papermill":{"duration":0.953479,"end_time":"2026-07-17T01:11:25.542572+00:00","exception":false,"start_time":"2026-07-17T01:11:24.589093+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# SIMCLR AUGMENTATION\n# SimCLR works by generating two DIFFERENT augmented views of the same image and then training the encoder to produce similar embeddings for both views (and dissimilar for different images).\n# The augmentations need to be strong enough that the two views look visually different — but weak enough that the underlying lesion content is preserved.\n#   - Strong ColorJitter: teaches the encoder to be colour-agnostic\n#   - RandomGrayscale:    forces representation in terms of texture/shape\n#   - GaussianBlur:       forces robustness to fine detail (reduces skin pore dependence)\n\nclass SimCLRAugment:\n    def __init__(self):\n        self.transform = transforms.Compose([\n            transforms.RandomResizedCrop(224),          # random crop + resize\n            transforms.RandomHorizontalFlip(),           # random flip\n            transforms.ColorJitter(                      # heavy colour variation\n                brightness=0.4,\n                contrast=0.4,\n                saturation=0.4,\n                hue=0.1\n            ),\n            transforms.RandomGrayscale(p=0.2),           # drop colour 20% of the time\n            transforms.GaussianBlur(kernel_size=3),      # slight blur\n            transforms.ToTensor(),                        # PIL → [0,1] tensor\n        ])\n\n    def __call__(self, x):\n        # Apply the same pipeline twice independently → two different views\n        return self.transform(x), self.transform(x)\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:25.613909Z","iopub.status.busy":"2026-07-17T01:11:25.613228Z","iopub.status.idle":"2026-07-17T01:11:25.618520Z","shell.execute_reply":"2026-07-17T01:11:25.617965Z"},"papermill":{"duration":0.041621,"end_time":"2026-07-17T01:11:25.619959+00:00","exception":false,"start_time":"2026-07-17T01:11:25.578338+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SimCLRDataset(torch.utils.data.Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.augment = SimCLRAugment()\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_path = self.df.iloc[idx]['image_path']\n        image = Image.open(img_path).convert(\"RGB\")\n\n        x1, x2 = self.augment(image)\n        return x1, x2","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:25.689611Z","iopub.status.busy":"2026-07-17T01:11:25.689252Z","iopub.status.idle":"2026-07-17T01:11:25.694174Z","shell.execute_reply":"2026-07-17T01:11:25.693412Z"},"papermill":{"duration":0.041059,"end_time":"2026-07-17T01:11:25.695512+00:00","exception":false,"start_time":"2026-07-17T01:11:25.654453+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"simclr_dataset = SimCLRDataset(train_df)\nsimclr_loader = DataLoader(simclr_dataset, batch_size=8, shuffle=True, generator=torch.Generator().manual_seed(SEED))\nprint(\"SimCLR loader ready:\", len(simclr_loader))","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:25.763528Z","iopub.status.busy":"2026-07-17T01:11:25.763161Z","iopub.status.idle":"2026-07-17T01:11:25.768024Z","shell.execute_reply":"2026-07-17T01:11:25.767231Z"},"papermill":{"duration":0.04085,"end_time":"2026-07-17T01:11:25.769576+00:00","exception":false,"start_time":"2026-07-17T01:11:25.728726+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# NT-XENT CONTRASTIVE LOSS (SimCLR)\n# This loss trains the encoder to:\n#   - Pull embeddings from the same image (different augmentations) together\n#   - Push embeddings from different images apart\n# Temperature τ controls the sharpness of the similarity distribution:\n#   Lower τ → harder negatives enforced more strictly\n#   Higher τ → softer contrast (more stable but slower convergence)\n\ndef contrastive_loss(z1, z2, temperature=0.5):\n    # L2-normalise so cosine similarity = dot product\n    z1 = F.normalize(z1, dim=1)  # [B, D]\n    z2 = F.normalize(z2, dim=1)  # [B, D]\n\n    batch_size = z1.shape[0]\n\n    # Stack both views: [2B, D]\n    representations = torch.cat([z1, z2], dim=0)\n\n    # Full similarity matrix: [2B, 2B]\n    # Entry (i, j) = cosine similarity between sample i and sample j\n    similarity_matrix = torch.matmul(representations, representations.T)  # [2B, 2B]\n\n    # Labels: for each view i, its positive pair is i + batch_size (or i - batch_size)\n    labels = torch.arange(batch_size, device=device)\n    labels = torch.cat([labels + batch_size, labels], dim=0)  # [2B]\n    # labels[i] points to the index of sample i's positive pair\n\n    # Scale by temperature and compute cross-entropy\n    loss = F.cross_entropy(similarity_matrix / temperature, labels)\n\n    return loss\n\n# SIMCLR PRETRAINING LOOP\n# We run a short pretraining phase so the backbone learns general skin-lesion representations before supervised training begins.\n# Important:\n#   - We use model(x) in pretraining too, but ONLY care about the embedding (third return value).\n#   - Lesion and skin-tone logits are ignored at this stage.\n#   - After pretraining, we save ONLY the backbone weights —the projection head is task-specific and will be re-trained.\n\nSIMCLR_EPOCHS = 10\noptimizer_pretrain = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-6)\n# Cosine annealing for smooth LR decay during pretraining\nscheduler_pretrain = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer_pretrain, T_max=SIMCLR_EPOCHS, eta_min=1e-6\n)\n\nprint(\"Starting SimCLR pretraining...\")\n\nfor epoch in range(SIMCLR_EPOCHS):\n    model.train()\n    total_loss = 0.0\n\n    loop = tqdm(simclr_loader, desc=f\"SimCLR Epoch {epoch+1}/{SIMCLR_EPOCHS}\")\n\n    for x1, x2 in loop:\n        x1 = x1.to(device)\n        x2 = x2.to(device)\n\n        # Concatenate both views → one forward pass instead of two (halves peak memory)\n        x_both = torch.cat([x1, x2], dim=0)          # [2B, C, H, W]\n        _, _, z_both = model(x_both)                  # [2B, 256]\n        z1, z2 = z_both.chunk(2, dim=0)               # each [B, 256]\n        loss = contrastive_loss(z1, z2)\n\n        optimizer_pretrain.zero_grad()\n        loss.backward()\n        optimizer_pretrain.step()\n\n        total_loss += loss.item()\n        loop.set_postfix(loss=f\"{loss.item():.4f}\")\n\n    avg_loss = total_loss / len(simclr_loader)\n    scheduler_pretrain.step()\n    print(f\"  SimCLR Epoch {epoch+1}/{SIMCLR_EPOCHS} — Avg Loss: {avg_loss:.4f}\")\n\n    os.environ[\"PYTORCH_ALLOC_CONF\"] = \"expandable_segments:True\"\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(\"\\n✅ SimCLR pretraining complete\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T01:11:25.840394Z","iopub.status.busy":"2026-07-17T01:11:25.840152Z","iopub.status.idle":"2026-07-17T02:06:52.374667Z","shell.execute_reply":"2026-07-17T02:06:52.373687Z"},"papermill":{"duration":3326.572834,"end_time":"2026-07-17T02:06:52.376341+00:00","exception":false,"start_time":"2026-07-17T01:11:25.803507+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# SAVE PRETRAINED BACKBONE\nprint(\"Saving pretrained backbone weights...\")\ntorch.save(model.backbone.state_dict(), \"simclr_backbone.pth\")\nprint(\"  Saved → simclr_backbone.pth\")\n\n# Free GPU memory before starting supervised training\ntorch.cuda.empty_cache()\ngc.collect()\nprint(\"  GPU memory cleared ✅\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T02:06:53.797612Z","iopub.status.busy":"2026-07-17T02:06:53.797193Z","iopub.status.idle":"2026-07-17T02:06:54.082927Z","shell.execute_reply":"2026-07-17T02:06:54.081910Z"},"papermill":{"duration":1.037865,"end_time":"2026-07-17T02:06:54.084406+00:00","exception":false,"start_time":"2026-07-17T02:06:53.046541+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FAIR TRAINING FUNCTION (one epoch)\n# Loss formula:\n#   L_total = L_lesion + λ * L_skin_tone\n# MixUp removed — it was destabilising learning on this heavily imbalanced dataset.\n\ndef train_one_epoch(model, loader, optimizer, scheduler,\n                    criterion_lesion, criterion_skin_tone,\n                    lambda_adv=0.2):\n    model.train()\n\n    total_loss_sum     = 0.0\n    lesion_loss_sum    = 0.0\n    skin_tone_loss_sum = 0.0\n\n    loop = tqdm(loader, desc=\"  Training\")\n\n    for images, masks, labels, fitzpatrick_groups in loop:\n\n        images             = images.to(device)\n        masks              = masks.to(device)\n        labels             = labels.to(device)\n        fitzpatrick_groups = fitzpatrick_groups.to(device)\n\n        # Apply lesion segmentation mask\n        masked_images = images * masks\n\n        optimizer.zero_grad()\n\n        # Forward pass\n        lesion_logits, skin_tone_logits, embedding = model(masked_images)\n\n        # Loss 1: Lesion classification with per-group weighting\n        # Pass fitzpatrick_groups so FairFocalLoss can upweight Groups 2 & 3\n        loss_lesion = criterion_lesion(lesion_logits, labels, fitzpatrick_groups)\n\n        # Loss 2: Adversarial skin-tone debiasing\n        # GradReverse already flipped the sign in backward — we ADD here\n        loss_skin_tone = criterion_skin_tone(skin_tone_logits, fitzpatrick_groups)\n\n        loss_total = loss_lesion + lambda_adv * loss_skin_tone\n\n        loss_total.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.5)\n        optimizer.step()\n        scheduler.step()  # OneCycleLR steps per batch\n\n        total_loss_sum     += loss_total.item()\n        lesion_loss_sum    += loss_lesion.item()\n        skin_tone_loss_sum += loss_skin_tone.item()\n\n        loop.set_postfix({\n            'L_total'  : f\"{loss_total.item():.3f}\",\n            'L_lesion' : f\"{loss_lesion.item():.3f}\",\n            'L_tone'   : f\"{loss_skin_tone.item():.3f}\",\n        })\n\n    n = len(loader)\n    return {\n        'total'     : total_loss_sum     / n,\n        'lesion'    : lesion_loss_sum    / n,\n        'skin_tone' : skin_tone_loss_sum / n,\n    }\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T02:06:55.554517Z","iopub.status.busy":"2026-07-17T02:06:55.553711Z","iopub.status.idle":"2026-07-17T02:06:55.561362Z","shell.execute_reply":"2026-07-17T02:06:55.560579Z"},"papermill":{"duration":0.796527,"end_time":"2026-07-17T02:06:55.562732+00:00","exception":false,"start_time":"2026-07-17T02:06:54.766205+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FOCAL LOSS FOR CLASS IMBALANCE\n# Focal Loss addresses dataset imbalance by DOWN-WEIGHTING easy examples:\n#   FL(p_t) = -α_t * (1 - p_t)^γ * log(p_t)\n#   where:\n#     p_t  = model confidence on the correct class\n#     α_t  = class weight  (higher for rare classes)\n#     γ    = focusing parameter (2.0 is standard; higher → more focus on hard examples)\n#     (1 - p_t)^γ = if model is confident (p_t → 1) this term → 0, so loss ≈ 0\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n    def forward(self, predictions, targets):\n        # label_smoothing=0.1 prevents overconfidence and improves calibration\n        ce_loss = F.cross_entropy(\n            predictions, targets,\n            reduction='none',\n            weight=self.alpha,\n            label_smoothing=0.1\n        )  # [B]\n        p_t = torch.exp(-ce_loss)  # [B]\n        focal_weight = (1.0 - p_t) ** self.gamma  # [B]\n        loss = (focal_weight * ce_loss).mean()\n        return loss\n\n# ── Compute per-class weights (sqrt inverse frequency) ────────\n# Softer than hard inverse frequency — reduces minority class boost\n# without suppressing majority classes so hard they stop being predicted\nclass_counts  = df['label'].value_counts().sort_index().values\nclass_weights = 1.0 / np.sqrt(class_counts)\nclass_weights = class_weights / class_weights.sum() * len(class_weights)\nclass_weights = torch.FloatTensor(class_weights).to(device)\n\nprint(\"Class weights (α for Focal Loss):\")\nfor cls_id, weight in enumerate(class_weights):\n    cls_name = le.classes_[cls_id]\n    print(f\"  [{cls_id}] {cls_name:<6} : {weight:.4f}\")\n\n# ── Per-Fitzpatrick-group loss weights ───────────────────────\n# Groups 2 and 3 are severely underrepresented (~63 and ~44 samples).\n# We tell the loss function that misclassifying these groups is more\n# costly — the model is penalised more heavily for getting them wrong.\n# Group 0 (Very Light) and 1 (Light) are well-represented → weight 1.0\n# Group 2 (Medium) → 2.5×  Group 3 (Dark) → 3.0×  Group 4 (Very Dark) → 1.2×\nGROUP_LOSS_WEIGHTS = {0: 1.0, 1: 1.0, 2: 2.5, 3: 3.0, 4: 1.2}\nprint(\"Per-Fitzpatrick-group loss weights:\")\nfor g, w in GROUP_LOSS_WEIGHTS.items():\n    print(f\"  Group {g}: {w}×\")\n\nclass FairFocalLoss(nn.Module):\n    \"\"\"\n    Focal Loss with two levels of weighting:\n      1. Per-class weights (α)  — handles class imbalance (e.g. DF >> MEL)\n      2. Per-group weights      — handles Fitzpatrick group imbalance\n                                  (Groups 2 and 3 are underrepresented)\n\n    For each sample in a batch, its loss contribution is scaled by:\n        α_class × group_weight × focal_weight\n    This means a misclassified Group 3 (Dark) sample contributes 3× more\n    to the gradient than the same mistake on a Group 0 (Very Light) sample.\n    \"\"\"\n    def __init__(self, alpha=None, gamma=2.0, group_weights=None):\n        super().__init__()\n        self.alpha         = alpha         # per-class weights tensor [num_classes]\n        self.gamma         = gamma         # focusing parameter\n        self.group_weights = group_weights # dict: fitzpatrick_group → float\n\n    def forward(self, predictions, targets, fitzpatrick_groups=None):\n        # Per-sample cross-entropy with class weights and label smoothing\n        ce_loss = F.cross_entropy(\n            predictions, targets,\n            reduction='none',        # keep per-sample so we can apply group weights\n            weight=self.alpha,\n            label_smoothing=0.1\n        )  # shape: [B]\n\n        # Focal weight: down-weight easy examples (high confidence = low loss)\n        p_t           = torch.exp(-ce_loss)\n        focal_weight  = (1.0 - p_t) ** self.gamma\n        loss          = focal_weight * ce_loss   # [B]\n\n        # Apply per-group weights if provided\n        if self.group_weights is not None and fitzpatrick_groups is not None:\n            # Build a weight tensor matching the batch\n            gw = torch.tensor(\n                [self.group_weights.get(int(g), 1.0) for g in fitzpatrick_groups],\n                dtype=torch.float32, device=predictions.device\n            )  # [B]\n            loss = loss * gw   # upweight underrepresented groups\n\n        return loss.mean()\n\n# ── Loss functions ────────────────────────────────────────────\n# FairFocalLoss replaces FocalLoss — same behaviour when no group weights\n# are passed, so the supervised training loop still works unchanged.\ncriterion_lesion     = FairFocalLoss(alpha=class_weights, gamma=2.0,\n                                     group_weights=GROUP_LOSS_WEIGHTS)\ncriterion_skin_tone  = nn.CrossEntropyLoss()\n\n# ── Optimiser & scheduler ─────────────────────────────────────\noptimizer = torch.optim.Adam(model.parameters(), lr=3e-4, weight_decay=1e-5)\nNUM_EPOCHS_SCHED = 25\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=3e-4,\n    steps_per_epoch=len(train_loader),\n    epochs=NUM_EPOCHS_SCHED,\n    pct_start=0.2,\n    div_factor=10,\n    final_div_factor=100\n)\n\nprint(\"\\n✅ Training setup ready\")\nprint(f\"   Optimiser  : Adam (lr=3e-4, wd=1e-5)\")\nprint(f\"   Scheduler  : OneCycleLR (max_lr=3e-4, epochs={NUM_EPOCHS_SCHED})\")\nprint(f\"   Lesion loss: FocalLoss (γ=2.0, α=sqrt inverse frequency)\")\nprint(f\"   Adv. loss  : CrossEntropyLoss\")","metadata":{"execution":{"iopub.execute_input":"2026-07-17T02:06:56.886636Z","iopub.status.busy":"2026-07-17T02:06:56.886194Z","iopub.status.idle":"2026-07-17T02:06:56.904743Z","shell.execute_reply":"2026-07-17T02:06:56.903748Z"},"papermill":{"duration":0.682318,"end_time":"2026-07-17T02:06:56.906286+00:00","exception":false,"start_time":"2026-07-17T02:06:56.223968+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#   EVALUATION FUNCTION\n#   - model.eval() disables Dropout and BatchNorm running-stat updates\n#   - torch.no_grad() skips gradient computation → faster + less memory\n#   - GradReverse still exists but has no effect (no backward pass)\n#   - We collect predictions, labels, Fitzpatrick groups, and probabilities so we can compute both overall and per-group fairness metrics.\n\ndef evaluate(model, loader, criterion_lesion):\n    \"\"\"\n    Evaluate model on a validation or test DataLoader.\n\n    Args:\n        model            : DermAidModel\n        loader           : DataLoader (val or test)\n        criterion_lesion : FocalLoss — used to track validation loss\n\n    Returns:\n        dict with keys:\n          'loss'         : average lesion loss\n          'accuracy'     : overall top-1 accuracy\n          'preds'        : np.array of predicted class indices\n          'labels'       : np.array of true class indices\n          'fitzpatrick'  : np.array of Fitzpatrick group IDs\n          'probs'        : np.array of softmax probabilities [N, num_classes]\n    \"\"\"\n    model.eval()\n\n    total_loss = 0.0\n    correct    = 0\n\n    all_preds       = []\n    all_labels      = []\n    all_fitzpatrick = []\n    all_probs       = []\n\n    with torch.no_grad():  # no gradient tracking needed for eval\n        for images, masks, labels, fitzpatrick_groups in loader:\n\n            images             = images.to(device)\n            masks              = masks.to(device)\n            labels             = labels.to(device)\n\n            # Apply segmentation mask (same preprocessing as training)\n            masked_images = images * masks\n\n            # Forward pass — unpack all three return values.\n            # We only use lesion_logits here; embedding is discarded.\n            lesion_logits, _, _ = model(masked_images)\n\n            # Compute lesion loss for monitoring\n            loss = criterion_lesion(lesion_logits, labels)\n            total_loss += loss.item()\n\n            # Predicted class = argmax of logits\n            preds = lesion_logits.argmax(dim=1)\n            correct += (preds == labels).sum().item()\n\n            # Softmax probabilities — needed for AUC computation\n            probs = F.softmax(lesion_logits, dim=1)\n\n            # Accumulate results (move to CPU for numpy conversion)\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            all_fitzpatrick.extend(fitzpatrick_groups.cpu().numpy())  # already on CPU\n            all_probs.extend(probs.cpu().numpy())\n\n    accuracy = correct / len(loader.dataset)\n\n    return {\n        'loss'       : total_loss / len(loader),\n        'accuracy'   : accuracy,\n        'preds'      : np.array(all_preds),\n        'labels'     : np.array(all_labels),\n        'fitzpatrick': np.array(all_fitzpatrick),\n        'probs'      : np.array(all_probs),\n    }\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T02:06:58.317987Z","iopub.status.busy":"2026-07-17T02:06:58.317611Z","iopub.status.idle":"2026-07-17T02:06:58.325297Z","shell.execute_reply":"2026-07-17T02:06:58.324404Z"},"papermill":{"duration":0.673737,"end_time":"2026-07-17T02:06:58.327103+00:00","exception":false,"start_time":"2026-07-17T02:06:57.653366+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Phase 2 — Supervised Training with Adversarial Debiasing\nWe load the SimCLR-pretrained backbone and train the full model for 25 epochs.\n\n**Three-component loss:**\n- `L_lesion` — Focal Loss with label smoothing for disease classification\n- `L_skin_tone` — adversarial loss via Gradient Reversal Layer (penalises the embedding for encoding skin tone)\n- `λ_adv` ramps from 0.05 → 0.4 over the first 10 epochs to stabilise early training\n\nThe backbone is **frozen for the first 5 epochs** to let the new heads stabilise, then unfrozen for end-to-end fine-tuning.","metadata":{"papermill":{"duration":0.671383,"end_time":"2026-07-17T02:06:59.760430+00:00","exception":false,"start_time":"2026-07-17T02:06:59.089047+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# LOAD PRETRAINED BACKBONE INTO SUPERVISED MODEL\nmodel = DermAidModel(\n    num_classes=len(df['label'].unique()),\n    num_fitzpatrick=5\n).to(device)\n\nif not os.path.exists(\"simclr_backbone.pth\"):\n    raise FileNotFoundError(\"simclr_backbone.pth not found — re-run SimCLR pretraining first (Cell 18)\")\nmodel.backbone.load_state_dict(torch.load(\"simclr_backbone.pth\"))\nprint(\"✅ Pretrained backbone weights loaded from simclr_backbone.pth\")\n\n# Freeze backbone BEFORE creating the optimizer so frozen params\n# are never tracked — saves memory and avoids wasted compute\nfor param in model.backbone.parameters():\n    param.requires_grad = False\nprint(\"Backbone frozen for first 5 epochs.\")\n\n# Optimizer created AFTER freezing — only trains projection + classifier heads\noptimizer = torch.optim.Adam(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=3e-4, weight_decay=1e-5\n)\n\nNUM_EPOCHS_SCHED = 25\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=3e-4,\n    steps_per_epoch=len(train_loader),\n    epochs=NUM_EPOCHS_SCHED,\n    pct_start=0.2,\n    div_factor=10,\n    final_div_factor=100\n)\nprint(\"✅ Optimizer and scheduler ready (tracking unfrozen params only)\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T02:07:01.178801Z","iopub.status.busy":"2026-07-17T02:07:01.178060Z","iopub.status.idle":"2026-07-17T02:07:01.608790Z","shell.execute_reply":"2026-07-17T02:07:01.608074Z"},"papermill":{"duration":1.099087,"end_time":"2026-07-17T02:07:01.610366+00:00","exception":false,"start_time":"2026-07-17T02:07:00.511279+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MAIN SUPERVISED TRAINING LOOP\n\nNUM_EPOCHS  = 25\nLAMBDA_MAX  = 0.4\nbest_val_loss = float('inf')\n\ndef get_lambda(epoch):\n    if epoch < 10:\n        return 0.05 + (LAMBDA_MAX - 0.05) * (epoch / 10)\n    return LAMBDA_MAX\n\nprint(f\"Starting supervised training ({NUM_EPOCHS} epochs, λ_adv=0.05→{LAMBDA_MAX})\")\nprint(\"=\"*70)\nhistory = {\n    'train_loss': [], 'val_loss': [],\n    'train_lesion_loss': [], 'val_accuracy': []\n}\n\nfor epoch in range(NUM_EPOCHS):\n\n    # At epoch 5: unfreeze backbone and recreate optimizer to include all params\n    if epoch == 5:\n        for param in model.backbone.parameters():\n            param.requires_grad = True\n        # Recreate optimizer now that all params are trainable\n        optimizer = torch.optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-5)\n        # Recreate scheduler for the remaining 20 epochs\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer,\n            max_lr=1e-4,\n            steps_per_epoch=len(train_loader),\n            epochs=NUM_EPOCHS - 5,\n            pct_start=0.1,\n            div_factor=10,\n            final_div_factor=100\n        )\n        print(\"Backbone unfrozen — optimizer reset with lower lr=1e-4.\")\n\n    lambda_adv = get_lambda(epoch)\n    train_losses = train_one_epoch(\n        model, train_loader, optimizer, scheduler,\n        criterion_lesion, criterion_skin_tone,\n        lambda_adv=lambda_adv\n    )\n\n    val_metrics = evaluate(model, val_loader, criterion_lesion)\n    current_lr  = optimizer.param_groups[0]['lr']\n\n    is_best = val_metrics['loss'] < best_val_loss\n    if is_best:\n        best_val_loss = val_metrics['loss']\n        torch.save(model.state_dict(), \"best_model_derm_aid.pth\")\n        checkpoint_tag = \" ← best ✅\"\n    else:\n        checkpoint_tag = \"\"\n\n    print(\n        f\"Epoch {epoch+1:02d}/{NUM_EPOCHS} | \"\n        f\"Train L_total={train_losses['total']:.4f} \"\n        f\"(lesion={train_losses['lesion']:.4f}, \"\n        f\"skin_tone={train_losses['skin_tone']:.4f}) | \"\n        f\"Val Loss={val_metrics['loss']:.4f}  \"\n        f\"Val Acc={val_metrics['accuracy']:.4f}  \"\n        f\"LR={current_lr:.2e}\"\n        f\"{checkpoint_tag}\"\n    )\n\n    history['train_loss'].append(train_losses['total'])\n    history['train_lesion_loss'].append(train_losses['lesion'])\n    history['val_loss'].append(val_metrics['loss'])\n    history['val_accuracy'].append(val_metrics['accuracy'])\n\nprint(\"\\n\" + \"=\"*70)\nprint(f\"Training complete. Best val loss: {best_val_loss:.4f}\")\nprint(\"Best model saved to best_model_derm_aid.pth ✅\")\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))\nax1.plot(history['train_loss'],        label='Train total loss')\nax1.plot(history['train_lesion_loss'],  label='Train lesion loss', linestyle='--')\nax1.plot(history['val_loss'],           label='Val loss')\nax1.set_xlabel('Epoch'); ax1.set_ylabel('Loss')\nax1.set_title('Training & Validation Loss')\nax1.legend(); ax1.grid(True)\n\nax2.plot(history['val_accuracy'], color='green', label='Val accuracy')\nax2.set_xlabel('Epoch'); ax2.set_ylabel('Accuracy')\nax2.set_title('Validation Accuracy over Epochs')\nax2.legend(); ax2.grid(True)\n\nplt.tight_layout()\nplt.savefig('training_curves.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint(\"Training curves saved ✅\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T02:07:03.046880Z","iopub.status.busy":"2026-07-17T02:07:03.046142Z","iopub.status.idle":"2026-07-17T03:40:28.966983Z","shell.execute_reply":"2026-07-17T03:40:28.966177Z"},"papermill":{"duration":5607.745266,"end_time":"2026-07-17T03:40:30.037032+00:00","exception":false,"start_time":"2026-07-17T02:07:02.291766+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# TEST SET EVALUATION\nprint(\"Loading best model checkpoint...\")\nmodel.load_state_dict(torch.load(\"best_model_derm_aid.pth\"))\nprint(\"  Loaded best_model_derm_aid.pth ✅\")\n\ntest_metrics = evaluate(model, test_loader, criterion_lesion)\n\nprint(f\"\\n{'='*70}\")\nprint(f\"TEST SET PERFORMANCE\")\nprint(f\"{'='*70}\")\nprint(f\"  Loss     : {test_metrics['loss']:.4f}\")\nprint(f\"  Accuracy : {test_metrics['accuracy']:.4f}\")\nprint(f\"{'='*70}\")\n\n# ── Unpack results for downstream fairness analysis ───────────\nall_preds       = test_metrics['preds']\nall_labels      = test_metrics['labels']\nall_fitzpatrick = test_metrics['fitzpatrick']\nall_probs       = test_metrics['probs']\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:40:33.237814Z","iopub.status.busy":"2026-07-17T03:40:33.237378Z","iopub.status.idle":"2026-07-17T03:41:00.835469Z","shell.execute_reply":"2026-07-17T03:41:00.834496Z"},"papermill":{"duration":29.192745,"end_time":"2026-07-17T03:41:00.837267+00:00","exception":false,"start_time":"2026-07-17T03:40:31.644522+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# LOAD EXPLAINABILITY MODULE\n# derm_aid_explainability.py contains:\n#   - GradCAM                : visual heatmap explanations\n#   - compute_metrics        : all fairness + performance metrics in one call\n#   - generate_patient_report: per-patient PDF (prediction + fairness)\n#   - export_for_mobile      : TorchScript .ptl for offline mobile use\n\n!pip install reportlab -q\n\nimport sys\nsys.path.append('/kaggle/input/datasets/jemimambabacha/explainability')\n\nfrom derm_aid_explainability import (\n    GradCAM,\n    compute_metrics,\n    generate_patient_report,\n    export_for_mobile,\n    FITZPATRICK_LABELS,\n    HAM10000_CLASSES,\n)\n\nprint(\"✅ Explainability module loaded\")","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:41:04.024741Z","iopub.status.busy":"2026-07-17T03:41:04.024186Z","iopub.status.idle":"2026-07-17T03:41:09.968423Z","shell.execute_reply":"2026-07-17T03:41:09.967410Z"},"papermill":{"duration":7.534202,"end_time":"2026-07-17T03:41:09.970202+00:00","exception":false,"start_time":"2026-07-17T03:41:02.436000+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# COMPUTE ALL METRICS\nfairness_summary = compute_metrics(\n    all_preds        = all_preds,\n    all_labels       = all_labels,\n    all_probs        = all_probs,\n    all_fitzpatrick  = all_fitzpatrick,\n    label_encoder    = le,\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:41:13.181009Z","iopub.status.busy":"2026-07-17T03:41:13.180327Z","iopub.status.idle":"2026-07-17T03:41:13.265172Z","shell.execute_reply":"2026-07-17T03:41:13.263925Z"},"papermill":{"duration":1.685768,"end_time":"2026-07-17T03:41:13.267350+00:00","exception":false,"start_time":"2026-07-17T03:41:11.581582+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Phase 3 — Skin-Tone-Conditioned Contrastive Learning\n### The Problem We Discovered\nAfter the masked model was trained, the **linear probe showed 60% skin-tone predictability** from the disease embedding — well above the 20% chance level. This revealed that skin tone information lives *inside the lesion itself* (not just in surrounding skin), because two kinds of features were being encoded:\n\n- **Kind 1 — Biologically entangled features** *(keep)*: melanin concentration genuinely affects how a disease presents. A melanoma on dark skin has different absolute colour values than the same disease on light skin. This is real biology we cannot and should not remove.\n- **Kind 2 — Spurious skin-tone-correlated features** *(remove)*: the model learned that certain colour ranges predict certain diagnoses simply because they co-occur with certain skin tones in the training data — not because they are biologically meaningful. This is **shortcut learning**.\n\n### The New Strategy\nRather than making the model *blind* to skin tone (which destroys Kind 1), we teach it that **the same disease should look similar across skin tones in representation space**.\n\nWe do this by:\n1. Building **cross-skin-tone positive pairs**: same disease class, different Fitzpatrick group\n2. Generating **synthetic Type IV images** via ITA-based skin tone transfer (solves data scarcity)\n3. Training with a **three-component loss**: disease classification + cross-skin-tone contrastive + mild adversarial\n\n**New target metrics:** Type IV TPR ≥ 40% | TPR gap < 5% | Overall accuracy ≥ 65% | Probe accuracy 35–45%","metadata":{"papermill":{"duration":1.607238,"end_time":"2026-07-17T03:41:16.517226+00:00","exception":false,"start_time":"2026-07-17T03:41:14.909988+00:00","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ── SYNTHETIC SKIN TONE TRANSFER ─────────────────────────────────────────────\n#\n# WHY THIS IS NEEDED:\n#   Group 3 (Dark, Type IV) has only ~44 training images across 7 disease classes.\n#   Some classes have zero Type IV examples — making cross-skin-tone contrastive\n#   pairs impossible to find in the real data.\n#\n# WHAT WE DO:\n#   Take images from well-represented groups (Very Light, Light, Very Dark),\n#   and synthetically shift their skin tone toward the Type IV ITA range.\n#   This gives us hundreds of synthetic Type IV images to pair with real ones.\n#\n# HOW IT WORKS:\n#   Images are converted to LAB colour space (L=lightness, A=green-red, B=blue-yellow).\n#   We shift the L (lightness) channel of skin pixels to match the target ITA value.\n#   Lesion pixels are shifted too, but at 30% of the skin shift — preserving the\n#   relative lesion-to-skin contrast while changing the absolute skin tone context.\n#\n# WHAT IS ITA?\n#   ITA = arctan((L_mean - 50) / b_mean) * 180/π\n#   Higher ITA = lighter skin. Lower ITA = darker skin.\n#   Type IV (Dark) ITA ≈ 5.0\n\nfrom skimage.color import rgb2lab, lab2rgb\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndef synthetic_skin_tone_transfer(image_np, mask_np, source_ita, target_ita):\n    \"\"\"\n    Transfer an image from one skin tone to another using ITA-based lightness shifting.\n\n    Args:\n        image_np   : HxWx3 uint8 RGB image\n        mask_np    : HxW binary mask — 1 = lesion, 0 = surrounding skin\n        source_ita : ITA value of the original image\n        target_ita : ITA value we want to simulate (e.g. 5.0 for Type IV)\n\n    Returns:\n        Transferred image as HxWx3 uint8 RGB\n    \"\"\"\n    image_float = image_np.astype(np.float32) / 255.0\n\n    # Convert to LAB so we can manipulate lightness independently\n    image_lab = rgb2lab(image_float)  # L in [0,100], A and B in [-128,127]\n\n    # Boolean masks for skin and lesion regions\n    skin_mask   = (mask_np == 0)   # surrounding skin — where skin tone lives\n    lesion_mask = (mask_np == 1)   # lesion — shifted less aggressively\n\n    if skin_mask.sum() < 10:\n        # Not enough skin pixels to estimate a reliable shift — return unchanged\n        return image_np\n\n    # How much do we need to shift ITA?\n    # ITA is mainly driven by L (lightness): darker skin = lower L\n    # We approximate the required lightness shift as half the ITA delta\n    ita_delta       = target_ita - source_ita\n    lightness_shift = ita_delta * 0.5   # approximate L-channel shift\n\n    # Apply full shift to skin pixels\n    image_lab[:, :, 0][skin_mask] = np.clip(\n        image_lab[:, :, 0][skin_mask] + lightness_shift,\n        0, 100  # L must stay in [0, 100]\n    )\n\n    # Apply 30% of the shift to lesion pixels\n    # The lesion's appearance IS influenced by skin tone (Kind 1 biology),\n    # so we shift it partially — not fully, to preserve the relative contrast\n    image_lab[:, :, 0][lesion_mask] = np.clip(\n        image_lab[:, :, 0][lesion_mask] + lightness_shift * 0.3,\n        0, 100\n    )\n\n    # Convert back to RGB\n    result = lab2rgb(image_lab)   # returns float [0,1]\n    return (result * 255).astype(np.uint8)\n\n\ndef augment_type_iv_with_synthetic(dataset_df, mask_dict, target_multiplier=10):\n    \"\"\"\n    Generate synthetic Type IV (Dark skin) images from well-represented groups.\n\n    For each disease class, we take images from Very Light, Light, and Very Dark groups\n    and transfer them to a Type IV ITA range (~5.0). This:\n      1. Solves the data scarcity problem for Group 3\n      2. Creates cross-skin-tone pairs for contrastive learning\n      3. Forces the model to encounter the same disease across a full range of skin tones\n\n    Args:\n        dataset_df        : DataFrame with image metadata (must have ita_score column)\n        mask_dict         : dict mapping image_id → mask file path\n        target_multiplier : how many synthetic images to generate per disease class\n\n    Returns:\n        List of dicts, each with keys: image (np array), label, fitzpatrick_group, is_synthetic\n    \"\"\"\n    TYPE_IV_ITA_TARGET = 5.0    # centre of Type IV ITA range\n    SOURCE_GROUPS      = [0, 1, 4]  # Very Light, Light, Very Dark — well-represented\n\n    synthetic_records = []\n\n    print(\"Generating synthetic Type IV images...\")\n\n    for disease_class in dataset_df['label'].unique():\n        class_df     = dataset_df[dataset_df['label'] == disease_class]\n        source_df    = class_df[class_df['fitzpatrick_group'].isin(SOURCE_GROUPS)]\n\n        if len(source_df) == 0:\n            print(f\"  Class {disease_class}: no source images from light groups — skipping\")\n            continue\n\n        # Sample up to target_multiplier source images for this class\n        n_sample = min(target_multiplier, len(source_df))\n        sample   = source_df.sample(n=n_sample, random_state=SEED)\n\n        for _, row in sample.iterrows():\n            try:\n                # Load original image\n                img = cv2.imread(str(row['image_path']))\n                if img is None:\n                    continue\n                img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n                # Load mask (or create full-image default)\n                mask_path = mask_dict.get(row['image_id'])\n                if mask_path and Path(mask_path).exists():\n                    mask = cv2.imread(str(mask_path), cv2.IMREAD_GRAYSCALE)\n                    mask = cv2.resize(mask, (img_rgb.shape[1], img_rgb.shape[0]))\n                    mask_binary = (mask > 128).astype(np.uint8)  # 1=lesion, 0=skin\n                else:\n                    mask_binary = np.zeros(img_rgb.shape[:2], dtype=np.uint8)  # all skin\n\n                # Apply synthetic skin tone transfer\n                source_ita = row.get('ita_score', 40.0)  # fallback to light ITA\n                synthetic  = synthetic_skin_tone_transfer(\n                    img_rgb, mask_binary, source_ita, TYPE_IV_ITA_TARGET\n                )\n\n                synthetic_records.append({\n                    'image'           : synthetic,\n                    'label'           : disease_class,\n                    'fitzpatrick_group': 3,      # assigned to Group 3 (Dark)\n                    'is_synthetic'    : True,\n                    'source_image_id' : row['image_id']\n                })\n\n            except Exception as e:\n                print(f\"  Warning: could not process {row['image_id']}: {e}\")\n                continue\n\n    print(f\"Generated {len(synthetic_records)} synthetic Type IV images ✅\")\n    print(f\"Distribution by class:\")\n    from collections import Counter\n    class_counts = Counter(r['label'] for r in synthetic_records)\n    for cls, count in sorted(class_counts.items()):\n        print(f\"  Class {cls}: {count} synthetic images\")\n\n    return synthetic_records\n\n# Run synthetic generation\nsynthetic_type_iv = augment_type_iv_with_synthetic(\n    train_df, mask_dict, target_multiplier=10\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:41:19.763713Z","iopub.status.busy":"2026-07-17T03:41:19.763378Z","iopub.status.idle":"2026-07-17T03:41:22.717066Z","shell.execute_reply":"2026-07-17T03:41:22.716017Z"},"papermill":{"duration":4.575475,"end_time":"2026-07-17T03:41:22.718739+00:00","exception":false,"start_time":"2026-07-17T03:41:18.143264+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── CROSS-SKIN-TONE CONTRASTIVE PAIR BUILDER ──────────────────────────────────\n#\n# CORE IDEA (from mentor):\n#   Standard SimCLR says: \"same image augmented differently → same representation\"\n#   Our extension says:   \"same DISEASE on different skin tones → same representation\"\n#\n# This directly attacks Kind 2 (spurious) features without destroying Kind 1\n# (biological) features — because we never tell the model to ignore skin tone\n# entirely. We just tell it that skin tone should not change what a disease LOOKS\n# LIKE in representation space.\n#\n# HOW PAIRS ARE BUILT:\n#   For each disease class, we find image pairs where:\n#     - Same disease label\n#     - Different Fitzpatrick group (including synthetic Type IV images)\n#   These become positive pairs: the model is rewarded for producing similar\n#   embeddings for both, regardless of the skin tone difference.\n\ndef build_cross_skin_tone_pairs(dataset_df, synthetic_records, n_pairs_per_class=500):\n    \"\"\"\n    Build cross-skin-tone positive pairs for contrastive training.\n\n    A positive pair = (image_A, image_B) where:\n      - image_A and image_B show the SAME disease\n      - image_A and image_B come from DIFFERENT Fitzpatrick groups\n\n    The model learns: same disease = same representation, across all skin tones.\n\n    Args:\n        dataset_df        : real training DataFrame\n        synthetic_records : list of synthetic Type IV dicts from augment_type_iv_with_synthetic()\n        n_pairs_per_class : how many pairs to generate per disease class\n\n    Returns:\n        DataFrame of pairs with columns: image_a, image_b, label, skin_tone_a, skin_tone_b\n    \"\"\"\n    pairs = []\n\n    # Convert synthetic records into a mini-DataFrame for easy sampling\n    if synthetic_records:\n        synth_df = pd.DataFrame([{\n            'image_id'        : f\"synth_{i}\",\n            'image_path'      : None,          # no path — image array stored directly\n            'image_array'     : r['image'],\n            'label'           : r['label'],\n            'fitzpatrick_group': r['fitzpatrick_group'],\n            'is_synthetic'    : True,\n        } for i, r in enumerate(synthetic_records)])\n    else:\n        synth_df = pd.DataFrame()\n\n    # Combine real and synthetic\n    real_df = dataset_df.copy()\n    real_df['is_synthetic'] = False\n    real_df['image_array']  = None\n\n    if not synth_df.empty:\n        combined_df = pd.concat([real_df, synth_df], ignore_index=True)\n    else:\n        combined_df = real_df\n\n    for disease_class in combined_df['label'].unique():\n        class_df = combined_df[combined_df['label'] == disease_class]\n\n        # Which Fitzpatrick groups are available for this class?\n        available_groups = class_df['fitzpatrick_group'].unique()\n\n        if len(available_groups) < 2:\n            # Can't build cross-skin-tone pairs with only one group\n            print(f\"  Class {disease_class}: only 1 skin tone group — skipping cross-tone pairs\")\n            continue\n\n        n_built = 0\n        max_attempts = n_pairs_per_class * 3  # allow retries to hit n_pairs_per_class\n\n        for _ in range(max_attempts):\n            if n_built >= n_pairs_per_class:\n                break\n\n            # Pick two DIFFERENT Fitzpatrick groups\n            group_a, group_b = np.random.choice(available_groups, size=2, replace=False)\n\n            group_a_df = class_df[class_df['fitzpatrick_group'] == group_a]\n            group_b_df = class_df[class_df['fitzpatrick_group'] == group_b]\n\n            if len(group_a_df) == 0 or len(group_b_df) == 0:\n                continue\n\n            img_a = group_a_df.sample(1).iloc[0]\n            img_b = group_b_df.sample(1).iloc[0]\n\n            pairs.append({\n                'image_id_a'   : img_a['image_id'],\n                'image_id_b'   : img_b['image_id'],\n                'image_path_a' : img_a['image_path'],\n                'image_path_b' : img_b['image_path'],\n                'image_array_a': img_a['image_array'],   # None if real, np array if synthetic\n                'image_array_b': img_b['image_array'],\n                'label'        : disease_class,\n                'skin_tone_a'  : group_a,\n                'skin_tone_b'  : group_b,\n                'synth_a'      : img_a['is_synthetic'],\n                'synth_b'      : img_b['is_synthetic'],\n            })\n            n_built += 1\n\n        print(f\"  Class {disease_class}: {n_built} cross-skin-tone pairs built \"\n              f\"(groups: {sorted(available_groups)})\")\n\n    pairs_df = pd.DataFrame(pairs)\n    print(f\"\\nTotal pairs: {len(pairs_df)} across {pairs_df['label'].nunique()} disease classes ✅\")\n    return pairs_df\n\n\n# Build the pairs — includes real + synthetic Type IV\ncross_tone_pairs = build_cross_skin_tone_pairs(\n    train_df,\n    synthetic_type_iv,\n    n_pairs_per_class=500\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:41:25.985464Z","iopub.status.busy":"2026-07-17T03:41:25.985193Z","iopub.status.idle":"2026-07-17T03:41:31.583442Z","shell.execute_reply":"2026-07-17T03:41:31.582420Z"},"papermill":{"duration":7.255597,"end_time":"2026-07-17T03:41:31.585181+00:00","exception":false,"start_time":"2026-07-17T03:41:24.329584+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── CROSS-SKIN-TONE CONTRASTIVE DATASET ───────────────────────────────────────\n#\n# This dataset serves each pair as (image_A, image_B, disease_label).\n# During training, the model learns:\n#   - image_A and image_B → embeddings that are CLOSE (same disease, different skin tone)\n#   - images from different disease classes → embeddings that are FAR APART\n#\n# The SkinDataset augmentations (colour jitter, contrast norm etc) are applied\n# to both images in the pair independently for extra variation.\n\nclass CrossSkinToneDataset(torch.utils.data.Dataset):\n    def __init__(self, pairs_df, mask_dict, input_size=224):\n        self.pairs    = pairs_df.reset_index(drop=True)\n        self.mask_dict = mask_dict\n\n        # Same preprocessing as SkinDataset (must be consistent)\n        self.resize    = transforms.Resize((input_size, input_size))\n        self.to_tensor = transforms.ToTensor()\n        self.normalize = transforms.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n        )\n        # Light augmentation for contrastive pairs — heavier than val, lighter than train\n        self.augment = transforms.Compose([\n            transforms.RandomHorizontalFlip(),\n            transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.1),\n            transforms.RandomGrayscale(p=0.1),\n        ])\n\n    def __len__(self):\n        return len(self.pairs)\n\n    def _load_image(self, image_path, image_array):\n        \"\"\"Load image from path (real) or numpy array (synthetic).\"\"\"\n        if image_array is not None:\n            # Synthetic image stored as numpy array\n            return Image.fromarray(image_array.astype(np.uint8))\n        else:\n            return Image.open(image_path).convert(\"RGB\")\n\n    def _apply_contrast_norm(self, image_tensor, mask_tensor):\n        \"\"\"Gentle contrast normalisation — same as SkinDataset.\"\"\"\n        skin_region = (mask_tensor < 0.5).expand_as(image_tensor)\n        skin_pixels = image_tensor[skin_region]\n        if skin_pixels.numel() > 0:\n            skin_mean   = skin_pixels.mean()\n            image_tensor = image_tensor - skin_mean\n        return image_tensor\n\n    def __getitem__(self, idx):\n        row = self.pairs.iloc[idx]\n\n        # Load both images in the pair\n        img_a = self._load_image(row['image_path_a'], row['image_array_a'])\n        img_b = self._load_image(row['image_path_b'], row['image_array_b'])\n\n        # Resize\n        img_a = self.resize(img_a)\n        img_b = self.resize(img_b)\n\n        # Light augmentation (independently per image)\n        img_a = self.augment(img_a)\n        img_b = self.augment(img_b)\n\n        # Convert to tensor\n        img_a = self.to_tensor(img_a)\n        img_b = self.to_tensor(img_b)\n\n        # Load masks for contrast normalisation\n        mask_a = self._load_mask(row['image_id_a'])\n        mask_b = self._load_mask(row['image_id_b'])\n\n        # Apply gentle contrast normalisation (remove absolute skin tone baseline)\n        img_a = self._apply_contrast_norm(img_a, mask_a)\n        img_b = self._apply_contrast_norm(img_b, mask_b)\n\n        # Final normalisation (ImageNet stats)\n        img_a = self.normalize(img_a)\n        img_b = self.normalize(img_b)\n\n        label = int(row['label'])\n        return img_a, img_b, label\n\n    def _load_mask(self, image_id):\n        \"\"\"Load segmentation mask or return full-image default.\"\"\"\n        mask_path = self.mask_dict.get(str(image_id))\n        if mask_path and Path(str(mask_path)).exists():\n            mask_pil = Image.open(mask_path).convert(\"L\").resize((224, 224))\n            return transforms.ToTensor()(mask_pil)\n        return torch.ones(1, 224, 224)  # default: treat whole image as lesion\n\n\n# Create the dataset and loader\ncross_tone_dataset = CrossSkinToneDataset(cross_tone_pairs, mask_dict)\n# Batch size increased from 16 → 32.\n# NT-Xent contrastive loss needs enough negatives per batch to learn\n# meaningful distinctions. With batch_size=16 there are only 14 negatives\n# per positive pair — too few for the loss to stay stable.\n# batch_size=32 gives 30 negatives, which is the practical minimum.\ncross_tone_loader  = DataLoader(\n    cross_tone_dataset,\n    batch_size=32,\n    shuffle=True,\n    num_workers=0,\n    generator=torch.Generator().manual_seed(SEED)\n)\nprint(f\"Cross-skin-tone contrastive loader ready: {len(cross_tone_loader)} batches ✅\")\nprint(f\"  Batch size: 32 (increased from 16 — NT-Xent needs more negatives per batch)\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:41:34.805107Z","iopub.status.busy":"2026-07-17T03:41:34.804152Z","iopub.status.idle":"2026-07-17T03:41:34.818900Z","shell.execute_reply":"2026-07-17T03:41:34.817970Z"},"papermill":{"duration":1.631438,"end_time":"2026-07-17T03:41:34.821134+00:00","exception":false,"start_time":"2026-07-17T03:41:33.189696+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── THREE-COMPONENT LOSS + CONDITIONED CONTRASTIVE TRAINING ──────────────────\n#\n# LOSS FORMULA (mentor's design):\n#   L_total = L_disease                          ← learn the correct diagnosis\n#           + 0.5 * L_contrastive               ← same disease = same representation across skin tones\n#           - λ_adv * L_skin_tone               ← mild adversarial (GRL already flips gradient)\n#\n# KEY DIFFERENCE FROM BEFORE:\n#   Previously: λ_adv was the main tool for disentanglement\n#   Now:        λ_adv is CAPPED at 0.3 (gentle) — the heavy lifting is done by\n#               the contrastive loss, which directly teaches cross-skin-tone invariance\n#               WITHOUT destroying Kind 1 (biological) features\n#\n# WHY 0.5 FOR CONTRASTIVE WEIGHT?\n#   The contrastive loss scale is typically larger than cross-entropy losses.\n#   0.5 balances it so neither loss overwhelms the other.\n#\n# NT-XENT at temperature=0.07 (lower than SimCLR's 0.5):\n#   Lower temperature = harder negatives enforced more strictly\n#   This is important here because our pairs ARE the same disease —\n#   we need the model to pull them together strongly\n\ndef compute_conditioned_loss(embeddings_a, embeddings_b,\n                              disease_logits, disease_labels,\n                              fitzpatrick_groups):\n    \"\"\"\n    Two-component loss for skin-tone-conditioned contrastive training.\n\n    WHY THE ADVERSARIAL TERM WAS REMOVED:\n    The previous three-component loss included - λ_adv * L_skin_tone.\n    In practice this caused the total loss to go negative from epoch 8,\n    meaning the optimiser was being rewarded for maximising skin-tone loss\n    at the expense of learning to classify diseases. The contrastive loss\n    already does the disentanglement work — the adversarial term was\n    redundant and destabilising. Adversarial debiasing remains in the\n    supervised training phase (Cell 28) where it works correctly.\n\n    NEW FORMULA:\n        L_total = L_disease + 0.5 * L_contrastive\n\n        L_disease     → learn to classify the correct lesion type\n        L_contrastive → pull same-disease embeddings together across skin tones\n\n    Args:\n        embeddings_a/b    : [B, 256] embeddings from the two images in each pair\n        disease_logits    : [B, 7] disease classification scores\n        disease_labels    : [B] ground truth disease labels\n        fitzpatrick_groups: [B] Fitzpatrick group labels (for FairFocalLoss weighting)\n\n    Returns:\n        total_loss, disease_loss, contrastive_loss\n    \"\"\"\n    # ── Component 1: Disease classification ───────────────────\n    # FairFocalLoss with per-group weights — Groups 2 & 3 contribute more\n    disease_loss = criterion_lesion(disease_logits, disease_labels, fitzpatrick_groups)\n\n    # ── Component 2: Cross-skin-tone contrastive loss ─────────\n    # L2-normalise embeddings so cosine similarity is used (not dot product)\n    z_a = F.normalize(embeddings_a, dim=1)\n    z_b = F.normalize(embeddings_b, dim=1)\n\n    # NT-Xent at temperature=0.07 — strictly penalises same-disease pairs\n    # that are far apart in embedding space, regardless of skin tone\n    contrastive_loss = contrastive_loss_fn(z_a, z_b, temperature=0.07)\n\n    # ── Total ─────────────────────────────────────────────────\n    # 0.5 weight on contrastive balances its larger magnitude vs Focal Loss\n    total = disease_loss + 0.5 * contrastive_loss\n\n    return total, disease_loss, contrastive_loss\n\n\ndef contrastive_loss_fn(z_a, z_b, temperature=0.07):\n    \"\"\"NT-Xent loss for a batch of positive pairs (z_a[i], z_b[i]).\"\"\"\n    batch_size = z_a.shape[0]\n    representations = torch.cat([z_a, z_b], dim=0)           # [2B, D]\n    similarity      = torch.matmul(representations, representations.T)  # [2B, 2B]\n    labels          = torch.arange(batch_size, device=device)\n    labels          = torch.cat([labels + batch_size, labels], dim=0)   # [2B]\n    return F.cross_entropy(similarity / temperature, labels)\n\n\n# ── CONDITIONED CONTRASTIVE FINE-TUNING LOOP ─────────────────────────────────\n#\n# We fine-tune the BEST SUPERVISED MODEL (best_model_derm_aid.pth) — not from scratch.\n# This builds on the representation already learned and refines it with cross-skin-tone alignment.\n#\n# Why not retrain from scratch?\n#   The supervised model already learned to classify lesions well.\n#   We only need to push the embedding toward cross-skin-tone invariance —\n#   that's faster and safer at low LR than full retraining.\n\nCONDITIONED_EPOCHS = 15\nbest_conditioned_loss = float('inf')\n\n# Load best supervised model as starting point\nmodel.load_state_dict(torch.load(\"best_model_derm_aid.pth\"))\nprint(\"✅ Loaded best_model_derm_aid.pth as starting point for conditioned fine-tuning\")\n\n# Low learning rate — we are refining, not retraining\noptimizer_cond = torch.optim.Adam(model.parameters(), lr=3e-5, weight_decay=1e-5)\nscheduler_cond = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer_cond,\n    max_lr=3e-5,\n    steps_per_epoch=len(cross_tone_loader),\n    epochs=CONDITIONED_EPOCHS,\n    pct_start=0.1,\n    div_factor=10,\n    final_div_factor=100\n)\n\nprint(f\"Conditioned contrastive fine-tuning for {CONDITIONED_EPOCHS} epochs\")\nprint(f\"Loss = L_disease + 0.5*L_contrastive  (adversarial term removed — contrastive does disentanglement)\")\nprint(\"=\"*70)\n\ncond_history = {'train_loss': [], 'val_loss': [], 'val_accuracy': []}\n\nfor epoch in range(CONDITIONED_EPOCHS):\n    model.train()\n\n    total_loss_sum       = 0.0\n    disease_loss_sum     = 0.0\n    contrastive_loss_sum = 0.0\n    # skin_tone_loss_sum removed — adversarial term no longer in contrastive phase\n\n    loop = tqdm(cross_tone_loader, desc=f\"  Cond. Epoch {epoch+1}/{CONDITIONED_EPOCHS}\")\n\n    for img_a, img_b, disease_labels in loop:\n        img_a          = img_a.to(device)\n        img_b          = img_b.to(device)\n        disease_labels = disease_labels.to(device)\n\n        optimizer_cond.zero_grad()\n\n        # Forward pass for both images in the pair\n        # We still get skin_tone_logits from the model (it's part of the architecture)\n        # but we don't use them in the contrastive loss — adversarial debiasing\n        # happens separately in the supervised fine-tuning phase (Cell 38)\n        disease_logits_a, _, emb_a = model(img_a)\n        disease_logits_b, _, emb_b = model(img_b)\n\n        # Use image A's logits for disease classification loss\n        # (both should predict the same class — they're the same disease)\n        # Note: skin_tone_logits no longer needed here — adversarial term removed\n        total, d_loss, c_loss = compute_conditioned_loss(\n            emb_a, emb_b,\n            disease_logits_a, disease_labels,\n            disease_labels   # fitzpatrick_groups not in pair loader — use disease labels\n                             # FairFocalLoss group weights are applied in supervised phase\n        )\n\n        total.backward()\n        # Gradient clipping — prevents spikes when contrastive loss is large\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer_cond.step()\n        scheduler_cond.step()\n\n        total_loss_sum       += total.item()\n        disease_loss_sum     += d_loss.item()\n        contrastive_loss_sum += c_loss.item()\n\n        loop.set_postfix({\n            'L_total'       : f\"{total.item():.3f}\",\n            'L_disease'     : f\"{d_loss.item():.3f}\",\n            'L_contrast'    : f\"{c_loss.item():.3f}\",\n        })\n\n    n = len(cross_tone_loader)\n    avg_total       = total_loss_sum       / n\n    avg_disease     = disease_loss_sum     / n\n    avg_contrastive = contrastive_loss_sum / n\n\n    # Evaluate on validation set after each epoch\n    val_metrics = evaluate(model, val_loader, criterion_lesion)\n\n    is_best = val_metrics['loss'] < best_conditioned_loss\n    if is_best:\n        best_conditioned_loss = val_metrics['loss']\n        torch.save(model.state_dict(), \"best_model_conditioned.pth\")\n        tag = \" ← best ✅\"\n    else:\n        tag = \"\"\n\n    print(\n        f\"Epoch {epoch+1:02d}/{CONDITIONED_EPOCHS} | \"\n        f\"L_total={avg_total:.4f} \"\n        f\"(disease={avg_disease:.4f}, contrastive={avg_contrastive:.4f}) | \"\n        f\"Val Loss={val_metrics['loss']:.4f}  Val Acc={val_metrics['accuracy']:.4f}\"\n        f\"{tag}\"\n    )\n\n    cond_history['train_loss'].append(avg_total)\n    cond_history['val_loss'].append(val_metrics['loss'])\n    cond_history['val_accuracy'].append(val_metrics['accuracy'])\n\nprint(f\"\\nConditioned contrastive training complete.\")\nprint(f\"Best val loss: {best_conditioned_loss:.4f}\")\nprint(\"Saved: best_model_conditioned.pth ✅\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T03:41:37.975615Z","iopub.status.busy":"2026-07-17T03:41:37.975354Z","iopub.status.idle":"2026-07-17T04:36:02.742251Z","shell.execute_reply":"2026-07-17T04:36:02.741456Z"},"papermill":{"duration":3268.116706,"end_time":"2026-07-17T04:36:04.506823+00:00","exception":false,"start_time":"2026-07-17T03:41:36.390117+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── EVALUATE CONDITIONED MODEL + UPDATED LINEAR PROBE ────────────────────────\n#\n# We now compare THREE checkpoints:\n#   1. best_model_derm_aid.pth      — supervised baseline (contrast norm + adversarial)\n#   2. best_model_conditioned.pth   — + cross-skin-tone contrastive fine-tuning\n#\n# For each we run:\n#   (a) Fairness metrics (TPR per Fitzpatrick group, TPR gap)\n#   (b) Linear probe — logistic regression to predict Fitzpatrick group\n#       from the frozen disease embedding\n#\n# WHAT THE PROBE TELLS US:\n#   Probe ≈ 20% (chance)  → skin tone NOT decodable from embedding → disentanglement ✅\n#   Probe ≈ 35-45%        → PARTIAL disentanglement (clinically acceptable target)\n#   Probe >> 45%          → skin tone still strongly encoded → needs more work ❌\n#\n# NOTE: We now accept 35-45% as the target (not 20%).\n# Complete disentanglement is impossible because Kind 1 features (biological entanglement\n# of melanin and lesion appearance) are real and should be preserved.\n\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.model_selection import cross_val_score\n\ndef run_linear_probe(model, loader, checkpoint_path, label=\"model\"):\n    \"\"\"\n    Freeze the encoder, extract 256-dim disease embeddings,\n    then train logistic regression to predict Fitzpatrick group.\n\n    Args:\n        model            : DermAidModel instance\n        loader           : test DataLoader\n        checkpoint_path  : path to .pth file to evaluate\n        label            : name to print in the report\n    \"\"\"\n    model.load_state_dict(torch.load(checkpoint_path))\n    model.eval()\n\n    embeddings, fitz_labels = [], []\n\n    with torch.no_grad():\n        for images, masks, labels, fitzpatrick_groups in loader:\n            images = images.to(device)\n            masks  = masks.to(device)\n\n            # Apply mask (consistent with training)\n            masked_images = images * masks\n\n            # Third return value is the 256-dim flat embedding\n            _, _, embedding = model(masked_images)\n            embeddings.append(embedding.cpu().numpy())\n            fitz_labels.extend(fitzpatrick_groups.numpy())\n\n    embeddings  = np.vstack(embeddings)\n    fitz_labels = np.array(fitz_labels)\n\n    probe  = LogisticRegression(max_iter=1000, random_state=SEED)\n    scores = cross_val_score(probe, embeddings, fitz_labels, cv=5, scoring='accuracy')\n\n    probe_acc = scores.mean()\n    print(f\"\\nLinear Probe — {label}\")\n    print(f\"  Probe accuracy     : {probe_acc:.4f} ± {scores.std():.4f}\")\n    print(f\"  Chance level       : 0.2000 (20%)\")\n    print(f\"  Target range       : 0.35 – 0.45 (partial disentanglement)\")\n\n    if probe_acc <= 0.35:\n        status = \"✅ Strong disentanglement (may be losing Kind 1 biology)\"\n    elif probe_acc <= 0.45:\n        status = \"✅ Partial disentanglement — clinically acceptable\"\n    elif probe_acc <= 0.60:\n        status = \"⚠️  Moderate — skin tone still somewhat encoded\"\n    else:\n        status = \"❌ Strong skin tone encoding — further work needed\"\n\n    print(f\"  Status             : {status}\")\n    return scores, probe_acc\n\n\n# ── Run fairness metrics on conditioned model ─────────────────────────────\nprint(\"=\"*70)\nprint(\"FAIRNESS EVALUATION — CONDITIONED CONTRASTIVE MODEL\")\nprint(\"=\"*70)\nmodel.load_state_dict(torch.load(\"best_model_conditioned.pth\"))\ncond_test_metrics = evaluate(model, test_loader, criterion_lesion)\ncond_fairness = compute_metrics(\n    all_preds       = cond_test_metrics['preds'],\n    all_labels      = cond_test_metrics['labels'],\n    all_probs       = cond_test_metrics['probs'],\n    all_fitzpatrick = cond_test_metrics['fitzpatrick'],\n    label_encoder   = le,\n)\n\n# ── Linear probe comparison ───────────────────────────────────────────────\nprint(\"\\n\" + \"=\"*70)\nprint(\"LINEAR PROBE COMPARISON\")\nprint(\"=\"*70)\n\nscores_baseline, acc_baseline = run_linear_probe(\n    model, test_loader,\n    \"best_model_derm_aid.pth\",\n    label=\"Supervised baseline (contrast norm + adversarial)\"\n)\n\nscores_conditioned, acc_conditioned = run_linear_probe(\n    model, test_loader,\n    \"best_model_conditioned.pth\",\n    label=\"Conditioned contrastive model\"\n)\n\n# ── Summary table ─────────────────────────────────────────────────────────\nprint(\"\\n\" + \"=\"*70)\nprint(\"DISENTANGLEMENT SUMMARY\")\nprint(\"=\"*70)\nprint(f\"  {'Checkpoint':<45}  {'Probe Acc':>10}  {'Status':>30}\")\nprint(f\"  {'─'*45}  {'─'*10}  {'─'*30}\")\nprint(f\"  {'Supervised baseline':<45}  {acc_baseline:>10.4f}  {'(see above)'}\")\nprint(f\"  {'Conditioned contrastive':<45}  {acc_conditioned:>10.4f}  {'(see above)'}\")\nprint(f\"  {'Chance level':<45}  {'0.2000':>10}  {'—'}\")\nprint(f\"  {'Clinical target':<45}  {'0.35–0.45':>10}  {'—'}\")\nprint(\"=\"*70)\ndelta = acc_baseline - acc_conditioned\nprint(f\"  Probe drop: {delta:+.4f} ({'improvement' if delta > 0 else 'no improvement'})\")\nprint(f\"  Interpretation: Kind 1 features (biological entanglement) set a floor\")\nprint(f\"  around 35-45%. Below that risks destroying real biology.\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T04:36:08.009406Z","iopub.status.busy":"2026-07-17T04:36:08.008983Z","iopub.status.idle":"2026-07-17T04:38:24.439268Z","shell.execute_reply":"2026-07-17T04:38:24.436155Z"},"papermill":{"duration":140.131932,"end_time":"2026-07-17T04:38:26.405051+00:00","exception":false,"start_time":"2026-07-17T04:36:06.273119+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── PHASE 2: ADVERSARIAL FINE-TUNING ON MASKED MODEL ────────────────────────\n# Builds on best_model_derm_aid.pth (trained with gentle contrast normalisation).\n# Adds adversarial skin-tone disentanglement for 10 epochs at low LR.\n# Goal: push Group 3 TPR further and close the TPR gap below 5%.\n\nFINETUNE_EPOCHS  = 10\nLAMBDA_ADV_FINE  = 0.4   # fixed — no warmup, model is already stable\nbest_ft_val_loss = float('inf')\n\n# Load the best masked model as the starting point\nmodel.load_state_dict(torch.load(\"best_model_derm_aid.pth\"))\nprint(\"✅ Loaded best_model_derm_aid.pth as fine-tuning starting point\")\n\n# Unfreeze everything — fine-tune end-to-end at low LR\nfor param in model.parameters():\n    param.requires_grad = True\n\noptimizer_ft = torch.optim.Adam(model.parameters(), lr=5e-5, weight_decay=1e-5)\nscheduler_ft = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer_ft,\n    max_lr=5e-5,\n    steps_per_epoch=len(train_loader),\n    epochs=FINETUNE_EPOCHS,\n    pct_start=0.1,\n    div_factor=10,\n    final_div_factor=100\n)\n\nprint(f\"Fine-tuning for {FINETUNE_EPOCHS} epochs (λ_adv={LAMBDA_ADV_FINE})\")\nprint(\"=\"*70)\n\nft_history = {'train_loss': [], 'val_loss': [], 'val_accuracy': []}\n\nfor epoch in range(FINETUNE_EPOCHS):\n    train_losses = train_one_epoch(\n        model, train_loader, optimizer_ft, scheduler_ft,\n        criterion_lesion, criterion_skin_tone,\n        lambda_adv=LAMBDA_ADV_FINE\n    )\n    val_metrics = evaluate(model, val_loader, criterion_lesion)\n    current_lr  = optimizer_ft.param_groups[0]['lr']\n\n    is_best = val_metrics['loss'] < best_ft_val_loss\n    if is_best:\n        best_ft_val_loss = val_metrics['loss']\n        torch.save(model.state_dict(), \"best_model_masked_adversarial.pth\")\n        checkpoint_tag = \" ← best ✅\"\n    else:\n        checkpoint_tag = \"\"\n\n    print(\n        f\"FT Epoch {epoch+1:02d}/{FINETUNE_EPOCHS} | \"\n        f\"L_total={train_losses['total']:.4f} \"\n        f\"(lesion={train_losses['lesion']:.4f}, \"\n        f\"skin_tone={train_losses['skin_tone']:.4f}) | \"\n        f\"Val Loss={val_metrics['loss']:.4f}  \"\n        f\"Val Acc={val_metrics['accuracy']:.4f}  \"\n        f\"LR={current_lr:.2e}\"\n        f\"{checkpoint_tag}\"\n    )\n\n    ft_history['train_loss'].append(train_losses['total'])\n    ft_history['val_loss'].append(val_metrics['loss'])\n    ft_history['val_accuracy'].append(val_metrics['accuracy'])\n\nprint(f\"\\nFine-tuning complete. Best val loss: {best_ft_val_loss:.4f}\")\nprint(\"Saved: best_model_masked_adversarial.pth ✅\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T04:38:29.797267Z","iopub.status.busy":"2026-07-17T04:38:29.796521Z","iopub.status.idle":"2026-07-17T05:17:31.247420Z","shell.execute_reply":"2026-07-17T05:17:31.246701Z"},"papermill":{"duration":2345.297528,"end_time":"2026-07-17T05:17:33.373002+00:00","exception":false,"start_time":"2026-07-17T04:38:28.075474+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── EVALUATE FINE-TUNED MODEL + LINEAR PROBE ─────────────────────────────────\n# Step 1: fairness metrics on the fine-tuned model\n# Step 2: linear probe on both checkpoints to measure skin-tone disentanglement\n#\n# Linear probe logic:\n#   Freeze the encoder. Train logistic regression to predict Fitzpatrick group\n#   from the 256-dim disease embedding.\n#   Probe accuracy ≈ 20% (chance) → skin tone NOT encoded → disentanglement ✅\n#   Probe accuracy >> 20%         → skin tone still encoded → needs more work ❌\n\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.model_selection import cross_val_score\n\n# ── Load fine-tuned model and evaluate fairness ───────────────\nmodel.load_state_dict(torch.load(\"best_model_masked_adversarial.pth\"))\nprint(\"✅ Loaded best_model_masked_adversarial.pth\")\n\nft_test_metrics = evaluate(model, test_loader, criterion_lesion)\nprint(\"\\nFairness metrics — Masked + Adversarial model:\")\nft_fairness = compute_metrics(\n    all_preds       = ft_test_metrics['preds'],\n    all_labels      = ft_test_metrics['labels'],\n    all_probs       = ft_test_metrics['probs'],\n    all_fitzpatrick = ft_test_metrics['fitzpatrick'],\n    label_encoder   = le,\n)\n\n# ── Linear probe helper ───────────────────────────────────────\ndef run_linear_probe(model, loader, label=\"model\"):\n    model.eval()\n    embeddings, fitz_labels = [], []\n\n    with torch.no_grad():\n        for images, masks, labels, fitzpatrick_groups in loader:\n            images = images.to(device)\n            masks  = masks.to(device)\n            masked_images = images * masks\n\n            # Third return value is the 256-dim flat embedding\n            _, _, embedding = model(masked_images)\n            embeddings.append(embedding.cpu().numpy())\n            fitz_labels.extend(fitzpatrick_groups.numpy())\n\n    embeddings  = np.vstack(embeddings)\n    fitz_labels = np.array(fitz_labels)\n\n    probe  = LogisticRegression(max_iter=1000, random_state=SEED)\n    scores = cross_val_score(probe, embeddings, fitz_labels, cv=5, scoring='accuracy')\n\n    print(f\"\\nLinear Probe — {label}\")\n    print(f\"  Probe accuracy : {scores.mean():.4f} ± {scores.std():.4f}\")\n    print(f\"  Chance level   : {1/5:.4f} (20.00%)\")\n    skin_tone_encoded = scores.mean() > 0.25\n    print(f\"  Skin tone encoded in representation: {skin_tone_encoded}\")\n    return scores\n\n# ── Run probe on both checkpoints ────────────────────────────\nmodel.load_state_dict(torch.load(\"best_model_derm_aid.pth\"))\nprobe_baseline = run_linear_probe(model, test_loader, label=\"Masked model (baseline)\")\n\nmodel.load_state_dict(torch.load(\"best_model_masked_adversarial.pth\"))\nprobe_adversarial = run_linear_probe(model, test_loader, label=\"Masked + Adversarial model\")\n\n# ── Summary table ─────────────────────────────────────────────\nprint(\"\\n\" + \"=\"*70)\nprint(\"DISENTANGLEMENT SUMMARY\")\nprint(\"=\"*70)\nprint(f\"{'':35s}  {'Baseline':>10}  {'Masked+Adv':>10}\")\nprint(f\"{'Probe accuracy':35s}  {probe_baseline.mean():>10.4f}  {probe_adversarial.mean():>10.4f}\")\nprint(f\"{'Chance level':35s}  {'0.2000':>10}  {'0.2000':>10}\")\nprint(\"=\"*70)\ndelta = probe_baseline.mean() - probe_adversarial.mean()\nprint(f\"  Probe drop after adversarial training: {delta:.4f}\")\nprint(f\"  {'✅ Disentanglement confirmed' if delta > 0.02 else '❌ Minimal disentanglement — consider increasing LAMBDA_ADV_FINE'}\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T05:17:37.460886Z","iopub.status.busy":"2026-07-17T05:17:37.460067Z","iopub.status.idle":"2026-07-17T05:20:03.081002Z","shell.execute_reply":"2026-07-17T05:20:03.080107Z"},"papermill":{"duration":149.768478,"end_time":"2026-07-17T05:20:05.279320+00:00","exception":false,"start_time":"2026-07-17T05:17:35.510842+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── GRAD-CAM DEMO + FULL FITZPATRICK AUDIT ────────────────────────────────────\n# Fixes the RecursionError:\n#   Root cause: the old _ModelWithFeaturesAttr overrode .train() and .eval(),\n#   but cam.explain() called self.model.eval() which called self._real_model.eval()\n#   which called .train(False) on the nn.Module, which iterates .children() and\n#   hit the wrapper's .train() again — infinite loop.\n#\n#   Fix: do NOT subclass or mimic nn.Module at all.  The wrapper is a plain Python\n#   object.  It never defines .train() or .eval() — PyTorch's machinery therefore\n#   never recurses into it.  Instead we call _real_model.eval() ourselves exactly\n#   once, before we hand the wrapper to GradCAM.\n\nfrom derm_aid_explainability import GradCAM, INFERENCE_TRANSFORM, FITZPATRICK_LABELS\nimport types, os\n\n# ─────────────────────────────────────────────────────────────────────────────\n# WRAPPER — plain Python object, NOT an nn.Module subclass.\n# Exposes .backbone.features so GradCAM can register its hook on the right layer.\n# Never defines .train() / .eval() to prevent any recursion.\n# ─────────────────────────────────────────────────────────────────────────────\nclass _GradCAMWrapper:\n    \"\"\"Thin adapter so GradCAM can find model.backbone.features without\n    touching nn.Module internals.\"\"\"\n    def __init__(self, real_model):\n        self._real_model = real_model\n        # SimpleNamespace is not an nn.Module — PyTorch won't iterate its children\n        self.backbone = types.SimpleNamespace(features=real_model.backbone)\n\n    def __call__(self, x):\n        return self._real_model(x)\n\n    def parameters(self):\n        return self._real_model.parameters()\n\n    def zero_grad(self):\n        self._real_model.zero_grad()\n\n    def eval(self):\n        self._real_model.eval()\n        return self\n\n    def train(self, mode=True):\n        self._real_model.train(mode)\n        return self\n\n\n# Put the real model in eval mode ONCE, before wrapping\nmodel.eval()\nwrapped_model = _GradCAMWrapper(model)\ncam = GradCAM(wrapped_model)\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# HELPER: preprocess one image+mask → masked tensor\n# ─────────────────────────────────────────────────────────────────────────────\ndef prepare_sample(img_path, mask_path_sample, device):\n    \"\"\"Return (orig_pil, masked_tensor) for one sample.\"\"\"\n    orig = Image.open(img_path).convert(\"RGB\").resize((224, 224))\n    if mask_path_sample and Path(mask_path_sample).exists():\n        mask_pil = Image.open(mask_path_sample).convert(\"L\").resize((224, 224))\n        mask_np  = np.array(mask_pil).astype(np.float32) / 255.0\n        mask_t   = torch.tensor(mask_np).unsqueeze(0).unsqueeze(0).expand(1, 3, -1, -1)\n    else:\n        mask_t = torch.ones(1, 3, 224, 224)\n    img_t  = INFERENCE_TRANSFORM(orig).unsqueeze(0).to(device)\n    masked = (img_t * mask_t.to(device)).requires_grad_(True)\n    return orig, masked\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# SECTION 1 — DEMO: one sample per Fitzpatrick group (original behaviour)\n# ─────────────────────────────────────────────────────────────────────────────\nprint(\"Generating demo Grad-CAM (1 sample per Fitzpatrick group)...\")\n\nfig, axes = plt.subplots(2, 5, figsize=(18, 8))\nfig.suptitle(\n    \"Grad-CAM Explanations per Fitzpatrick Skin-Tone Group\\n\"\n    \"Top: original masked image   |   Bottom: Grad-CAM heatmap overlay\",\n    fontsize=12, fontweight='bold'\n)\n\nfor col, group_id in enumerate(range(5)):\n    group_mask = test_df['fitzpatrick_group'] == group_id\n    if group_mask.sum() == 0:\n        for row in range(2):\n            axes[row][col].axis('off')\n        continue\n\n    sample   = test_df[group_mask].iloc[0]\n    img_path = sample['image_path']\n    msk_path = mask_dict.get(Path(img_path).stem, None)\n\n    orig, masked = prepare_sample(img_path, msk_path, device)\n    _, overlay   = cam.explain(masked)\n\n    axes[0][col].imshow(orig)\n    axes[0][col].set_title(FITZPATRICK_LABELS[group_id], fontsize=8)\n    axes[0][col].axis('off')\n\n    axes[1][col].imshow(overlay)\n    with torch.no_grad():\n        pred_label = le.classes_[wrapped_model(masked)[0].argmax().item()]\n    axes[1][col].set_title(f\"Pred: {pred_label}\", fontsize=8, color='#1D4ED8')\n    axes[1][col].axis('off')\n\ncam.remove_hooks()\nplt.tight_layout()\nplt.savefig('gradcam_per_group.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint(\"Saved: gradcam_per_group.png ✅\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# SECTION 2 — FULL GRAD-CAM AUDIT\n#\n# For each Fitzpatrick group (0–4):\n#   • Sample up to 20 CORRECTLY classified test images\n#   • Sample up to 20 MISCLASSIFIED test images\n#   • Generate Grad-CAM heatmaps for all of them\n#   • Save a 4×5 grid (= 20 images) per group × split combination\n#\n# This lets you answer:\n#   \"Where is the model looking when it gets it right?\"\n#   \"Where is it looking when it gets it wrong?\"\n#   \"Does the attention pattern change across skin tones?\"\n# ─────────────────────────────────────────────────────────────────────────────\nprint(\"\\n\" + \"=\"*70)\nprint(\"GRAD-CAM AUDIT — 20 correct + 20 misclassified per Fitzpatrick group\")\nprint(\"=\"*70)\n\n# Build a DataFrame of test-set predictions (reuse all_preds / all_labels / all_fitzpatrick)\naudit_df = test_df.copy().reset_index(drop=True)\naudit_df['pred']    = all_preds\naudit_df['correct'] = (audit_df['pred'] == audit_df['label']).astype(int)\n\nAUDIT_N   = 20        # samples per (group, split) cell\nAUDIT_DIR = 'gradcam_audit'\nos.makedirs(AUDIT_DIR, exist_ok=True)\n\nFITZ_NAMES = {\n    0: 'Very_Light',\n    1: 'Light',\n    2: 'Medium',\n    3: 'Dark',\n    4: 'Very_Dark',\n}\n\nfor group_id in range(5):\n    group_df = audit_df[audit_df['fitzpatrick_group'] == group_id]\n    if len(group_df) == 0:\n        print(f\"  Fitzpatrick {group_id}: no samples — skipped\")\n        continue\n\n    for split_label, correct_flag in [('correct', 1), ('misclassified', 0)]:\n        subset = group_df[group_df['correct'] == correct_flag]\n        n_avail = len(subset)\n\n        if n_avail == 0:\n            print(f\"  Fitzpatrick {group_id} / {split_label}: 0 samples — skipped\")\n            continue\n\n        sample_n = min(AUDIT_N, n_avail)\n        subset   = subset.sample(n=sample_n, random_state=42).reset_index(drop=True)\n\n        # Determine grid shape\n        ncols = 5\n        nrows = max(1, (sample_n + ncols - 1) // ncols)  # ceil division\n\n        fig, axes = plt.subplots(nrows, ncols,\n                                 figsize=(ncols * 4, nrows * 4))\n        # Normalise axes to always be 2-D list\n        if nrows == 1 and ncols == 1:\n            axes = [[axes]]\n        elif nrows == 1:\n            axes = [axes]\n        elif ncols == 1:\n            axes = [[ax] for ax in axes]\n\n        fitz_name = FITZ_NAMES[group_id]\n        fig.suptitle(\n            f\"Fitzpatrick {group_id} ({fitz_name}) — {split_label.capitalize()} ({sample_n} samples)\\n\"\n            f\"Model attention: where does it look when it gets the diagnosis {'right' if correct_flag else 'wrong'}?\",\n            fontsize=11, fontweight='bold'\n        )\n\n        # Reinitialise cam for each grid (fresh hooks)\n        cam_audit = GradCAM(wrapped_model)\n\n        for idx, (_, row) in enumerate(subset.iterrows()):\n            r, c = divmod(idx, ncols)\n            ax   = axes[r][c]\n\n            img_path = row['image_path']\n            msk_path = mask_dict.get(Path(img_path).stem, None)\n\n            try:\n                orig, masked = prepare_sample(img_path, msk_path, device)\n                _, overlay   = cam_audit.explain(masked)\n\n                true_name = le.classes_[row['label']]\n                pred_name = le.classes_[row['pred']]\n                colour    = '#16A34A' if correct_flag else '#DC2626'\n\n                ax.imshow(overlay)\n                ax.set_title(\n                    f\"GT:{true_name}\\nP:{pred_name}\",\n                    fontsize=7, color=colour, pad=2\n                )\n            except Exception as e:\n                ax.text(0.5, 0.5, f\"Error:\\n{e}\", ha='center', va='center',\n                        fontsize=6, transform=ax.transAxes, wrap=True)\n\n            ax.axis('off')\n\n        # Hide any unused subplot slots\n        for empty_idx in range(sample_n, nrows * ncols):\n            r, c = divmod(empty_idx, ncols)\n            axes[r][c].axis('off')\n\n        cam_audit.remove_hooks()\n\n        fname = f\"{AUDIT_DIR}/fitz{group_id}_{fitz_name}_{split_label}.png\"\n        plt.tight_layout()\n        plt.savefig(fname, dpi=130, bbox_inches='tight')\n        plt.close()\n        print(f\"  Saved: {fname}  ({sample_n} images)\")\n\nprint(\"\\n✅ Grad-CAM audit complete.\")\nprint(f\"   Files saved to: {AUDIT_DIR}/\")\nprint(\"  ✔ Heatmap on the LESION (not bare skin)       → model is reasoning correctly\")\nprint(\"  ✗ Heatmap on SURROUNDING SKIN                 → skin tone may be a shortcut\")\nprint(\"  ✗ Attention SHIFTS across Fitzpatrick groups  → bias signal\")\nprint(\"  ✗ Misclassifications cluster on darker groups → high-priority fairness flag\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T05:20:09.398083Z","iopub.status.busy":"2026-07-17T05:20:09.397703Z","iopub.status.idle":"2026-07-17T05:20:44.606850Z","shell.execute_reply":"2026-07-17T05:20:44.605840Z"},"papermill":{"duration":37.269679,"end_time":"2026-07-17T05:20:44.608537+00:00","exception":false,"start_time":"2026-07-17T05:20:07.338858+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── QUANTITATIVE GRAD-CAM EVALUATION — IoU per Fitzpatrick Group ─────────────\n#\n# WHY THIS MATTERS (supervisor feedback):\n#   Visual inspection of Grad-CAM heatmaps is qualitative.\n#   We need a number: how much of the model's attention actually falls\n#   on the lesion vs the surrounding skin?\n#\n# HOW IoU WORKS HERE:\n#   - Grad-CAM produces a heatmap: high values = model is attending there\n#   - We threshold it at 0.5 → binary attention map\n#   - We compare that binary map to the ground-truth segmentation mask\n#   - IoU = intersection / union of the two binary maps\n#\n# WHAT THE RESULT MEANS:\n#   IoU close to 1.0 → model is looking exactly at the lesion ✅\n#   IoU close to 0.0 → model is looking at skin, not the lesion ❌\n#   If IoU drops for darker Fitzpatrick groups → quantitative bias signal\n\nimport glob\n\ndef gradcam_iou(cam_np, mask_pil, threshold=0.5):\n    \"\"\"\n    Compute IoU between a Grad-CAM heatmap and a ground-truth segmentation mask.\n\n    Args:\n        cam_np    : HxW float array in [0,1] — raw Grad-CAM output\n        mask_pil  : PIL Image (greyscale) — ground-truth lesion mask\n        threshold : binarisation threshold for the heatmap (default 0.5)\n\n    Returns:\n        iou : float in [0, 1]\n    \"\"\"\n    # Resize CAM to match mask size (224×224)\n    cam_resized = cv2.resize(cam_np, (224, 224))\n\n    # Binarise the CAM — anything above threshold = model attended here\n    cam_binary  = (cam_resized > threshold).astype(np.uint8)\n\n    # Binarise the ground-truth mask — 1 = lesion, 0 = skin\n    mask_np     = np.array(mask_pil.resize((224, 224)).convert('L'))\n    mask_binary = (mask_np > 128).astype(np.uint8)\n\n    intersection = (cam_binary & mask_binary).sum()\n    union        = (cam_binary | mask_binary).sum()\n\n    if union == 0:\n        # Edge case: both maps are empty — undefined IoU, return 0\n        return 0.0\n\n    return float(intersection) / float(union)\n\n\n# ── Run IoU across test set, grouped by Fitzpatrick type ─────\nprint(\"Computing Grad-CAM IoU per Fitzpatrick group...\")\nprint(\"(This re-uses the wrapped_model and cam objects from the Grad-CAM cell above)\")\nprint(\"=\"*70)\n\n# Re-initialise cam on the best model checkpoint\nmodel.load_state_dict(torch.load(\"best_model_derm_aid.pth\"))\nmodel.eval()\nwrapped_model_iou = _GradCAMWrapper(model)\ncam_iou           = GradCAM(wrapped_model_iou)\n\n# We sample up to N_SAMPLE images per Fitzpatrick group to keep runtime reasonable\nN_SAMPLE    = 30\niou_results = {g: [] for g in range(5)}\n\nFITZ_NAMES  = {0:'Very Light', 1:'Light', 2:'Medium', 3:'Dark', 4:'Very Dark'}\n\nfor group_id in range(5):\n    group_df = test_df[test_df['fitzpatrick_group'] == group_id]\n    sample   = group_df.sample(n=min(N_SAMPLE, len(group_df)), random_state=SEED)\n\n    for _, row in sample.iterrows():\n        img_path  = row['image_path']\n        mask_path = mask_dict.get(Path(img_path).stem)\n\n        # Skip if no ground-truth mask — IoU is meaningless without it\n        if not mask_path or not Path(mask_path).exists():\n            continue\n\n        try:\n            orig, masked = prepare_sample(img_path, mask_path, device)\n\n            # Generate Grad-CAM — returns (raw_cam_np, overlay_pil)\n            raw_cam, _ = cam_iou.explain(masked)\n\n            # Load ground-truth mask\n            mask_pil = Image.open(mask_path).convert('L')\n\n            iou = gradcam_iou(raw_cam, mask_pil)\n            iou_results[group_id].append(iou)\n\n        except Exception as e:\n            print(f\"  Warning: {Path(img_path).stem} — {e}\")\n            continue\n\ncam_iou.remove_hooks()\n\n# ── Print results ─────────────────────────────────────────────\nprint(\"\\nGrad-CAM IoU Results (higher = model looks at lesion, not skin):\")\nprint(\"-\"*55)\nall_ious = []\nfor group_id in range(5):\n    ious = iou_results[group_id]\n    if len(ious) == 0:\n        print(f\"  Group {group_id} ({FITZ_NAMES[group_id]:<12}): no masked samples\")\n        continue\n    mean_iou = np.mean(ious)\n    std_iou  = np.std(ious)\n    all_ious.extend(ious)\n    flag = \"✅\" if mean_iou > 0.4 else \"⚠️ \" if mean_iou > 0.25 else \"❌\"\n    print(f\"  Group {group_id} ({FITZ_NAMES[group_id]:<12}): \"\n          f\"IoU = {mean_iou:.3f} ± {std_iou:.3f}  (n={len(ious)})  {flag}\")\n\nprint(\"-\"*55)\nif all_ious:\n    print(f\"  Overall mean IoU : {np.mean(all_ious):.3f}\")\n    # IoU gap between best and worst group — like the TPR gap but for attention\n    group_means = {g: np.mean(v) for g, v in iou_results.items() if len(v) > 0}\n    best_g  = max(group_means, key=group_means.get)\n    worst_g = min(group_means, key=group_means.get)\n    iou_gap = group_means[best_g] - group_means[worst_g]\n    print(f\"  IoU gap          : {iou_gap:.3f}  \"\n          f\"(best: Group {best_g} / worst: Group {worst_g})\")\n    print(f\"\\n  Interpretation:\")\n    print(f\"  IoU gap > 0.10 → model attends differently by skin tone → attention bias\")\n    print(f\"  IoU gap < 0.10 → attention is consistent across skin tones → fair attention\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T05:20:48.914071Z","iopub.status.busy":"2026-07-17T05:20:48.913727Z","iopub.status.idle":"2026-07-17T05:20:58.486472Z","shell.execute_reply":"2026-07-17T05:20:58.485547Z"},"papermill":{"duration":11.666657,"end_time":"2026-07-17T05:20:58.488122+00:00","exception":false,"start_time":"2026-07-17T05:20:46.821465+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import Image, display\n\n# Demo overview\ndisplay(Image('gradcam_per_group.png'))\n # Audit grids — run this to flip through all of them\nimport glob\n\nfor path in sorted(glob.glob('gradcam_audit/*.png')):\n    print(f\"\\n{'='*50}\\n{path}\\n{'='*50}\")\n    display(Image(path))","metadata":{"execution":{"iopub.execute_input":"2026-07-17T05:21:02.588425Z","iopub.status.busy":"2026-07-17T05:21:02.588000Z","iopub.status.idle":"2026-07-17T05:21:02.895923Z","shell.execute_reply":"2026-07-17T05:21:02.895228Z"},"papermill":{"duration":2.34789,"end_time":"2026-07-17T05:21:02.935517+00:00","exception":false,"start_time":"2026-07-17T05:21:00.587627+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# GENERATE PER-PATIENT PDF REPORT — DEMO\n# The PDF contains:\n#   Page 1: lesion image + Grad-CAM + prediction + class probabilities + plain-language disclaimer for patients\n#   Page 2: model-level fairness metrics embedded as static context so the clinician knows how reliable the model is across skin tones when they receive the report.\n\nsample     = test_df.iloc[0]\nimg_path   = sample['image_path']\nmask_path_demo  = mask_dict.get(Path(img_path).stem, img_path)  # fallback to img if no mask\npatient_id = f\"DEMO-{Path(img_path).stem[:8].upper()}\"\n\nreport_path = generate_patient_report(\n    image_path       = img_path,\n    mask_path        = mask_path_demo,\n    model            = wrapped_model,\n    label_encoder    = le,\n    fairness_summary = fairness_summary,\n    output_path      = f\"DermAid_Report_{patient_id}.pdf\",\n    patient_id       = patient_id,\n    device           = device,\n)\n\nprint(f\"\\nReport generated: {report_path}\")\nprint(\"Open it to review the layout before deploying to the app.\")\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T05:21:07.797133Z","iopub.status.busy":"2026-07-17T05:21:07.796436Z","iopub.status.idle":"2026-07-17T05:21:08.342583Z","shell.execute_reply":"2026-07-17T05:21:08.341727Z"},"papermill":{"duration":2.914939,"end_time":"2026-07-17T05:21:08.344221+00:00","exception":false,"start_time":"2026-07-17T05:21:05.429282+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# EXPORT MODEL FOR MOBILE (OFFLINE DEPLOYMENT)\n# Produces a TorchScript .ptl file that PyTorch Mobilecan load on iOS and Android with no internet connection.\n# What gets stripped:\n#   - GradientReversalLayer (training-only, no-op at inference)\n#   - Skin-tone adversary head (not needed for patient-facing prediction)\n#\n# What is kept:\n#   - EfficientNet-B0 backbone\n#   - Projection head (512 → 256)\n#   - Lesion classifier (256 → 7)\n#\n# The app runs Grad-CAM separately on the full Python model (or implements it natively); this export is for prediction only.\n\nexport_for_mobile(\n    model       = model,\n    output_path = \"derm_aid_mobile.ptl\",\n    device      = device,\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-07-17T05:21:13.057505Z","iopub.status.busy":"2026-07-17T05:21:13.057002Z","iopub.status.idle":"2026-07-17T05:21:17.178301Z","shell.execute_reply":"2026-07-17T05:21:17.177413Z"},"papermill":{"duration":6.479237,"end_time":"2026-07-17T05:21:17.179833+00:00","exception":false,"start_time":"2026-07-17T05:21:10.700596+00:00","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. External Validation — ISIC 2020\nTo test generalisability beyond HAM10000, the best-performing model is evaluated on the **ISIC 2020** dataset. Since ISIC 2020 does not include Fitzpatrick labels, skin tone groups are estimated using the same **ITA scoring pipeline** already used for HAM10000. This section:\n1. Loads and preprocesses ISIC 2020 images\n2. Maps ISIC 2020 disease labels to the HAM10000 7-class taxonomy\n3. Computes ITA scores to assign Fitzpatrick groups\n4. Runs the same `evaluate()` and `compute_metrics()` functions used throughout\n\n**To add the dataset on Kaggle:** go to *Add Data* → search `siim-isic-melanoma-classification` → Add.","metadata":{}},{"cell_type":"code","source":"# ── ISIC 2020: LOAD METADATA AND MAP IMAGE PATHS ─────────────────────────────\n#\n# ISIC 2020 is a Kaggle competition dataset for melanoma detection.\n# It contains 33,126 dermoscopy images with diagnosis labels and image IDs.\n# We use it here for EXTERNAL VALIDATION ONLY — the model is not retrained.\n#\n# Dataset path on Kaggle after adding siim-isic-melanoma-classification:\nISIC2020_PATH = Path('/kaggle/input/competitions/siim-isic-melanoma-classification')\nISIC2020_CSV  = ISIC2020_PATH / 'train.csv'\nISIC2020_IMGS = ISIC2020_PATH / 'jpeg' / 'train'\n\nisic_df = pd.read_csv(ISIC2020_CSV)\nprint(f\"ISIC 2020 metadata loaded: {len(isic_df)} images\")\nprint(f\"Columns: {list(isic_df.columns)}\")\nprint(f\"\\nDiagnosis distribution:\")\nprint(isic_df['diagnosis'].value_counts())\n\n# ── Map image paths ───────────────────────────────────────────\n# ISIC 2020 stores images as {image_name}.jpg\nisic_df['image_path'] = isic_df['image_name'].apply(\n    lambda x: str(ISIC2020_IMGS / f\"{x}.jpg\")\n)\n\n# Drop images whose file doesn't exist on disk\nisic_df = isic_df[isic_df['image_path'].apply(lambda p: Path(p).exists())]\nprint(f\"\\nImages found on disk: {len(isic_df)}\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(list(le.classes_))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}