{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042},{"sourceType":"datasetVersion","sourceId":23812,"datasetId":17810,"databundleVersionId":23851},{"sourceType":"datasetVersion","sourceId":1493513,"datasetId":876960,"databundleVersionId":1527499},{"sourceType":"datasetVersion","sourceId":6717213,"datasetId":1317048,"databundleVersionId":6801677},{"sourceType":"datasetVersion","sourceId":18613,"datasetId":5839,"databundleVersionId":18613},{"sourceType":"datasetVersion","sourceId":15374689,"datasetId":9834806,"databundleVersionId":16286847}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Chest X-Ray SSL Feature Extraction Pipeline\n\nA unified pipeline that:\n- Loads and samples a multi-source chest X-ray parquet dataset with demographic metadata\n- **Applies CLAHE pre-processing and aggressive border-cropping** to eliminate scanner bias (shortcut learning fix)\n- Domain-adapts a frozen DINOv2 ViT-S/14 backbone via NT-Xent contrastive (SSL) training\n- Extracts both raw 384-d and SSL-adapted 128-d embeddings in a single forward pass\n- Applies unsupervised severity indexing across pathological classes\n- **Builds fused conditioning vectors (SSL 128-d ⊕ demographic 64-d = 192-d)** ready for conditional cGAN input\n- **Outputs a Parquet file** containing image relative paths, severity labels, demographics, raw + adapted embeddings, and fused conditioning vectors\n\n**Phases covered:** 0 → 13, Phase 11 (scanner-bias audit), V0 → V10, T1 → T6  \n**Shortcut-learning fixes:** CLAHE normalisation, aggressive RandomResizedCrop (scale 0.5–0.85), demographic fallback from original parquet  \n**Primary output:** `cgan_input_full.parquet` — one row per image, all cGAN conditioning columns included","metadata":{}},{"cell_type":"markdown","source":"## Phase 0 — Imports & Environment","metadata":{}},{"cell_type":"code","source":"# ── PHASE 0: IMPORTS & ENVIRONMENT ──────────────────────────────────────────\nimport os\nimport time\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport cv2                          # for CLAHE (scanner-bias fix)\nfrom pathlib import Path\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.cuda.amp import GradScaler, autocast\nimport torchvision.transforms as T\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.cluster import KMeans\nfrom scipy.spatial.distance import euclidean\n\nwarnings.filterwarnings(\"ignore\")\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device : {DEVICE}\")\nprint(f\"PyTorch: {torch.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:45.498854Z","iopub.execute_input":"2026-04-01T06:43:45.499476Z","iopub.status.idle":"2026-04-01T06:43:45.506570Z","shell.execute_reply.started":"2026-04-01T06:43:45.499437Z","shell.execute_reply":"2026-04-01T06:43:45.505729Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 1 — Configuration & Constants","metadata":{}},{"cell_type":"markdown","source":"### 1A — Paths","metadata":{}},{"cell_type":"code","source":"# ── 1A: PATHS ────────────────────────────────────────────────────────────────\nPARQUET_PATH = \"/kaggle/input/datasets/rohanpesuecbtech2023/train-parquet/unchanged_train.parquet\"\nOUTPUT_DIR   = Path(\"/kaggle/working/outputs\")\nDCM_OUT_DIR  = Path(\"/kaggle/working/converted_pngs\")\n\nDATASET_ROOTS = {\n    \"CheXpert\"  : \"/kaggle/input/datasets/willarevalo/chexpert-v10-small/CheXpert-v1.0-small\",\n    \"Pediatric\" : \"/kaggle/input/datasets/paultimothymooney/chest-xray-pneumonia/chest_xray\",\n    \"NIH\"       : \"/kaggle/input/datasets/organizations/nih-chest-xrays/data\",\n    \"COVIDx\"    : \"/kaggle/input/datasets/andyczhao/covidx-cxr2\",\n    \"RSNA\"      : \"/kaggle/input/competitions/rsna-pneumonia-detection-challenge/stage_2_train_images\",\n}\n\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\nDCM_OUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprint(\"Paths configured.\")\nprint(f\"  Output directory  : {OUTPUT_DIR}\")\nprint(f\"  DICOM cache dir   : {DCM_OUT_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:45.507795Z","iopub.execute_input":"2026-04-01T06:43:45.508042Z","iopub.status.idle":"2026-04-01T06:43:45.520075Z","shell.execute_reply.started":"2026-04-01T06:43:45.508021Z","shell.execute_reply":"2026-04-01T06:43:45.519407Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1B — Hyperparameters","metadata":{}},{"cell_type":"code","source":"# ── 1B: HYPERPARAMETERS ──────────────────────────────────────────────────────\nSAMPLES_PER_CLASS = 1_000\nRANDOM_STATE      = 42\nBATCH_SIZE        = 64\nSSL_BATCH_SIZE    = 128\nNUM_WORKERS       = 2\nSSL_EPOCHS        = 30\nSSL_LR            = 1e-3\nSSL_WEIGHT_DECAY  = 1e-4\nDINO_DIM          = 384   # DINOv2 ViT-S/14 CLS-token dimension\nSSL_DIM           = 128   # Projection head output dimension\n\n# ── Scanner-bias / shortcut-learning suppression ──────────────────────────────\n# RandomResizedCrop: scale reduced from (0.7,1.0) → (0.5,0.85) so the model\n# rarely sees the full image border (where hospital markers / padding live).\nSSL_CROP_SCALE    = (0.5, 0.85)\n\n# CLAHE (Contrast Limited Adaptive Histogram Equalisation):\n# Unifies per-scanner contrast before DINOv2 sees the pixel values.\n# clip_limit=2.0  tile_grid=(8,8) are standard chest-X-ray settings.\nCLAHE_ENABLED     = True\nCLAHE_CLIP_LIMIT  = 2.0\nCLAHE_TILE_GRID   = (8, 8)\n\nprint(\"Hyperparameters configured.\")\nprint(f\"  Samples per class : {SAMPLES_PER_CLASS:,}\")\nprint(f\"  SSL epochs        : {SSL_EPOCHS}  |  LR: {SSL_LR}  |  Batch: {SSL_BATCH_SIZE}\")\nprint(f\"  Embedding dims    : DINOv2={DINO_DIM}-d  |  SSL={SSL_DIM}-d\")\nprint(f\"  Shortcut-fix crop : scale={SSL_CROP_SCALE}\")\nprint(f\"  CLAHE             : {'enabled' if CLAHE_ENABLED else 'disabled'}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:45.520900Z","iopub.execute_input":"2026-04-01T06:43:45.521124Z","iopub.status.idle":"2026-04-01T06:43:45.534606Z","shell.execute_reply.started":"2026-04-01T06:43:45.521104Z","shell.execute_reply":"2026-04-01T06:43:45.533874Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1C — Demographic Config","metadata":{}},{"cell_type":"code","source":"# ── 1C: DEMOGRAPHIC CONFIG ────────────────────────────────────────────────────\n# Columns to attempt loading from the parquet; missing ones are silently skipped.\nDEMOGRAPHIC_COLS = [\"age\", \"sex\", \"view_position\", \"ap_pa\", \"patient_id\"]\n\nprint(f\"Demographic columns to probe: {DEMOGRAPHIC_COLS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:45.535519Z","iopub.execute_input":"2026-04-01T06:43:45.535848Z","iopub.status.idle":"2026-04-01T06:43:45.550573Z","shell.execute_reply.started":"2026-04-01T06:43:45.535782Z","shell.execute_reply":"2026-04-01T06:43:45.550000Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 2 — Data","metadata":{}},{"cell_type":"markdown","source":"### 2A — Path Resolution Helpers","metadata":{}},{"cell_type":"code","source":"# ── 2A: PATH RESOLUTION HELPERS ──────────────────────────────────────────────\ndef resolve_path(row):\n    \"\"\"Maps a parquet-relative path to its absolute Kaggle dataset path.\"\"\"\n    root = DATASET_ROOTS.get(row[\"dataset\"])\n    if root is None:\n        return row[\"path\"]\n    p = row[\"path\"]\n    if row[\"dataset\"] == \"RSNA\":\n        return str(Path(root) / (Path(p).stem + \".dcm\"))\n    return str(Path(root) / p)\n\n\ndef make_relative(abs_path):\n    \"\"\"Strips the dataset-root prefix to produce a portable relative path.\"\"\"\n    for root in DATASET_ROOTS.values():\n        if abs_path.startswith(root):\n            return abs_path.replace(root + \"/\", \"\")\n    return abs_path\n\n\ndef check_paths(df):\n    \"\"\"Returns (valid_df, missing_paths) after verifying each file exists on disk.\"\"\"\n    valid_indices, missing = [], []\n    for idx, row in df.iterrows():\n        if Path(row[\"path\"]).exists():\n            valid_indices.append(idx)\n        else:\n            missing.append(row[\"path\"])\n    return df.loc[valid_indices].reset_index(drop=True), missing\n\n\nprint(\"Path helpers defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:45.552257Z","iopub.execute_input":"2026-04-01T06:43:45.552667Z","iopub.status.idle":"2026-04-01T06:43:45.562886Z","shell.execute_reply.started":"2026-04-01T06:43:45.552647Z","shell.execute_reply":"2026-04-01T06:43:45.562270Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2B — Load & Sample","metadata":{}},{"cell_type":"code","source":"# ── 2B: LOAD & SAMPLE ────────────────────────────────────────────────────────\nprint(\"[ 2B ] Loading parquet and applying stratified sampling ...\")\n\ncore_cols = [\"path\", \"disease\", \"dataset\"]\n\n# Detect which demographic columns actually exist in the parquet schema.\n# Using pyarrow schema avoids the empty-column probe issue seen on Kaggle.\nimport pyarrow.parquet as pq\nschema_cols = set(pq.read_schema(PARQUET_PATH).names)\navailable_demo_cols = [c for c in DEMOGRAPHIC_COLS if c in schema_cols]\nload_cols = core_cols + available_demo_cols\n\ndf_raw = pd.read_parquet(PARQUET_PATH, columns=load_cols)\ndf_raw[\"path\"] = df_raw.apply(resolve_path, axis=1)\n\nprint(f\"  Raw parquet rows         : {len(df_raw):,}\")\nprint(f\"  Demographic cols loaded  : {available_demo_cols if available_demo_cols else 'none found'}\")\nprint(f\"\\n  Disease distribution:\\n{df_raw['disease'].value_counts().to_string()}\")\nprint(f\"\\n  Source dataset counts:\\n{df_raw['dataset'].value_counts().to_string()}\")\n\n# Stratified sample — up to SAMPLES_PER_CLASS per disease class\ndf_sample = (\n    df_raw\n    .groupby(\"disease\", group_keys=False)\n    .apply(lambda g: g.sample(n=min(len(g), SAMPLES_PER_CLASS), random_state=RANDOM_STATE))\n    .reset_index(drop=True)\n)\n\nprint(f\"\\n  Sampled rows (pre-verification): {len(df_sample):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:45.563647Z","iopub.execute_input":"2026-04-01T06:43:45.563903Z","iopub.status.idle":"2026-04-01T06:43:47.616738Z","shell.execute_reply.started":"2026-04-01T06:43:45.563876Z","shell.execute_reply":"2026-04-01T06:43:47.616007Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2C — Verify","metadata":{}},{"cell_type":"code","source":"# ── 2C: VERIFY ───────────────────────────────────────────────────────────────\nprint(\"[ 2C ] Verifying sampled paths on disk ...\")\n\ndf_sample, missing_sample = check_paths(df_sample)\n\nprint(f\"  Valid sampled rows : {len(df_sample):,}\")\nif missing_sample:\n    print(f\"  Dropped            : {len(missing_sample)} files not found on disk.\")\nelse:\n    print(\"  All sampled paths verified — no missing files.\")\n\n# Quick per-source sanity check (5 samples each)\nprint(\"\\n  Per-source sanity check (5 random paths each):\")\nfor ds, grp in df_raw.groupby(\"dataset\"):\n    s = grp.sample(n=min(5, len(grp)), random_state=RANDOM_STATE)\n    bad = [p for p in s[\"path\"] if not Path(p).exists()]\n    status = \"OK\" if not bad else f\"FAIL ({len(bad)} missing)\"\n    print(f\"    {ds:<12}  {status}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:47.617577Z","iopub.execute_input":"2026-04-01T06:43:47.617786Z","iopub.status.idle":"2026-04-01T06:43:50.542680Z","shell.execute_reply.started":"2026-04-01T06:43:47.617766Z","shell.execute_reply":"2026-04-01T06:43:50.541990Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 3 — Transforms & Datasets","metadata":{}},{"cell_type":"markdown","source":"### 3A — DICOM Loader","metadata":{}},{"cell_type":"code","source":"# ── 3A: DICOM LOADER + CLAHE PRE-PROCESSOR ──────────────────────────────────\ndef apply_clahe(pil_img: Image.Image) -> Image.Image:\n    \"\"\"\n    Applies CLAHE to a PIL image (any mode).\n    Unifies per-scanner contrast so DINOv2 sees anatomy, not scanner brightness.\n    This is one of the two primary fixes for the shortcut-learning / scanner-bias problem.\n    \"\"\"\n    if not CLAHE_ENABLED:\n        return pil_img\n    clahe = cv2.createCLAHE(clipLimit=CLAHE_CLIP_LIMIT, tileGridSize=CLAHE_TILE_GRID)\n    gray  = np.array(pil_img.convert(\"L\"), dtype=np.uint8)\n    eq    = clahe.apply(gray)\n    return Image.fromarray(eq, mode=\"L\")\n\n\ndef safe_dcm_to_png(dcm_path):\n    \"\"\"Converts a DICOM file to a normalised PNG, caching to DCM_OUT_DIR. Applies CLAHE.\"\"\"\n    out_path = DCM_OUT_DIR / (Path(dcm_path).stem + \".png\")\n    if out_path.exists():\n        return apply_clahe(Image.open(out_path))\n    ds  = pydicom.dcmread(dcm_path)\n    arr = ds.pixel_array.astype(np.float32)\n    arr = (arr - arr.min()) / (arr.max() - arr.min() + 1e-8) * 255.0\n    img = Image.fromarray(arr.astype(np.uint8))\n    img.save(out_path)\n    return apply_clahe(img)\n\n\ndef load_image(path: str) -> Image.Image:\n    \"\"\"Loads any supported image format (DICOM or standard). Applies CLAHE.\"\"\"\n    p = Path(path)\n    if p.suffix.lower() == \".dcm\":\n        return safe_dcm_to_png(str(p))\n    img = Image.open(p)\n    return apply_clahe(img)\n\n\nprint(\"DICOM loader + CLAHE pre-processor defined.\")\nprint(f\"  CLAHE: {'enabled' if CLAHE_ENABLED else 'disabled'} | clip={CLAHE_CLIP_LIMIT} | tile={CLAHE_TILE_GRID}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:50.544404Z","iopub.execute_input":"2026-04-01T06:43:50.544863Z","iopub.status.idle":"2026-04-01T06:43:50.552853Z","shell.execute_reply.started":"2026-04-01T06:43:50.544837Z","shell.execute_reply":"2026-04-01T06:43:50.552031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3B — Clean Transform (Extraction)","metadata":{}},{"cell_type":"code","source":"# ── 3B: CLEAN TRANSFORM (Extraction) ─────────────────────────────────────────\n# Deterministic grayscale pipeline used at inference / feature extraction.\n# CenterCrop(200) trims ~10% of edges where scanner markers and padding sit,\n# then Resize(224) brings it back to the DINOv2 expected input size.\nclean_transform = T.Compose([\n    T.Grayscale(num_output_channels=3),\n    T.CenterCrop(200),                          # ← crop border artefacts (shortcut-fix)\n    T.Resize((224, 224)),\n    T.ToTensor(),\n    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\n\nprint(\"Clean (extraction) transform compiled — CenterCrop(200)→Resize(224) ImageNet-normalised.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:50.553796Z","iopub.execute_input":"2026-04-01T06:43:50.554073Z","iopub.status.idle":"2026-04-01T06:43:50.567900Z","shell.execute_reply.started":"2026-04-01T06:43:50.554051Z","shell.execute_reply":"2026-04-01T06:43:50.567111Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3C — SSL Transform (Contrastive Augmentation)","metadata":{}},{"cell_type":"code","source":"from torchvision.transforms import v2\nimport torch\n\n# Define the image size used by your DINOv2 backbone (usually 224)\nIMG_SIZE = 224 \n\nprint(\"[ Phase 3 ] Building Augmented SSL Transforms (with Bias Fixes)...\")\n\nssl_transform = v2.Compose([\n    # 1. Convert to tensor first so math operations can run\n    v2.ToImage(),\n    v2.ToDtype(torch.uint8, scale=True),\n    \n    # 2. THE BIAS FIXES\n    # Aggressive Crop: Cuts off the outer 20-50% of the image to hide hospital borders\n    v2.RandomResizedCrop(size=(IMG_SIZE, IMG_SIZE), scale=(0.5, 0.8)), \n    # CLAHE/Equalization: Forces 100% of images to have the exact same contrast distribution\n    v2.RandomEqualize(p=1.0), \n    \n    # 3. Standard SSL Spatial Augmentations\n    v2.RandomHorizontalFlip(p=0.5),\n    v2.RandomRotation(degrees=15),\n    \n    # 4. Standard SSL Color Jitter (Applied after equalization)\n    v2.ColorJitter(brightness=0.2, contrast=0.2),\n    \n    # 5. Final Normalization for DINOv2 (ImageNet standards)\n    v2.ToDtype(torch.float32, scale=True),\n    v2.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nprint(\" SSL Transforms successfully updated with Aggressive Cropping & Equalization.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:50.568772Z","iopub.execute_input":"2026-04-01T06:43:50.568960Z","iopub.status.idle":"2026-04-01T06:43:50.683704Z","shell.execute_reply.started":"2026-04-01T06:43:50.568940Z","shell.execute_reply":"2026-04-01T06:43:50.682966Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3D — Datasets & Collators","metadata":{}},{"cell_type":"code","source":"# ── 3D: DATASETS & COLLATORS ─────────────────────────────────────────────────\nclass CleanDataset(Dataset):\n    \"\"\"\n    Applies the clean extraction transform and carries full metadata\n    (disease, dataset, relative path, any available demographic fields).\n    \"\"\"\n    def __init__(self, df):\n        self.df = df.reset_index(drop=True)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        try:\n            img    = load_image(row[\"path\"]).convert(\"L\")\n            tensor = clean_transform(img)\n        except Exception:\n            return None\n\n        item = {\n            \"image\"   : tensor,\n            \"disease\" : row[\"disease\"],\n            \"dataset\" : row[\"dataset\"],\n            \"path\"    : row[\"path\"],\n        }\n        for col in available_demo_cols:\n            item[col] = str(row[col]) if pd.notna(row.get(col)) else \"unknown\"\n        return item\n\n\nclass SSLDataset(Dataset):\n    \"\"\"Produces two independently augmented RGB views per image for NT-Xent training.\"\"\"\n    def __init__(self, df):\n        self.paths = df[\"path\"].tolist()\n\n    def __len__(self):\n        return len(self.paths)\n\n    def __getitem__(self, idx):\n        try:\n            img = load_image(self.paths[idx]).convert(\"RGB\")\n            return {\"view1\": ssl_transform(img), \"view2\": ssl_transform(img)}\n        except Exception:\n            return None\n\n\ndef clean_collate(batch):\n    batch = [b for b in batch if b is not None]\n    if not batch:\n        return None\n    out = {\n        \"image\"   : torch.stack([b[\"image\"] for b in batch]),\n        \"disease\" : [b[\"disease\"] for b in batch],\n        \"dataset\" : [b[\"dataset\"] for b in batch],\n        \"path\"    : [b[\"path\"]    for b in batch],\n    }\n    for col in available_demo_cols:\n        out[col] = [b.get(col, \"unknown\") for b in batch]\n    return out\n\n\ndef ssl_collate(batch):\n    batch = [b for b in batch if b is not None]\n    if not batch:\n        return None\n    return {\n        \"view1\": torch.stack([b[\"view1\"] for b in batch]),\n        \"view2\": torch.stack([b[\"view2\"] for b in batch]),\n    }\n\n\nprint(\"CleanDataset, SSLDataset, and collation functions defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:50.684597Z","iopub.execute_input":"2026-04-01T06:43:50.684846Z","iopub.status.idle":"2026-04-01T06:43:50.695380Z","shell.execute_reply.started":"2026-04-01T06:43:50.684824Z","shell.execute_reply":"2026-04-01T06:43:50.694617Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3E — DataLoaders","metadata":{}},{"cell_type":"code","source":"# ── 3E: DATALOADERS ──────────────────────────────────────────────────────────\nclean_dataloader = DataLoader(\n    CleanDataset(df_sample),\n    batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, collate_fn=clean_collate, pin_memory=True,\n)\nssl_dataloader = DataLoader(\n    SSLDataset(df_sample),\n    batch_size=SSL_BATCH_SIZE, shuffle=True,\n    num_workers=NUM_WORKERS, collate_fn=ssl_collate, pin_memory=True,\n)\n\nprint(f\"Clean DataLoader  : {len(clean_dataloader)} batches  (batch size {BATCH_SIZE})\")\nprint(f\"SSL DataLoader    : {len(ssl_dataloader)} batches  (batch size {SSL_BATCH_SIZE})\")\nprint(f\"Demographic fields carried through collation: {available_demo_cols}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:50.696336Z","iopub.execute_input":"2026-04-01T06:43:50.697574Z","iopub.status.idle":"2026-04-01T06:43:50.711806Z","shell.execute_reply.started":"2026-04-01T06:43:50.697551Z","shell.execute_reply":"2026-04-01T06:43:50.711096Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 4 — Unified Encoder","metadata":{}},{"cell_type":"markdown","source":"### 4A — DINOv2 Backbone","metadata":{}},{"cell_type":"code","source":"# ── 4A: DINOV2 BACKBONE ──────────────────────────────────────────────────────\nprint(\"[ 4A ] Loading DINOv2 ViT-S/14 pretrained backbone ...\")\n\n_backbone = torch.hub.load(\n    \"facebookresearch/dinov2\", \"dinov2_vits14\",\n    pretrained=True, verbose=False,\n)\n\nprint(f\"  DINOv2 ViT-S/14 loaded — CLS-token dim: {DINO_DIM}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:50.712694Z","iopub.execute_input":"2026-04-01T06:43:50.713080Z","iopub.status.idle":"2026-04-01T06:43:51.145299Z","shell.execute_reply.started":"2026-04-01T06:43:50.713059Z","shell.execute_reply":"2026-04-01T06:43:51.144588Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4B — UnifiedEncoder Class","metadata":{}},{"cell_type":"code","source":"# ── 4B: UNIFIED ENCODER ──────────────────────────────────────────────────────\nclass UnifiedEncoder(nn.Module):\n    \"\"\"\n    Wraps the frozen DINOv2 backbone and a trainable SSL projection head.\n\n    Public extraction API\n    ---------------------\n    encoder.extract_raw(x)      ->  L2-normalised 384-d DINOv2 embedding\n    encoder.extract_adapted(x)  ->  L2-normalised 128-d SSL-projected embedding\n    encoder(x)                  ->  adapted embedding (default; used during training)\n\n    Only ssl_head parameters are trainable; the backbone is permanently frozen.\n    \"\"\"\n    def __init__(self, backbone, dino_dim: int = DINO_DIM, ssl_dim: int = SSL_DIM):\n        super().__init__()\n        self.backbone = backbone\n\n        for param in self.backbone.parameters():\n            param.requires_grad = False\n\n        self.ssl_head = nn.Sequential(\n            nn.Linear(dino_dim, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(inplace=True),\n            nn.Linear(256, ssl_dim),\n        )\n\n    @torch.no_grad()\n    def _backbone_features(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Extract raw backbone CLS features with gradient tracking disabled.\"\"\"\n        return self.backbone(x)\n\n    def extract_raw(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Returns L2-normalised 384-d DINOv2 embeddings (backbone only).\"\"\"\n        return F.normalize(self._backbone_features(x), dim=-1)\n\n    def extract_adapted(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"Returns L2-normalised 128-d SSL-projected embeddings.\"\"\"\n        projected = self.ssl_head(self._backbone_features(x))\n        return F.normalize(projected, dim=-1)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return self.extract_adapted(x)\n\n\nprint(\"UnifiedEncoder class defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:51.146263Z","iopub.execute_input":"2026-04-01T06:43:51.146562Z","iopub.status.idle":"2026-04-01T06:43:51.154139Z","shell.execute_reply.started":"2026-04-01T06:43:51.146530Z","shell.execute_reply":"2026-04-01T06:43:51.153405Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4C — NT-Xent Loss","metadata":{}},{"cell_type":"code","source":"# ── 4C: NT-XENT LOSS ─────────────────────────────────────────────────────────\ndef nt_xent_loss(z1: torch.Tensor, z2: torch.Tensor, temperature: float = 0.5) -> torch.Tensor:\n    \"\"\"\n    Normalised Temperature-scaled Cross-Entropy Loss (SimCLR formulation).\n\n    Parameters\n    ----------\n    z1, z2      : L2-normalised embeddings of shape (N, D)\n    temperature : softmax temperature (default 0.5)\n\n    Returns\n    -------\n    Scalar loss tensor.\n    \"\"\"\n    N        = z1.size(0)\n    features = torch.cat([z1, z2], dim=0)                       # (2N, D)\n    sim      = torch.matmul(features, features.T) / temperature  # (2N, 2N)\n\n    labels = torch.cat([torch.arange(N)] * 2).to(z1.device)\n    labels = (labels.unsqueeze(0) == labels.unsqueeze(1)).float()\n    mask   = torch.eye(2 * N, dtype=torch.bool, device=z1.device)\n\n    labels = labels[~mask].view(2 * N, -1)\n    sim    = sim[~mask].view(2 * N, -1)\n\n    positives   = sim[labels.bool()].view(2 * N, -1)\n    numerator   = torch.exp(positives)\n    denominator = torch.exp(sim).sum(dim=1, keepdim=True)\n\n    return -torch.log(numerator / denominator).mean()\n\n\nprint(\"NT-Xent contrastive loss defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:51.157275Z","iopub.execute_input":"2026-04-01T06:43:51.157591Z","iopub.status.idle":"2026-04-01T06:43:51.172699Z","shell.execute_reply.started":"2026-04-01T06:43:51.157568Z","shell.execute_reply":"2026-04-01T06:43:51.172044Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 4D — Verify","metadata":{}},{"cell_type":"code","source":"# ── 4D: VERIFY ───────────────────────────────────────────────────────────────\nencoder = UnifiedEncoder(_backbone).to(DEVICE)\n\nwith torch.no_grad():\n    _dummy     = torch.zeros(2, 3, 224, 224, device=DEVICE)\n    _raw_out   = encoder.extract_raw(_dummy)\n    _adapt_out = encoder.extract_adapted(_dummy)\n\nprint(f\"UnifiedEncoder on {DEVICE}\")\nprint(f\"  extract_raw()     -> shape {tuple(_raw_out.shape)}    (DINOv2 384-d)\")\nprint(f\"  extract_adapted() -> shape {tuple(_adapt_out.shape)}   (SSL 128-d)\")\nprint(f\"  Trainable params  : {sum(p.numel() for p in encoder.parameters() if p.requires_grad):,}  (ssl_head only)\")\nprint(f\"  Frozen params     : {sum(p.numel() for p in encoder.parameters() if not p.requires_grad):,}  (DINOv2 backbone)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:51.173585Z","iopub.execute_input":"2026-04-01T06:43:51.173946Z","iopub.status.idle":"2026-04-01T06:43:51.234168Z","shell.execute_reply.started":"2026-04-01T06:43:51.173901Z","shell.execute_reply":"2026-04-01T06:43:51.233598Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 5 — SSL Training","metadata":{}},{"cell_type":"markdown","source":"### 5A — Optimiser Setup","metadata":{}},{"cell_type":"code","source":"# ── 5A: OPTIMISER SETUP ──────────────────────────────────────────────────────\n# Only ssl_head parameters are passed to the optimiser; backbone stays frozen.\noptimizer = optim.AdamW(\n    encoder.ssl_head.parameters(),\n    lr=SSL_LR,\n    weight_decay=SSL_WEIGHT_DECAY,\n)\nscaler = GradScaler()   # mixed-precision gradient scaler\n\nprint(f\"Optimiser : AdamW  |  LR: {SSL_LR}  |  Weight decay: {SSL_WEIGHT_DECAY}\")\nprint(f\"Mixed precision: GradScaler enabled\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:51.234982Z","iopub.execute_input":"2026-04-01T06:43:51.235274Z","iopub.status.idle":"2026-04-01T06:43:51.240346Z","shell.execute_reply.started":"2026-04-01T06:43:51.235232Z","shell.execute_reply":"2026-04-01T06:43:51.239723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 5B — Training Loop","metadata":{}},{"cell_type":"code","source":"# ── 5B: TRAINING LOOP ────────────────────────────────────────────────────────\nprint(f\"[ 5B ] Starting SSL contrastive training — {SSL_EPOCHS} epochs\")\nprint(f\"       Batch size: {SSL_BATCH_SIZE}  |  Batches/epoch: {len(ssl_dataloader)}\")\nprint()\n\nencoder.train()\nt0_train    = time.time()\nepoch_losses = []                      # ← persisted for V3 loss curve\n\nfor epoch in range(SSL_EPOCHS):\n    epoch_loss    = 0.0\n    valid_batches = 0\n\n    for batch in ssl_dataloader:\n        if batch is None:\n            continue\n\n        v1 = batch[\"view1\"].to(DEVICE, non_blocking=True)\n        v2 = batch[\"view2\"].to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad()\n\n        with autocast():\n            z1   = encoder(v1)\n            z2   = encoder(v2)\n            loss = nt_xent_loss(z1, z2)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        epoch_loss    += loss.item()\n        valid_batches += 1\n\n    avg_loss = epoch_loss / max(valid_batches, 1)\n    epoch_losses.append(avg_loss)       # ← cache per-epoch mean loss\n    print(f\"  Epoch [{epoch + 1:>2}/{SSL_EPOCHS}]  Contrastive Loss: {avg_loss:.4f}\")\n\ntrain_mins = (time.time() - t0_train) / 60\nprint(f\"\\nTraining complete in {train_mins:.1f} minutes.\")\nprint(\"SSL head is now domain-adapted to the chest X-ray feature space.\")\n\n# Persist loss log so V3 can reload it in future kernel restarts\n_loss_csv = OUTPUT_DIR / \"ssl_epoch_losses.csv\"\nimport pandas as _pd_tmp\n_pd_tmp.DataFrame({\"epoch\": range(1, len(epoch_losses) + 1),\n                   \"loss\": epoch_losses}).to_csv(_loss_csv, index=False)\nprint(f\"  Loss log saved → {_loss_csv}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T06:43:51.241196Z","iopub.execute_input":"2026-04-01T06:43:51.241698Z","iopub.status.idle":"2026-04-01T07:11:43.470640Z","shell.execute_reply.started":"2026-04-01T06:43:51.241668Z","shell.execute_reply":"2026-04-01T07:11:43.469751Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 6 — Joint Extraction — Sample Set","metadata":{}},{"cell_type":"markdown","source":"### 6A — Extract","metadata":{}},{"cell_type":"code","source":"# ── 6A: EXTRACT (SAMPLE SET) ─────────────────────────────────────────────────\nprint(f\"[ 6A ] Joint extraction — {len(clean_dataloader.dataset)} samples, {len(clean_dataloader)} batches\")\nprint()\n\nencoder.eval()\n\nall_raw_features     = []\nall_adapted_features = []\nall_diseases, all_datasets, all_paths = [], [], []\nall_demo = {col: [] for col in available_demo_cols}\n\ntotal_batches = len(clean_dataloader)\nt0 = time.time()\n\nwith torch.no_grad():\n    for i, batch in enumerate(clean_dataloader, 1):\n        if batch is None:\n            continue\n\n        images = batch[\"image\"].to(DEVICE, non_blocking=True)\n\n        # Single backbone forward — branch into both heads\n        raw_emb     = encoder.extract_raw(images)      # 384-d, L2-normalised\n        adapted_emb = encoder.extract_adapted(images)  # 128-d, L2-normalised\n\n        all_raw_features.append(raw_emb.cpu().numpy())\n        all_adapted_features.append(adapted_emb.cpu().numpy())\n\n        all_diseases.extend(batch[\"disease\"])\n        all_datasets.extend(batch[\"dataset\"])\n        all_paths.extend(make_relative(p) for p in batch[\"path\"])\n\n        for col in available_demo_cols:\n            all_demo[col].extend(batch.get(col, [\"unknown\"] * len(batch[\"disease\"])))\n\n        if i % 10 == 0 or i == total_batches:\n            print(f\"  Extracted {i:>3}/{total_batches} batches | {time.time() - t0:.1f}s elapsed\", end=\"\\r\")\n\nprint(f\"\\n  Extraction complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:11:43.472140Z","iopub.execute_input":"2026-04-01T07:11:43.472505Z","iopub.status.idle":"2026-04-01T07:12:21.296183Z","shell.execute_reply.started":"2026-04-01T07:11:43.472475Z","shell.execute_reply":"2026-04-01T07:12:21.295354Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 6B — Consolidate & Save","metadata":{}},{"cell_type":"code","source":"# ── 6B: CONSOLIDATE & SAVE (SAMPLE SET) ──────────────────────────────────────\nfeatures_384d = np.concatenate(all_raw_features,     axis=0)\nfeatures_128d = np.concatenate(all_adapted_features, axis=0)\nlabels_arr    = np.array(all_diseases)\ndatasets_arr  = np.array(all_datasets)\npaths_arr     = np.array(all_paths)\n\nmeta_df = pd.DataFrame({\"path\": paths_arr, \"disease\": labels_arr, \"dataset\": datasets_arr})\nfor col in available_demo_cols:\n    meta_df[col] = all_demo[col]\n\nprint(f\"  Raw DINOv2 features  : {features_384d.shape}  (384-d, L2-normalised)\")\nprint(f\"  SSL-adapted features : {features_128d.shape}  (128-d, L2-normalised)\")\nprint(f\"  Metadata columns     : {list(meta_df.columns)}\")\n\n# Persist sample-set arrays\nnp.save(OUTPUT_DIR / \"features_384d.npy\", features_384d)\nnp.save(OUTPUT_DIR / \"features_128d.npy\", features_128d)\nnp.save(OUTPUT_DIR / \"labels.npy\",        labels_arr)\nnp.save(OUTPUT_DIR / \"datasets.npy\",      datasets_arr)\nnp.save(OUTPUT_DIR / \"paths.npy\",         paths_arr)\nmeta_df.to_csv(OUTPUT_DIR / \"metadata.csv\", index=False)\n\nprint(f\"\\n  Sample-set outputs saved to {OUTPUT_DIR}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.297688Z","iopub.execute_input":"2026-04-01T07:12:21.298074Z","iopub.status.idle":"2026-04-01T07:12:21.338448Z","shell.execute_reply.started":"2026-04-01T07:12:21.298044Z","shell.execute_reply":"2026-04-01T07:12:21.337549Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 7 — Validation","metadata":{}},{"cell_type":"markdown","source":"### 7A — Structural Checks","metadata":{}},{"cell_type":"code","source":"# ── 7A: STRUCTURAL CHECKS ────────────────────────────────────────────────────\nprint(\"[ 7A ] Validating extracted feature arrays ...\")\nprint()\n\nN = features_384d.shape[0]\nchecks = [\n    (\"Raw features dimension     == 384\",     features_384d.shape[1] == 384),\n    (\"Adapted features dimension == 128\",     features_128d.shape[1] == 128),\n    (\"Labels aligned with N\",                 labels_arr.shape[0]   == N),\n    (\"Datasets aligned with N\",               datasets_arr.shape[0] == N),\n    (\"Paths aligned with N\",                  paths_arr.shape[0]    == N),\n    (\"Metadata rows aligned with N\",          len(meta_df)          == N),\n    (\"No NaNs in raw features\",               not np.isnan(features_384d).any()),\n    (\"No Infs in raw features\",               not np.isinf(features_384d).any()),\n    (\"No NaNs in adapted features\",           not np.isnan(features_128d).any()),\n    (\"No Infs in adapted features\",           not np.isinf(features_128d).any()),\n    (\"Raw features L2-normalised\",            np.isclose(np.linalg.norm(features_384d, axis=1).mean(), 1.0, atol=1e-3)),\n    (\"Adapted features L2-normalised\",        np.isclose(np.linalg.norm(features_128d, axis=1).mean(), 1.0, atol=1e-3)),\n    (\"Paths are relative (no /kaggle/ prefix)\", not any(p.startswith(\"/kaggle/\") for p in paths_arr)),\n    (\"Samples extracted > 0\",                 N > 0),\n]\nall_passed = True\nfor name, passed in checks:\n    status = \"PASS\" if passed else \"FAIL\"\n    print(f\"  [{status}]  {name}\")\n    if not passed:\n        all_passed = False\n\nprint()\nif all_passed:\n    print(\"  All structural checks passed — feature banks are sound.\")\nelse:\n    print(\"  WARNING: One or more checks failed. Review the pipeline before proceeding.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.339567Z","iopub.execute_input":"2026-04-01T07:12:21.339820Z","iopub.status.idle":"2026-04-01T07:12:21.355458Z","shell.execute_reply.started":"2026-04-01T07:12:21.339797Z","shell.execute_reply":"2026-04-01T07:12:21.354725Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 7B — Demographic Coverage","metadata":{}},{"cell_type":"code","source":"# ── 7B: DEMOGRAPHIC COVERAGE ─────────────────────────────────────────────────\nprint(\"[ 7B ] Demographic coverage summary:\")\nprint()\n\ntotal_samples = len(meta_df)\ncandidate_demo_cols = DEMOGRAPHIC_COLS if \"DEMOGRAPHIC_COLS\" in globals() else [\"age\", \"sex\", \"view_position\", \"ap_pa\", \"patient_id\"]\ndemo_cols = [c for c in candidate_demo_cols if c in meta_df.columns]\n\nif not demo_cols:\n    print(\"  No demographic columns available in metadata.\")\n    print(\"  Check Phase 2 loading and dataset schema for demographic fields.\")\nelse:\n    missing_tokens = {\"\", \"unknown\", \"nan\", \"none\", \"na\", \"n/a\"}\n    for col in demo_cols:\n        if col == \"age\":\n            known = pd.to_numeric(meta_df[col], errors=\"coerce\").notna().sum()\n        else:\n            col_series = meta_df[col].astype(str).str.strip().str.lower()\n            known = (~col_series.isin(missing_tokens)).sum()\n\n        pct = 100.0 * known / total_samples if total_samples else 0.0\n        print(f\"  {col:<20}: {known:>6} / {total_samples} samples have data  ({pct:.1f}% coverage)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.356457Z","iopub.execute_input":"2026-04-01T07:12:21.356689Z","iopub.status.idle":"2026-04-01T07:12:21.369432Z","shell.execute_reply.started":"2026-04-01T07:12:21.356668Z","shell.execute_reply":"2026-04-01T07:12:21.368636Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 8 — Severity Indexing","metadata":{}},{"cell_type":"markdown","source":"### 8A — Healthy Anchor","metadata":{}},{"cell_type":"code","source":"# ── 8A: HEALTHY ANCHOR ───────────────────────────────────────────────────────\nprint(\"[ 8A ] Computing healthy (No Finding) anchor centroid ...\")\n\nhealthy_mask     = (labels_arr == \"No Finding\")\nhealthy_centroid = np.mean(features_128d[healthy_mask], axis=0)\n\nprint(f\"  Healthy anchor computed from {healthy_mask.sum()} 'No Finding' samples.\")\nprint(\"  'No Finding' samples will be labelled 'Normal' — no sub-clustering applied.\")\n\n# Cast to object dtype so per-element writes of variable-length severity\n# strings (e.g. \"Moderate Pleural Effusion\") are never silently truncated.\ndetailed_labels = np.where(healthy_mask, \"Normal\", labels_arr).astype(object)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.370256Z","iopub.execute_input":"2026-04-01T07:12:21.370552Z","iopub.status.idle":"2026-04-01T07:12:21.386799Z","shell.execute_reply.started":"2026-04-01T07:12:21.370523Z","shell.execute_reply":"2026-04-01T07:12:21.386198Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 8B — Grade Function","metadata":{}},{"cell_type":"code","source":"# ── 8B: GRADE FUNCTION ───────────────────────────────────────────────────────\ndef grade_severity(disease_name: str) -> None:\n    \"\"\"\n    Applies 3-cluster KMeans to the 128-d SSL embeddings of one disease class.\n\n    Clusters are ranked by Euclidean distance from the healthy centroid:\n      Closest  ->  Mild <disease>\n      Middle   ->  Moderate <disease>\n      Furthest ->  Severe <disease>\n\n    'No Finding' samples are never passed to this function.\n    \"\"\"\n    global detailed_labels\n\n    mask  = (labels_arr == disease_name)\n    feats = features_128d[mask]\n\n    if len(feats) < 3:\n        print(f\"  Skipping '{disease_name}' — insufficient samples ({len(feats)}).\")\n        return\n\n    kmeans    = KMeans(n_clusters=3, random_state=RANDOM_STATE, n_init=10)\n    sub_lbls  = kmeans.fit_predict(feats)\n    centroids = kmeans.cluster_centers_\n\n    distances      = [euclidean(c, healthy_centroid) for c in centroids]\n    sorted_indices = np.argsort(distances)   # ascending: closest -> furthest\n\n    grade_map = {\n        sorted_indices[0]: f\"Mild {disease_name}\",\n        sorted_indices[1]: f\"Moderate {disease_name}\",\n        sorted_indices[2]: f\"Severe {disease_name}\",\n    }\n\n    # Use explicit integer indices — avoids boolean-mask write issues on\n    # fixed-dtype string arrays produced by np.where()\n    indices = np.where(mask)[0]\n    for pos, cluster_id in zip(indices, sub_lbls):\n        detailed_labels[pos] = grade_map[cluster_id]\n\n    print(f\"  {disease_name}\")\n    for rank, cluster_idx in enumerate(sorted_indices):\n        grade = (\"Mild\", \"Moderate\", \"Severe\")[rank]\n        count = int((sub_lbls == cluster_idx).sum())\n        print(f\"    {grade:<10}: {count} cases\")\n\n\nprint(\"Severity grading function defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.387701Z","iopub.execute_input":"2026-04-01T07:12:21.387951Z","iopub.status.idle":"2026-04-01T07:12:21.397971Z","shell.execute_reply.started":"2026-04-01T07:12:21.387925Z","shell.execute_reply":"2026-04-01T07:12:21.397436Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 8C — Apply & Save","metadata":{}},{"cell_type":"code","source":"# ── 8C: APPLY & SAVE ─────────────────────────────────────────────────────────\npathology_classes = sorted(c for c in np.unique(labels_arr) if c != \"No Finding\")\nprint(f\"[ 8C ] Applying severity grading to {len(pathology_classes)} pathological classes ...\")\nprint()\n\nfor disease in pathology_classes:\n    grade_severity(disease)\n\nmeta_df[\"severity_label\"] = detailed_labels\n\n# Final distribution\nprint(\"\\n  Final severity label distribution:\")\nunique_lbls, counts = np.unique(detailed_labels, return_counts=True)\nfor lbl, cnt in sorted(zip(unique_lbls, counts)):\n    print(f\"    {lbl:<35}: {cnt:>6}\")\n\nnp.save(OUTPUT_DIR / \"labels_detailed.npy\", detailed_labels)\nmeta_df.to_csv(OUTPUT_DIR / \"metadata.csv\", index=False)   # re-save with severity column\n\nprint(f\"\\n  labels_detailed.npy saved to {OUTPUT_DIR}\")\nprint(\"  metadata.csv updated with 'severity_label' column.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.399261Z","iopub.execute_input":"2026-04-01T07:12:21.399534Z","iopub.status.idle":"2026-04-01T07:12:21.648198Z","shell.execute_reply.started":"2026-04-01T07:12:21.399513Z","shell.execute_reply":"2026-04-01T07:12:21.647620Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 9 — Demographic Encoder — Option B","metadata":{}},{"cell_type":"markdown","source":"### 9A — Encode Demographics","metadata":{}},{"cell_type":"code","source":"# ── 9A: ENCODE DEMOGRAPHICS ──────────────────────────────────────────────────\nprint(\"[ 9A ] Encoding demographic metadata into a numeric conditioning vector ...\")\nprint()\n\n# Build per-column category maps from the sample-set metadata\ncategory_maps = {}\nfor col in available_demo_cols:\n    if col == \"age\":\n        continue  # handled numerically\n    unique_vals = sorted(meta_df[col].dropna().unique())\n    category_maps[col] = {v: i for i, v in enumerate(unique_vals)}\n\ndef encode_demographic_row(row: pd.Series) -> np.ndarray:\n    \"\"\"\n    Encodes a single metadata row into a fixed-length float32 vector:\n      - age         : normalised float in [0, 1]  (unknown -> 0.5)\n      - cat. cols   : one-hot per column\n    \"\"\"\n    parts = []\n\n    # Age\n    if \"age\" in available_demo_cols:\n        try:\n            age_val = float(row[\"age\"])\n            parts.append(np.array([np.clip(age_val / 100.0, 0.0, 1.0)], dtype=np.float32))\n        except (ValueError, TypeError):\n            parts.append(np.array([0.5], dtype=np.float32))\n\n    # Categorical columns\n    for col, cmap in category_maps.items():\n        vec     = np.zeros(len(cmap), dtype=np.float32)\n        val     = row.get(col, \"unknown\")\n        if val in cmap:\n            vec[cmap[val]] = 1.0\n        parts.append(vec)\n\n    return np.concatenate(parts) if parts else np.zeros(1, dtype=np.float32)\n\n\n# Apply to entire metadata table\ndemo_matrix = np.vstack([encode_demographic_row(row) for _, row in meta_df.iterrows()])\nDEMO_DIM    = demo_matrix.shape[1]\n\nprint(f\"  Demographic feature matrix : {demo_matrix.shape}  ({DEMO_DIM}-d per sample)\")\nprint(f\"  Category maps constructed  : {list(category_maps.keys())}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.649048Z","iopub.execute_input":"2026-04-01T07:12:21.649331Z","iopub.status.idle":"2026-04-01T07:12:21.863933Z","shell.execute_reply.started":"2026-04-01T07:12:21.649309Z","shell.execute_reply":"2026-04-01T07:12:21.863397Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9B — DemographicEncoder MLP","metadata":{}},{"cell_type":"code","source":"# ── 9B: DEMOGRAPHIC ENCODER MLP ──────────────────────────────────────────────\nclass DemographicEncoder(nn.Module):\n    \"\"\"\n    Lightweight MLP that projects a raw demographic vector into a\n    dense conditioning embedding of dimension `out_dim`.\n\n    Architecture:  demo_dim -> 64 -> ReLU -> out_dim -> L2-normalise\n    \"\"\"\n    def __init__(self, in_dim: int, out_dim: int = 64):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(in_dim, 64),\n            nn.ReLU(inplace=True),\n            nn.Linear(64, out_dim),\n        )\n        self.out_dim = out_dim\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        return F.normalize(self.net(x), dim=-1)\n\n\nCONDITIONING_DIM   = 64\ndemo_encoder       = DemographicEncoder(in_dim=DEMO_DIM, out_dim=CONDITIONING_DIM).to(DEVICE)\n\nprint(f\"DemographicEncoder instantiated on {DEVICE}\")\nprint(f\"  Input dim  : {DEMO_DIM}\")\nprint(f\"  Output dim : {CONDITIONING_DIM}  (L2-normalised)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.864826Z","iopub.execute_input":"2026-04-01T07:12:21.865062Z","iopub.status.idle":"2026-04-01T07:12:21.873701Z","shell.execute_reply.started":"2026-04-01T07:12:21.865040Z","shell.execute_reply":"2026-04-01T07:12:21.873066Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9C — Fused Conditioning Vector","metadata":{}},{"cell_type":"code","source":"# ── 9C: FUSED CONDITIONING VECTOR ────────────────────────────────────────────\n# The cGAN needs THREE explicit conditioning signals:\n#   1. Disease class  — one-hot over all disease labels\n#   2. Severity       — one-hot [Normal, Mild, Moderate, Severe]\n#   3. Demographics   — 64-d MLP embedding (age + sex + view_position + ap_pa)\n#\n# These are concatenated with the SSL 128-d visual embedding to form the\n# full conditioning vector the cGAN generator and discriminator both receive.\n#\n# Final layout (sample set):\n#   [SSL 128-d] ⊕ [disease one-hot D-d] ⊕ [severity one-hot 4-d] ⊕ [demo 64-d]\n#   Total = 128 + D + 4 + 64  (D = number of disease classes)\n\nprint(\"[ 9C ] Building full cGAN conditioning vectors ...\")\nprint()\n\n# ── 1. Disease one-hot ────────────────────────────────────────────────────────\ndisease_classes   = sorted(np.unique(labels_arr))\ndisease_to_idx    = {d: i for i, d in enumerate(disease_classes)}\nD                 = len(disease_classes)\ndisease_onehot    = np.zeros((len(labels_arr), D), dtype=np.float32)\nfor i, lbl in enumerate(labels_arr):\n    disease_onehot[i, disease_to_idx[lbl]] = 1.0\nprint(f\"  Disease one-hot    : {disease_onehot.shape}  ({D} classes)\")\nprint(f\"  Classes            : {disease_classes}\")\n\n# ── 2. Severity one-hot ───────────────────────────────────────────────────────\nSEV_GRADES     = [\"Normal\", \"Mild\", \"Moderate\", \"Severe\"]\nsev_to_idx     = {s: i for i, s in enumerate(SEV_GRADES)}\n\ndef sev_label_to_grade(label: str) -> str:\n    l = str(label).lower()\n    if l == \"normal\":       return \"Normal\"\n    if l.startswith(\"mild\"):     return \"Mild\"\n    if l.startswith(\"moderate\"): return \"Moderate\"\n    if l.startswith(\"severe\"):   return \"Severe\"\n    return \"Normal\"  # safe fallback\n\nseverity_onehot = np.zeros((len(detailed_labels), 4), dtype=np.float32)\nfor i, lbl in enumerate(detailed_labels):\n    severity_onehot[i, sev_to_idx[sev_label_to_grade(lbl)]] = 1.0\nprint(f\"  Severity one-hot   : {severity_onehot.shape}  (Normal/Mild/Moderate/Severe)\")\n\n# ── 3. Demographic 64-d embedding ────────────────────────────────────────────\ndemo_encoder.eval()\ndemo_tensor = torch.tensor(demo_matrix, dtype=torch.float32, device=DEVICE)\nwith torch.no_grad():\n    demo_embeddings = demo_encoder(demo_tensor).cpu().numpy()   # (N, 64)\nprint(f\"  Demographic 64-d   : {demo_embeddings.shape}\")\n\n# ── 4. Concatenate all conditioning signals ───────────────────────────────────\nfused_vectors = np.concatenate(\n    [features_128d, disease_onehot, severity_onehot, demo_embeddings], axis=1\n)\nFUSED_TOTAL = 128 + D + 4 + 64\nprint()\nprint(f\"  Full conditioning vector : {fused_vectors.shape}\")\nprint(f\"    128-d SSL visual embedding\")\nprint(f\"    {D}-d disease class one-hot\")\nprint(f\"    4-d severity one-hot\")\nprint(f\"    64-d demographic embedding\")\nprint(f\"    ─────────────────────────\")\nprint(f\"    {FUSED_TOTAL}-d total\")\n\n# Store for Phase 9D verify and later use\nCONDITIONING_DISEASE_CLASSES = disease_classes\nCONDITIONING_DIM_TOTAL       = FUSED_TOTAL","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.874568Z","iopub.execute_input":"2026-04-01T07:12:21.874914Z","iopub.status.idle":"2026-04-01T07:12:21.902508Z","shell.execute_reply.started":"2026-04-01T07:12:21.874884Z","shell.execute_reply":"2026-04-01T07:12:21.901900Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9D — Verify","metadata":{}},{"cell_type":"code","source":"# ── 9D: VERIFY ───────────────────────────────────────────────────────────────\nprint(\"[ 9D ] Verifying full cGAN conditioning vectors ...\")\nprint()\n\n_FUSED_DIM_EXPECTED = 128 + len(CONDITIONING_DISEASE_CLASSES) + 4 + 64\n\np9_checks = [\n    (f\"Fused dim == {_FUSED_DIM_EXPECTED} (128+{len(CONDITIONING_DISEASE_CLASSES)}+4+64)\",\n     fused_vectors.shape[1] == _FUSED_DIM_EXPECTED),\n    (\"Fused rows aligned with N\",\n     fused_vectors.shape[0] == features_128d.shape[0]),\n    (\"No NaNs in fused vectors\",\n     not np.isnan(fused_vectors).any()),\n    (\"No Infs in fused vectors\",\n     not np.isinf(fused_vectors).any()),\n    (\"Disease one-hot sums to 1 per row\",\n     np.allclose(disease_onehot.sum(axis=1), 1.0)),\n    (\"Severity one-hot sums to 1 per row\",\n     np.allclose(severity_onehot.sum(axis=1), 1.0)),\n    (\"Demo embeddings L2-normalised\",\n     np.isclose(np.linalg.norm(demo_embeddings, axis=1).mean(), 1.0, atol=1e-3)),\n]\n\nall_passed_9 = True\nfor name, passed in p9_checks:\n    status = \"PASS\" if passed else \"FAIL\"\n    print(f\"  [{status}]  {name}\")\n    if not passed:\n        all_passed_9 = False\n\nprint()\nif all_passed_9:\n    print(\"  All Phase 9 checks passed.\")\n    print(f\"  Conditioning vector ({fused_vectors.shape[1]}-d) carries:\")\n    print(f\"    ✅ Visual pathology signal  (SSL 128-d)\")\n    print(f\"    ✅ Disease class            ({len(CONDITIONING_DISEASE_CLASSES)}-d one-hot)\")\n    print(f\"    ✅ Severity grade           (4-d one-hot: Normal/Mild/Moderate/Severe)\")\n    print(f\"    ✅ Patient demographics     (64-d: age + sex + view + ap_pa)\")\nelse:\n    print(\"  WARNING: One or more Phase 9 checks failed — review above.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.903278Z","iopub.execute_input":"2026-04-01T07:12:21.903564Z","iopub.status.idle":"2026-04-01T07:12:21.912122Z","shell.execute_reply.started":"2026-04-01T07:12:21.903544Z","shell.execute_reply":"2026-04-01T07:12:21.911230Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 11 — Scanner-Bias Validation\n\nQuantifies whether the shortcut-learning fix (CLAHE + aggressive crop) reduced\ndataset-source leakage relative to disease-class separability in the sample-set embeddings.\n\n| Outcome | Interpretation |\n|---------|----------------|\n| Disease sil. > Dataset sil. | ✅ Fix working — embeddings encode pathology |\n| Dataset sil. > Disease sil. | ⚠️ Scanner bias still dominant |","metadata":{}},{"cell_type":"code","source":"# ── PHASE 11: SCANNER-BIAS VALIDATION ───────────────────────────────────────\nfrom sklearn.metrics import silhouette_score\nfrom sklearn.preprocessing import LabelEncoder as _LE\n\nprint(\"[ 11 ] Scanner-bias silhouette audit on SSL 128-d embeddings (sample set) ...\")\nprint()\n\n_SIL_MAX = 5_000\n_f_sil   = features_128d\n_l_sil   = labels_arr\n_d_sil   = datasets_arr\n\n# Sub-sample if large\nif len(_f_sil) > _SIL_MAX:\n    _rng   = np.random.default_rng(42)\n    _sidx  = _rng.choice(len(_f_sil), _SIL_MAX, replace=False)\n    _f_sil, _l_sil, _d_sil = _f_sil[_sidx], _l_sil[_sidx], _d_sil[_sidx]\n\n_sil_disease = silhouette_score(_f_sil, _LE().fit_transform(_l_sil), metric=\"cosine\")\n_sil_dataset = silhouette_score(_f_sil, _LE().fit_transform(_d_sil), metric=\"cosine\")\n\nprint(f\"  Silhouette score (by disease class) : {_sil_disease:+.4f}\")\nprint(f\"  Silhouette score (by dataset source): {_sil_dataset:+.4f}\")\nprint()\n\nif _sil_disease > _sil_dataset:\n    print(\"  ✅  Disease silhouette > Dataset silhouette\")\n    print(\"     → CLAHE + aggressive crop fix is working. Pathology > Scanner bias.\")\nelif abs(_sil_disease - _sil_dataset) < 0.01:\n    print(\"  ⚠️   Scores are nearly equal — marginal improvement.\")\n    print(\"     → Consider increasing SSL_EPOCHS or tightening SSL_CROP_SCALE further.\")\nelse:\n    print(\"  ⚠️   Dataset silhouette ≥ Disease silhouette\")\n    print(\"     → Scanner bias still present. Try: (a) more SSL epochs,\")\n    print(\"        (b) domain-adversarial training, or (c) supervised fine-tuning.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:21.913140Z","iopub.execute_input":"2026-04-01T07:12:21.913512Z","iopub.status.idle":"2026-04-01T07:12:22.165362Z","shell.execute_reply.started":"2026-04-01T07:12:21.913481Z","shell.execute_reply":"2026-04-01T07:12:22.164645Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 13 — Full-Scale Extraction\n\n> **Note:** Phases 10–12 (cGAN architecture, training, and synthetic generation) are handled by a separate team.  \n> This phase produces the four primary `.npy` arrays that serve as inputs to that pipeline.","metadata":{}},{"cell_type":"markdown","source":"### 13A — DataLoader (Full Dataset)","metadata":{}},{"cell_type":"code","source":"# ── 13A: FULL-SCALE DATALOADER ───────────────────────────────────────────────\nprint(\"[ 13A ] Building DataLoader for the full dataset ...\")\n\ndf_full, missing_full = check_paths(df_raw)\nprint(f\"  Total raw rows  : {len(df_raw):,}\")\nprint(f\"  Valid paths     : {len(df_full):,}\")\nif missing_full:\n    print(f\"  Dropped         : {len(missing_full)} missing files\")\n\nfull_dataloader = DataLoader(\n    CleanDataset(df_full),\n    batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=NUM_WORKERS, collate_fn=clean_collate, pin_memory=True,\n)\nprint(f\"  Total batches   : {len(full_dataloader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:12:22.166581Z","iopub.execute_input":"2026-04-01T07:12:22.167364Z","iopub.status.idle":"2026-04-01T07:15:30.929524Z","shell.execute_reply.started":"2026-04-01T07:12:22.167330Z","shell.execute_reply":"2026-04-01T07:15:30.928675Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 13B — Joint Extract (Full Scale)","metadata":{}},{"cell_type":"code","source":"# ── 13B: JOINT EXTRACTION (FULL SCALE) ───────────────────────────────────────\nprint(\"[ 13B ] Extracting DINOv2 384-d + SSL 128-d embeddings across the full dataset ...\")\nprint()\n\nencoder.eval()\n\nall_raw_full      = []\nall_adapted_full  = []\nall_diseases_full = []\nall_datasets_full = []\nall_paths_full    = []\nall_demo_full     = {col: [] for col in available_demo_cols}\n\ntotal_batches = len(full_dataloader)\nt0 = time.time()\n\nwith torch.no_grad():\n    for i, batch in enumerate(full_dataloader, 1):\n        if batch is None:\n            continue\n\n        images = batch[\"image\"].to(DEVICE, non_blocking=True)\n\n        raw_emb     = encoder.extract_raw(images)\n        adapted_emb = encoder.extract_adapted(images)\n\n        all_raw_full.append(raw_emb.cpu().numpy())\n        all_adapted_full.append(adapted_emb.cpu().numpy())\n\n        all_diseases_full.extend(batch[\"disease\"])\n        all_datasets_full.extend(batch[\"dataset\"])\n        all_paths_full.extend(make_relative(p) for p in batch[\"path\"])\n\n        for col in available_demo_cols:\n            all_demo_full[col].extend(batch.get(col, [\"unknown\"] * len(batch[\"disease\"])))\n\n        if i % 50 == 0 or i == total_batches:\n            elapsed = time.time() - t0\n            n_done  = sum(len(f) for f in all_raw_full)\n            eta     = (elapsed / i) * (total_batches - i)\n            print(\n                f\"  Batch {i:>4}/{total_batches} | \"\n                f\"Samples: {n_done:>7,} | \"\n                f\"Elapsed: {elapsed / 60:.1f}m | \"\n                f\"ETA: {eta / 60:.1f}m\",\n                end=\"\\r\",\n            )\n\nprint(\"\\n  Full-scale extraction complete.\")\n\nfeatures_384d_full = np.concatenate(all_raw_full,     axis=0)\nfeatures_128d_full = np.concatenate(all_adapted_full, axis=0)\nlabels_full        = np.array(all_diseases_full)\ndatasets_full      = np.array(all_datasets_full)\npaths_full         = np.array(all_paths_full)\n\nmeta_df_full = pd.DataFrame({\"path\": paths_full, \"disease\": labels_full, \"dataset\": datasets_full})\nfor col in available_demo_cols:\n    meta_df_full[col] = all_demo_full[col]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:15:30.930449Z","iopub.execute_input":"2026-04-01T07:15:30.930732Z","iopub.status.idle":"2026-04-01T07:54:32.532631Z","shell.execute_reply.started":"2026-04-01T07:15:30.930708Z","shell.execute_reply":"2026-04-01T07:54:32.531946Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 13C — Severity Index (Full Scale)","metadata":{}},{"cell_type":"code","source":"# ── 13C: SEVERITY INDEXING (FULL SCALE) ──────────────────────────────────────\nprint(\"[ 13C ] Severity indexing on the full dataset ...\")\nprint()\n\nhealthy_mask_full     = (labels_full == \"No Finding\")\nhealthy_centroid_full = np.mean(features_128d_full[healthy_mask_full], axis=0)\n\n# Cast to object dtype — prevents string truncation on long disease names\ndetailed_labels_full = np.where(\n    healthy_mask_full, \"Normal\", labels_full\n).astype(object)\n\nprint(f\"  Healthy anchor from {healthy_mask_full.sum()} 'No Finding' samples.\")\n\n\ndef grade_severity_full(disease_name: str) -> None:\n    \"\"\"Severity grading on the full feature bank (mirrors Phase 8B logic).\"\"\"\n    global detailed_labels_full\n\n    mask  = (labels_full == disease_name)\n    feats = features_128d_full[mask]\n\n    if len(feats) < 3:\n        print(f\"  Skipping '{disease_name}' — insufficient samples.\")\n        return\n\n    kmeans    = KMeans(n_clusters=3, random_state=RANDOM_STATE, n_init=10)\n    sub_lbls  = kmeans.fit_predict(feats)\n    centroids = kmeans.cluster_centers_\n\n    distances      = [euclidean(c, healthy_centroid_full) for c in centroids]\n    sorted_indices = np.argsort(distances)\n\n    grade_map = {\n        sorted_indices[0]: f\"Mild {disease_name}\",\n        sorted_indices[1]: f\"Moderate {disease_name}\",\n        sorted_indices[2]: f\"Severe {disease_name}\",\n    }\n\n    indices = np.where(mask)[0]\n    for pos, cluster_id in zip(indices, sub_lbls):\n        detailed_labels_full[pos] = grade_map[cluster_id]\n\n    print(f\"  {disease_name}\")\n    for rank, cluster_idx in enumerate(sorted_indices):\n        grade = (\"Mild\", \"Moderate\", \"Severe\")[rank]\n        count = int((sub_lbls == cluster_idx).sum())\n        print(f\"    {grade:<10}: {count} cases\")\n\n\npathology_classes_full = sorted(c for c in np.unique(labels_full) if c != \"No Finding\")\nfor disease in pathology_classes_full:\n    grade_severity_full(disease)\n\nmeta_df_full[\"severity_label\"] = detailed_labels_full\n\n# ── Numeric severity index (0=Normal, 1=Mild, 2=Moderate, 3=Severe) ───────────\ndef severity_to_numeric(label: str) -> int:\n    l = str(label).lower()\n    if l == \"normal\":\n        return 0\n    if l.startswith(\"mild\"):\n        return 1\n    if l.startswith(\"moderate\"):\n        return 2\n    if l.startswith(\"severe\"):\n        return 3\n    return -1   # unexpected\n\nmeta_df_full[\"severity_index\"] = [severity_to_numeric(l) for l in detailed_labels_full]\n\nprint()\nprint(\"  Severity label distribution (full dataset):\")\nfor lbl in [\"Normal\", \"Mild\", \"Moderate\", \"Severe\"]:\n    mask_lbl = [str(l).startswith(lbl) or l == lbl for l in detailed_labels_full]\n    print(f\"    {lbl:<10}: {sum(mask_lbl):>7,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:54:32.533799Z","iopub.execute_input":"2026-04-01T07:54:32.534143Z","iopub.status.idle":"2026-04-01T07:54:34.258893Z","shell.execute_reply.started":"2026-04-01T07:54:32.534113Z","shell.execute_reply":"2026-04-01T07:54:34.258191Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 13D — Save","metadata":{}},{"cell_type":"code","source":"# ── 13D: SAVE FULL-SCALE ARRAYS + cGAN PARQUET ──────────────────────────────\nprint(\"\\n[ 13D ] Saving full-scale arrays and cGAN-ready Parquet ...\")\nprint()\n\n# ── Primary outputs (4 key .npy files) ────────────────────────────────────────\nnp.save(OUTPUT_DIR / \"features_384d_full.npy\",   features_384d_full)\nnp.save(OUTPUT_DIR / \"features_128d_full.npy\",   features_128d_full)\nnp.save(OUTPUT_DIR / \"labels_full.npy\",          labels_full)\nnp.save(OUTPUT_DIR / \"labels_detailed_full.npy\", detailed_labels_full)\nnp.save(OUTPUT_DIR / \"datasets_full.npy\",        datasets_full)\nnp.save(OUTPUT_DIR / \"paths_full.npy\",           paths_full)\n\n# ── Demographic fallback ───────────────────────────────────────────────────────\nprint(\"[ 13D ] Demographic fallback — checking columns ...\")\n_REQUIRED_DEMO = [\"age\", \"sex\", \"view_position\", \"ap_pa\"]\n_missing_demo  = [c for c in _REQUIRED_DEMO if c not in meta_df_full.columns\n                  or meta_df_full[c].eq(\"unknown\").mean() > 0.9]\n\nif _missing_demo:\n    print(f\"  Missing/mostly-unknown demographics: {_missing_demo}\")\n    try:\n        import pyarrow.parquet as _pq_fb\n        _schema_fb = set(_pq_fb.read_schema(PARQUET_PATH).names)\n        _avail_fb  = [c for c in _missing_demo if c in _schema_fb]\n        if _avail_fb:\n            _df_fb = pd.read_parquet(PARQUET_PATH, columns=[\"path\"] + _avail_fb)\n            _df_fb = _df_fb.drop_duplicates(subset=[\"path\"])\n            meta_df_full = meta_df_full.merge(_df_fb, on=\"path\", how=\"left\", suffixes=(\"\", \"_fb\"))\n            for col in _avail_fb:\n                fb_col = col + \"_fb\"\n                if fb_col in meta_df_full.columns:\n                    meta_df_full[col] = meta_df_full[col].where(\n                        meta_df_full[col].notna() & (meta_df_full[col] != \"unknown\"),\n                        meta_df_full[fb_col]\n                    )\n                    meta_df_full.drop(columns=[fb_col], inplace=True)\n            print(f\"  Fallback merged: {_avail_fb}\")\n        else:\n            print(\"  Columns not in original parquet — keeping unknowns.\")\n    except Exception as _e_fb:\n        print(f\"  Fallback failed ({_e_fb}) — continuing.\")\nelse:\n    print(\"  All demographic columns present — no fallback needed.\")\n\n# ── Build the THREE conditioning signals ──────────────────────────────────────\nprint()\nprint(\"[ 13D ] Building full-scale conditioning matrix ...\")\nprint(\"        Signals: disease one-hot + severity one-hot + demographic 64-d\")\nprint()\n\n# ── Signal 1: Disease class one-hot ──────────────────────────────────────────\n_disease_classes_full = sorted(np.unique(labels_full))\n_disease_to_idx_full  = {d: i for i, d in enumerate(_disease_classes_full)}\n_D_full               = len(_disease_classes_full)\n_disease_onehot_full  = np.zeros((len(labels_full), _D_full), dtype=np.float32)\nfor i, lbl in enumerate(labels_full):\n    _disease_onehot_full[i, _disease_to_idx_full[lbl]] = 1.0\nprint(f\"  Disease one-hot     : {_disease_onehot_full.shape}  ({_D_full} classes)\")\nprint(f\"  Classes             : {_disease_classes_full}\")\n\n# ── Signal 2: Severity one-hot ────────────────────────────────────────────────\n_SEV_GRADES    = [\"Normal\", \"Mild\", \"Moderate\", \"Severe\"]\n_sev_to_idx    = {s: i for i, s in enumerate(_SEV_GRADES)}\n\ndef _sev_label_to_grade(label: str) -> str:\n    l = str(label).lower()\n    if l == \"normal\":            return \"Normal\"\n    if l.startswith(\"mild\"):     return \"Mild\"\n    if l.startswith(\"moderate\"): return \"Moderate\"\n    if l.startswith(\"severe\"):   return \"Severe\"\n    return \"Normal\"\n\n_severity_onehot_full = np.zeros((len(detailed_labels_full), 4), dtype=np.float32)\nfor i, lbl in enumerate(detailed_labels_full):\n    _severity_onehot_full[i, _sev_to_idx[_sev_label_to_grade(lbl)]] = 1.0\nprint(f\"  Severity one-hot    : {_severity_onehot_full.shape}  (Normal/Mild/Moderate/Severe)\")\n\n# ── Signal 3: Demographic 64-d embedding ─────────────────────────────────────\n_cat_maps_full = {}\nfor col in _REQUIRED_DEMO:\n    if col == \"age\":\n        continue\n    if col in meta_df_full.columns:\n        unique_vals = sorted(meta_df_full[col].dropna().unique())\n        _cat_maps_full[col] = {v: i for i, v in enumerate(unique_vals)}\n\ndef _encode_demo_full(row):\n    parts = []\n    if \"age\" in meta_df_full.columns:\n        try:\n            parts.append(np.array([np.clip(float(row[\"age\"]) / 100.0, 0.0, 1.0)], dtype=np.float32))\n        except (ValueError, TypeError):\n            parts.append(np.array([0.5], dtype=np.float32))\n    for col, cmap in _cat_maps_full.items():\n        vec = np.zeros(len(cmap), dtype=np.float32)\n        val = row.get(col, \"unknown\")\n        if val in cmap:\n            vec[cmap[val]] = 1.0\n        parts.append(vec)\n    return np.concatenate(parts) if parts else np.zeros(1, dtype=np.float32)\n\n_demo_matrix_full = np.vstack([_encode_demo_full(row) for _, row in meta_df_full.iterrows()])\n_demo_enc         = DemographicEncoder(in_dim=_demo_matrix_full.shape[1], out_dim=64).to(DEVICE)\n_demo_enc.eval()\nwith torch.no_grad():\n    _demo_emb_full = _demo_enc(\n        torch.tensor(_demo_matrix_full, dtype=torch.float32, device=DEVICE)\n    ).cpu().numpy()   # (N, 64)\nprint(f\"  Demographic 64-d    : {_demo_emb_full.shape}\")\n\n# ── Concatenate all three signals + SSL visual embedding ──────────────────────\n_fused_full = np.concatenate(\n    [features_128d_full, _disease_onehot_full, _severity_onehot_full, _demo_emb_full],\n    axis=1\n)\n_FUSED_DIM_FULL = 128 + _D_full + 4 + 64\nprint()\nprint(f\"  Full conditioning vector: {_fused_full.shape}\")\nprint(f\"    [SSL 128-d] ⊕ [disease {_D_full}-d one-hot] ⊕ [severity 4-d one-hot] ⊕ [demo 64-d]\")\nprint(f\"    = {_FUSED_DIM_FULL}-d total\")\n\n# Validate conditioning signals\nassert not np.isnan(_fused_full).any(),  \"NaN in fused vector\"\nassert not np.isinf(_fused_full).any(),  \"Inf in fused vector\"\nassert np.allclose(_disease_onehot_full.sum(axis=1),  1.0), \"Disease one-hot invalid\"\nassert np.allclose(_severity_onehot_full.sum(axis=1), 1.0), \"Severity one-hot invalid\"\nprint(\"  ✅ All conditioning signal assertions passed.\")\n\nnp.save(OUTPUT_DIR / \"fused_conditioning_full.npy\", _fused_full)\n\n# ── Save conditioning metadata (maps for the cGAN to decode labels) ───────────\nimport json as _json\n_cond_meta = {\n    \"disease_classes\":    _disease_classes_full,\n    \"disease_to_idx\":     _disease_to_idx_full,\n    \"severity_grades\":    _SEV_GRADES,\n    \"severity_to_idx\":    _sev_to_idx,\n    \"fused_layout\": {\n        \"ssl_128d\":         [0, 128],\n        \"disease_onehot\":   [128, 128 + _D_full],\n        \"severity_onehot\":  [128 + _D_full, 128 + _D_full + 4],\n        \"demo_64d\":         [128 + _D_full + 4, _FUSED_DIM_FULL],\n        \"total_dim\":        _FUSED_DIM_FULL,\n    }\n}\nwith open(OUTPUT_DIR / \"conditioning_meta.json\", \"w\") as _f:\n    _json.dump(_cond_meta, _f, indent=2)\nprint(f\"  conditioning_meta.json saved — slice indices for each signal included.\")\n\n# ── Build and save the cGAN Parquet ───────────────────────────────────────────\nprint()\nprint(\"[ 13D ] Assembling cGAN-ready Parquet ...\")\n\ncgan_df = meta_df_full.copy()\ncgan_df[\"embedding_384d\"]          = list(features_384d_full)\ncgan_df[\"embedding_128d\"]          = list(features_128d_full)\ncgan_df[\"disease_onehot\"]          = list(_disease_onehot_full)\ncgan_df[\"severity_onehot\"]         = list(_severity_onehot_full)\ncgan_df[\"demo_embedding_64d\"]      = list(_demo_emb_full)\ncgan_df[\"cgan_conditioning\"]       = list(_fused_full)\n\n# Ensure severity columns present\nif \"severity_label\" not in cgan_df.columns:\n    cgan_df[\"severity_label\"] = detailed_labels_full\nif \"severity_index\" not in cgan_df.columns:\n    cgan_df[\"severity_index\"] = [\n        {\"Normal\": 0, \"Mild\": 1, \"Moderate\": 2, \"Severe\": 3}.get(\n            _sev_label_to_grade(str(l)), -1)\n        for l in detailed_labels_full\n    ]\n\n_col_order = (\n    [\"path\", \"disease\", \"dataset\", \"severity_label\", \"severity_index\"]\n    + [c for c in [\"age\", \"sex\", \"view_position\", \"ap_pa\", \"patient_id\"] if c in cgan_df.columns]\n    + [\"embedding_384d\", \"embedding_128d\",\n       \"disease_onehot\", \"severity_onehot\", \"demo_embedding_64d\",\n       \"cgan_conditioning\"]\n)\n_col_order += [c for c in cgan_df.columns if c not in _col_order]\ncgan_df = cgan_df[_col_order]\n\nPARQUET_OUT = OUTPUT_DIR / \"cgan_input_full.parquet\"\ncgan_df.to_parquet(PARQUET_OUT, index=False, engine=\"pyarrow\")\n\nprint(f\"  cgan_input_full.parquet : {len(cgan_df):,} rows × {len(cgan_df.columns)} cols\")\nprint(f\"  Columns: {list(cgan_df.columns)}\")\nprint()\n\n# ── Summary ────────────────────────────────────────────────────────────────────\nprint(f\"All outputs saved to {OUTPUT_DIR}\")\nprint()\nprint(f\"  ╔══ PRIMARY OUTPUTS ═══════════════════════════════════════════════════╗\")\nprint(f\"  ║  features_384d_full.npy     raw DINOv2 384-d            {features_384d_full.shape}\")\nprint(f\"  ║  features_128d_full.npy     SSL-adapted 128-d           {features_128d_full.shape}\")\nprint(f\"  ║  fused_conditioning_full.npy  {_FUSED_DIM_FULL}-d cGAN input         {_fused_full.shape}\")\nprint(f\"  ║  cgan_input_full.parquet    all-in-one cGAN Parquet     {cgan_df.shape}\")\nprint(f\"  ║  conditioning_meta.json     label maps + slice indices\")\nprint(f\"  ╠══ CONDITIONING BREAKDOWN ════════════════════════════════════════════╣\")\nprint(f\"  ║  [  0:128 ]  SSL visual embedding    (128-d)\")\nprint(f\"  ║  [128:{128+_D_full:<4}]  Disease class one-hot  ({_D_full}-d, classes: {_disease_classes_full})\")\nprint(f\"  ║  [{128+_D_full}:{128+_D_full+4:<4}]  Severity one-hot      (4-d: Normal/Mild/Moderate/Severe)\")\nprint(f\"  ║  [{128+_D_full+4}:{_FUSED_DIM_FULL:<4}]  Demographic embedding (64-d: age+sex+view+ap_pa)\")\nprint(f\"  ╚══════════════════════════════════════════════════════════════════════╝\")\nprint()\nprint(\"  Pipeline complete. cgan_input_full.parquet is ready for the cGAN team.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:54:34.259818Z","iopub.execute_input":"2026-04-01T07:54:34.260026Z","iopub.status.idle":"2026-04-01T07:54:44.140440Z","shell.execute_reply.started":"2026-04-01T07:54:34.260006Z","shell.execute_reply":"2026-04-01T07:54:44.139692Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 14 — Extended Validation & Visualisation\n\nDeep quality checks and diagnostic plots covering:\n- **V1** Feature space statistics & norm audits\n- **V2** Class & dataset balance\n- **V3** SSL training loss curve\n- **V4** UMAP projections (raw vs adapted, disease, dataset, severity, demographics)\n- **V5** Inter-class cosine similarity heat-map\n- **V6** Severity gradient analysis\n- **V7** Demographic conditioning coverage\n- **V8** Feature drift (raw → adapted)\n- **V9** KNN purity & intra/inter-class distance audit\n- **V10** GAN-readiness checklist","metadata":{}},{"cell_type":"markdown","source":"### V0 — Validation Imports","metadata":{}},{"cell_type":"code","source":"# ── V0: VALIDATION / VIZ IMPORTS ─────────────────────────────────────────────\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nimport seaborn as sns\nimport umap.umap_ as umap\nfrom sklearn.neighbors import KNeighborsClassifier\nfrom sklearn.metrics import classification_report\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom scipy.spatial.distance import cdist\nimport warnings\n\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\n\nsns.set_theme(style=\"whitegrid\", palette=\"muted\")\nPALETTE = sns.color_palette(\"tab20\", 20)\nPLOT_DIR = OUTPUT_DIR / \"plots\"\nPLOT_DIR.mkdir(parents=True, exist_ok=True)\n\n# ── Convenience: pick full-scale arrays when available, else sample-set ───────\n_f384   = features_384d_full   if \"features_384d_full\"   in dir() else features_384d\n_f128   = features_128d_full   if \"features_128d_full\"   in dir() else features_128d\n_labels = labels_full          if \"labels_full\"           in dir() else labels_arr\n_detail = detailed_labels_full if \"detailed_labels_full\"  in dir() else detailed_labels\n_dsets  = datasets_full        if \"datasets_full\"         in dir() else datasets_arr\n_meta   = meta_df_full         if \"meta_df_full\"          in dir() else meta_df\n\nN_VIZ = len(_labels)\nprint(f\"Validation operating on {N_VIZ:,} samples.\")\nprint(f\"Plots will be saved to: {PLOT_DIR}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:54:44.141505Z","iopub.execute_input":"2026-04-01T07:54:44.141810Z","iopub.status.idle":"2026-04-01T07:55:21.557980Z","shell.execute_reply.started":"2026-04-01T07:54:44.141779Z","shell.execute_reply":"2026-04-01T07:55:21.557280Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V1 — Feature Space Statistics & Norm Audit","metadata":{}},{"cell_type":"code","source":"# ── V1: FEATURE SPACE STATISTICS (FIXED) ──────────────────────────────────────\nprint(\"[ V1 ] Feature space statistics\")\nprint()\n\nfor tag, arr in [(\"Raw 384-d\", _f384), (\"SSL 128-d\", _f128)]:\n    norms = np.linalg.norm(arr, axis=1)\n    print(f\"  {tag}\")\n    print(f\"    Shape         : {arr.shape}\")\n    print(f\"    dtype         : {arr.dtype}\")\n    print(f\"    Mean value    : {arr.mean():.4f}  ± {arr.std():.4f}\")\n    print(f\"    L2-norm  mean : {norms.mean():.6f}  ± {norms.std():.2e}\")\n    print(f\"    L2-norm  min  : {norms.min():.6f}\")\n    print(f\"    L2-norm  max  : {norms.max():.6f}\")\n    print(f\"    NaN count     : {np.isnan(arr).sum()}\")\n    print(f\"    Inf count     : {np.isinf(arr).sum()}\")\n    print()\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\nfor ax, (tag, arr) in zip(axes, [(\"Raw 384-d\", _f384), (\"SSL 128-d\", _f128)]):\n    norms = np.linalg.norm(arr, axis=1)\n\n    all_same = np.allclose(norms, norms[0], atol=1e-5)\n\n    if all_same:\n        # All norms are identical — show a confirmation panel instead of histogram\n        ax.set_xlim(0.99, 1.01)\n        ax.set_ylim(0, 1)\n        ax.axvline(1.0, color=\"crimson\", linestyle=\"--\", linewidth=2, label=\"L2 = 1\")\n        ax.bar([1.0], [0.85], width=0.002, color=\"steelblue\",\n               alpha=0.85, label=f\"All {len(norms):,} norms = {norms[0]:.6f}\")\n        ax.set_xlabel(\"L2 norm\")\n        ax.set_ylabel(\"Density (all values identical)\")\n        ax.set_title(f\"L2-norm distribution — {tag}\\n\"\n                     f\"✅ Perfectly normalised: mean={norms.mean():.6f}, std={norms.std():.2e}\")\n        ax.legend(fontsize=9)\n\n        # Add annotation box\n        ax.text(0.5, 0.5,\n                f\"All {len(norms):,} embeddings\\nperfectly L2-normalised\\nnorm = {norms[0]:.8f}\",\n                transform=ax.transAxes,\n                ha=\"center\", va=\"center\", fontsize=11,\n                bbox=dict(boxstyle=\"round,pad=0.5\",\n                          facecolor=\"lightgreen\", alpha=0.7,\n                          edgecolor=\"green\", linewidth=1.5))\n    else:\n        # Normal case — norms vary, show histogram\n        ax.hist(norms, bins=60, color=\"steelblue\",\n                edgecolor=\"white\", linewidth=0.4, density=True)\n        ax.axvline(1.0, color=\"crimson\", linestyle=\"--\",\n                   linewidth=1.5, label=\"L2 = 1\")\n        ax.axvline(norms.mean(), color=\"orange\", linestyle=\"-\",\n                   linewidth=1.5, label=f\"mean={norms.mean():.4f}\")\n        ax.set_xlabel(\"L2 norm\")\n        ax.set_ylabel(\"Density\")\n        ax.set_title(f\"L2-norm distribution — {tag}\\n\"\n                     f\"mean={norms.mean():.4f}, std={norms.std():.4f}\")\n        ax.legend(fontsize=9)\n\nplt.suptitle(\"V1 · Feature Norm Audit\", fontweight=\"bold\", fontsize=13)\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v1_norm_audit.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → v1_norm_audit.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:55:21.561912Z","iopub.execute_input":"2026-04-01T07:55:21.562437Z","iopub.status.idle":"2026-04-01T07:55:22.900425Z","shell.execute_reply.started":"2026-04-01T07:55:21.562413Z","shell.execute_reply":"2026-04-01T07:55:22.899665Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V2 — Class & Dataset Balance","metadata":{}},{"cell_type":"code","source":"# ── V2: CLASS & DATASET BALANCE ───────────────────────────────────────────────\nprint(\"[ V2 ] Class and dataset balance\")\n\nlabel_counts   = pd.Series(_labels).value_counts()\ndataset_counts = pd.Series(_dsets).value_counts()\n\nfig, axes = plt.subplots(1, 2, figsize=(16, 5))\n\n# Disease distribution\naxes[0].barh(label_counts.index[::-1], label_counts.values[::-1],\n             color=PALETTE[:len(label_counts)])\naxes[0].set_xlabel(\"Sample count\")\naxes[0].set_title(\"Disease class distribution\")\nfor i, v in enumerate(label_counts.values[::-1]):\n    axes[0].text(v + 5, i, f\"{v:,}\", va=\"center\", fontsize=8)\n\n# Dataset distribution\naxes[1].bar(dataset_counts.index, dataset_counts.values,\n            color=PALETTE[:len(dataset_counts)])\naxes[1].set_xlabel(\"Source dataset\")\naxes[1].set_ylabel(\"Sample count\")\naxes[1].set_title(\"Source dataset distribution\")\naxes[1].tick_params(axis=\"x\", rotation=30)\nfor i, (k, v) in enumerate(dataset_counts.items()):\n    axes[1].text(i, v + 5, f\"{v:,}\", ha=\"center\", fontsize=8)\n\nplt.suptitle(\"V2 · Sample Distribution\", fontweight=\"bold\")\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v2_class_dataset_balance.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\n# Class balance table\nprint()\nprint(\"  Disease counts:\")\nfor lbl, cnt in label_counts.items():\n    pct = 100 * cnt / N_VIZ\n    bar = \"█\" * int(pct / 2)\n    print(f\"    {lbl:<30} {cnt:>6,}  ({pct:5.1f}%)  {bar}\")\nprint()\nprint(\"  Saved → v2_class_dataset_balance.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:55:22.901476Z","iopub.execute_input":"2026-04-01T07:55:22.902076Z","iopub.status.idle":"2026-04-01T07:55:23.478962Z","shell.execute_reply.started":"2026-04-01T07:55:22.902041Z","shell.execute_reply":"2026-04-01T07:55:23.478088Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V3 — SSL Training Loss Curve","metadata":{}},{"cell_type":"code","source":"# ── V3: SSL TRAINING LOSS CURVE ─────────────────────────────────────────────\n# epoch_losses is saved to CSV in Phase 5B. This cell loads from memory if\n# available, or falls back to the CSV so it works on kernel restart too.\n_loss_csv = OUTPUT_DIR / \"ssl_epoch_losses.csv\"\n\nif \"epoch_losses\" in dir() and len(epoch_losses) > 0:\n    _losses = epoch_losses\n    print(\"[ V3 ] Using in-memory epoch_losses.\")\nelif _loss_csv.exists():\n    _df_loss = pd.read_csv(_loss_csv)\n    _losses  = _df_loss[\"loss\"].tolist()\n    print(f\"[ V3 ] Loaded loss log from {_loss_csv}.\")\nelse:\n    print(\"[ V3 ] No loss history found — run Phase 5B first.\")\n    _losses = []\n\nif _losses:\n    fig, ax = plt.subplots(figsize=(9, 4))\n    epochs  = range(1, len(_losses) + 1)\n    ax.plot(epochs, _losses, marker=\"o\", linewidth=2,\n            color=\"darkorange\", markersize=5, label=\"NT-Xent loss\")\n\n    # Rolling mean smoothing if enough epochs\n    if len(_losses) >= 3:\n        _smoothed = pd.Series(_losses).rolling(3, center=True).mean()\n        ax.plot(epochs, _smoothed, linewidth=1.5, linestyle=\"--\",\n                color=\"crimson\", alpha=0.7, label=\"3-epoch rolling mean\")\n\n    best_ep = int(np.argmin(_losses)) + 1\n    ax.axvline(best_ep, color=\"steelblue\", linestyle=\":\", linewidth=1.5,\n               label=f\"Best epoch {best_ep}  (loss={min(_losses):.4f})\")\n\n    ax.set_xlabel(\"Epoch\", fontsize=11)\n    ax.set_ylabel(\"NT-Xent Loss\", fontsize=11)\n    ax.set_title(\"V3 · SSL Contrastive Training Loss Curve\\n\"\n                 \"(Lower = better feature alignment across augmented views)\",\n                 fontweight=\"bold\")\n    ax.set_xticks(list(epochs))\n    ax.grid(True, linestyle=\"--\", alpha=0.4)\n    ax.legend(fontsize=9)\n    plt.tight_layout()\n    plt.savefig(PLOT_DIR / \"v3_ssl_loss_curve.png\", dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(f\"  Epochs trained : {len(_losses)}\")\n    print(f\"  Best epoch     : {best_ep}  |  Loss: {min(_losses):.4f}\")\n    print(f\"  Final loss     : {_losses[-1]:.4f}\")\n    print(f\"  Δ loss         : {_losses[0] - _losses[-1]:+.4f}  (improvement over training)\")\n    print(\"  Saved → v3_ssl_loss_curve.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:55:23.480040Z","iopub.execute_input":"2026-04-01T07:55:23.480315Z","iopub.status.idle":"2026-04-01T07:55:24.010022Z","shell.execute_reply.started":"2026-04-01T07:55:23.480292Z","shell.execute_reply":"2026-04-01T07:55:24.009276Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V4 — UMAP Projections","metadata":{}},{"cell_type":"code","source":"# ── V4A: UMAP FIT — run once, reuse all sub-plots ────────────────────────────\nprint(\"[ V4 ] Fitting UMAP on 128-d SSL embeddings ...\")\nprint(\"       (This may take ~60-120 s for large datasets)\")\n\n_N_UMAP = min(8_000, N_VIZ)   # cap for speed; remove cap for full-scale runs\nrng = np.random.default_rng(42)\n_idx = rng.choice(N_VIZ, _N_UMAP, replace=False)\n\n_emb    = _f128[_idx]\n_ulbls  = _labels[_idx]\n_udsets = _dsets[_idx]\n_udetail= _detail[_idx]\n\nreducer = umap.UMAP(n_neighbors=30, min_dist=0.1, metric=\"cosine\",\n                    random_state=42, verbose=False)\n_umap_2d = reducer.fit_transform(_emb)\nprint(f\"  UMAP projection complete — shape: {_umap_2d.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:55:24.010886Z","iopub.execute_input":"2026-04-01T07:55:24.011079Z","iopub.status.idle":"2026-04-01T07:56:00.733949Z","shell.execute_reply.started":"2026-04-01T07:55:24.011060Z","shell.execute_reply":"2026-04-01T07:56:00.733288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── V4B: UMAP coloured by disease class ──────────────────────────────────────\nunique_diseases = sorted(set(_ulbls))\n_cmap_d = {d: PALETTE[i % 20] for i, d in enumerate(unique_diseases)}\n\nfig, ax = plt.subplots(figsize=(11, 8))\nfor disease in unique_diseases:\n    mask = (_ulbls == disease)\n    ax.scatter(_umap_2d[mask, 0], _umap_2d[mask, 1],\n               s=4, alpha=0.55, color=_cmap_d[disease], label=disease)\nax.set_title(\"V4B · UMAP — SSL 128-d  (colour = disease class)\", fontweight=\"bold\")\nax.set_xlabel(\"UMAP-1\"); ax.set_ylabel(\"UMAP-2\")\nax.legend(markerscale=3, bbox_to_anchor=(1.01, 1), loc=\"upper left\",\n          fontsize=8, framealpha=0.7)\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v4b_umap_disease.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → v4b_umap_disease.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:00.735155Z","iopub.execute_input":"2026-04-01T07:56:00.735577Z","iopub.status.idle":"2026-04-01T07:56:01.354947Z","shell.execute_reply.started":"2026-04-01T07:56:00.735549Z","shell.execute_reply":"2026-04-01T07:56:01.354259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── V4C: UMAP coloured by source dataset ─────────────────────────────────────\nunique_dsets = sorted(set(_udsets))\n_cmap_ds = {d: PALETTE[i % 20] for i, d in enumerate(unique_dsets)}\n\nfig, ax = plt.subplots(figsize=(10, 7))\nfor ds in unique_dsets:\n    mask = (_udsets == ds)\n    ax.scatter(_umap_2d[mask, 0], _umap_2d[mask, 1],\n               s=4, alpha=0.55, color=_cmap_ds[ds], label=ds)\nax.set_title(\"V4C · UMAP — SSL 128-d  (colour = source dataset)\", fontweight=\"bold\")\nax.set_xlabel(\"UMAP-1\"); ax.set_ylabel(\"UMAP-2\")\nax.legend(markerscale=3, bbox_to_anchor=(1.01, 1), loc=\"upper left\",\n          fontsize=8, framealpha=0.7)\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v4c_umap_dataset.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → v4c_umap_dataset.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:01.355708Z","iopub.execute_input":"2026-04-01T07:56:01.355905Z","iopub.status.idle":"2026-04-01T07:56:02.015064Z","shell.execute_reply.started":"2026-04-01T07:56:01.355885Z","shell.execute_reply":"2026-04-01T07:56:02.014279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── V4D: UMAP coloured by severity grade ──────────────────────────────────────\nunique_severity = sorted(set(_udetail))\n_cmap_sev = {s: PALETTE[i % 20] for i, s in enumerate(unique_severity)}\n\nfig, ax = plt.subplots(figsize=(12, 8))\nfor sev in unique_severity:\n    mask = (_udetail == sev)\n    ax.scatter(_umap_2d[mask, 0], _umap_2d[mask, 1],\n               s=4, alpha=0.55, color=_cmap_sev[sev], label=sev)\nax.set_title(\"V4D · UMAP — SSL 128-d  (colour = severity grade)\", fontweight=\"bold\")\nax.set_xlabel(\"UMAP-1\"); ax.set_ylabel(\"UMAP-2\")\nax.legend(markerscale=3, bbox_to_anchor=(1.01, 1), loc=\"upper left\",\n          fontsize=8, framealpha=0.7)\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v4d_umap_severity.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → v4d_umap_severity.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:02.016112Z","iopub.execute_input":"2026-04-01T07:56:02.016384Z","iopub.status.idle":"2026-04-01T07:56:02.823982Z","shell.execute_reply.started":"2026-04-01T07:56:02.016362Z","shell.execute_reply":"2026-04-01T07:56:02.823288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── V4E: Side-by-side Raw 384-d vs SSL 128-d UMAP ────────────────────────────\nprint(\"[ V4E ] Fitting UMAP on raw 384-d embeddings for comparison ...\")\n\nreducer_raw = umap.UMAP(n_neighbors=30, min_dist=0.1, metric=\"cosine\",\n                        random_state=42, verbose=False)\n_umap_raw = reducer_raw.fit_transform(_f384[_idx])\n\nfig, axes = plt.subplots(1, 2, figsize=(18, 7))\nfor ax, coords, title in zip(axes,\n                              [_umap_raw, _umap_2d],\n                              [\"Raw DINOv2 384-d\", \"SSL-adapted 128-d\"]):\n    for disease in unique_diseases:\n        mask = (_ulbls == disease)\n        ax.scatter(coords[mask, 0], coords[mask, 1],\n                   s=4, alpha=0.5, color=_cmap_d[disease], label=disease)\n    ax.set_title(f\"V4E · UMAP — {title}\", fontweight=\"bold\")\n    ax.set_xlabel(\"UMAP-1\"); ax.set_ylabel(\"UMAP-2\")\naxes[1].legend(markerscale=3, bbox_to_anchor=(1.01, 1), loc=\"upper left\",\n               fontsize=8, framealpha=0.7)\nplt.suptitle(\"Feature Space Comparison: Raw vs SSL-adapted\", fontsize=13)\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v4e_umap_raw_vs_ssl.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → v4e_umap_raw_vs_ssl.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:02.825433Z","iopub.execute_input":"2026-04-01T07:56:02.825659Z","iopub.status.idle":"2026-04-01T07:56:23.845986Z","shell.execute_reply.started":"2026-04-01T07:56:02.825637Z","shell.execute_reply":"2026-04-01T07:56:23.845192Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V5 — Inter-Class Cosine Similarity Heat-Map","metadata":{}},{"cell_type":"code","source":"# ── V5: INTER-CLASS COSINE SIMILARITY ────────────────────────────────────────\nprint(\"[ V5 ] Computing per-class mean embeddings and cosine similarity matrix ...\")\n\nunique_classes = sorted(set(_labels))\nclass_centroids_128 = np.stack([\n    _f128[_labels == c].mean(axis=0) for c in unique_classes\n])   # (C, 128)\n\n# L2-normalise centroids for cosine similarity\nclass_centroids_norm = class_centroids_128 / (\n    np.linalg.norm(class_centroids_128, axis=1, keepdims=True) + 1e-9\n)\ncos_sim_matrix = class_centroids_norm @ class_centroids_norm.T   # (C, C)\n\nfig, ax = plt.subplots(figsize=(11, 9))\nsns.heatmap(\n    cos_sim_matrix,\n    xticklabels=unique_classes,\n    yticklabels=unique_classes,\n    annot=True, fmt=\".2f\", cmap=\"RdYlGn\",\n    vmin=-0.2, vmax=1.0,\n    linewidths=0.4,\n    ax=ax,\n)\nax.set_title(\"V5 · Inter-class Cosine Similarity (SSL 128-d centroids)\",\n             fontweight=\"bold\", pad=12)\nax.set_xticklabels(ax.get_xticklabels(), rotation=45, ha=\"right\", fontsize=8)\nax.set_yticklabels(ax.get_yticklabels(), rotation=0, fontsize=8)\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v5_interclass_cosine_sim.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\n# Flag highly similar classes (potential overlap risk for cGAN conditioning)\nprint()\nprint(\"  High-similarity pairs (cosine > 0.90, excluding diagonal):\")\nfound = False\nfor i in range(len(unique_classes)):\n    for j in range(i + 1, len(unique_classes)):\n        if cos_sim_matrix[i, j] > 0.90:\n            print(f\"    {unique_classes[i]:<25}  ↔  {unique_classes[j]:<25}  \"\n                  f\"cosine = {cos_sim_matrix[i,j]:.3f}\")\n            found = True\nif not found:\n    print(\"    None — all class centroids are well-separated.\")\nprint()\nprint(\"  Saved → v5_interclass_cosine_sim.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:23.847147Z","iopub.execute_input":"2026-04-01T07:56:23.847475Z","iopub.status.idle":"2026-04-01T07:56:24.293110Z","shell.execute_reply.started":"2026-04-01T07:56:23.847443Z","shell.execute_reply":"2026-04-01T07:56:24.292276Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V6 — Severity Gradient Analysis","metadata":{}},{"cell_type":"code","source":"# ── V6: SEVERITY GRADIENT ANALYSIS ──────────────────────────────────────────\nprint(\"[ V6 ] Severity gradient: distance from healthy anchor per grade ...\")\nprint()\n\n# ── Resolve which arrays to use ───────────────────────────────────────────────\n_sev_f128   = features_128d_full   if \"features_128d_full\"   in dir() else features_128d\n_sev_labels = labels_full          if \"labels_full\"           in dir() else labels_arr\n_sev_detail = (detailed_labels_full if \"detailed_labels_full\" in dir()\n               else detailed_labels)\n\n# ── Diagnostic: inspect what's actually in _sev_detail ───────────────────────\nunique_detail = pd.Series(_sev_detail).value_counts()\nprint(f\"  _sev_detail has {len(unique_detail)} unique values. Sample:\")\nfor lbl, cnt in unique_detail.head(15).items():\n    print(f\"    {lbl!r:<40}  n={cnt}\")\nprint()\n\n# ── Guard: check severity grading was actually applied ───────────────────────\n_has_graded = any(\n    str(v).startswith((\"Mild \", \"Moderate \", \"Severe \"))\n    for v in unique_detail.index\n)\n\nif not _has_graded:\n    print(\"  ⚠  Severity grading labels not found in _sev_detail.\")\n    print(\"     Re-applying Phase 8 grading logic on-the-fly ...\")\n    print()\n\n    _healthy_mask = (\n        (_sev_labels == \"No Finding\") |\n        (_sev_detail == \"Normal\")     |\n        (_sev_detail == \"No Finding\")\n    )\n    if _healthy_mask.sum() == 0:\n        raise RuntimeError(\n            \"Cannot find any 'No Finding' / 'Normal' samples to build healthy anchor.\"\n        )\n\n    _healthy_centroid_v6 = np.mean(_sev_f128[_healthy_mask], axis=0)\n    print(f\"  Healthy anchor built from {_healthy_mask.sum()} samples.\")\n\n    from sklearn.cluster import KMeans\n    from scipy.spatial.distance import euclidean\n\n    # Object dtype prevents silent truncation of long grade strings\n    _sev_detail_v6 = np.array(_sev_detail, dtype=object)\n    pathology_classes_v6 = [\n        c for c in np.unique(_sev_labels)\n        if c not in (\"No Finding\", \"Normal\", \"\")\n    ]\n\n    for disease in pathology_classes_v6:\n        mask  = (_sev_labels == disease)\n        feats = _sev_f128[mask]\n        if len(feats) < 3:\n            print(f\"  Skipping '{disease}' — only {len(feats)} samples.\")\n            continue\n        km        = KMeans(n_clusters=3, random_state=42, n_init=10)\n        sub_lbls  = km.fit_predict(feats)\n        centroids = km.cluster_centers_\n        dists_to_healthy = [euclidean(c, _healthy_centroid_v6) for c in centroids]\n        sorted_idx = np.argsort(dists_to_healthy)\n        grade_map  = {\n            sorted_idx[0]: f\"Mild {disease}\",\n            sorted_idx[1]: f\"Moderate {disease}\",\n            sorted_idx[2]: f\"Severe {disease}\",\n        }\n        # Use pos (not i) to avoid shadowing the outer enumerate variable\n        indices = np.where(mask)[0]\n        for pos, lbl in zip(indices, sub_lbls):\n            _sev_detail_v6[pos] = grade_map[lbl]\n        print(f\"  Graded: {disease}\")\n\nelse:\n    _healthy_mask = (\n        (_sev_labels == \"No Finding\") |\n        (_sev_detail == \"Normal\")     |\n        (_sev_detail == \"No Finding\")\n    )\n    _healthy_centroid_v6 = np.mean(_sev_f128[_healthy_mask], axis=0)\n    # Cast to object so any downstream per-element writes are safe\n    _sev_detail_v6 = np.array(_sev_detail, dtype=object)\n\nprint()\n\n# ── Build records ─────────────────────────────────────────────────────────────\nsev_records = []\nfor disease in sorted(set(_sev_labels)):\n    if disease in (\"No Finding\", \"Normal\", \"\"):\n        continue\n    for grade_prefix in (\"Mild\", \"Moderate\", \"Severe\"):\n        grade_label = f\"{grade_prefix} {disease}\"\n        mask = (_sev_detail_v6 == grade_label)\n        if mask.sum() == 0:\n            continue\n        feats = _sev_f128[mask]\n        dists = np.linalg.norm(feats - _healthy_centroid_v6, axis=1)\n        sev_records.append({\n            \"disease\":   disease,\n            \"grade\":     grade_prefix,\n            \"count\":     int(mask.sum()),\n            \"mean_dist\": float(dists.mean()),\n            \"std_dist\":  float(dists.std()),\n        })\n\nif not sev_records:\n    print(\"  Still no severity records found. Check that Phase 8 / 13C ran correctly.\")\nelse:\n    sev_df = pd.DataFrame(sev_records)\n    print(f\"  Built severity records for {sev_df['disease'].nunique()} diseases.\")\n    print()\n\n    # ── Plot ──────────────────────────────────────────────────────────────────\n    diseases_with_severity = sev_df[\"disease\"].unique()\n    ncols = 3\n    nrows = int(np.ceil(len(diseases_with_severity) / ncols))\n    fig, axes = plt.subplots(nrows, ncols,\n                             figsize=(ncols * 5, nrows * 3.5), squeeze=False)\n    grade_colors = {\"Mild\": \"#2ecc71\", \"Moderate\": \"#f39c12\", \"Severe\": \"#e74c3c\"}\n    grade_order  = [\"Mild\", \"Moderate\", \"Severe\"]\n\n    for ax_idx, disease in enumerate(diseases_with_severity):\n        ax  = axes[ax_idx // ncols][ax_idx % ncols]\n        sub = sev_df[sev_df[\"disease\"] == disease]\n        sub = sub.set_index(\"grade\").reindex(\n            [g for g in grade_order if g in sub[\"grade\"].values]\n        ).reset_index()\n        grades = sub[\"grade\"].tolist()\n        means  = sub[\"mean_dist\"].tolist()\n        stds   = sub[\"std_dist\"].tolist()\n        colors = [grade_colors[g] for g in grades]\n        bars   = ax.bar(grades, means, yerr=stds, capsize=4,\n                        color=colors, edgecolor=\"white\", linewidth=0.5)\n        ax.set_title(disease, fontsize=9, fontweight=\"bold\")\n        ax.set_ylabel(\"Dist from healthy\", fontsize=8)\n        ax.set_ylim(0, None)\n        max_mean = max(means) if means else 1.0\n        for bar, n in zip(bars, sub[\"count\"].tolist()):\n            ax.text(bar.get_x() + bar.get_width() / 2,\n                    bar.get_height() + max_mean * 0.02,\n                    f\"n={n}\", ha=\"center\", va=\"bottom\", fontsize=7)\n\n    for ax_idx in range(len(diseases_with_severity), nrows * ncols):\n        axes[ax_idx // ncols][ax_idx % ncols].set_visible(False)\n\n    plt.suptitle(\n        \"V6 · Severity Gradient — Distance from Healthy Anchor (SSL 128-d)\",\n        fontweight=\"bold\",\n    )\n    plt.tight_layout()\n    plt.savefig(PLOT_DIR / \"v6_severity_gradient.png\", dpi=150, bbox_inches=\"tight\")\n    plt.show()\n\n    print(sev_df.to_string(index=False))\n    print()\n    print(\"  Saved → v6_severity_gradient.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:24.294197Z","iopub.execute_input":"2026-04-01T07:56:24.294508Z","iopub.status.idle":"2026-04-01T07:56:25.465340Z","shell.execute_reply.started":"2026-04-01T07:56:24.294483Z","shell.execute_reply":"2026-04-01T07:56:25.464610Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V7 — Demographic Conditioning Coverage","metadata":{}},{"cell_type":"code","source":"# ── V7: DEMOGRAPHIC COVERAGE (DETAIL) ────────────────────────────────────────\nprint(\"[ V7 ] Demographic coverage per disease class ...\")\n\n_demo_cols = [c for c in [\"sex\", \"age\", \"view_position\", \"ap_pa\"]\n              if c in _meta.columns]\n\nif not _demo_cols:\n    print(\"  No demographic columns available in metadata — skipping V7.\")\nelse:\n    fig, axes = plt.subplots(1, len(_demo_cols),\n                             figsize=(5 * len(_demo_cols), 5), squeeze=False)\n    axes = axes[0]\n\n    for ax, col in zip(axes, _demo_cols):\n        if col == \"age\":\n            # Age distribution violin per disease\n            _age_df = _meta[[\"disease\", \"age\"]].copy() if \"disease\" in _meta.columns                       else _meta[[col]].assign(disease=_labels)\n            _age_df = _age_df.rename(columns={0: \"age\"})\n            _age_df[\"age\"] = pd.to_numeric(_age_df[\"age\"], errors=\"coerce\")\n            _age_df = _age_df.dropna(subset=[\"age\"])\n            if len(_age_df) > 0:\n                sns.violinplot(data=_age_df, x=\"disease\", y=\"age\",\n                               palette=\"muted\", ax=ax, inner=\"quartile\",\n                               scale=\"width\")\n                ax.set_xticklabels(ax.get_xticklabels(), rotation=45, ha=\"right\",\n                                   fontsize=7)\n                ax.set_title(f\"Age distribution by disease\")\n        else:\n            # Stacked bar for categorical\n            _cat = pd.DataFrame({\"disease\": _labels,\n                                  col: _meta[col].values if col in _meta.columns\n                                       else np.full(N_VIZ, \"unknown\")})\n            _piv = (_cat.groupby([\"disease\", col])\n                        .size()\n                        .unstack(fill_value=0))\n            _piv_pct = _piv.div(_piv.sum(axis=1), axis=0) * 100\n            _piv_pct.plot(kind=\"bar\", stacked=True, ax=ax, colormap=\"tab10\",\n                          edgecolor=\"white\", linewidth=0.3)\n            ax.set_xticklabels(ax.get_xticklabels(), rotation=45, ha=\"right\",\n                                fontsize=7)\n            ax.set_ylabel(\"% of class\")\n            ax.set_title(f\"{col} coverage by disease\")\n            ax.legend(fontsize=7, bbox_to_anchor=(1, 1), loc=\"upper left\")\n\n    plt.suptitle(\"V7 · Demographic Conditioning Coverage\", fontweight=\"bold\")\n    plt.tight_layout()\n    plt.savefig(PLOT_DIR / \"v7_demographic_coverage.png\", dpi=150,\n                bbox_inches=\"tight\")\n    plt.show()\n    print(\"  Saved → v7_demographic_coverage.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:25.466244Z","iopub.execute_input":"2026-04-01T07:56:25.466466Z","iopub.status.idle":"2026-04-01T07:56:26.198963Z","shell.execute_reply.started":"2026-04-01T07:56:25.466444Z","shell.execute_reply":"2026-04-01T07:56:26.198170Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V8 — Feature Drift: Raw 384-d → SSL 128-d","metadata":{}},{"cell_type":"code","source":"# ── V8: FEATURE DRIFT ─────────────────────────────────────────────────────────\nprint(\"[ V8 ] Measuring feature drift: raw DINOv2 384-d vs SSL-adapted 128-d ...\")\n\n# Project 384-d to 128-d via PCA for a fair pair-wise comparison\nfrom sklearn.decomposition import PCA\n\npca = PCA(n_components=128, random_state=42)\n_f384_proj = pca.fit_transform(_f384)\n_f384_proj /= (np.linalg.norm(_f384_proj, axis=1, keepdims=True) + 1e-9)\n\n# Per-class centroid drift\ndrift_records = []\nfor cls in sorted(set(_labels)):\n    mask = (_labels == cls)\n    c_raw = _f384_proj[mask].mean(axis=0)\n    c_ssl = _f128[mask].mean(axis=0)\n    # Cosine similarity between centroids\n    cos = float(np.dot(c_raw, c_ssl) / (\n        np.linalg.norm(c_raw) * np.linalg.norm(c_ssl) + 1e-9))\n    dist = float(np.linalg.norm(c_raw - c_ssl))\n    drift_records.append({\"class\": cls, \"cosine_sim\": cos, \"l2_drift\": dist})\n\ndrift_df = pd.DataFrame(drift_records).sort_values(\"cosine_sim\")\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\n# Cosine similarity\ncolor_cos = [\"#e74c3c\" if c < 0.5 else \"#2ecc71\" for c in drift_df[\"cosine_sim\"]]\naxes[0].barh(drift_df[\"class\"], drift_df[\"cosine_sim\"], color=color_cos)\naxes[0].axvline(0.5, linestyle=\"--\", color=\"gray\", linewidth=1)\naxes[0].set_xlabel(\"Cosine similarity (PCA-384 centroid vs SSL-128 centroid)\")\naxes[0].set_title(\"Centroid alignment after SSL adaptation\")\naxes[0].set_xlim(0, 1.05)\n\n# L2 drift\naxes[1].barh(drift_df[\"class\"], drift_df[\"l2_drift\"], color=\"steelblue\")\naxes[1].set_xlabel(\"L2 distance between centroids\")\naxes[1].set_title(\"Centroid L2 drift after SSL adaptation\")\n\nplt.suptitle(\"V8 · Feature Drift — Raw DINOv2 vs SSL-adapted Embeddings\",\n             fontweight=\"bold\")\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v8_feature_drift.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\nprint()\nprint(drift_df.to_string(index=False))\nprint()\nprint(\"  Saved → v8_feature_drift.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:26.199971Z","iopub.execute_input":"2026-04-01T07:56:26.200194Z","iopub.status.idle":"2026-04-01T07:56:27.160487Z","shell.execute_reply.started":"2026-04-01T07:56:26.200171Z","shell.execute_reply":"2026-04-01T07:56:27.159460Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V9 — KNN Purity & Intra/Inter-class Distance Audit","metadata":{}},{"cell_type":"code","source":"# ── V9: KNN PURITY ────────────────────────────────────────────────────────────\nprint(\"[ V9 ] KNN purity and intra/inter-class distance audit ...\")\n\n_N_KNN = min(5_000, N_VIZ)\nrng_knn = np.random.default_rng(0)\n_kidx   = rng_knn.choice(N_VIZ, _N_KNN, replace=False)\n_kX     = _f128[_kidx]\n_ky     = _labels[_kidx]\n\nle   = LabelEncoder()\n_ky_enc = le.fit_transform(_ky)\n\n# 5-NN leave-one-out approximate (train on full, eval stratified 20% split)\nfrom sklearn.model_selection import StratifiedShuffleSplit\nsss = StratifiedShuffleSplit(n_splits=1, test_size=0.2, random_state=42)\ntrain_idx, test_idx = next(sss.split(_kX, _ky_enc))\n\nknn = KNeighborsClassifier(n_neighbors=5, metric=\"cosine\", n_jobs=-1)\nknn.fit(_kX[train_idx], _ky_enc[train_idx])\n_kpred = knn.predict(_kX[test_idx])\n\nprint(\"  5-NN Classification Report (SSL 128-d, cosine metric):\")\nprint()\nprint(classification_report(\n    _ky_enc[test_idx], _kpred,\n    target_names=le.classes_, zero_division=0\n))\n\n# ── Intra vs inter-class distance bar chart ────────────────────────────────\nintra_dists, inter_dists = [], []\nunique_cls = sorted(set(_ky))\n\nfor cls in unique_cls:\n    cls_feats = _kX[_ky == cls]\n    other_feats = _kX[_ky != cls]\n    if len(cls_feats) < 2 or len(other_feats) < 2:\n        continue\n    # Sample for speed\n    _samp = min(200, len(cls_feats))\n    _osamp = min(200, len(other_feats))\n    ci = cls_feats[np.random.choice(len(cls_feats), _samp, replace=False)]\n    oi = other_feats[np.random.choice(len(other_feats), _osamp, replace=False)]\n    intra = cdist(ci, ci, metric=\"cosine\").mean()\n    inter = cdist(ci, oi, metric=\"cosine\").mean()\n    intra_dists.append((cls, intra))\n    inter_dists.append((cls, inter))\n\n_cls_names = [x[0] for x in intra_dists]\n_intra_v   = [x[1] for x in intra_dists]\n_inter_v   = [x[1] for x in inter_dists]\n\nx_pos = np.arange(len(_cls_names))\nfig, ax = plt.subplots(figsize=(13, 5))\nax.bar(x_pos - 0.2, _intra_v, 0.4, label=\"Intra-class\", color=\"#3498db\", alpha=0.85)\nax.bar(x_pos + 0.2, _inter_v, 0.4, label=\"Inter-class\", color=\"#e74c3c\", alpha=0.85)\nax.set_xticks(x_pos)\nax.set_xticklabels(_cls_names, rotation=45, ha=\"right\", fontsize=8)\nax.set_ylabel(\"Mean cosine distance\")\nax.set_title(\"V9 · Intra vs Inter-class Cosine Distance (SSL 128-d)\",\n             fontweight=\"bold\")\nax.legend()\nplt.tight_layout()\nplt.savefig(PLOT_DIR / \"v9_knn_intra_inter_dist.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → v9_knn_intra_inter_dist.png\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:27.161574Z","iopub.execute_input":"2026-04-01T07:56:27.161874Z","iopub.status.idle":"2026-04-01T07:56:27.726381Z","shell.execute_reply.started":"2026-04-01T07:56:27.161842Z","shell.execute_reply":"2026-04-01T07:56:27.725584Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### V10 — GAN-Readiness Checklist","metadata":{}},{"cell_type":"code","source":"# ── V10: GAN-READINESS CHECKLIST ──────────────────────────────────────────────\nprint(\"[ V10 ] Evaluating GAN-readiness of all output arrays ...\")\nprint()\n\ndef _chk(name, val, note=\"\"):\n    sym = \"✅\" if val else \"❌\"\n    print(f\"  {sym}  {name}\" + (f\"  [{note}]\" if note else \"\"))\n    return int(val)\n\n_f128_chk   = features_128d_full if \"features_128d_full\" in dir() else features_128d\n_f384_chk   = features_384d_full if \"features_384d_full\" in dir() else features_384d\n_lbls_chk   = labels_full        if \"labels_full\"        in dir() else labels_arr\n_detail_chk = (detailed_labels_full if \"detailed_labels_full\" in dir()\n               else (detailed_labels if \"detailed_labels\" in dir() else np.array([])))\n_fused_chk  = _fused_full if \"_fused_full\" in dir() else None\n_meta_chk   = meta_df_full if \"meta_df_full\" in dir() else meta_df\n_dis_oh     = _disease_onehot_full if \"_disease_onehot_full\" in dir() else None\n_sev_oh     = _severity_onehot_full if \"_severity_onehot_full\" in dir() else None\n\nchecks_passed = 0\nchecks_total  = 0\n\nprint(\"  ── Array integrity ──────────────────────────────────────────────────\")\nchecks_total += 4\nchecks_passed += _chk(\"features_128d: no NaN/Inf\",\n    not (np.isnan(_f128_chk).any() or np.isinf(_f128_chk).any()))\nchecks_passed += _chk(\"features_384d: no NaN/Inf\",\n    not (np.isnan(_f384_chk).any() or np.isinf(_f384_chk).any()))\nchecks_passed += _chk(\"features_128d L2-normalised\",\n    np.isclose(np.linalg.norm(_f128_chk, axis=1).mean(), 1.0, atol=5e-3),\n    f\"mean norm={np.linalg.norm(_f128_chk, axis=1).mean():.4f}\")\nchecks_passed += _chk(\"features_384d L2-normalised\",\n    np.isclose(np.linalg.norm(_f384_chk, axis=1).mean(), 1.0, atol=5e-3),\n    f\"mean norm={np.linalg.norm(_f384_chk, axis=1).mean():.4f}\")\n\nprint()\nprint(\"  ── Disease class conditioning ───────────────────────────────────────\")\nchecks_total += 3\nif _dis_oh is not None:\n    _D = _dis_oh.shape[1]\n    checks_passed += _chk(f\"Disease one-hot present ({_D} classes)\", True, f\"shape={_dis_oh.shape}\")\n    checks_passed += _chk(\"Disease one-hot sums to 1.0 per row\",\n        np.allclose(_dis_oh.sum(axis=1), 1.0))\n    checks_passed += _chk(\"Disease one-hot no NaN/Inf\",\n        not (np.isnan(_dis_oh).any() or np.isinf(_dis_oh).any()))\nelse:\n    for _ in range(3): _chk(\"Disease one-hot present\", False, \"Phase 13D not completed\")\n\nprint()\nprint(\"  ── Severity conditioning ────────────────────────────────────────────\")\nchecks_total += 3\nif _sev_oh is not None:\n    checks_passed += _chk(\"Severity one-hot present (4-d)\", True, f\"shape={_sev_oh.shape}\")\n    checks_passed += _chk(\"Severity one-hot sums to 1.0 per row\",\n        np.allclose(_sev_oh.sum(axis=1), 1.0))\n    _sev_dist = dict(zip([\"Normal\",\"Mild\",\"Moderate\",\"Severe\"], _sev_oh.sum(axis=0).astype(int).tolist()))\n    checks_passed += _chk(f\"All 4 severity grades present {_sev_dist}\",\n        all(v > 0 for v in _sev_dist.values()))\nelse:\n    for _ in range(3): _chk(\"Severity one-hot present\", False, \"Phase 13D not completed\")\n\nprint()\nprint(\"  ── Demographic conditioning ─────────────────────────────────────────\")\nchecks_total += 3\n_demo_cols_present = [c for c in [\"age\",\"sex\",\"view_position\",\"ap_pa\"] if c in _meta_chk.columns]\nchecks_passed += _chk(f\"Demographic columns present {_demo_cols_present}\",\n    len(_demo_cols_present) > 0)\nif _demo_cols_present:\n    _unk_pct = _meta_chk[_demo_cols_present].isin([\"unknown\"]).mean().mean()\n    checks_passed += _chk(f\"Unknown demographics < 50% (actual {_unk_pct:.0%})\", _unk_pct < 0.5)\nelse:\n    _chk(\"Unknown demographics < 50%\", False, \"no demo cols found\")\n    checks_total -= 1\nif _fused_chk is not None:\n    checks_passed += _chk(\"demo_embedding_64d present in fused vector\",\n        _fused_chk.shape[1] >= 192)\nelse:\n    _chk(\"demo_embedding_64d present\", False, \"fused not built\")\n\nprint()\nprint(\"  ── Full conditioning vector ─────────────────────────────────────────\")\nchecks_total += 3\nif _fused_chk is not None:\n    _expected_dim = 128 + (_dis_oh.shape[1] if _dis_oh is not None else 0) + 4 + 64\n    checks_passed += _chk(f\"Fused dim == {_expected_dim} (128+D+4+64)\",\n        _fused_chk.shape[1] == _expected_dim, f\"actual={_fused_chk.shape[1]}\")\n    checks_passed += _chk(\"Fused vector no NaN/Inf\",\n        not (np.isnan(_fused_chk).any() or np.isinf(_fused_chk).any()))\n    checks_passed += _chk(\"Fused rows == N\",\n        _fused_chk.shape[0] == _f128_chk.shape[0])\nelse:\n    for _ in range(3): _chk(\"Fused conditioning present\", False, \"Phase 13D not completed\")\n\nprint()\nprint(\"  ── Saved artefacts ──────────────────────────────────────────────────\")\nexpected_files = [\n    \"features_384d_full.npy\", \"features_128d_full.npy\",\n    \"labels_full.npy\",        \"labels_detailed_full.npy\",\n    \"fused_conditioning_full.npy\", \"conditioning_meta.json\",\n    \"cgan_input_full.parquet\", \"metadata_full.csv\",\n]\nchecks_total += len(expected_files)\nfor fname in expected_files:\n    fp = OUTPUT_DIR / fname\n    exists = fp.exists()\n    checks_passed += _chk(f\"{fname} exists\", exists,\n        f\"{fp.stat().st_size/1e6:.1f} MB\" if exists else \"MISSING\")\n\nprint()\nprint(\"=\" * 66)\nprint(f\"  GAN-READINESS SCORE : {checks_passed} / {checks_total}\")\nif checks_passed == checks_total:\n    print(\"  🎉  All checks passed — arrays are ready for the cGAN pipeline.\")\nelif checks_passed >= checks_total * 0.85:\n    print(\"  ⚠️   Minor issues — review ❌ items above before handoff.\")\nelse:\n    print(\"  🚨  Critical issues — do not hand off until ❌ items are resolved.\")\nprint(\"=\" * 66)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:27.727762Z","iopub.execute_input":"2026-04-01T07:56:27.728458Z","iopub.status.idle":"2026-04-01T07:56:28.118838Z","shell.execute_reply.started":"2026-04-01T07:56:27.728429Z","shell.execute_reply":"2026-04-01T07:56:28.118109Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase T — t-SNE Evaluation\n\nEvaluates learned representations via PCA reduction followed by t-SNE projection to 2D. \nTwo diagnostic plots are generated:\n- **Plot T1** — coloured by disease label → checks for pathology clustering\n- **Plot T2** — coloured by source dataset → checks for dataset-bias leakage\n\nA silhouette score on the PCA-50 space gives a concise quantitative summary.","metadata":{}},{"cell_type":"markdown","source":"### T1 — Feature Selection & PCA-50 Reduction","metadata":{}},{"cell_type":"code","source":"# ── T1: FEATURE SELECTION & PCA-50 REDUCTION ─────────────────────────────────\n# Uses SSL-adapted 128-d embeddings. Prefers full-scale arrays (Phase 13),\n# falls back to sample-set (Phase 6). Stratified sub-sampling ensures every\n# disease class AND severity grade is represented in the plot.\nfrom sklearn.decomposition import PCA\nfrom sklearn.preprocessing import normalize\n\nprint(\"[ T1 ] Preparing features for t-SNE ...\")\nprint()\n\n# ── Resolve arrays ─────────────────────────────────────────────────────────────\nif \"features_128d_full\" in dir() and features_128d_full is not None and len(features_128d_full) > 0:\n    _tsne_feats    = features_128d_full.copy()\n    _tsne_labels   = labels_full.copy()\n    _tsne_detail   = detailed_labels_full.copy() if \"detailed_labels_full\" in dir() else _tsne_labels.copy()\n    _tsne_datasets = datasets_full.copy() if \"datasets_full\" in dir() else None\n    _tsne_sev_idx  = np.array([\n        {\"Normal\":0,\"Mild\":1,\"Moderate\":2,\"Severe\":3}.get(\n            \"Normal\" if str(l).lower()==\"normal\"\n            else \"Mild\" if str(l).lower().startswith(\"mild\")\n            else \"Moderate\" if str(l).lower().startswith(\"moderate\")\n            else \"Severe\" if str(l).lower().startswith(\"severe\")\n            else \"Normal\", 0)\n        for l in _tsne_detail\n    ])\n    print(f\"  Source : full-scale SSL 128-d  (shape {_tsne_feats.shape})\")\nelse:\n    _tsne_feats    = features_128d.copy()\n    _tsne_labels   = labels_arr.copy()\n    _tsne_detail   = detailed_labels.copy() if \"detailed_labels\" in dir() else _tsne_labels.copy()\n    _tsne_datasets = datasets_arr.copy() if \"datasets_arr\" in dir() else None\n    _tsne_sev_idx  = np.zeros(len(_tsne_labels), dtype=int)\n    print(f\"  Source : sample-set SSL 128-d  (shape {_tsne_feats.shape})\")\n\n# ── L2 normalisation ──────────────────────────────────────────────────────────\n_tsne_feats = normalize(_tsne_feats, norm=\"l2\")\nprint(\"  L2 normalisation applied.\")\n\n# ── Stratified sub-sample — balanced across both disease classes AND severity ──\nTSNE_MAX_SAMPLES = 5_000   # increased from 3000 for better coverage\nN_total = len(_tsne_feats)\n\nif N_total > TSNE_MAX_SAMPLES:\n    rng_tsne    = np.random.default_rng(42)\n    unique_cls  = np.unique(_tsne_labels)\n    per_class   = max(10, TSNE_MAX_SAMPLES // len(unique_cls))\n    sel_idx     = []\n    for cls in unique_cls:\n        cls_idx = np.where(_tsne_labels == cls)[0]\n        chosen  = rng_tsne.choice(cls_idx, min(per_class, len(cls_idx)), replace=False)\n        sel_idx.append(chosen)\n    sel_idx = np.concatenate(sel_idx)\n    rng_tsne.shuffle(sel_idx)\n\n    _tsne_feats    = _tsne_feats[sel_idx]\n    _tsne_labels   = _tsne_labels[sel_idx]\n    _tsne_detail   = np.array(_tsne_detail)[sel_idx]\n    _tsne_sev_idx  = _tsne_sev_idx[sel_idx]\n    if _tsne_datasets is not None:\n        _tsne_datasets = np.array(_tsne_datasets)[sel_idx]\n    print(f\"  Stratified sub-sample: {N_total:,} → {len(_tsne_feats):,}\")\n    print(f\"  Per class cap: {per_class}  |  Classes: {len(unique_cls)}\")\nelse:\n    print(f\"  Using all {N_total:,} samples\")\n\n# ── PCA-50 ────────────────────────────────────────────────────────────────────\nn_components = min(50, _tsne_feats.shape[1], len(_tsne_feats) - 1)\npca_tsne   = PCA(n_components=n_components, random_state=42)\n_tsne_pca50 = pca_tsne.fit_transform(_tsne_feats)\nexplained   = pca_tsne.explained_variance_ratio_.sum()\n\nprint()\nprint(f\"  PCA-{n_components} output shape : {_tsne_pca50.shape}\")\nprint(f\"  Variance explained  : {explained:.1%}\")\nprint()\nprint(\"  Label summary:\")\nfor cls in np.unique(_tsne_labels):\n    n = (np.array(_tsne_labels) == cls).sum()\n    print(f\"    {cls:<30}: {n:>5} samples\")\nprint()\nprint(\"  ✔  T1 complete — PCA features ready for t-SNE\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:28.120591Z","iopub.execute_input":"2026-04-01T07:56:28.120882Z","iopub.status.idle":"2026-04-01T07:56:28.288401Z","shell.execute_reply.started":"2026-04-01T07:56:28.120860Z","shell.execute_reply":"2026-04-01T07:56:28.287665Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T2 — t-SNE Projection (2D)","metadata":{}},{"cell_type":"code","source":"# ── T2: t-SNE 2D PROJECTION ──────────────────────────────────────────────────\nfrom sklearn.manifold import TSNE\n\nprint(\"[ T2 ] Running t-SNE on PCA features ...\")\nprint()\n\n# Perplexity scales with dataset size. Rule of thumb: sqrt(N)/3, clamped 15-50.\n_N_tsne    = len(_tsne_pca50)\n_perplexity = int(np.clip(np.sqrt(_N_tsne) / 3, 15, 50))\nprint(f\"  Samples       : {_N_tsne:,}\")\nprint(f\"  Perplexity    : {_perplexity}  (auto-scaled to sqrt(N)/3, clamped 15-50)\")\nprint(f\"  n_iter        : 1500  (increased from 1000 for better convergence)\")\nprint(f\"  init          : pca   (deterministic, avoids random restarts)\")\nprint()\n\n_t0 = time.time()\n\ntsne = TSNE(\n    n_components       = 2,\n    perplexity         = _perplexity,\n    n_iter             = 1500,        # was 1000 — more iterations = better convergence\n    early_exaggeration = 12,          # default=12, keeps well-separated clusters\n    init               = \"pca\",       # deterministic + faster than random\n    learning_rate      = \"auto\",\n    random_state       = 42,\n    n_jobs             = -1,\n)\n_tsne_2d = tsne.fit_transform(_tsne_pca50)\n\n_elapsed = time.time() - _t0\nprint(f\"  t-SNE output shape : {_tsne_2d.shape}\")\nprint(f\"  KL divergence      : {tsne.kl_divergence_:.4f}\")\nprint(f\"  Wall time          : {_elapsed:.1f} s\")\nprint()\n\n# KL divergence quality guide\nif tsne.kl_divergence_ < 0.5:\n    print(\"  ✅  KL divergence < 0.5  — very good projection quality\")\nelif tsne.kl_divergence_ < 1.0:\n    print(\"  ✅  KL divergence < 1.0  — good projection quality\")\nelif tsne.kl_divergence_ < 2.0:\n    print(\"  ⚠️   KL divergence < 2.0  — acceptable; consider more iterations\")\nelse:\n    print(\"  ❌  KL divergence > 2.0  — increase n_iter or adjust perplexity\")\nprint()\nprint(\"  ✔  T2 complete — 2-D embedding ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:56:28.289423Z","iopub.execute_input":"2026-04-01T07:56:28.289695Z","iopub.status.idle":"2026-04-01T07:57:01.163535Z","shell.execute_reply.started":"2026-04-01T07:56:28.289654Z","shell.execute_reply":"2026-04-01T07:57:01.162662Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T3 — Plot 1: Disease-Coloured t-SNE","metadata":{}},{"cell_type":"code","source":"# ── T3: PLOT 1 — DISEASE-COLOURED t-SNE ─────────────────────────────────────\nprint(\"[ T3 ] Generating disease-coloured t-SNE plot ...\")\n\n_DISEASE_PALETTE = [\n    \"#4C72B0\",\"#DD8452\",\"#55A868\",\"#C44E52\",\"#8172B2\",\n    \"#937860\",\"#DA8BC3\",\"#8C8C8C\",\"#CCB974\",\"#64B5CD\",\n]\nunique_diseases = sorted(np.unique(_tsne_labels))\n_d_color_map    = {d: _DISEASE_PALETTE[i % len(_DISEASE_PALETTE)]\n                   for i, d in enumerate(unique_diseases)}\n\nfig, ax = plt.subplots(figsize=(11, 8))\nfor disease in unique_diseases:\n    mask = np.array(_tsne_labels) == disease\n    ax.scatter(_tsne_2d[mask, 0], _tsne_2d[mask, 1],\n               c=_d_color_map[disease],\n               label=f\"{disease} (n={mask.sum():,})\",\n               s=12, alpha=0.65, linewidths=0)\n\nax.set_title(\"T3 · t-SNE — Disease Label Clustering\\n\"\n             \"DINOv2 SSL 128-d → PCA → t-SNE 2D\",\n             fontsize=13, fontweight=\"bold\", pad=12)\nax.set_xlabel(\"t-SNE Dim 1\", fontsize=10)\nax.set_ylabel(\"t-SNE Dim 2\", fontsize=10)\nax.legend(title=\"Disease class\", bbox_to_anchor=(1.02, 1), loc=\"upper left\",\n          framealpha=0.9, fontsize=8, title_fontsize=9)\nax.grid(True, linestyle=\"--\", alpha=0.25)\nplt.tight_layout()\nplt.savefig(OUTPUT_DIR / \"tsne_disease.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → tsne_disease.png\")\nprint(\"  ✅ Good: clusters separate by disease  |  ⚠️ Bad: all mixed together\")\nprint()\n\n# ── T3b: SEVERITY-COLOURED t-SNE ─────────────────────────────────────────────\nprint(\"[ T3b ] Generating severity-coloured t-SNE plot ...\")\n\n_SEV_COLOR = {\"Normal\":\"#2ecc71\",\"Mild\":\"#f1c40f\",\"Moderate\":\"#e67e22\",\"Severe\":\"#e74c3c\"}\n\ndef _label_to_grade(lbl):\n    l = str(lbl).lower()\n    if l == \"normal\":            return \"Normal\"\n    if l.startswith(\"mild\"):     return \"Mild\"\n    if l.startswith(\"moderate\"): return \"Moderate\"\n    if l.startswith(\"severe\"):   return \"Severe\"\n    return \"Normal\"\n\n_tsne_grades = np.array([_label_to_grade(l) for l in _tsne_detail])\n\nfig, ax = plt.subplots(figsize=(11, 8))\nfor grade in [\"Normal\",\"Mild\",\"Moderate\",\"Severe\"]:\n    mask = _tsne_grades == grade\n    if mask.sum() == 0:\n        continue\n    ax.scatter(_tsne_2d[mask, 0], _tsne_2d[mask, 1],\n               c=_SEV_COLOR[grade],\n               label=f\"{grade} (n={mask.sum():,})\",\n               s=12, alpha=0.65, linewidths=0)\n\nax.set_title(\"T3b · t-SNE — Severity Grade Clustering\\n\"\n             \"DINOv2 SSL 128-d → PCA → t-SNE 2D\",\n             fontsize=13, fontweight=\"bold\", pad=12)\nax.set_xlabel(\"t-SNE Dim 1\", fontsize=10)\nax.set_ylabel(\"t-SNE Dim 2\", fontsize=10)\nax.legend(title=\"Severity\", bbox_to_anchor=(1.02, 1), loc=\"upper left\",\n          framealpha=0.9, fontsize=9, title_fontsize=9)\nax.grid(True, linestyle=\"--\", alpha=0.25)\nplt.tight_layout()\nplt.savefig(OUTPUT_DIR / \"tsne_severity.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"  Saved → tsne_severity.png\")\nprint(\"  ✅ Good: severity grades gradient from Normal → Severe within each cluster\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:57:01.164744Z","iopub.execute_input":"2026-04-01T07:57:01.165156Z","iopub.status.idle":"2026-04-01T07:57:02.507927Z","shell.execute_reply.started":"2026-04-01T07:57:01.165126Z","shell.execute_reply":"2026-04-01T07:57:02.507248Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T4 — Plot 2: Dataset-Coloured t-SNE","metadata":{}},{"cell_type":"code","source":"# ── T4: PLOT 2 — DATASET-COLOURED t-SNE ──────────────────────────────────────\nprint(\"[ T4 ] Generating dataset-coloured t-SNE plot ...\")\n\n# Known source datasets\n_DATASET_PALETTE = [\n    '#E6194B', '#3CB44B', '#4363D8', '#F58231', '#911EB4',\n    '#42D4F4', '#F032E6', '#BFEF45', '#FABED4', '#469990',\n]\n\nif _tsne_datasets is not None:\n    unique_dsets    = sorted(np.unique(_tsne_datasets))\n    _ds_color_map   = {d: _DATASET_PALETTE[i % len(_DATASET_PALETTE)]\n                       for i, d in enumerate(unique_dsets)}\n\n    fig, ax = plt.subplots(figsize=(10, 8))\n\n    for dset in unique_dsets:\n        mask = _tsne_datasets == dset\n        ax.scatter(\n            _tsne_2d[mask, 0], _tsne_2d[mask, 1],\n            c     = _ds_color_map[dset],\n            label = f\"{dset} (n={mask.sum():,})\",\n            s     = 12,\n            alpha = 0.65,\n            linewidths = 0,\n        )\n\n    ax.set_title(\"t-SNE — Source Dataset Clustering\\n\"\n                 \"(DINOv2 SSL 128-d → PCA-50 → t-SNE 2D)\",\n                 fontsize=14, fontweight='bold', pad=12)\n    ax.set_xlabel(\"t-SNE Dimension 1\", fontsize=11)\n    ax.set_ylabel(\"t-SNE Dimension 2\", fontsize=11)\n    ax.legend(\n        title='Source dataset',\n        bbox_to_anchor=(1.02, 1), loc='upper left',\n        framealpha=0.9, fontsize=9, title_fontsize=10,\n    )\n    ax.tick_params(labelsize=9)\n    ax.grid(True, linestyle='--', alpha=0.3)\n\n    fig.tight_layout()\n    plt.savefig(OUTPUT_DIR / 'tsne_dataset.png', dpi=150, bbox_inches='tight')\n    plt.show()\n    print(\"  Saved → tsne_dataset.png\")\n    print()\n    print(\"  Interpretation guide:\")\n    print(\"    ✅ Good: dataset colours are mixed within disease clusters → no bias\")\n    print(\"    ⚠  Bad : clusters align by dataset source               → dataset bias\")\nelse:\n    print(\"  ⚠  _tsne_datasets is None — skipping dataset plot.\")\n    print(\"     Ensure 'datasets_arr' or '_datasets' is available from Phase 6/13.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:57:02.508804Z","iopub.execute_input":"2026-04-01T07:57:02.509020Z","iopub.status.idle":"2026-04-01T07:57:03.229426Z","shell.execute_reply.started":"2026-04-01T07:57:02.508999Z","shell.execute_reply":"2026-04-01T07:57:03.228653Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T5 — Silhouette Score (Quantitative Check)","metadata":{}},{"cell_type":"code","source":"# ── T5: SILHOUETTE SCORE ──────────────────────────────────────────────────────\nfrom sklearn.metrics import silhouette_score\nfrom sklearn.preprocessing import LabelEncoder as _LE\n\nprint(\"[ T5 ] Computing silhouette score on PCA-50 features ...\")\nprint(\"       (Supporting evidence only — not the primary claim)\")\nprint()\n\n# Encode string labels to integers\n_le_tsne = _LE()\n_y_enc   = _le_tsne.fit_transform(_tsne_labels)\n\n# Sub-sample for speed if very large\n_SIL_MAX = 5_000\n_N_sil   = len(_tsne_pca50)\nif _N_sil > _SIL_MAX:\n    rng_sil  = np.random.default_rng(42)\n    _sil_idx = rng_sil.choice(_N_sil, _SIL_MAX, replace=False)\n    _X_sil   = _tsne_pca50[_sil_idx]\n    _y_sil   = _y_enc[_sil_idx]\n    print(f\"  Sub-sampled to {_SIL_MAX:,} for silhouette computation.\")\nelse:\n    _X_sil, _y_sil = _tsne_pca50, _y_enc\n\n_sil_score = silhouette_score(_X_sil, _y_sil, metric='euclidean', random_state=42)\n\nprint(f\"  Silhouette score (PCA-50, euclidean) : {_sil_score:.4f}\")\nprint()\nprint(\"  Scale reference:\")\nprint(\"    > 0.30  Strong clustering — pathology separability is high\")\nprint(\"    0.10–0.30  Moderate — some class structure present\")\nprint(\"    < 0.10  Weak — overlapping representations; inspect t-SNE for bias\")\nprint()\n\n# ── Per-class silhouette breakdown ────────────────────────────────────────────\nfrom sklearn.metrics import silhouette_samples\n\n_sil_samples = silhouette_samples(_X_sil, _y_sil, metric='euclidean')\n\nprint(\"  Per-class silhouette breakdown:\")\nprint(f\"  {'Class':<22}  {'n':>5}  {'Mean sil':>9}  {'Median sil':>11}\")\nprint(\"  \" + \"-\" * 53)\nfor cls_idx, cls_name in enumerate(_le_tsne.classes_):\n    mask_cls = _y_sil == cls_idx\n    sil_cls  = _sil_samples[mask_cls]\n    print(f\"  {cls_name:<22}  {mask_cls.sum():>5}  \"\n          f\"{sil_cls.mean():>+9.4f}  {np.median(sil_cls):>+11.4f}\")\n\nprint()\nprint(\"  ✔  T5 complete\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:57:03.230462Z","iopub.execute_input":"2026-04-01T07:57:03.230707Z","iopub.status.idle":"2026-04-01T07:57:03.823050Z","shell.execute_reply.started":"2026-04-01T07:57:03.230678Z","shell.execute_reply":"2026-04-01T07:57:03.822527Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### T6 — t-SNE Summary & Interpretation","metadata":{}},{"cell_type":"code","source":"# ── T6: SUMMARY & INTERPRETATION ────────────────────────────────────────────\nprint(\"═\" * 65)\nprint(\" t-SNE EVALUATION SUMMARY\")\nprint(\"═\" * 65)\nprint()\nprint(f\"  Samples visualised : {len(_tsne_2d):,}\")\nprint(f\"  Feature source     : SSL 128-d → L2-norm → PCA → t-SNE 2D\")\nprint(f\"  PCA variance kept  : {explained:.1%}\")\nprint(f\"  Perplexity used    : {_perplexity}  (auto-scaled)\")\nprint(f\"  t-SNE n_iter       : 1500\")\nprint(f\"  KL divergence      : {tsne.kl_divergence_:.4f}\")\nprint(f\"  Silhouette score   : {_sil_score:.4f}  (by disease class, PCA space)\")\nprint()\n\n# ── Scanner-bias comparison ───────────────────────────────────────────────────\nif \"_sil_disease\" in dir() and \"_sil_dataset\" in dir():\n    print(f\"  Scanner-bias audit (Phase 11):\")\n    print(f\"    Disease sil.   : {_sil_disease:+.4f}\")\n    print(f\"    Dataset sil.   : {_sil_dataset:+.4f}\")\n    _bias_ok = _sil_disease > _sil_dataset\n    print(f\"    Bias status    : {'✅ Disease > Dataset — CLAHE fix working' if _bias_ok else '⚠️  Dataset ≥ Disease — bias may remain'}\")\n    print()\n\nprint(\"  Diagnosis (t-SNE silhouette):\")\nif _sil_score > 0.30:\n    print(\"    ✅ Strong pathology clustering — model captures disease structure well.\")\n    print(\"       cGAN conditioning on disease + severity is well-grounded.\")\nelif _sil_score > 0.10:\n    print(\"    ⚠️  Moderate clustering — partial separation present.\")\n    print(\"       Inspect t-SNE plots: are disease clusters visible even if not tight?\")\n    print(\"       The cGAN conditioning is still usable but may benefit from more SSL epochs.\")\nelse:\n    print(\"    ❌ Weak clustering — representations overlap significantly.\")\n    print(\"       Check Phase 11 scanner-bias audit. Consider more SSL_EPOCHS (e.g. 20).\")\n    print(\"       The cGAN conditioning may not generalise well at this separation level.\")\n\nprint()\nprint(\"  Output plots:\")\n_plot_files = [\n    (\"tsne_disease.png\",   \"Disease class clustering  (T3)\"),\n    (\"tsne_severity.png\",  \"Severity grade clustering (T3b)\"),\n    (\"tsne_dataset.png\",   \"Scanner/dataset bias      (T4)\"),\n]\nfor fname, desc in _plot_files:\n    fp = OUTPUT_DIR / fname\n    sym = \"✅\" if fp.exists() else \"⚠️ \"\n    print(f\"    {sym} {fp}  [{desc}]\")\n\nprint()\nprint(\"═\" * 65)\nprint(\" t-SNE EVALUATION COMPLETE\")\nprint(\"═\" * 65)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T07:57:03.823695Z","iopub.execute_input":"2026-04-01T07:57:03.823902Z","iopub.status.idle":"2026-04-01T07:57:03.833267Z","shell.execute_reply.started":"2026-04-01T07:57:03.823882Z","shell.execute_reply":"2026-04-01T07:57:03.832637Z"}},"outputs":[],"execution_count":null}]}