{"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-17T11:53:37.059555Z","iopub.execute_input":"2026-04-17T11:53:37.059960Z","iopub.status.idle":"2026-04-17T11:53:52.690986Z","shell.execute_reply.started":"2026-04-17T11:53:37.059930Z","shell.execute_reply":"2026-04-17T11:53:52.690062Z"}},"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-17T11:53:52.692393Z","iopub.execute_input":"2026-04-17T11:53:52.692821Z","iopub.status.idle":"2026-04-17T11:53:52.700452Z","shell.execute_reply.started":"2026-04-17T11:53:52.692796Z","shell.execute_reply":"2026-04-17T11:53:52.699687Z"}},"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        = 100\nSSL_LR            = 5e-4\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-17T11:53:52.701542Z","iopub.execute_input":"2026-04-17T11:53:52.701886Z","iopub.status.idle":"2026-04-17T11:53:52.724686Z","shell.execute_reply.started":"2026-04-17T11:53:52.701853Z","shell.execute_reply":"2026-04-17T11:53:52.724020Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1C — Demographic Config","metadata":{}},{"cell_type":"code","source":"# ── 1C: DEMOGRAPHIC CONFIG ────────────────────────────────────────────────────\n# Core columns used for the 64-d embedding\nDEMOGRAPHIC_COLS = [\n    \"age\", \n    \"gender\",           \n    \"view\",             \n    \"original_view\",    \n    \"identity_id\"       \n]\n\n# Validation columns used to check data integrity (handling missing/noisy data)\nVALIDATION_COLS = [\n    \"has_age\",\n    \"has_gender\",\n    \"age_group\",\n    \"disease\",      # To verify pathological labels\n    \"raw_labels\"    # To see the exact original text before dataset merging\n]\n\n# Load both sets when reading the parquet\nALL_COLS_TO_LOAD = DEMOGRAPHIC_COLS + VALIDATION_COLS\n\nprint(f\"Demographic features for MLP: {DEMOGRAPHIC_COLS}\")\nprint(f\"Validation columns for QC: {VALIDATION_COLS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T11:53:52.726249Z","iopub.execute_input":"2026-04-17T11:53:52.726461Z","iopub.status.idle":"2026-04-17T11:53:52.741504Z","shell.execute_reply.started":"2026-04-17T11:53:52.726442Z","shell.execute_reply":"2026-04-17T11:53:52.740819Z"}},"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-17T11:53:52.742471Z","iopub.execute_input":"2026-04-17T11:53:52.742766Z","iopub.status.idle":"2026-04-17T11:53:52.763426Z","shell.execute_reply.started":"2026-04-17T11:53:52.742735Z","shell.execute_reply":"2026-04-17T11:53:52.762677Z"}},"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-17T11:53:52.764353Z","iopub.execute_input":"2026-04-17T11:53:52.764637Z","iopub.status.idle":"2026-04-17T11:53:55.225572Z","shell.execute_reply.started":"2026-04-17T11:53:52.764593Z","shell.execute_reply":"2026-04-17T11:53:55.224746Z"}},"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-17T11:53:55.226854Z","iopub.execute_input":"2026-04-17T11:53:55.227201Z","iopub.status.idle":"2026-04-17T11:54:14.386528Z","shell.execute_reply.started":"2026-04-17T11:53:55.227160Z","shell.execute_reply":"2026-04-17T11:54:14.385865Z"}},"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-17T11:54:14.387534Z","iopub.execute_input":"2026-04-17T11:54:14.387806Z","iopub.status.idle":"2026-04-17T11:54:14.396117Z","shell.execute_reply.started":"2026-04-17T11:54:14.387783Z","shell.execute_reply":"2026-04-17T11:54:14.395430Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3B — Clean Transform (Extraction)","metadata":{}},{"cell_type":"code","source":"# ── 3B: OPTIMISED CLEAN TRANSFORM (Extraction) ──────────────────────────────\nimport cv2\nfrom PIL import Image\n\nclass ApplyCLAHE:\n    def __call__(self, img):\n        # Convert PIL to CV2 (grayscale)\n        img_np = np.array(img.convert(\"L\"))\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        img_clahe = clahe.apply(img_np)\n        return Image.fromarray(img_clahe)\n\nclean_transform = T.Compose([\n    T.Resize(256),              # Resize first so CenterCrop is relative\n    ApplyCLAHE(),               # MANDATORY: Matches your SSL training logic\n    T.Grayscale(num_output_channels=3),\n    T.CenterCrop(224),          # Trims borders while keeping the 224 input size\n    T.ToTensor(),\n    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\n\nprint(\"Clean transform optimized: CLAHE integrated + Resolution aligned.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T11:54:14.397268Z","iopub.execute_input":"2026-04-17T11:54:14.397749Z","iopub.status.idle":"2026-04-17T11:54:14.417571Z","shell.execute_reply.started":"2026-04-17T11:54:14.397725Z","shell.execute_reply":"2026-04-17T11:54:14.416896Z"}},"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-17T11:54:14.419831Z","iopub.execute_input":"2026-04-17T11:54:14.420235Z","iopub.status.idle":"2026-04-17T11:54:14.629226Z","shell.execute_reply.started":"2026-04-17T11:54:14.420213Z","shell.execute_reply":"2026-04-17T11:54:14.628596Z"}},"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-17T11:54:14.630151Z","iopub.execute_input":"2026-04-17T11:54:14.630503Z","iopub.status.idle":"2026-04-17T11:54:14.641713Z","shell.execute_reply.started":"2026-04-17T11:54:14.630478Z","shell.execute_reply":"2026-04-17T11:54:14.640984Z"}},"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-17T11:54:14.642555Z","iopub.execute_input":"2026-04-17T11:54:14.642794Z","iopub.status.idle":"2026-04-17T11:54:14.659706Z","shell.execute_reply.started":"2026-04-17T11:54:14.642774Z","shell.execute_reply":"2026-04-17T11:54:14.659119Z"}},"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-17T11:54:14.660681Z","iopub.execute_input":"2026-04-17T11:54:14.661000Z","iopub.status.idle":"2026-04-17T11:54:16.756769Z","shell.execute_reply.started":"2026-04-17T11:54:14.660978Z","shell.execute_reply":"2026-04-17T11:54:16.756104Z"}},"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-17T11:54:16.757658Z","iopub.execute_input":"2026-04-17T11:54:16.757918Z","iopub.status.idle":"2026-04-17T11:54:16.765952Z","shell.execute_reply.started":"2026-04-17T11:54:16.757897Z","shell.execute_reply":"2026-04-17T11:54:16.765509Z"}},"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-17T11:54:16.766694Z","iopub.execute_input":"2026-04-17T11:54:16.766877Z","iopub.status.idle":"2026-04-17T11:54:16.804392Z","shell.execute_reply.started":"2026-04-17T11:54:16.766860Z","shell.execute_reply":"2026-04-17T11:54:16.803665Z"}},"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-17T11:54:16.805360Z","iopub.execute_input":"2026-04-17T11:54:16.805633Z","iopub.status.idle":"2026-04-17T11:54:18.673819Z","shell.execute_reply.started":"2026-04-17T11:54:16.805608Z","shell.execute_reply":"2026-04-17T11:54:18.673120Z"}},"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-17T11:54:18.674700Z","iopub.execute_input":"2026-04-17T11:54:18.674984Z","iopub.status.idle":"2026-04-17T11:54:18.680165Z","shell.execute_reply.started":"2026-04-17T11:54:18.674961Z","shell.execute_reply":"2026-04-17T11:54:18.679582Z"}},"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-17T11:54:18.681122Z","iopub.execute_input":"2026-04-17T11:54:18.681417Z"}},"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},"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# Dynamically create the 'sample' subdirectory if it does not exist\nsample_dir = OUTPUT_DIR / \"sample\"\nsample_dir.mkdir(parents=True, exist_ok=True)\n\n# Persist sample-set arrays into the newly verified directory\nnp.save(sample_dir / \"features_384d.npy\", features_384d)\nnp.save(sample_dir / \"features_128d.npy\", features_128d)\nnp.save(sample_dir / \"labels.npy\",        labels_arr)\nnp.save(sample_dir / \"datasets.npy\",      datasets_arr)\nnp.save(sample_dir / \"paths.npy\",         paths_arr)\nmeta_df.to_csv(sample_dir / \"metadata.csv\", index=False)\n\nprint(f\"\\n  Sample-set outputs saved to {sample_dir}\")","metadata":{"trusted":true},"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},"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)\n\n# We use the DEMOGRAPHIC_COLS variable defined in Phase 1C\ndemo_cols = [c for c in DEMOGRAPHIC_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    \n    for col in demo_cols:\n        if col == \"age\":\n            # 1. Must be a valid number\n            is_numeric = pd.to_numeric(meta_df[col], errors=\"coerce\").notna()\n            \n            # 2. Must be verified by the dataset's native boolean flag (if it exists)\n            if \"has_age\" in meta_df.columns:\n                known = (is_numeric & (meta_df[\"has_age\"] == True)).sum()\n            else:\n                # Fallback: Treat 0 or negative ages as imputed/missing\n                valid_range = pd.to_numeric(meta_df[col], errors=\"coerce\") > 0\n                known = (is_numeric & valid_range).sum()\n                \n        elif col == \"gender\" and \"has_gender\" in meta_df.columns:\n            # Trust the native flag over text matching if available\n            known = (meta_df[\"has_gender\"] == True).sum()\n            \n        else:\n            # Standard string parsing for 'view', 'original_view', 'identity_id'\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},"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\")\n\n# Defensive check to ensure the dataset isn't pathology-only\nassert healthy_mask.sum() > 0, \"CRITICAL ERROR: No 'No Finding' samples available to anchor the severity index.\"\n\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\ndetailed_labels = np.where(healthy_mask, \"Normal\", labels_arr).astype(object)","metadata":{"trusted":true},"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},"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\n# Assuming sample_dir was defined in Phase 6B\nnp.save(sample_dir / \"labels_detailed.npy\", detailed_labels)\nmeta_df.to_csv(sample_dir / \"metadata.csv\", index=False)   # re-save with severity column\n\n# Corrected the print statement to reflect the actual save location\nprint(f\"\\n  labels_detailed.npy saved to {sample_dir}\")\nprint(\"  metadata.csv updated with 'severity_label' column.\")","metadata":{"trusted":true},"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 ──────────────────────────────────────────────────\nimport pandas as pd\nimport numpy as np\n\nprint(\"[ 9A ] Encoding demographic metadata into a numeric conditioning vector ...\")\n\n# 1. Define which columns actually contribute to the MLP vector\n# We exclude identity_id and validation columns from the one-hot maps\nmlp_feature_cols = [\"gender\", \"view\", \"original_view\"] \nvalidation_cols = [\"has_age\", \"has_gender\", \"age_group\", \"disease\", \"raw_labels\"]\n\ncategory_maps = {}\nfor col in mlp_feature_cols:\n    if col in meta_df.columns:\n        unique_vals = sorted(meta_df[col].dropna().astype(str).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    parts = []\n\n    # Age feature (Normalized to [0, 1])\n    if \"age\" in meta_df.columns:\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 features (One-hot encoded)\n    for col, cmap in category_maps.items():\n        vec = np.zeros(len(cmap), dtype=np.float32)\n        val = str(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# Generate the numeric matrix for the MLP\ndemo_matrix = np.vstack([encode_demographic_row(row) for _, row in meta_df.iterrows()])\nDEMO_DIM = demo_matrix.shape[1]\n\nprint(f\"  MLP Input Dimension      : {DEMO_DIM}-d\")\nprint(f\"  Features in matrix       : age + {list(category_maps.keys())}\")\nprint(f\"  QC columns in DataFrame  : {validation_cols}\")","metadata":{"trusted":true},"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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 9C — Final Vector","metadata":{}},{"cell_type":"code","source":"# ── 9C: FINAL CONSOLIDATION & EXPORT ─────────────────────────────────────────\nimport torch\nimport pandas as pd\nimport numpy as np\n\nprint(\"[ 9C ] Finalizing embeddings and metadata merge...\")\n\n# 1. GENERATE THE 4 SEPARATE VECTORS\n# ─────────────────────────────────────────────────────────────────────────────\n\n# A: Visual 128-d \n# (Already exists in memory as 'features_128d' from Phase 6B)\n\n# B: Disease One-Hot (N x Number of Diseases)\n# We use get_dummies to create a binary matrix for categorical disease names\ndisease_onehot = pd.get_dummies(meta_df['disease']).values.astype(np.float32)\n\n# C: Severity One-Hot (Normal, Mild, Moderate, Severe)\n# We extract the first word (e.g., \"Mild Pneumonia\" -> \"Mild\") to standardise\nseverity_only = meta_df['severity_label'].astype(str).str.split().str[0]\nseverity_onehot = pd.get_dummies(severity_only).values.astype(np.float32)\n\n# D: Demographic 64-d (The MLP Latent Space)\n# We pass the raw numeric matrix through our trained MLP to get the embedding\ndemo_encoder.eval() \nwith torch.no_grad():\n    demo_tensor = torch.from_numpy(demo_matrix).to(DEVICE)\n    # The output is L2-normalised inside the model's forward pass\n    demo_embeddings = demo_encoder(demo_tensor).cpu().numpy()\n\n# 2. SAVE SEPARATE NUMPY FILES\n# ─────────────────────────────────────────────────────────────────────────────\n# Creating a dedicated folder for these vectors to keep the GAN input clean\nsave_dir = OUTPUT_DIR / \"final_conditions\"\nsave_dir.mkdir(parents=True, exist_ok=True)\n\nnp.save(save_dir / \"image_embeddings_128d.npy\",      features_128d)\nnp.save(save_dir / \"disease_onehot.npy\",             disease_onehot)\nnp.save(save_dir / \"severity_onehot.npy\",            severity_onehot)\nnp.save(save_dir / \"demographic_embeddings_64d.npy\", demo_embeddings)\n\n# 3. UPDATE MASTER PARQUET\n# ─────────────────────────────────────────────────────────────────────────────\n# Loading the original file to perform the merge\noriginal_parquet_path = \"/kaggle/input/datasets/rohanpesuecbtech2023/train-parquet/unchanged_train.parquet\"\nfull_df = pd.read_parquet(original_parquet_path)\n\n# Map the sampled severity back to the full dataset using the 'path' column\nseverity_mapping = meta_df[['path', 'severity_label']]\n\n# A 'left' join ensures all 112,679 rows from full_df are kept.\n# Unprocessed rows will simply have 'NaN' in the severity_label column.\nupdated_df = full_df.merge(severity_mapping, on='path', how='left')\n\n# 4. EXPORT UPDATED PARQUET\n# ─────────────────────────────────────────────────────────────────────────────\noutput_path = OUTPUT_DIR / \"train_with_severity_separated.parquet\"\nupdated_df.to_parquet(output_path, index=False)\n\n# 5. SUMMARY REPORT\n# ─────────────────────────────────────────────────────────────────────────────\nprint(f\"\\n  Success:\")\nprint(f\"  - 4 Condition Arrays saved to: {save_dir}\")\nprint(f\"  - Vector Dimensions:\")\nprint(f\"      Visual (DINOv2+SSL) : {features_128d.shape}\")\nprint(f\"      Disease (One-Hot)   : {disease_onehot.shape}\")\nprint(f\"      Severity (One-Hot)  : {severity_onehot.shape}\")\nprint(f\"      Demographic (MLP)   : {demo_embeddings.shape}\")\nprint(f\"  - Master Parquet updated:\")\nprint(f\"      Save Location       : {output_path}\")\nprint(f\"      Total Rows Kept     : {len(updated_df):,}\")\nprint(f\"      Processed Samples   : {updated_df['severity_label'].notna().sum():,}\")","metadata":{"trusted":true},"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# Create the fused vector for verification and GAN input\nfused_vectors = np.concatenate([\n    features_128d, \n    disease_onehot, \n    severity_onehot, \n    demo_embeddings\n], axis=1)\n\n# Ensure this variable is defined if not already in your Phase 1/2\nCONDITIONING_DISEASE_CLASSES = sorted(meta_df['disease'].unique())\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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Save weights","metadata":{}},{"cell_type":"code","source":"# ── 10: SAVE ALL WEIGHTS ─────────────────────────────────────────────────────\nimport torch\n\n# Ensure the output directory exists\nOUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\n# 1. Save SSL Head (Attribute of the 'encoder' object)\nssl_weights_path = OUTPUT_DIR / \"ssl_projection_head_weights.pth\"\n# Access via encoder.ssl_head\ntorch.save(encoder.ssl_head.state_dict(), ssl_weights_path)\nprint(f\"✅ Saved SSL weights to: {ssl_weights_path}\")\n\n# 2. Save Demographic MLP (Variable named 'demo_encoder' in Phase 9B)\ndemo_weights_path = OUTPUT_DIR / \"demographic_mlp_weights.pth\"\ntorch.save(demo_encoder.state_dict(), demo_weights_path)\nprint(f\"✅ Saved MLP weights to: {demo_weights_path}\")","metadata":{"trusted":true},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── PRE-EXTRACTION: LOAD INTELLIGENCE ────────────────────────────────────────\nprint(\"[ SETUP ] Loading trained weights for full-scale inference...\")\n\n# 1. Load SSL Head Weights into the Encoder\ntry:\n    ssl_weights_path = OUTPUT_DIR / \"ssl_projection_head_weights.pth\"\n    encoder.ssl_head.load_state_dict(torch.load(ssl_weights_path))\n    encoder.eval()\n    print(f\"   SSL Head weights loaded from {ssl_weights_path}\")\nexcept Exception as e:\n    print(f\"  Could not load SSL weights ({e}) - check if weights exist!\")\n\n# 2. We don't load the Demo Encoder yet because Phase 13D \n#    will re-instantiate it based on the final column count.","metadata":{"trusted":true},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 13D — Save","metadata":{}},{"cell_type":"code","source":"# ── 13D: SAVE SEPARATE ARRAYS + cGAN PARQUET (FINAL) ─────────────────────────\nprint(\"\\n[ 13D ] Saving 4 separate conditioning vectors and final Parquet ...\")\n\n# 1. DYNAMIC COLUMN DISCOVERY (Fixes the 'sex' vs 'gender' KeyError)\nexisting_cols = meta_df_full.columns.tolist()\n# We look for ANY demographic-style columns available\ndemo_cols_present = [c for c in [\"age\", \"sex\", \"gender\", \"view_position\", \"view\", \"ap_pa\"] if c in existing_cols]\nprint(f\"  Encoding demographics from: {demo_cols_present}\")\n\n# ── SIGNAL 1: Disease One-Hot ────────────────────────────────────────────────\n_disease_onehot = pd.get_dummies(meta_df_full['disease']).values.astype(np.float32)\nnp.save(OUTPUT_DIR / \"disease_onehot_full.npy\", _disease_onehot)\nprint(f\"  - Saved Disease One-Hot   : {_disease_onehot.shape}\")\n\n# ── SIGNAL 2: Severity One-Hot ───────────────────────────────────────────────\ndef _to_sev_grade(label):\n    l = str(label).lower()\n    if \"normal\" in l:   return \"Normal\"\n    if \"mild\" in l:     return \"Mild\"\n    if \"moderate\" in l: return \"Moderate\"\n    if \"severe\" in l:   return \"Severe\"\n    return \"Normal\"\n\n_sev_grades = meta_df_full['severity_label'].apply(_to_sev_grade)\n_severity_onehot = pd.get_dummies(_sev_grades).values.astype(np.float32)\nnp.save(OUTPUT_DIR / \"severity_onehot_full.npy\", _severity_onehot)\nprint(f\"  - Saved Severity One-Hot  : {_severity_onehot.shape}\")\n\n# ── SIGNAL 3: Demographic 64-d Embedding (LOADS WEIGHTS) ──────────────────────\n_cat_maps = {}\nfor col in demo_cols_present:\n    if col == \"age\": continue\n    _cat_maps[col] = {v: i for i, v in enumerate(sorted(meta_df_full[col].dropna().unique()))}\n\ndef _encode_row_safe(row):\n    parts = []\n    try:\n        age_val = float(row.get(\"age\", 50))\n        parts.append(np.array([np.clip(age_val / 100.0, 0.0, 1.0)], dtype=np.float32))\n    except:\n        parts.append(np.array([0.5], dtype=np.float32))\n    for col, cmap in _cat_maps.items():\n        vec = np.zeros(len(cmap), dtype=np.float32)\n        val = row.get(col, \"unknown\")\n        if val in cmap: vec[cmap[val]] = 1.0\n        parts.append(vec)\n    return np.concatenate(parts)\n\n_demo_matrix = np.vstack([_encode_row_safe(row) for _, row in meta_df_full.iterrows()])\n\n# Re-init and Load Weights\ndemo_encoder = DemographicEncoder(in_dim=_demo_matrix.shape[1], out_dim=64).to(DEVICE)\ntry:\n    demo_weights_path = OUTPUT_DIR / \"demographic_mlp_weights.pth\"\n    demo_encoder.load_state_dict(torch.load(demo_weights_path))\n    print(f\"  ✅ MLP Weights loaded from {demo_weights_path}\")\nexcept:\n    print(\"  ⚠️ MLP weights not found - using random init (Not Recommended!)\")\n\ndemo_encoder.eval()\nwith torch.no_grad():\n    _demo_emb = demo_encoder(torch.tensor(_demo_matrix, device=DEVICE)).cpu().numpy()\n\nnp.save(OUTPUT_DIR / \"demographic_embeddings_64d_full.npy\", _demo_emb)\nprint(f\"  - Saved Demographic 64-d  : {_demo_emb.shape}\")\n\n# ── SIGNAL 4: SSL Visual Embedding ────────────────────────────────────────────\nnp.save(OUTPUT_DIR / \"features_128d_full.npy\", features_128d_full)\nprint(f\"  - Saved SSL Visual 128-d  : {features_128d_full.shape}\")\n\n# ── FINAL PARQUET ─────────────────────────────────────────────────────────────\ncgan_df = meta_df_full.copy()\ncgan_df[\"disease_onehot\"]     = list(_disease_onehot)\ncgan_df[\"severity_onehot\"]    = list(_severity_onehot)\ncgan_df[\"demo_embedding_64d\"] = list(_demo_emb)\ncgan_df[\"ssl_embedding_128d\"] = list(features_128d_full)\n\noutput_path = OUTPUT_DIR / \"cgan_input_full.parquet\"\ncgan_df.to_parquet(output_path, index=False)\n\nprint(f\"\\n✅ All 4 separate vectors and master Parquet saved to {OUTPUT_DIR}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}