{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":20270,"databundleVersionId":1222630,"sourceType":"competition"},{"sourceId":13205039,"sourceType":"datasetVersion","datasetId":8369086},{"sourceId":268190984,"sourceType":"kernelVersion"},{"sourceId":271785730,"sourceType":"kernelVersion"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\n\n# Retrieve the Hugging Face token from Kaggle Secrets and set it as an environment variable\nhf_token = user_secrets.get_secret(\"HUGGING_FACE_TOKEN\")\nos.environ[\"HF_TOKEN\"] = hf_token","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:02.698222Z","iopub.execute_input":"2025-10-28T09:28:02.698522Z","iopub.status.idle":"2025-10-28T09:28:02.981921Z","shell.execute_reply.started":"2025-10-28T09:28:02.698497Z","shell.execute_reply":"2025-10-28T09:28:02.981323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import requests\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\n\nfrom typing import Any\nfrom PIL import Image\nimport numpy as np\nimport math\nimport re\nimport base64\nimport random\nfrom google.cloud import storage\nimport pandas as pd\nimport io\nimport json\nimport ast\n\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\nfrom sklearn.metrics import f1_score, roc_auc_score, roc_curve, confusion_matrix\nfrom sklearn.preprocessing import label_binarize\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ndef seed_everything(seed: int):\n    import random, os\n    import numpy as np\n    import torch\n    \n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    \nseed_everything(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:07.420626Z","iopub.execute_input":"2025-10-28T09:28:07.421179Z","iopub.status.idle":"2025-10-28T09:28:31.122887Z","shell.execute_reply.started":"2025-10-28T09:28:07.421148Z","shell.execute_reply":"2025-10-28T09:28:31.122260Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Globals:\n  gcp_project = 'dx-scin-public' #@param\n  gcs_bucket_name = 'dx-scin-public-data' #@param\n  cases_csv = 'dataset/scin_cases.csv' #@param\n  labels_csv = 'dataset/scin_labels.csv' #@param\n  gcs_images_dir = 'dataset/images/' #@param\n\n  ### Key column names\n  image_path_columns = ['image_1_path', 'image_2_path', 'image_3_path']\n  weighted_skin_condition_label = \"weighted_skin_condition_label\"\n  skin_condition_label = \"dermatologist_skin_condition_on_label_name\"\n\n  gcs_storage_client = None\n  gcs_bucket = None\n  cases_df = None\n  cases_and_labels_df = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:31.123881Z","iopub.execute_input":"2025-10-28T09:28:31.124342Z","iopub.status.idle":"2025-10-28T09:28:31.128180Z","shell.execute_reply.started":"2025-10-28T09:28:31.124317Z","shell.execute_reply":"2025-10-28T09:28:31.127458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.environ[\"GOOGLE_APPLICATION_CREDENTIALS\"] = \"/kaggle/input/gcs-json/gen-lang-client-0182740461-c4c482d65f84.json\"\n\ndef list_blobs(storage_client, bucket_name):\n  \"\"\"Helper to list blobs in a bucket (useful for debugging).\"\"\"\n  blobs = storage_client.list_blobs(bucket_name)\n  for blob in blobs:\n    print(blob)\n\ndef initialize_df_with_metadata(bucket, csv_path):\n  \"\"\"Loads the given CSV into a pd.DataFrame.\"\"\"\n  df = pd.read_csv(io.BytesIO(bucket.blob(csv_path).download_as_string()), dtype={'case_id': str})\n  df['case_id'] = df['case_id'].astype(str)\n  return df\n\ndef augment_metadata_with_labels(df, bucket, csv_path):\n  \"\"\"Loads the given CSV into a pd.DataFrame.\"\"\"\n  labels_df = pd.read_csv(io.BytesIO(bucket.blob(csv_path).download_as_string()), dtype={'case_id': str})\n  labels_df['case_id'] = labels_df['case_id'].astype(str)\n  merged_df = pd.merge(df, labels_df, on='case_id')\n  return merged_df\n\nGlobals.gcs_storage_client = storage.Client(Globals.gcp_project)\nGlobals.gcs_bucket = Globals.gcs_storage_client.bucket(\n    Globals.gcs_bucket_name\n)\nGlobals.cases_df = initialize_df_with_metadata(Globals.gcs_bucket, Globals.cases_csv)\nGlobals.cases_and_labels_df = augment_metadata_with_labels(Globals.cases_df, Globals.gcs_bucket, Globals.labels_csv)\nprint(len(Globals.cases_and_labels_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:31.129039Z","iopub.execute_input":"2025-10-28T09:28:31.129569Z","iopub.status.idle":"2025-10-28T09:28:34.863227Z","shell.execute_reply.started":"2025-10-28T09:28:31.129540Z","shell.execute_reply":"2025-10-28T09:28:34.862453Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_top_condition_and_prob(label_str):\n    try:\n        label_dict = ast.literal_eval(label_str)\n        if not isinstance(label_dict, dict) or len(label_dict) == 0:\n            return None, None\n        cond, prob = max(label_dict.items(), key=lambda x: x[1])\n        return cond, prob\n    except Exception:\n        return None, None\n\nGlobals.cases_and_labels_df[[\"top_condition\", \"top_prob\"]] = (\n    Globals.cases_and_labels_df[\"weighted_skin_condition_label\"]\n    .apply(lambda s: pd.Series(get_top_condition_and_prob(s)))\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:34.864207Z","iopub.execute_input":"2025-10-28T09:28:34.864479Z","iopub.status.idle":"2025-10-28T09:28:35.279268Z","shell.execute_reply.started":"2025-10-28T09:28:34.864460Z","shell.execute_reply":"2025-10-28T09:28:35.278679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = Globals.cases_and_labels_df\ndf[\"image_paths\"] = df[Globals.image_path_columns].values.tolist()\ndf[\"weighted_skin_condition_label\"] = df[\"weighted_skin_condition_label\"].apply(\n    lambda x: ast.literal_eval(x) if isinstance(x, str) else x\n)\n# Drop rows with empty dicts\ndf = df[df[\"weighted_skin_condition_label\"].astype(str) != \"{}\"]\n# checking dict type\ndf = df[df[\"weighted_skin_condition_label\"].apply(lambda x: bool(x))]\nprint(df.shape)\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:35.281240Z","iopub.execute_input":"2025-10-28T09:28:35.281473Z","iopub.status.idle":"2025-10-28T09:28:35.380476Z","shell.execute_reply.started":"2025-10-28T09:28:35.281455Z","shell.execute_reply":"2025-10-28T09:28:35.379807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df[\"main_label\"] = df[\"weighted_skin_condition_label\"].apply(lambda d: max(d, key=d.get))\n# Count frequencies\nlabel_counts = df[\"main_label\"].value_counts()\n\n# Labels with less than 10 members\nrare_labels = label_counts[label_counts < 10].index.tolist()\ndf = df[~df[\"main_label\"].isin(rare_labels)].reset_index(drop=True)\ndf = df[[\"case_id\",\n        \"age_group\",\n        \"sex_at_birth\",\n        \"fitzpatrick_skin_type\",\n        \"combined_race\",\n        \"image_paths\",\n        \"main_label\"]]\nlabel_counts = df[\"main_label\"].value_counts()\nprint(f'New label counts: {label_counts}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:35.381306Z","iopub.execute_input":"2025-10-28T09:28:35.381608Z","iopub.status.idle":"2025-10-28T09:28:35.404226Z","shell.execute_reply.started":"2025-10-28T09:28:35.381582Z","shell.execute_reply":"2025-10-28T09:28:35.403629Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_image_lists = list(df[\"image_paths\"])\n\n# Flatten nested lists (in case each row has multiple images)\nall_image_paths = []\nfor item in all_image_lists:\n    if isinstance(item, (list, tuple)):\n        all_image_paths.extend(item)\n    elif isinstance(item, str):\n        all_image_paths.append(item)\n\n# Clean: remove None / nan / invalid\nall_image_paths = [\n    p for p in all_image_paths\n    if isinstance(p, str) and p.lower() != \"none\" and p.strip() != \"\"\n]\n\n# Deduplicate\nall_image_paths = list(set(all_image_paths))\nprint(f\"Total unique image paths: {len(all_image_paths)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:35.404870Z","iopub.execute_input":"2025-10-28T09:28:35.405085Z","iopub.status.idle":"2025-10-28T09:28:35.414998Z","shell.execute_reply.started":"2025-10-28T09:28:35.405069Z","shell.execute_reply":"2025-10-28T09:28:35.414218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # CONFIG\n# bucket_name = \"dx-scin-public-data\"\nlocal_root = \"./data_cache\" \n# os.makedirs(local_root, exist_ok=True)\n\n# def download_one(obj_path: str):\n#     \"\"\"Download one file from GCS to local cache.\"\"\"\n#     if obj_path is None or str(obj_path).lower() == \"none\":\n#         return None\n\n#     # Build local path\n#     local_path = os.path.join(local_root, obj_path)\n#     os.makedirs(os.path.dirname(local_path), exist_ok=True)\n\n#     # Skip if already exists\n#     if os.path.exists(local_path):\n#         return local_path\n\n#     # Build GCS public URL\n#     url = f\"https://storage.googleapis.com/download/storage/v1/b/{bucket_name}/o/{obj_path.replace('/', '%2F')}?alt=media\"\n\n#     for attempt in range(3):\n#         try:\n#             r = requests.get(url, timeout=15)\n#             if r.status_code == 200:\n#                 with open(local_path, \"wb\") as f:\n#                     f.write(r.content)\n#                 return local_path\n#             else:\n#                 raise RuntimeError(f\"HTTP {r.status_code}\")\n#         except Exception as e:\n#             if attempt == 2:\n#                 print(f\"[FAIL] {obj_path}: {e}\")\n#     return None\n\n\n# # Parallel download\n# results = []\n# with ThreadPoolExecutor(max_workers=4) as ex:\n#     futures = {ex.submit(download_one, p): p for p in all_image_paths}\n#     for fut in tqdm(as_completed(futures), total=len(futures), desc=\"Downloading\"):\n#         res = fut.result()\n#         if res:\n#             results.append(res)\n\n# print(f\"Downloaded {len(results)} / {len(all_image_paths)} files to {local_root}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:35.415720Z","iopub.execute_input":"2025-10-28T09:28:35.415951Z","iopub.status.idle":"2025-10-28T09:28:35.430973Z","shell.execute_reply.started":"2025-10-28T09:28:35.415936Z","shell.execute_reply":"2025-10-28T09:28:35.430300Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import zipfile\n\nzip_path = \"/kaggle/input/scin-dataset-download/images_backup.zip\"\nextract_path = \"/kaggle/working/data_cache/dataset/images\"\n\nos.makedirs(extract_path, exist_ok=True)\n\nwith zipfile.ZipFile(zip_path, 'r') as zip_ref:\n    zip_ref.extractall(extract_path)\n\nprint(f\"✅ Extracted {zip_path} to {extract_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:28:35.431748Z","iopub.execute_input":"2025-10-28T09:28:35.431976Z","iopub.status.idle":"2025-10-28T09:29:35.828288Z","shell.execute_reply.started":"2025-10-28T09:28:35.431961Z","shell.execute_reply":"2025-10-28T09:29:35.827680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(os.listdir('/kaggle/working/data_cache/dataset/images'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.828947Z","iopub.execute_input":"2025-10-28T09:29:35.829128Z","iopub.status.idle":"2025-10-28T09:29:35.837179Z","shell.execute_reply.started":"2025-10-28T09:29:35.829114Z","shell.execute_reply":"2025-10-28T09:29:35.836449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Split \ntrain_df, test_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df[\"main_label\"],\n    random_state=42\n)\ntrain_df = train_df.reset_index(drop=True)\ntest_df = test_df.reset_index(drop=True)\nprint(f\"Dropped labels with less than samples: {rare_labels}\")\nprint(train_df.shape, test_df.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.837983Z","iopub.execute_input":"2025-10-28T09:29:35.838621Z","iopub.status.idle":"2025-10-28T09:29:35.853760Z","shell.execute_reply.started":"2025-10-28T09:29:35.838601Z","shell.execute_reply":"2025-10-28T09:29:35.852971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CLASSES = sorted(df[\"main_label\"].unique().tolist())\n\n# Create mapping label -> index\nlabel2id = {label: idx for idx, label in enumerate(CLASSES)}\nid2label = {idx: label for label, idx in label2id.items()}\n\n# Combine both splits\nall_image_lists = list(train_df['image_paths']) + list(test_df['image_paths'])\n\n# Flatten nested lists (in case each row has multiple images)\nall_image_paths = []\nfor item in all_image_lists:\n    if isinstance(item, (list, tuple)):\n        all_image_paths.extend(item)\n    elif isinstance(item, str):\n        all_image_paths.append(item)\n\n# Clean: remove None / nan / invalid\nall_image_paths = [\n    p for p in all_image_paths\n    if isinstance(p, str) and p.lower() != \"none\" and p.strip() != \"\"\n]\nprint(len(all_image_paths))\n\nif not isinstance(train_df, pd.DataFrame):\n    train_df = train_df.to_pandas()\nif not isinstance(test_df, pd.DataFrame):\n    test_df = test_df.to_pandas()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.854561Z","iopub.execute_input":"2025-10-28T09:29:35.854822Z","iopub.status.idle":"2025-10-28T09:29:35.871379Z","shell.execute_reply.started":"2025-10-28T09:29:35.854797Z","shell.execute_reply":"2025-10-28T09:29:35.870826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Build a mapping from GCS path → local path\npath_map = {p: os.path.join(local_root, p) for p in all_image_paths}\n\ndef map_local_paths(image_list):\n    return [path_map.get(p, None) for p in image_list if isinstance(p, str)]\n\ntrain_df[\"local_image_paths\"] = train_df[\"image_paths\"].apply(map_local_paths)\ntest_df[\"local_image_paths\"] = test_df[\"image_paths\"].apply(map_local_paths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.872114Z","iopub.execute_input":"2025-10-28T09:29:35.872336Z","iopub.status.idle":"2025-10-28T09:29:35.892614Z","shell.execute_reply.started":"2025-10-28T09:29:35.872312Z","shell.execute_reply":"2025-10-28T09:29:35.891902Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, val_df = train_test_split(\n    train_df,\n    test_size=0.2,\n    stratify=train_df[\"main_label\"],\n    random_state=42\n)\ntrain_df = train_df.reset_index(drop=True)\nval_df = val_df.reset_index(drop=True)\ntrain_df.shape, val_df.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.896016Z","iopub.execute_input":"2025-10-28T09:29:35.896261Z","iopub.status.idle":"2025-10-28T09:29:35.913192Z","shell.execute_reply.started":"2025-10-28T09:29:35.896240Z","shell.execute_reply":"2025-10-28T09:29:35.912596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.914112Z","iopub.execute_input":"2025-10-28T09:29:35.914354Z","iopub.status.idle":"2025-10-28T09:29:35.925003Z","shell.execute_reply.started":"2025-10-28T09:29:35.914328Z","shell.execute_reply":"2025-10-28T09:29:35.924452Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Melanoma data","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/siim-isic-melanoma-classification/train.csv')\nprint('Data distribution')\nprint(train['target'].value_counts())\ntrain.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:35.925729Z","iopub.execute_input":"2025-10-28T09:29:35.925988Z","iopub.status.idle":"2025-10-28T09:29:36.011235Z","shell.execute_reply.started":"2025-10-28T09:29:35.925966Z","shell.execute_reply":"2025-10-28T09:29:36.010627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Step 1. Sample 10% stratified by target ---\nsampled_train, _ = train_test_split(\n    train,\n    test_size=0.9,\n    stratify=train[\"target\"],\n    random_state=42,\n)\n\n# --- Step 2. Split by patient_id to avoid leakage ---\n\n# First, get unique patient IDs and their corresponding label (mode of their samples)\npatient_target = (\n    sampled_train.groupby(\"patient_id\")[\"target\"]\n    .agg(lambda x: x.mode().iloc[0] if not x.mode().empty else 0)\n    .reset_index()\n)\n\n# Split patient_ids stratified by patient-level label\ntrain_patients, temp_patients = train_test_split(\n    patient_target,\n    test_size=0.3,  # 70% train, 30% temp (val+test)\n    stratify=patient_target[\"target\"],\n    random_state=42,\n)\nval_patients, test_patients = train_test_split(\n    temp_patients,\n    test_size=0.5,  # 15% val, 15% test\n    stratify=temp_patients[\"target\"],\n    random_state=42,\n)\n\n# --- Step 3. Filter the rows by patient_id ---\ntrain_df_mel = sampled_train[sampled_train[\"patient_id\"].isin(train_patients[\"patient_id\"])]\nval_df_mel = sampled_train[sampled_train[\"patient_id\"].isin(val_patients[\"patient_id\"])]\ntest_df_mel = sampled_train[sampled_train[\"patient_id\"].isin(test_patients[\"patient_id\"])]\n\n# --- Step 4. Check splits ---\nprint(f\"Total sampled: {len(sampled_train)}\")\nprint(f\"Train: {len(train_df)}, Val: {len(val_df_mel)}, Test: {len(test_df_mel)}\")\nprint(f\"Target ratio (train): {train_df_mel['target'].mean():.4f}\")\nprint(f\"Target ratio (val): {val_df_mel['target'].mean():.4f}\")\nprint(f\"Target ratio (test): {test_df_mel['target'].mean():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:36.011946Z","iopub.execute_input":"2025-10-28T09:29:36.012146Z","iopub.status.idle":"2025-10-28T09:29:36.245225Z","shell.execute_reply.started":"2025-10-28T09:29:36.012130Z","shell.execute_reply":"2025-10-28T09:29:36.244419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"base_image_dir = \"/kaggle/input/siim-isic-melanoma-classification/jpeg/train\"\n\ndef convert_melanoma_df(train_df_mel: pd.DataFrame) -> pd.DataFrame:\n    df = train_df_mel.copy()\n\n    # --- map fields ---\n    df[\"case_id\"] = df[\"patient_id\"]                         # use patient_id as case_id\n    df[\"sex_at_birth\"] = df[\"sex\"].fillna(\"unknown\")\n\n    # Age group binning\n    def age_to_group(age):\n        if pd.isna(age):\n            return \"unknown\"\n        age = float(age)\n        if age < 20:\n            return \"0-19\"\n        elif age < 40:\n            return \"20-39\"\n        elif age < 60:\n            return \"40-59\"\n        elif age < 80:\n            return \"60-79\"\n        else:\n            return \"80+\"\n\n    df[\"age_group\"] = df[\"age_approx\"].apply(age_to_group)\n\n    # Fitzpatrick and race — not available\n    df[\"fitzpatrick_skin_type\"] = None\n    df[\"combined_race\"] = None\n\n    # Image paths\n    import os\nimport pandas as pd\n\nbase_image_dir = \"/kaggle/input/siim-isic-melanoma-classification/jpeg/train\"\n\ndef convert_melanoma_df(train_df_mel: pd.DataFrame) -> pd.DataFrame:\n    df = train_df_mel.copy()\n\n    # --- map fields ---\n    df[\"case_id\"] = df[\"patient_id\"]\n    df[\"sex_at_birth\"] = df[\"sex\"].fillna(\"unknown\")\n\n    # --- Age group binning ---\n    def age_to_group(age):\n        if pd.isna(age):\n            return \"unknown\"\n        age = float(age)\n        if age < 20:\n            return \"0-19\"\n        elif age < 40:\n            return \"20-39\"\n        elif age < 60:\n            return \"40-59\"\n        elif age < 80:\n            return \"60-79\"\n        else:\n            return \"80+\"\n\n    df[\"age_group\"] = df[\"age_approx\"].apply(age_to_group)\n\n    # --- Fields not available ---\n    df[\"fitzpatrick_skin_type\"] = None\n    df[\"combined_race\"] = None\n\n    # --- image paths as STRINGIFIED LIST ---\n    df[\"image_paths\"] = df[\"image_name\"].apply(\n        lambda x: str([os.path.join(\"jpeg/train\", f\"{x}.jpg\")])\n    )\n    df[\"local_image_paths\"] = df[\"image_name\"].apply(\n        lambda x: str([os.path.join(base_image_dir, f\"{x}.jpg\")])\n    )\n\n    # --- main_label: use target column ---\n    df[\"main_label\"] = df[\"target\"].map({0: \"benign\", 1: \"cancerous\"}).fillna(\"unknown\")\n\n    # --- select final columns ---\n    final_df = df[\n        [\n            \"case_id\",\n            \"age_group\",\n            \"sex_at_birth\",\n            \"fitzpatrick_skin_type\",\n            \"combined_race\",\n            \"image_paths\",\n            \"main_label\",\n            \"local_image_paths\",\n        ]\n    ].reset_index(drop=True)\n\n    return final_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:36.246096Z","iopub.execute_input":"2025-10-28T09:29:36.246365Z","iopub.status.idle":"2025-10-28T09:29:36.255347Z","shell.execute_reply.started":"2025-10-28T09:29:36.246339Z","shell.execute_reply":"2025-10-28T09:29:36.254786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df_ready = convert_melanoma_df(train_df_mel)\nval_df_ready = convert_melanoma_df(val_df_mel)\ntest_df_ready = convert_melanoma_df(test_df_mel)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:29:36.256032Z","iopub.execute_input":"2025-10-28T09:29:36.256268Z","iopub.status.idle":"2025-10-28T09:29:36.298129Z","shell.execute_reply.started":"2025-10-28T09:29:36.256251Z","shell.execute_reply":"2025-10-28T09:29:36.297503Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# extra data","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n# --- Root directory of your deduplicated images ---\nextra_dir = \"/kaggle/input/dedup-extra-skin-images\"\n\n# --- Collect all image paths recursively ---\nimg_exts = {\".jpg\", \".jpeg\", \".png\", \".bmp\", \".webp\"}\nimage_paths = [str(p) for p in Path(extra_dir).rglob(\"*\") if p.suffix.lower() in img_exts]\n\nprint(f\"Found {len(image_paths)} images\")\n\n# --- Extract labels (main_label = parent folder name) ---\nrecords = []\nfor i, img_path in enumerate(image_paths):\n    path = Path(img_path)\n    main_label = path.parent.name  # e.g. 'Abscess'\n    records.append({\n        \"case_id\": f\"extra_{i:05d}\",\n        \"age_group\": None,\n        \"sex_at_birth\": None,\n        \"fitzpatrick_skin_type\": None,\n        \"combined_race\": None,\n        \"image_paths\": str(path),\n        \"main_label\": main_label,\n        \"local_image_paths\": str(path.relative_to(extra_dir))\n    })\n\nextra_df = pd.DataFrame(records)[[\n    \"case_id\",\n    \"age_group\",\n    \"sex_at_birth\",\n    \"fitzpatrick_skin_type\",\n    \"combined_race\",\n    \"image_paths\",\n    \"main_label\",\n    \"local_image_paths\"\n]]\n\nprint(\"extra_df created:\", extra_df.shape)\nextra_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:30:08.879429Z","iopub.execute_input":"2025-10-28T09:30:08.879709Z","iopub.status.idle":"2025-10-28T09:30:08.931053Z","shell.execute_reply.started":"2025-10-28T09:30:08.879688Z","shell.execute_reply":"2025-10-28T09:30:08.930502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"extra_df['main_label'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:30:55.576453Z","iopub.execute_input":"2025-10-28T09:30:55.576739Z","iopub.status.idle":"2025-10-28T09:30:55.583155Z","shell.execute_reply.started":"2025-10-28T09:30:55.576717Z","shell.execute_reply":"2025-10-28T09:30:55.582381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef stratified_split_with_min(df, label_col=\"main_label\", test_size=0.1, val_size=0.1, random_state=42):\n    \"\"\"\n    Stratified split ensuring that even classes with small sample counts\n    get at least 1 in val/test when possible.\n    \"\"\"\n    train_parts, val_parts, test_parts = [], [], []\n\n    for label, group in df.groupby(label_col):\n        n = len(group)\n        if n < 3:\n            # Too small, assign all to train\n            train_parts.append(group)\n            continue\n\n        # --- compute per-class test/val sizes ---\n        n_test = max(1, int(round(test_size * n)))\n        n_val = max(1, int(round(val_size * n)))\n        n_train = max(0, n - n_test - n_val)\n\n        # --- sample without replacement ---\n        group = group.sample(frac=1, random_state=random_state)\n        test = group.iloc[:n_test]\n        val = group.iloc[n_test:n_test + n_val]\n        train = group.iloc[n_test + n_val:]\n\n        test_parts.append(test)\n        val_parts.append(val)\n        train_parts.append(train)\n\n    train_df = pd.concat(train_parts).reset_index(drop=True)\n    val_df = pd.concat(val_parts).reset_index(drop=True)\n    test_df = pd.concat(test_parts).reset_index(drop=True)\n    return train_df, val_df, test_df\n\n\n# --- Run the function ---\ntrain_extra, val_extra, test_extra = stratified_split_with_min(\n    extra_df,\n    label_col=\"main_label\",\n    test_size=0.1,\n    val_size=0.1,\n    random_state=42\n)\n\nprint(f\"Extra split: train={len(train_extra)}, val={len(val_extra)}, test={len(test_extra)}\")\nprint(\"Label balance check:\")\nprint(pd.Series({\n    \"train\": train_extra[\"main_label\"].nunique(),\n    \"val\": val_extra[\"main_label\"].nunique(),\n    \"test\": test_extra[\"main_label\"].nunique(),\n}))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:02.278127Z","iopub.execute_input":"2025-10-28T09:33:02.278491Z","iopub.status.idle":"2025-10-28T09:33:02.300539Z","shell.execute_reply.started":"2025-10-28T09:33:02.278443Z","shell.execute_reply":"2025-10-28T09:33:02.299775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- concat each split ---\ntrain_df_combined = pd.concat([train_df, train_df_ready, train_extra], ignore_index=True)\nval_df_combined   = pd.concat([val_df, val_df_ready, val_extra], ignore_index=True)\ntest_df_combined  = pd.concat([test_df, test_df_ready, test_extra], ignore_index=True)\n\n# --- sanity check ---\nfor name, df in [(\"train\", train_df_combined), (\"val\", val_df_combined), (\"test\", test_df_combined)]:\n    print(f\"{name}: {len(df)} samples, {df['main_label'].nunique()} unique labels\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.181535Z","iopub.execute_input":"2025-10-28T09:33:27.181806Z","iopub.status.idle":"2025-10-28T09:33:27.190568Z","shell.execute_reply.started":"2025-10-28T09:33:27.181785Z","shell.execute_reply":"2025-10-28T09:33:27.189939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df_combined['main_label'].value_counts()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_df_combined['main_label'].value_counts()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df_combined['main_label'].value_counts()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.Resize(384, 384),\n\n    # --- Geometric augmentations ---\n    A.OneOf([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.3),\n        A.Affine( \n            scale=(0.9, 1.1),\n            translate_percent=(0.05, 0.05),\n            rotate=(-30, 30),\n            shear=(-10, 10),\n            p=0.5\n        ),\n    ], p=0.8),\n\n    # --- Color / brightness ---\n    A.OneOf([\n        A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=1.0),\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=1.0),\n        A.HueSaturationValue(hue_shift_limit=10, sat_shift_limit=15, val_shift_limit=10, p=1.0),\n    ], p=0.5),\n\n    # --- Blur / noise ---\n    A.OneOf([\n        A.GaussianBlur(blur_limit=(1, 5), p=0.5),\n        A.MotionBlur(blur_limit=(2, 7), p=0.5),\n        A.MedianBlur(blur_limit=3, p=0.5),\n        A.ISONoise(color_shift=(0.01, 0.05), intensity=(0.1, 0.5), p=0.5), \n    ], p=0.5),\n\n    # --- Distortion ---\n    A.OneOf([\n        A.OpticalDistortion(distort_limit=0.05, p=0.5),\n        A.GridDistortion(num_steps=5, distort_limit=0.1, p=0.5),\n        A.ElasticTransform(alpha=1.0, sigma=50.0, p=0.5),\n    ], p=0.5),\n\n    # --- Dropout / cutout ---\n    A.CoarseDropout(\n        num_holes_range=(5, 12),                   \n        hole_height_range=(0.05, 0.15),              # each hole 5–15% of height (~20–60 px)\n        hole_width_range=(0.05, 0.15),               # each hole 5–15% of width\n        fill=(0, 0, 0),                             \n        fill_mask=None,                             \n        p=0.3\n    ),\n\n    # --- Normalize ---\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2(),\n])\n\nval_transform = A.Compose([\n    A.Resize(384, 384),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2(),\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.191699Z","iopub.execute_input":"2025-10-28T09:33:27.191953Z","iopub.status.idle":"2025-10-28T09:33:27.218336Z","shell.execute_reply.started":"2025-10-28T09:33:27.191937Z","shell.execute_reply":"2025-10-28T09:33:27.217689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rand_bbox(size, lam):\n    \"\"\"Generate random bounding box.\"\"\"\n    W = size[2]\n    H = size[3]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = int(W * cut_rat)\n    cut_h = int(H * cut_rat)\n\n    # uniform center\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    x1 = np.clip(cx - cut_w // 2, 0, W)\n    y1 = np.clip(cy - cut_h // 2, 0, H)\n    x2 = np.clip(cx + cut_w // 2, 0, W)\n    y2 = np.clip(cy + cut_h // 2, 0, H)\n\n    return x1, y1, x2, y2\n\n\ndef cutmix(images, labels, alpha=1.0):\n    \"\"\"Apply CutMix to a batch of images and labels.\"\"\"\n    if alpha <= 0:\n        return images, labels, labels, 1.0  # skip CutMix\n    \n    lam = np.random.beta(alpha, alpha)\n    batch_size = images.size(0)\n    index = torch.randperm(batch_size)\n\n    shuffled_images = images[index]\n    shuffled_labels = labels[index]\n\n    x1, y1, x2, y2 = rand_bbox(images.size(), lam)\n    images[:, :, y1:y2, x1:x2] = shuffled_images[:, :, y1:y2, x1:x2]\n    \n    # adjust lambda to exactly match pixel ratio\n    lam = 1 - ((x2 - x1) * (y2 - y1) / (images.size(-1) * images.size(-2)))\n    \n    return images, labels, shuffled_labels, lam","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.219032Z","iopub.execute_input":"2025-10-28T09:33:27.219235Z","iopub.status.idle":"2025-10-28T09:33:27.226043Z","shell.execute_reply.started":"2025-10-28T09:33:27.219218Z","shell.execute_reply":"2025-10-28T09:33:27.225237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, df, local_root=\"./data_cache\",\n                 melanoma_root=\"/kaggle/input/siim-isic-melanoma-classification/jpeg/train\",\n                 transform=None):\n        self.df = df.reset_index(drop=True)\n        self.local_root = local_root\n        self.melanoma_root = melanoma_root\n        self.transform = transform\n        self.labels = sorted(df[\"main_label\"].unique())\n        self.label2idx = {l: i for i, l in enumerate(self.labels)}\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[int(idx)]\n        label = self.label2idx[str(row[\"main_label\"])]\n\n        # Parse image paths\n        image_paths = row[\"image_paths\"]\n        if isinstance(image_paths, float) or image_paths is None:\n            image_paths = []\n        elif isinstance(image_paths, str):\n            try:\n                image_paths = ast.literal_eval(image_paths)\n            except Exception:\n                image_paths = [image_paths]\n        elif not isinstance(image_paths, list):\n            image_paths = [str(image_paths)]\n\n        # Filter invalid\n        image_paths = [\n            str(p) for p in image_paths\n            if isinstance(p, (str, bytes)) and p != \"nan\" and p.strip() != \"\"\n        ]\n\n        img_list = []\n        for image_path in image_paths:\n            # Try primary dataset root first\n            local_path = os.path.join(self.local_root, image_path)\n            # If missing, try melanoma dataset root\n            if not os.path.exists(local_path):\n                # Only keep filename if relative path includes extra dirs\n                filename = os.path.basename(image_path)\n                local_path = os.path.join(self.melanoma_root, filename)\n\n            # Try load\n            if os.path.exists(local_path):\n                try:\n                    img = np.array(Image.open(local_path).convert(\"RGB\"))\n                    img_list.append(img)\n                except Exception as e:\n                    print(f\"[WARN] Error opening {local_path}: {e}\")\n\n        # If no valid image found\n        if len(img_list) == 0:\n            raise FileNotFoundError(f\"No valid images found for index {idx} (paths: {image_paths})\")\n\n        # Pick one randomly\n        img = random.choice(img_list)\n\n        if self.transform:\n            img = self.transform(image=img)[\"image\"]\n\n        return img, torch.tensor(label, dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.227349Z","iopub.execute_input":"2025-10-28T09:33:27.227611Z","iopub.status.idle":"2025-10-28T09:33:27.248813Z","shell.execute_reply.started":"2025-10-28T09:33:27.227596Z","shell.execute_reply":"2025-10-28T09:33:27.248200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = ImageDataset(train_df_combined, transform = train_transform)\nval_dataset = ImageDataset(val_df_combined, transform = val_transform)\ntest_dataset = ImageDataset(test_df_combined, transform = val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=4)\ntest_loader = DataLoader(test_dataset, batch_size=8, shuffle=False, num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.249530Z","iopub.execute_input":"2025-10-28T09:33:27.250190Z","iopub.status.idle":"2025-10-28T09:33:27.268116Z","shell.execute_reply.started":"2025-10-28T09:33:27.250168Z","shell.execute_reply":"2025-10-28T09:33:27.267449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n        \n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n        \n    def __repr__(self):\n        return self.__class__.__name__ + \\\n                '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + \\\n                ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.268901Z","iopub.execute_input":"2025-10-28T09:33:27.269650Z","iopub.status.idle":"2025-10-28T09:33:27.279012Z","shell.execute_reply.started":"2025-10-28T09:33:27.269632Z","shell.execute_reply":"2025-10-28T09:33:27.278439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EfficientNetClassifier(nn.Module):\n    \"\"\"\n    EfficientNet wrapper for classification with optional GeM pooling.\n    \"\"\"\n\n    def __init__(\n        self,\n        num_classes: int,\n        pretrained: bool = True,\n        model_name: str = \"efficientnet_b4\",\n        use_gem: bool = True,   \n    ):\n        super().__init__()\n        self.num_classes = num_classes\n        self.model_name = model_name\n        self.use_gem = use_gem\n\n        # Load pretrained backbone with no head\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            num_classes=0,         \n            global_pool=''           \n        )\n\n        # Optional: replace pooling with GeM\n        if use_gem:\n            self.pooling = GeM()\n        else:\n            self.pooling = nn.AdaptiveAvgPool2d(1)\n\n        # Classifier head\n        feat_dim = self.backbone.num_features\n        self.head = nn.Linear(feat_dim, num_classes)\n\n    def forward(self, x: torch.Tensor, return_features: bool = False):\n        # Extract features before pooling\n        if hasattr(self.backbone, \"forward_features\"):\n            feat_map = self.backbone.forward_features(x)\n        else:\n            feat_map = self.backbone(x)\n\n        # Apply pooling (GeM or Avg)\n        pooled = self.pooling(feat_map).flatten(1)\n\n        # Classify\n        logits = self.head(pooled)\n\n        if return_features:\n            return logits, pooled\n        return logits\n\n    def param_groups(self, lr_backbone: float, lr_head: float):\n        return [\n            {\"params\": self.backbone.parameters(), \"lr\": lr_backbone},\n            {\"params\": self.head.parameters(),     \"lr\": lr_head},\n        ]\n\nclass EfficientNetB7Classifier(nn.Module):\n    \"\"\"\n    EfficientNet-B7 wrapper for classification with optional GeM pooling.\n    \"\"\"\n\n    def __init__(\n        self,\n        num_classes: int,\n        pretrained: bool = True,\n        use_gem: bool = True,\n    ):\n        super().__init__()\n        self.num_classes = num_classes\n        self.model_name = \"tf_efficientnet_b7.aa_in1k\"\n        self.use_gem = use_gem\n\n        # Load pretrained backbone without classification head\n        self.backbone = timm.create_model(\n            self.model_name,\n            pretrained=pretrained,\n            num_classes=0,\n            global_pool=\"\"\n        )\n\n        # Optional: GeM or Average Pooling\n        self.pooling = GeM() if use_gem else nn.AdaptiveAvgPool2d(1)\n\n        # Classification head\n        feat_dim = self.backbone.num_features\n        self.head = nn.Linear(feat_dim, num_classes)\n\n    def forward(self, x: torch.Tensor, return_features: bool = False):\n        feat_map = (\n            self.backbone.forward_features(x)\n            if hasattr(self.backbone, \"forward_features\")\n            else self.backbone(x)\n        )\n\n        pooled = self.pooling(feat_map).flatten(1)\n        logits = self.head(pooled)\n\n        if return_features:\n            return logits, pooled\n        return logits\n\n    def param_groups(self, lr_backbone: float, lr_head: float):\n        return [\n            {\"params\": self.backbone.parameters(), \"lr\": lr_backbone},\n            {\"params\": self.head.parameters(),     \"lr\": lr_head},\n        ]\n\nclass EvaModel(nn.Module):\n    def __init__(self,\n                 model_name=\"eva02_small_patch14_336.mim_in22k_ft_in1k\",\n                 num_classes=49,\n                 pretrained=True,\n                 checkpoint_path=None):\n        super(EvaModel, self).__init__()\n        self.num_classes = num_classes\n        self.model_name = model_name\n\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            checkpoint_path=checkpoint_path,\n            num_classes=0  # remove classifier, returns 2D feature vector\n        )\n\n        # Classifier head\n        feat_dim = self.backbone.num_features\n        self.head = nn.Linear(feat_dim, num_classes)\n\n    def forward(self, x: torch.Tensor, return_features: bool = False):\n        features = self.backbone(x)\n\n        # Classify\n        logits = self.head(features)\n\n        if return_features:\n            return logits, features\n        return logits\n\n    def param_groups(self, lr_backbone: float, lr_head: float):\n        \"\"\"\n        Helper function to set different learning rates for backbone and head.\n        \"\"\"\n        return [\n            {\"params\": self.backbone.parameters(), \"lr\": lr_backbone},\n            {\"params\": self.head.parameters(),     \"lr\": lr_head},\n        ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.291179Z","iopub.execute_input":"2025-10-28T09:33:27.291423Z","iopub.status.idle":"2025-10-28T09:33:27.303791Z","shell.execute_reply.started":"2025-10-28T09:33:27.291380Z","shell.execute_reply":"2025-10-28T09:33:27.303214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nnum_classes = len(train_dataset.labels)   \nmodel = EfficientNetClassifier(num_classes=num_classes, pretrained=True, use_gem=False)\n# model = EvaModel(num_classes=num_classes, pretrained=True)\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:27.304895Z","iopub.execute_input":"2025-10-28T09:33:27.305118Z","iopub.status.idle":"2025-10-28T09:33:31.376287Z","shell.execute_reply.started":"2025-10-28T09:33:27.305094Z","shell.execute_reply":"2025-10-28T09:33:31.375668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tpr_weighted_multiclass_auc(\n    y_true,\n    y_pred,\n    tpr_mode=\"mean\",\n    fpr_target=0.1,\n    custom_tpr_fn=None,\n    tpr_threshold=0.8,\n    require_all=True,  \n):\n    \"\"\"\n    Compute a TPR-weighted multiclass AUC.\n\n    If require_all is True (default) then the function returns (np.nan, aucs, tprs, {})\n    whenever any class has TPR <= tpr_threshold or has undefined AUC.\n    If require_all is False the function falls back to the previous behavior:\n    average only over classes with TPR > tpr_threshold.\n    \"\"\"\n    classes = np.unique(y_true)\n    y_true_bin = label_binarize(y_true, classes=classes)\n\n    aucs, tprs = {}, {}\n\n    for i, cls in enumerate(classes):\n        y_true_cls = y_true_bin[:, i]\n        y_score_cls = y_pred[:, i]\n\n        # not enough labels to compute ROC for this class\n        if len(np.unique(y_true_cls)) < 2:\n            aucs[cls], tprs[cls] = np.nan, 0.0\n            continue\n\n        aucs[cls] = roc_auc_score(y_true_cls, y_score_cls)\n        fpr, tpr, _ = roc_curve(y_true_cls, y_score_cls)\n\n        if callable(custom_tpr_fn):\n            tpr_val = custom_tpr_fn(fpr, tpr)\n        elif tpr_mode == \"mean\":\n            tpr_val = np.mean(tpr)\n        elif tpr_mode == \"at_fpr\":\n            tpr_val = np.interp(fpr_target, fpr, tpr)\n        elif tpr_mode == \"max\":\n            tpr_val = np.max(tpr)\n        elif tpr_mode == \"youden\":\n            idx = np.argmax(tpr - fpr)\n            tpr_val = tpr[idx]\n        elif tpr_mode == \"auc\":\n            tpr_val = np.trapz(tpr, fpr)\n        else:\n            raise ValueError(f\"Unknown tpr_mode: {tpr_mode}\")\n\n        tprs[cls] = float(tpr_val)\n\n    # --- Enforce \"every class must be > threshold\" when require_all=True ---\n    if require_all:\n        failed = [cls for cls in classes if (np.isnan(aucs[cls]) or tprs[cls] <= tpr_threshold)]\n        if len(failed) > 0:\n            # you can also return a tuple with more information if you prefer\n            return np.nan, aucs, tprs, {}\n\n        filtered_classes = list(classes)  # all classes passed\n    else:\n        # previous behavior: include only classes with tpr > threshold and valid auc\n        filtered_classes = [cls for cls in classes if (not np.isnan(aucs[cls]) and tprs[cls] > tpr_threshold)]\n        if len(filtered_classes) == 0:\n            return np.nan, aucs, tprs, {}\n\n    total_tpr = sum(tprs[cls] for cls in filtered_classes)\n    # total_tpr must be > 0 here (because require_all ensured tprs > threshold > 0 or filtered list non-empty)\n    weights = {cls: (tprs[cls] / total_tpr) for cls in filtered_classes}\n    weighted_auc = np.nansum([aucs[cls] * weights[cls] for cls in filtered_classes])\n\n    return weighted_auc, aucs, tprs, weights\n\n# ---------------- TRAINING ----------------\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.param_groups(lr_backbone=1e-3, lr_head=1e-3), eps=1e-8)\nscheduler = CosineAnnealingWarmRestarts(optimizer, T_0=50, T_mult=1, eta_min=1e-6)\nbest_auc = -np.inf\nbest_f1 = -np.inf\nbest_path = \"best_model.pth\"\nnum_epochs = 100\npatience = 7\npatience_counter = 0\n\nfor epoch in range(num_epochs):\n    model.train()\n    total_loss = 0.0\n\n    for imgs, labels in tqdm(train_loader):\n        imgs, labels = imgs.cuda(), labels.cuda()\n        # --- 30% chance to apply CutMix ---\n        apply_cutmix = np.random.rand() < 0.3\n        if apply_cutmix:\n            imgs, labels_a, labels_b, lam = cutmix(imgs, labels, alpha=1.0)\n            outputs = model(imgs)\n            loss = lam * criterion(outputs, labels_a) + (1 - lam) * criterion(outputs, labels_b)\n        else:\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n\n    scheduler.step()\n    avg_train_loss = total_loss / len(train_loader)\n    print(f\"Epoch [{epoch+1}/{num_epochs}] | Train Loss: {avg_train_loss:.4f}\")\n\n    # ---------------- VALIDATION ----------------\n    model.eval()\n    preds_all, labels_all, probs_all = [], [], []\n    with torch.no_grad():\n        for imgs, labels in val_loader:\n            imgs, labels = imgs.cuda(), labels.cuda()\n            outputs = model(imgs)\n            probs = torch.softmax(outputs, dim=1)\n            preds_all.extend(probs.argmax(dim=1).cpu().numpy())\n            labels_all.extend(labels.cpu().numpy())\n            probs_all.extend(probs.cpu().numpy())\n\n    labels_all = np.array(labels_all)\n    probs_all = np.array(probs_all)\n    preds_all = np.array(preds_all)\n\n    f1 = f1_score(labels_all, preds_all, average=\"macro\")\n    weighted_auc, aucs, tprs, weights = tpr_weighted_multiclass_auc(\n        labels_all,\n        probs_all,\n        tpr_mode=\"at_fpr\",  # or \"youden\" / \"mean\"\n        fpr_target=0.1,\n        tpr_threshold=0.8,\n    )\n\n    print(f\"Epoch {epoch+1}: \"\n          f\"train_loss={avg_train_loss:.4f}, \"\n          f\"F1={f1:.4f}, \"\n          f\"WeightedAUC@TPR>0.8={weighted_auc:.4f}\")\n\n    # Save best\n    improved = False\n\n    if not np.isnan(weighted_auc) and weighted_auc > best_auc:\n        best_auc = weighted_auc\n        improved = True\n    \n    if f1 > best_f1:\n        best_f1 = f1\n        improved = True\n    \n    if improved:\n        patience_counter = 0\n        torch.save(model.state_dict(), best_path)\n        print(f\"✅ Saved new best model (AUC={best_auc:.4f}, F1={best_f1:.4f})\")\n    else:\n        patience_counter += 1\n        print(f\"No improvement. Patience counter: {patience_counter}/{patience}\")\n    \n        if patience_counter >= patience:\n            print(\"⏹️ Early stopping triggered.\")\n            break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T09:33:31.376992Z","iopub.execute_input":"2025-10-28T09:33:31.377181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, matthews_corrcoef\n\n# ---------------- TEST EVALUATION ----------------\nprint(\"\\nLoading best model and evaluating on test set...\")\nmodel.load_state_dict(torch.load(best_path))\nmodel.eval()\n\npreds_all, labels_all, probs_all = [], [], []\nwith torch.no_grad():\n    for imgs, labels in test_loader:\n        imgs, labels = imgs.cuda(), labels.cuda()\n        outputs = model(imgs)\n        probs = torch.softmax(outputs, dim=1)\n        preds_all.extend(probs.argmax(dim=1).cpu().numpy())\n        labels_all.extend(labels.cpu().numpy())\n        probs_all.extend(probs.cpu().numpy())\n\nlabels_all = np.array(labels_all)\nprobs_all = np.array(probs_all)\npreds_all = np.array(preds_all)\n\nf1_test = f1_score(labels_all, preds_all, average=\"macro\")\nweighted_auc_test, auc_test, tpr_test, weights = tpr_weighted_multiclass_auc(\n    labels_all, probs_all,\n    tpr_mode=\"at_fpr\", fpr_target=0.1,\n    tpr_threshold=0.8,\n)\nprint(f\"\\n[Test Results] F1={f1_test:.4f}, WeightedAUC@TPR>0.8={weighted_auc_test:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T10:31:25.325229Z","iopub.execute_input":"2025-10-28T10:31:25.325542Z","iopub.status.idle":"2025-10-28T10:32:12.471098Z","shell.execute_reply.started":"2025-10-28T10:31:25.325518Z","shell.execute_reply":"2025-10-28T10:32:12.470231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ---------------- CONFUSION MATRIX ----------------\ncm = confusion_matrix(labels_all, preds_all)\n\n# Retrieve readable label names from dataset\nclass_names = test_loader.dataset.labels  # <-- from your ImageDataset definition\n\nplt.figure(figsize=(12, 10))\nsns.heatmap(\n    cm,\n    annot=True,\n    fmt=\"d\",\n    cmap=\"Blues\",\n    xticklabels=class_names,\n    yticklabels=class_names,\n    cbar=False,\n    square=True,\n    linewidths=0.5,\n    linecolor=\"gray\"\n)\nplt.xlabel(\"Predicted Label\", fontsize=12)\nplt.ylabel(\"True Label\", fontsize=12)\nplt.title(\"Confusion Matrix\", fontsize=14, fontweight=\"bold\")\nplt.xticks(rotation=45, ha=\"right\")\nplt.yticks(rotation=0)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T10:35:49.578775Z","iopub.execute_input":"2025-10-28T10:35:49.579285Z","iopub.status.idle":"2025-10-28T10:35:53.317522Z","shell.execute_reply.started":"2025-10-28T10:35:49.579255Z","shell.execute_reply":"2025-10-28T10:35:53.316737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import Dict, List\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_recall_fscore_support,\n    confusion_matrix,\n    roc_auc_score,\n    average_precision_score,\n    matthews_corrcoef,\n    roc_curve,\n    auc\n)\n\ndef compute_partial_auroc(y_true: np.ndarray, y_score: np.ndarray, min_tpr: float = 0.8) -> float:\n    \"\"\"\n    Compute partial AUROC at minimum TPR threshold\n    \n    Args:\n        y_true: True binary labels\n        y_score: Predicted probabilities\n        min_tpr: Minimum TPR threshold\n    \n    Returns:\n        Partial AUROC value\n    \"\"\"\n    try:\n        fpr, tpr, _ = roc_curve(y_true, y_score)\n        \n        # Find index where TPR >= min_tpr\n        valid_indices = np.where(tpr >= min_tpr)[0]\n        \n        if len(valid_indices) == 0:\n            return 0.0\n        \n        # Calculate partial AUC\n        start_idx = valid_indices[0]\n        partial_fpr = fpr[start_idx:]\n        partial_tpr = tpr[start_idx:]\n        \n        if len(partial_fpr) < 2:\n            return 0.0\n        \n        # Normalize to [0, 1]\n        partial_fpr_norm = (partial_fpr - partial_fpr[0]) / (partial_fpr[-1] - partial_fpr[0] + 1e-10)\n        partial_tpr_norm = (partial_tpr - min_tpr) / (1.0 - min_tpr + 1e-10)\n        \n        return auc(partial_fpr_norm, partial_tpr_norm)\n        \n    except:\n        return 0.0\n\ndef compute_comprehensive_metrics(\n    y_true: np.ndarray,\n    y_pred: np.ndarray,\n    y_prob: np.ndarray,\n    class_names: List[str],\n    tpr_thresholds: List[float] = [0.8, 0.9]\n) -> Dict:\n    \"\"\"\n    Compute ALL metrics for wound classification\n    \n    Args:\n        y_true: Ground truth labels (N,)\n        y_pred: Predicted labels (N,)\n        y_prob: Predicted probabilities (N, num_classes)\n        class_names: List of class names\n        tpr_thresholds: TPR thresholds for partial AUROC\n    \n    Returns:\n        Dictionary with all metrics\n    \"\"\"\n    num_classes = len(class_names)\n    \n    # ==================== OVERALL METRICS ====================\n    accuracy = accuracy_score(y_true, y_pred)\n    mcc = matthews_corrcoef(y_true, y_pred)\n    \n    # Macro-averaged metrics\n    macro_prec, macro_rec, macro_f1, _ = precision_recall_fscore_support(\n        y_true, y_pred, average='macro', zero_division=0\n    )\n    \n    # Weighted-averaged metrics\n    weighted_prec, weighted_rec, weighted_f1, _ = precision_recall_fscore_support(\n        y_true, y_pred, average='weighted', zero_division=0\n    )\n    \n    # ==================== PER-CLASS METRICS ====================\n    per_class_prec, per_class_rec, per_class_f1, per_class_support = precision_recall_fscore_support(\n        y_true, y_pred, average=None, zero_division=0, labels=range(num_classes)\n    )\n    \n    # One-vs-rest for AUROC/AUPRC\n    per_class_auroc = []\n    per_class_auprc = []\n    per_class_partial_aurocs = {tpr: [] for tpr in tpr_thresholds}\n    \n    for i in range(num_classes):\n        try:\n            # Binary labels (class i vs rest)\n            y_true_binary = (y_true == i).astype(int)\n            y_score = y_prob[:, i]\n            \n            # AUROC\n            if len(np.unique(y_true_binary)) > 1:\n                auroc = roc_auc_score(y_true_binary, y_score)\n                auprc = average_precision_score(y_true_binary, y_score)\n            else:\n                auroc = 0.0\n                auprc = 0.0\n            \n            per_class_auroc.append(auroc)\n            per_class_auprc.append(auprc)\n            \n            # Partial AUROC at different TPR thresholds\n            for tpr in tpr_thresholds:\n                pauroc = compute_partial_auroc(y_true_binary, y_score, min_tpr=tpr)\n                per_class_partial_aurocs[tpr].append(pauroc)\n                \n        except Exception as e:\n            per_class_auroc.append(0.0)\n            per_class_auprc.append(0.0)\n            for tpr in tpr_thresholds:\n                per_class_partial_aurocs[tpr].append(0.0)\n    \n    per_class_auroc = np.array(per_class_auroc)\n    per_class_auprc = np.array(per_class_auprc)\n    \n    # Macro-averaged AUROC/AUPRC\n    macro_auroc = np.mean(per_class_auroc)\n    macro_auprc = np.mean(per_class_auprc)\n    \n    # Weighted-averaged AUROC/AUPRC\n    total_support = per_class_support.sum()\n    weighted_auroc = np.sum(per_class_auroc * per_class_support) / total_support\n    weighted_auprc = np.sum(per_class_auprc * per_class_support) / total_support\n    \n    # Macro/weighted partial AUROC\n    macro_partial_aurocs = {}\n    weighted_partial_aurocs = {}\n    \n    for tpr in tpr_thresholds:\n        pauroc_array = np.array(per_class_partial_aurocs[tpr])\n        macro_partial_aurocs[f'pAUROC_TPR{tpr}'] = np.mean(pauroc_array)\n        weighted_partial_aurocs[f'pAUROC_TPR{tpr}'] = np.sum(pauroc_array * per_class_support) / total_support\n    \n    # ==================== CONFUSION MATRIX ====================\n    cm = confusion_matrix(y_true, y_pred, labels=range(num_classes))\n    \n    # ==================== COMPILE ALL METRICS ====================\n    metrics = {\n        # Overall\n        'accuracy': float(accuracy),\n        'mcc': float(mcc),\n        \n        # Macro-averaged\n        'macro_precision': float(macro_prec),\n        'macro_recall': float(macro_rec),\n        'macro_f1': float(macro_f1),\n        'macro_auroc': float(macro_auroc),\n        'macro_auprc': float(macro_auprc),\n        \n        # Weighted-averaged\n        'weighted_precision': float(weighted_prec),\n        'weighted_recall': float(weighted_rec),\n        'weighted_f1': float(weighted_f1),\n        'weighted_auroc': float(weighted_auroc),\n        'weighted_auprc': float(weighted_auprc),\n        \n        # # Per-class (as lists)\n        # 'per_class_precision': per_class_prec.tolist(),\n        # 'per_class_recall': per_class_rec.tolist(),\n        # 'per_class_f1': per_class_f1.tolist(),\n        # 'per_class_support': per_class_support.tolist(),\n        # 'per_class_auroc': per_class_auroc.tolist(),\n        # 'per_class_auprc': per_class_auprc.tolist(),\n        \n        # # Confusion matrix\n        # 'confusion_matrix': cm.tolist(),\n        \n        # # Class names\n        # 'class_names': class_names\n    }\n    \n    # Add partial AUROC metrics\n    for key, value in macro_partial_aurocs.items():\n        metrics[f'macro_{key}'] = float(value)\n    for key, value in weighted_partial_aurocs.items():\n        metrics[f'weighted_{key}'] = float(value)\n    \n    # Add per-class partial AUROC\n    for tpr in tpr_thresholds:\n        metrics[f'per_class_pAUROC_TPR{tpr}'] = per_class_partial_aurocs[tpr]\n    \n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T10:32:16.624790Z","iopub.execute_input":"2025-10-28T10:32:16.625313Z","iopub.status.idle":"2025-10-28T10:32:16.632874Z","shell.execute_reply.started":"2025-10-28T10:32:16.625289Z","shell.execute_reply":"2025-10-28T10:32:16.632178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = sorted(set(test_loader.dataset.labels))\ncompute_comprehensive_metrics(labels_all, preds_all, probs_all, class_names=class_names)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    n_classes = probs_all.shape[1]\n    aucs_per_class = []\n    for c in range(n_classes):\n        y_true_bin = (labels_all == c).astype(int)\n        auc_c = roc_auc_score(y_true_bin, probs_all[:, c])\n        aucs_per_class.append(auc_c)\n\n    class_names = test_loader.dataset.labels\n    aucs_per_class = np.array(aucs_per_class)\n    weights_norm = np.bincount(labels_all) / len(labels_all)\n\n    weighted_avg_auc = np.sum(aucs_per_class * weights_norm)\n    macro_avg_auc = np.mean(aucs_per_class)\n\n    print(\"\\nPer-Class AUCs:\")\n    for name, auc_val in zip(class_names, aucs_per_class):\n        print(f\"  {name:<20} {auc_val:.4f}\")\n\n    print(f\"\\nWeighted Avg AUC: {weighted_avg_auc:.4f}\")\n    print(f\"Macro Avg AUC:    {macro_avg_auc:.4f}\")\n\nexcept Exception as e:\n    print(f\"[WARN] Could not compute per-class AUCs: {e}\")\n\nprint(f\"\\n[Test Results] F1={f1_test:.4f}, WeightedAUC@TPR>0.8={weighted_auc_test:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T10:32:42.532076Z","iopub.execute_input":"2025-10-28T10:32:42.532889Z","iopub.status.idle":"2025-10-28T10:32:42.599216Z","shell.execute_reply.started":"2025-10-28T10:32:42.532860Z","shell.execute_reply":"2025-10-28T10:32:42.598445Z"}},"outputs":[],"execution_count":null}]}