{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":2169393,"datasetId":1302315,"databundleVersionId":2210641},{"sourceType":"datasetVersion","sourceId":18613,"datasetId":5839,"databundleVersionId":18613},{"sourceType":"datasetVersion","sourceId":1800067,"datasetId":1069810,"databundleVersionId":1837524}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q -U transformers\n!pip install -q scikit-learn matplotlib pandas tqdm\nprint(\"Setup complete\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:39:44.214507Z","iopub.execute_input":"2026-05-10T04:39:44.214816Z","iopub.status.idle":"2026-05-10T04:39:51.850439Z","shell.execute_reply.started":"2026-05-10T04:39:44.214793Z","shell.execute_reply":"2026-05-10T04:39:51.849531Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom pathlib import Path\n\nCONFIG = {\n    # === Datasets ===\n    \"datasets\": [\"nih\", \"vindr\", \"chexpert\"],   # Tier 1 + 2: 3 datasets\n\n    # === Backbones ===\n    \"backbones\": [\n        \"facebook/dinov2-base\",\n        \"microsoft/rad-dino\",\n        \"microsoft/BiomedCLIP\",   # placeholder name; handled by extractor\n    ],\n\n    # === Seeds ===\n    \"seeds\": [42, 123, 7, 2024, 555],\n\n    # === RQ2 primary ===\n    \"rq2_primary_dataset\": \"vindr\",\n    \"rq2_primary_backbone\": \"microsoft/rad-dino\",\n\n    # === Dataset paths ===\n    \"nih_data_dir\": None,\n    \"vindr_data_dir\": \"/kaggle/input/datasets/awsaf49/vinbigdata-512-image-dataset/vinbigdata\",\n    \"chexpert_data_dir\": None,    # Set manually if auto-detect fails\n\n    # === Output paths ===\n    \"cache_dir\": \"/kaggle/working/cache\",\n    \"output_dir\": \"/kaggle/working/output\",\n    \"checkpoint_dir\": \"/kaggle/working/checkpoints\",\n    \"figures_dir\": \"/kaggle/working/output/figures\",\n    \"tables_dir\": \"/kaggle/working/output/tables\",\n\n    # === Backbone settings ===\n    \"feature_dim\": 768,           # Default for DINOv2/RAD-DINO; BiomedCLIP=512 auto-handled\n    \"image_size\": 224,\n\n    # === Dataset settings ===\n    \"selected_classes\": [\"Atelectasis\", \"Cardiomegaly\", \"Consolidation\",\n                         \"Effusion\", \"Pneumothorax\"],\n    \"subset_size\": 30000,\n    \"test_ratio\": 0.20,\n    \"val_ratio\": 0.10,\n\n    # === Federation ===\n    \"num_clients\": 5,\n    \"non_iid_alpha\": 0.5,\n\n    # === Ridge / upgrades ===\n    \"ridge_gamma\": 1.0,\n    \"nystrom_anchors\": 1024,\n    \"label_smoothing_eps\": 0.1,\n    \"gamma_grid\": [0.1, 0.3, 1.0, 3.0, 10.0, 30.0],\n\n    # === Numerical ===\n    \"use_fp64_server\": True,\n\n    # === RQ2 ===\n    \"n_patient_deletions\": 100,\n    \"rare_class\": \"Pneumothorax\",\n    \"withdrawn_site_idx\": 0,\n\n    # === SISA baseline (Tier 1 item 3) ===\n    \"sisa_num_shards\": 5,\n\n    # === Speed ===\n    \"batch_size\": 64,\n    \"num_workers\": 2,\n\n    # === Quick mode ===\n    \"quick_mode\": False,\n}\n\nif CONFIG[\"quick_mode\"]:\n    CONFIG[\"subset_size\"] = 5000\n    CONFIG[\"n_patient_deletions\"] = 30\n    CONFIG[\"seeds\"] = [42, 123]\n\nfor d in [CONFIG[\"cache_dir\"], CONFIG[\"output_dir\"], CONFIG[\"checkpoint_dir\"],\n          CONFIG[\"figures_dir\"], CONFIG[\"tables_dir\"]]:\n    Path(d).mkdir(parents=True, exist_ok=True)\n\n# Auto-detect dataset paths\ndef find_nih():\n    for c in [\"/kaggle/input/data\", \"/kaggle/input/nih-chest-xrays/data\",\n              \"/kaggle/input/datasets/organizations/nih-chest-xrays/data\"]:\n        if os.path.exists(os.path.join(c, \"Data_Entry_2017.csv\")):\n            return c\n    return None\n\ndef find_vindr():\n    for c in [\"/kaggle/input/vinbigdata-512-image-dataset\",\n              \"/kaggle/input/vinbigdata-chest-xray-resized-png-256x256\",\n              \"/kaggle/input/datasets/awsaf49/vinbigdata-512-image-dataset/vinbigdata\"]:\n        if os.path.exists(os.path.join(c, \"train.csv\")):\n            return c\n    return None\n\ndef find_chexpert():\n    for c in [\"/kaggle/input/chexpert\", \"/kaggle/input/chexpert-v10\",\n              \"/kaggle/input/chexpert-small\",\n              \"/kaggle/input/datasets/ashery/chexpert\"]:\n        if os.path.exists(c):\n            for sub in [\"\", \"CheXpert-v1.0/\", \"CheXpert-v1.0-small/\"]:\n                if os.path.exists(os.path.join(c, sub, \"train.csv\")):\n                    return c\n    # Recursive deep search\n    for root, dirs, files in os.walk(\"/kaggle/input\"):\n        if \"train.csv\" in files and \"chexpert\" in root.lower():\n            return root\n    return None\n\nif CONFIG[\"nih_data_dir\"] is None:\n    CONFIG[\"nih_data_dir\"] = find_nih()\nif CONFIG[\"vindr_data_dir\"] is None:\n    CONFIG[\"vindr_data_dir\"] = find_vindr()\nif CONFIG[\"chexpert_data_dir\"] is None:\n    CONFIG[\"chexpert_data_dir\"] = find_chexpert()\n\nprint(f\"NIH dir:      {CONFIG['nih_data_dir']}\")\nprint(f\"VinDr dir:    {CONFIG['vindr_data_dir']}\")\nprint(f\"CheXpert dir: {CONFIG['chexpert_data_dir']}\")\nprint(f\"Datasets to run: {CONFIG['datasets']}\")\nprint(f\"Backbones to run: {CONFIG['backbones']}\")\nprint(f\"Seeds: {CONFIG['seeds']}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:39:51.852272Z","iopub.execute_input":"2026-05-10T04:39:51.852542Z","iopub.status.idle":"2026-05-10T04:39:51.872113Z","shell.execute_reply.started":"2026-05-10T04:39:51.852512Z","shell.execute_reply":"2026-05-10T04:39:51.871344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json, time, pickle, warnings\nfrom copy import deepcopy\nfrom collections import defaultdict\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom sklearn.metrics import roc_auc_score, average_precision_score\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings(\"ignore\")\n\ndef set_seed(seed):\n    np.random.seed(seed); torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nDTYPE_SERVER = torch.float64 if CONFIG[\"use_fp64_server\"] else torch.float32\nprint(f\"Device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:39:51.873334Z","iopub.execute_input":"2026-05-10T04:39:51.873712Z","iopub.status.idle":"2026-05-10T04:40:02.619447Z","shell.execute_reply.started":"2026-05-10T04:39:51.873685Z","shell.execute_reply":"2026-05-10T04:40:02.618594Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_nih(config):\n    cache_path = Path(config[\"cache_dir\"]) / \"metadata_nih.pkl\"\n    if cache_path.exists():\n        with open(cache_path, \"rb\") as f:\n            return pickle.load(f)\n\n    base = config[\"nih_data_dir\"]\n    df = pd.read_csv(os.path.join(base, \"Data_Entry_2017.csv\"))\n\n    image_path_map = {}\n    for sub in os.listdir(base):\n        sp = os.path.join(base, sub)\n        if not os.path.isdir(sp): continue\n        for cand in [os.path.join(sp, \"images\"), sp]:\n            if os.path.isdir(cand):\n                for fn in os.listdir(cand):\n                    if fn.endswith(\".png\"):\n                        image_path_map[fn] = os.path.join(cand, fn)\n\n    df = df[df[\"Image Index\"].isin(image_path_map.keys())].copy()\n    df[\"path\"] = df[\"Image Index\"].map(image_path_map)\n    df[\"patient_id\"] = df[\"Patient ID\"].astype(int)\n\n    selected = config[\"selected_classes\"]\n    label_matrix = np.zeros((len(df), len(selected)), dtype=np.float32)\n    for i, ls in enumerate(df[\"Finding Labels\"].values):\n        labs = ls.split(\"|\")\n        for j, c in enumerate(selected):\n            if c in labs:\n                label_matrix[i, j] = 1.0\n    df = df.reset_index(drop=True)\n\n    rng = np.random.default_rng(config[\"seed\"] if \"seed\" in config else 42)\n    pids = df[\"patient_id\"].unique()\n    rng.shuffle(pids)\n    selected_p, n = [], 0\n    for pid in pids:\n        np_pid = (df[\"patient_id\"] == pid).sum()\n        if n + np_pid > config[\"subset_size\"]:\n            if n == 0: selected_p.append(pid); n += np_pid\n            break\n        selected_p.append(pid); n += np_pid\n\n    mask = df[\"patient_id\"].isin(selected_p).values\n    df = df.loc[mask].reset_index(drop=True)\n    label_matrix = label_matrix[mask]\n\n    print(f\"[NIH] Total: {len(df)}, patients: {len(selected_p)}\")\n    print(\"Class prevalence:\")\n    for j, c in enumerate(selected):\n        print(f\"  {c}: {label_matrix[:, j].mean():.4f} ({int(label_matrix[:, j].sum())} pos)\")\n\n    result = {\"df\": df, \"labels\": label_matrix, \"dataset\": \"nih\"}\n    with open(cache_path, \"wb\") as f:\n        pickle.dump(result, f)\n    return result\n\n\ndef prepare_vindr(config):\n    cache_path = Path(config[\"cache_dir\"]) / \"metadata_vindr.pkl\"\n    if cache_path.exists():\n        with open(cache_path, \"rb\") as f:\n            return pickle.load(f)\n\n    base = config[\"vindr_data_dir\"]\n    if base is None:\n        raise FileNotFoundError(\"VinDr-CXR not found. Attach 'vinbigdata-512-image-dataset' on Kaggle.\")\n\n    df_raw = pd.read_csv(os.path.join(base, \"train.csv\"))\n\n    # VinDr class id mapping (Kaggle competition uses 14 + No finding=14)\n    vindr_classes = [\n        \"Aortic enlargement\", \"Atelectasis\", \"Calcification\", \"Cardiomegaly\",\n        \"Consolidation\", \"ILD\", \"Infiltration\", \"Lung Opacity\", \"Nodule/Mass\",\n        \"Other lesion\", \"Pleural effusion\", \"Pleural thickening\", \"Pneumothorax\",\n        \"Pulmonary fibrosis\"\n    ]\n    name_map = {\n        \"Atelectasis\": \"Atelectasis\",\n        \"Cardiomegaly\": \"Cardiomegaly\",\n        \"Consolidation\": \"Consolidation\",\n        \"Effusion\": \"Pleural effusion\",  # VinDr uses different term\n        \"Pneumothorax\": \"Pneumothorax\",\n    }\n    selected = config[\"selected_classes\"]\n    cls_indices = [vindr_classes.index(name_map[c]) for c in selected]\n\n    # Aggregate multi-rater: positive if ≥2 of 3 radiologists agree\n    image_class_counts = defaultdict(lambda: defaultdict(int))\n    for _, row in df_raw.iterrows():\n        cid = row[\"class_id\"]\n        if 0 <= cid < len(vindr_classes):\n            image_class_counts[row[\"image_id\"]][cid] += 1\n\n    image_ids = sorted(image_class_counts.keys())\n    label_matrix = np.zeros((len(image_ids), len(selected)), dtype=np.float32)\n    for i, iid in enumerate(image_ids):\n        for j, cid in enumerate(cls_indices):\n            if image_class_counts[iid][cid] >= 2:\n                label_matrix[i, j] = 1.0\n\n    # Build paths (try common subdirs)\n    train_dirs = [os.path.join(base, \"train\"), base]\n    paths = []\n    for iid in image_ids:\n        p = None\n        for td in train_dirs:\n            for ext in [\".png\", \".jpg\"]:\n                cand = os.path.join(td, f\"{iid}{ext}\")\n                if os.path.exists(cand):\n                    p = cand; break\n            if p: break\n        paths.append(p)\n\n    valid = [i for i, p in enumerate(paths) if p is not None]\n    image_ids = [image_ids[i] for i in valid]\n    paths = [paths[i] for i in valid]\n    label_matrix = label_matrix[valid]\n\n    df = pd.DataFrame({\n        \"Image Index\": [f\"{iid}.png\" for iid in image_ids],\n        \"patient_id\": list(range(len(image_ids))),  # 1 image = 1 \"patient\" (no patient ID in VinDr)\n        \"path\": paths,\n    })\n\n    print(f\"[VinDr] Total: {len(df)}\")\n    print(\"Class prevalence:\")\n    for j, c in enumerate(selected):\n        print(f\"  {c}: {label_matrix[:, j].mean():.4f} ({int(label_matrix[:, j].sum())} pos)\")\n\n    result = {\"df\": df, \"labels\": label_matrix, \"dataset\": \"vindr\"}\n    with open(cache_path, \"wb\") as f:\n        pickle.dump(result, f)\n    return result\n\n\ndef prepare_metadata(dataset_name, config):\n    \"\"\"Dispatcher for dataset preparation.\"\"\"\n    if dataset_name == \"nih\":\n        return prepare_nih(config)\n    elif dataset_name == \"vindr\":\n        return prepare_vindr(config)\n    elif dataset_name == \"chexpert\":\n        return prepare_chexpert(config)\n    else:\n        raise ValueError(f\"Unknown dataset: {dataset_name}\")\n\n\ndef split_train_val_test(meta, config, seed):\n    \"\"\"Patient-level split into train, val, test.\"\"\"\n    df = meta[\"df\"]\n    labels = meta[\"labels\"]\n    pids = df[\"patient_id\"].values\n    unique_pids = np.unique(pids)\n\n    rng = np.random.default_rng(seed)\n    rng.shuffle(unique_pids)\n\n    n_test = int(len(unique_pids) * config[\"test_ratio\"])\n    n_val = int(len(unique_pids) * config[\"val_ratio\"])\n    test_pids = set(unique_pids[:n_test])\n    val_pids = set(unique_pids[n_test:n_test+n_val])\n\n    test_mask = np.isin(pids, list(test_pids))\n    val_mask = np.isin(pids, list(val_pids))\n    train_mask = ~(test_mask | val_mask)\n\n    return {\n        \"train_idx\": np.where(train_mask)[0],\n        \"val_idx\": np.where(val_mask)[0],\n        \"test_idx\": np.where(test_mask)[0],\n    }\n\n\ndef prepare_chexpert(config):\n    \"\"\"Prepare CheXpert metadata: frontal-only, U-Ones uncertainty.\"\"\"\n    cache_path = Path(config[\"cache_dir\"]) / \"metadata_chexpert.pkl\"\n    if cache_path.exists():\n        with open(cache_path, \"rb\") as f:\n            cached = pickle.load(f)\n        if abs(len(cached[\"df\"]) - config[\"subset_size\"]) / config[\"subset_size\"] <= 0.2:\n            return cached\n\n    base = config[\"chexpert_data_dir\"]\n    if base is None:\n        raise FileNotFoundError(\"CheXpert not found. Set CONFIG['chexpert_data_dir'].\")\n\n    csv_path = None\n    for cand in [os.path.join(base, \"train.csv\"),\n                 os.path.join(base, \"CheXpert-v1.0/train.csv\"),\n                 os.path.join(base, \"CheXpert-v1.0-small/train.csv\")]:\n        if os.path.exists(cand):\n            csv_path = cand\n            break\n    if csv_path is None:\n        # Try recursive\n        for root, dirs, files in os.walk(base):\n            if \"train.csv\" in files:\n                csv_path = os.path.join(root, \"train.csv\")\n                break\n    if csv_path is None:\n        raise FileNotFoundError(f\"CheXpert train.csv not found under {base}\")\n\n    df_raw = pd.read_csv(csv_path)\n    print(f\"[CheXpert] Raw rows: {len(df_raw)}\")\n\n    # Frontal only\n    if \"Frontal/Lateral\" in df_raw.columns:\n        df_raw = df_raw[df_raw[\"Frontal/Lateral\"] == \"Frontal\"].copy()\n    print(f\"[CheXpert] After frontal filter: {len(df_raw)}\")\n\n    # Map class names to CheXpert columns\n    name_map = {\n        \"Atelectasis\": \"Atelectasis\",\n        \"Cardiomegaly\": \"Cardiomegaly\",\n        \"Consolidation\": \"Consolidation\",\n        \"Effusion\": \"Pleural Effusion\",\n        \"Pneumothorax\": \"Pneumothorax\",\n    }\n    selected = config[\"selected_classes\"]\n\n    # Build label matrix: U-Ones (uncertain → positive)\n    label_matrix = np.zeros((len(df_raw), len(selected)), dtype=np.float32)\n    for j, c in enumerate(selected):\n        col = name_map[c]\n        if col not in df_raw.columns:\n            print(f\"  WARNING: column {col} missing in CheXpert\")\n            continue\n        vals = df_raw[col].fillna(0).values\n        label_matrix[:, j] = ((vals == 1) | (vals == -1)).astype(np.float32)\n\n    # Patient ID extraction from path\n    def extract_pid(path_str):\n        parts = str(path_str).split(\"/\")\n        for p in parts:\n            if p.startswith(\"patient\"):\n                return p\n        return str(path_str)\n\n    df_raw[\"patient_id\"] = df_raw[\"Path\"].apply(extract_pid)\n    df_raw = df_raw.reset_index(drop=True)\n\n    # Subsample at patient level\n    rng = np.random.default_rng(42)\n    pids = df_raw[\"patient_id\"].unique()\n    rng.shuffle(pids)\n    selected_p, n = [], 0\n    for pid in pids:\n        np_pid = (df_raw[\"patient_id\"] == pid).sum()\n        if n + np_pid > config[\"subset_size\"]:\n            if n == 0:\n                selected_p.append(pid); n += np_pid\n            break\n        selected_p.append(pid); n += np_pid\n\n    mask = df_raw[\"patient_id\"].isin(selected_p).values\n    df_raw = df_raw.loc[mask].reset_index(drop=True)\n    label_matrix = label_matrix[mask]\n\n    # Resolve image paths\n    paths = []\n    csv_dir = os.path.dirname(csv_path)\n    for p_str in df_raw[\"Path\"].values:\n        full = None\n        for base_cand in [base, csv_dir, os.path.dirname(csv_dir)]:\n            cand_path = os.path.join(base_cand, str(p_str))\n            if os.path.exists(cand_path):\n                full = cand_path; break\n        paths.append(full)\n\n    valid = [i for i, p in enumerate(paths) if p is not None]\n    print(f\"[CheXpert] Valid images: {len(valid)} / {len(paths)}\")\n    if len(valid) == 0:\n        raise FileNotFoundError(\"No CheXpert images found. Check CONFIG['chexpert_data_dir'] structure.\")\n\n    df = pd.DataFrame({\n        \"Image Index\": [paths[i] for i in valid],\n        \"patient_id\": df_raw[\"patient_id\"].values[valid],\n        \"path\": [paths[i] for i in valid],\n    })\n    label_matrix = label_matrix[valid]\n\n    print(f\"[CheXpert] Final: {len(df)} samples, {df['patient_id'].nunique()} patients\")\n    print(\"Class prevalence:\")\n    for j, c in enumerate(selected):\n        pos = int(label_matrix[:, j].sum())\n        print(f\"  {c}: {pos} positives ({pos/len(df):.2%})\")\n\n    out = {\"dataset\": \"chexpert\", \"df\": df, \"labels\": label_matrix}\n    with open(cache_path, \"wb\") as f:\n        pickle.dump(out, f)\n    return out\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.621179Z","iopub.execute_input":"2026-05-10T04:40:02.621547Z","iopub.status.idle":"2026-05-10T04:40:02.657197Z","shell.execute_reply.started":"2026-05-10T04:40:02.621523Z","shell.execute_reply":"2026-05-10T04:40:02.656484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ChestXrayDataset(Dataset):\n    def __init__(self, paths, transform):\n        self.paths, self.transform = paths, transform\n    def __len__(self): return len(self.paths)\n    def __getitem__(self, idx):\n        try:\n            img = Image.open(self.paths[idx]).convert(\"RGB\")\n        except Exception:\n            img = Image.new(\"RGB\", (224, 224))\n        return self.transform(img), idx\n\n\ndef extract_features(meta, backbone_name, config):\n    safe_bb = backbone_name.replace(\"/\", \"_\")\n    cache_path = Path(config[\"cache_dir\"]) / f\"features_{meta['dataset']}_{safe_bb}_{len(meta['df'])}.npz\"\n    if cache_path.exists():\n        print(f\"[CACHED] {cache_path.name}\")\n        return np.load(cache_path)[\"features\"]\n\n    print(f\"Extracting: {meta['dataset']} × {backbone_name}\")\n    from transformers import AutoModel, AutoImageProcessor\n    processor = AutoImageProcessor.from_pretrained(backbone_name)\n    model = AutoModel.from_pretrained(backbone_name).to(DEVICE).eval()\n\n    # === Multi-GPU support ===\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        print(f\"  Using {n_gpus} GPUs via DataParallel\")\n        model = nn.DataParallel(model)\n        effective_bs = config[\"batch_size\"] * n_gpus      # scale batch\n        effective_workers = max(config[\"num_workers\"] * 2, 4)  # more loaders\n    else:\n        print(f\"  Using single GPU\")\n        effective_bs = config[\"batch_size\"]\n        effective_workers = config[\"num_workers\"]\n\n    transform = transforms.Compose([\n        transforms.Resize((config[\"image_size\"], config[\"image_size\"])),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=processor.image_mean, std=processor.image_std),\n    ])\n\n    df = meta[\"df\"]\n    dataset = ChestXrayDataset(df[\"path\"].tolist(), transform)\n    loader = DataLoader(\n        dataset, batch_size=effective_bs,\n        num_workers=effective_workers, shuffle=False,\n        pin_memory=True, persistent_workers=(effective_workers > 0)\n    )\n\n    feat_dim = config[\"feature_dim\"]\n    features = np.zeros((len(df), feat_dim), dtype=np.float32)\n    with torch.no_grad():\n        for imgs, idx in tqdm(loader, desc=f\"{meta['dataset']}/{safe_bb}\"):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            out = model(imgs)\n            # Handle DataParallel wrap\n            if hasattr(out, \"pooler_output\") and out.pooler_output is not None:\n                feat = out.pooler_output\n            elif hasattr(out, \"last_hidden_state\"):\n                feat = out.last_hidden_state[:, 0, :]\n            else:\n                # Fallback: assume tensor or first element\n                feat = out[0][:, 0, :] if isinstance(out, tuple) else out[:, 0, :]\n            features[idx.numpy()] = feat.float().cpu().numpy()\n\n    np.savez_compressed(cache_path, features=features)\n    print(f\"Saved: {cache_path.name}\")\n    del model\n    torch.cuda.empty_cache()\n    return features\n\n\n# ===================== BiomedCLIP support ======================\ndef extract_features_biomedclip(meta, config):\n    \"\"\"Special handler for BiomedCLIP via open_clip.\"\"\"\n    safe_bb = \"biomedclip\"\n    cache_path = Path(config[\"cache_dir\"]) / f\"features_{meta['dataset']}_{safe_bb}_{len(meta['df'])}.npz\"\n    if cache_path.exists():\n        print(f\"[CACHED] {cache_path.name}\")\n        return np.load(cache_path)[\"features\"]\n\n    print(f\"Extracting: {meta['dataset']} × BiomedCLIP\")\n    try:\n        import open_clip\n    except ImportError:\n        print(\"  Installing open_clip_torch...\")\n        import subprocess\n        subprocess.run([\"pip\", \"install\", \"-q\", \"open_clip_torch\"], check=True)\n        import open_clip\n\n    model, preprocess = open_clip.create_model_from_pretrained(\n        \"hf-hub:microsoft/BiomedCLIP-PubMedBERT_256-vit_base_patch16_224\"\n    )\n    model = model.to(DEVICE).eval()\n\n    n_gpus = torch.cuda.device_count()\n    if n_gpus > 1:\n        print(f\"  Using {n_gpus} GPUs via DataParallel\")\n        model = nn.DataParallel(model)\n        effective_bs = config[\"batch_size\"] * n_gpus\n    else:\n        effective_bs = config[\"batch_size\"]\n\n    df = meta[\"df\"]\n\n    class _BiomedDS(Dataset):\n        def __init__(self, paths): self.paths = paths\n        def __len__(self): return len(self.paths)\n        def __getitem__(self, idx):\n            try:\n                img = Image.open(self.paths[idx]).convert(\"RGB\")\n            except Exception:\n                img = Image.new(\"RGB\", (224, 224))\n            return preprocess(img), idx\n\n    loader = DataLoader(_BiomedDS(df[\"path\"].tolist()),\n                        batch_size=effective_bs,\n                        num_workers=config[\"num_workers\"],\n                        shuffle=False, pin_memory=True)\n\n    # Probe feature dim\n    with torch.no_grad():\n        dummy = torch.zeros(1, 3, 224, 224).to(DEVICE)\n        m_inner = model.module if isinstance(model, nn.DataParallel) else model\n        feat_test = m_inner.encode_image(dummy)\n        feat_dim = feat_test.shape[-1]\n    print(f\"  BiomedCLIP feature dim: {feat_dim}\")\n\n    features = np.zeros((len(df), feat_dim), dtype=np.float32)\n    with torch.no_grad():\n        for imgs, idx in tqdm(loader, desc=f\"{meta['dataset']}/biomedclip\"):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            m_inner = model.module if isinstance(model, nn.DataParallel) else model\n            feat = m_inner.encode_image(imgs)\n            features[idx.numpy()] = feat.float().cpu().numpy()\n\n    np.savez_compressed(cache_path, features=features)\n    print(f\"Saved: {cache_path.name}\")\n    del model\n    torch.cuda.empty_cache()\n    return features\n\n\n# Wrapper that dispatches based on backbone\n_extract_features_orig = extract_features\ndef extract_features(meta, backbone_name, config):\n    if \"biomedclip\" in backbone_name.lower():\n        return extract_features_biomedclip(meta, config)\n    return _extract_features_orig(meta, backbone_name, config)\n\n\n# ===================== Contamination tracking ======================\ndef is_clean_config(dataset_name, backbone_name):\n    \"\"\"Return True if (dataset, backbone) combination is contamination-free.\n\n    RAD-DINO was pre-trained on MIMIC-CXR + CheXpert + PadChest + CXR14 + BRAX.\n    CXR14 = NIH ChestX-ray14, so NIH × RAD-DINO and CheXpert × RAD-DINO are\n    contaminated.\n    \"\"\"\n    bb_lower = backbone_name.lower()\n    if \"rad-dino\" in bb_lower:\n        if dataset_name in [\"nih\", \"chexpert\"]:\n            return False\n    return True\n\nprint(\"\\nClean config matrix:\")\nfor d in CONFIG[\"datasets\"]:\n    for b in CONFIG[\"backbones\"]:\n        clean = \"✓\" if is_clean_config(d, b) else \"✗ contaminated\"\n        print(f\"  {d:<10} × {b}: {clean}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.658461Z","iopub.execute_input":"2026-05-10T04:40:02.6588Z","iopub.status.idle":"2026-05-10T04:40:02.685906Z","shell.execute_reply.started":"2026-05-10T04:40:02.658776Z","shell.execute_reply":"2026-05-10T04:40:02.684978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dirichlet_partition(labels, num_clients, alpha, seed):\n    rng = np.random.default_rng(seed)\n    n, n_cls = labels.shape\n    primary = np.full(n, n_cls, dtype=int)\n    for i in range(n):\n        pos = np.where(labels[i] > 0)[0]\n        if len(pos) > 0: primary[i] = pos[0]\n\n    client_idx = [[] for _ in range(num_clients)]\n    for c in range(n_cls + 1):\n        ic = np.where(primary == c)[0]\n        if len(ic) == 0: continue\n        rng.shuffle(ic)\n        props = rng.dirichlet(np.repeat(alpha, num_clients))\n        cuts = (np.cumsum(props) * len(ic)).astype(int)[:-1]\n        splits = np.split(ic, cuts)\n        for k, s in enumerate(splits):\n            client_idx[k].extend(s.tolist())\n\n    out = []\n    for k in range(num_clients):\n        a = np.array(client_idx[k], dtype=int)\n        rng.shuffle(a)\n        out.append(a)\n    return out\n\nprint(\"dirichlet_partition defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.687262Z","iopub.execute_input":"2026-05-10T04:40:02.687658Z","iopub.status.idle":"2026-05-10T04:40:02.704538Z","shell.execute_reply.started":"2026-05-10T04:40:02.687632Z","shell.execute_reply":"2026-05-10T04:40:02.703771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SufficientStats:\n    def __init__(self, d, c, dtype=DTYPE_SERVER):\n        self.d, self.c, self.dtype = d, c, dtype\n        self.S = torch.zeros(d, d, dtype=dtype)\n        self.G = torch.zeros(d, c, dtype=dtype)\n        self.n = 0\n    def add(self, F, Y):\n        F = torch.as_tensor(F, dtype=self.dtype); Y = torch.as_tensor(Y, dtype=self.dtype)\n        self.S += F.T @ F; self.G += F.T @ Y; self.n += F.shape[0]\n    def remove(self, F, Y):\n        F = torch.as_tensor(F, dtype=self.dtype); Y = torch.as_tensor(Y, dtype=self.dtype)\n        self.S -= F.T @ F; self.G -= F.T @ Y; self.n -= F.shape[0]\n    def solve(self, gamma):\n        H = self.S + gamma * torch.eye(self.d, dtype=self.dtype)\n        try:\n            L = torch.linalg.cholesky(H)\n            return torch.cholesky_solve(self.G, L)\n        except Exception:\n            return torch.linalg.solve(H, self.G)\n    def solve_per_class(self, gamma_vec):\n        W = torch.zeros(self.d, self.c, dtype=self.dtype)\n        I_d = torch.eye(self.d, dtype=self.dtype)\n        for j in range(self.c):\n            H = self.S + gamma_vec[j] * I_d\n            try:\n                L = torch.linalg.cholesky(H)\n                W[:, j:j+1] = torch.cholesky_solve(self.G[:, j:j+1], L)\n            except Exception:\n                W[:, j:j+1] = torch.linalg.solve(H, self.G[:, j:j+1])\n        return W\n\n\ndef make_nystrom(features_pool, n_anchors, gamma_rbf, seed):\n    rng = np.random.default_rng(seed)\n    idx = rng.choice(features_pool.shape[0], min(n_anchors, len(features_pool)), replace=False)\n    anchors = features_pool[idx].astype(np.float64)\n    m = anchors.shape[0]\n    XX = np.sum(anchors**2, axis=1, keepdims=True)\n    K_mm = np.exp(-gamma_rbf * (XX + XX.T - 2 * anchors @ anchors.T)) + 1e-6 * np.eye(m)\n    eigvals, eigvecs = np.linalg.eigh(K_mm)\n    eigvals = np.maximum(eigvals, 1e-10)\n    inv_sqrt = eigvecs / np.sqrt(eigvals)\n    anchor_norms = np.sum(anchors**2, axis=1, keepdims=True).T\n    def transform(F, batch_size=2048):\n        F_np = np.asarray(F, dtype=np.float64); n = F_np.shape[0]\n        out = np.zeros((n, m), dtype=np.float64)\n        for i in range(0, n, batch_size):\n            b = F_np[i:i+batch_size]\n            b_norm = np.sum(b**2, axis=1, keepdims=True)\n            sqd = b_norm + anchor_norms - 2 * b @ anchors.T\n            out[i:i+batch_size] = np.exp(-gamma_rbf * sqd) @ inv_sqrt\n        return out\n    return transform, anchors\n\n\n# === Variants ===\nclass VariantV0:\n    name = \"V0_plain\"\n    def __init__(self, d, c, gamma):\n        self.d, self.c, self.gamma = d, c, gamma\n        self.stats = SufficientStats(d, c); self.W = None\n    def _phi(self, F): return torch.as_tensor(F, dtype=DTYPE_SERVER)\n    def _psi(self, Y): return torch.as_tensor(Y, dtype=DTYPE_SERVER)\n    def add(self, F, Y): self.stats.add(self._phi(F), self._psi(Y))\n    def remove(self, F, Y): self.stats.remove(self._phi(F), self._psi(Y))\n    def fit(self): self.W = self.stats.solve(self.gamma); return self.W\n    def predict(self, F): return (self._phi(F) @ self.W).cpu().numpy()\n    def get_W(self): return self.W\n\n\nclass VariantV1(VariantV0):\n    name = \"V1_nystrom\"\n    def __init__(self, c, gamma, nystrom_fn, n_anchors):\n        super().__init__(n_anchors, c, gamma)\n        self.nystrom_fn = nystrom_fn\n    def _phi(self, F):\n        return torch.as_tensor(self.nystrom_fn(F), dtype=DTYPE_SERVER)\n\n\nclass VariantV2(VariantV1):\n    \"\"\"Nyström + label smoothing + per-class γ (validation-tuned).\"\"\"\n    name = \"V2_valgamma\"\n    def __init__(self, c, gamma_vec, nystrom_fn, n_anchors, eps):\n        # Use first gamma as placeholder for parent class\n        super().__init__(c, float(gamma_vec[0]), nystrom_fn, n_anchors)\n        self.eps = eps\n        self.gamma_vec = torch.tensor(gamma_vec, dtype=DTYPE_SERVER)\n    def _psi(self, Y):\n        Y_t = torch.as_tensor(Y, dtype=DTYPE_SERVER)\n        return (1 - self.eps) * Y_t + 0.5 * self.eps\n    def fit(self):\n        self.W = self.stats.solve_per_class(self.gamma_vec); return self.W\n\n\nclass VariantV2_freq(VariantV1):\n    \"\"\"Frequency-based per-class γ (the original failing version).\"\"\"\n    name = \"V2_freqgamma\"\n    def __init__(self, c, gamma_base, nystrom_fn, n_anchors, eps, class_freq):\n        super().__init__(c, gamma_base, nystrom_fn, n_anchors)\n        self.eps = eps\n        n_max = max(class_freq.max(), 1)\n        self.gamma_vec = torch.tensor(\n            [gamma_base * np.sqrt(n_max / max(class_freq[j], 1)) for j in range(c)],\n            dtype=DTYPE_SERVER)\n    def _psi(self, Y):\n        Y_t = torch.as_tensor(Y, dtype=DTYPE_SERVER)\n        return (1 - self.eps) * Y_t + 0.5 * self.eps\n    def fit(self):\n        self.W = self.stats.solve_per_class(self.gamma_vec); return self.W\n\n\nprint(\"Variants ready: V0, V1, V2_freq, V2_val\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.705739Z","iopub.execute_input":"2026-05-10T04:40:02.706446Z","iopub.status.idle":"2026-05-10T04:40:02.730323Z","shell.execute_reply.started":"2026-05-10T04:40:02.706418Z","shell.execute_reply":"2026-05-10T04:40:02.72937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def tune_per_class_gamma(train_features, train_labels, val_features, val_labels,\n                         classes, gamma_grid, nystrom_fn, n_anchors, eps=0.0):\n    \"\"\"For each class, find γ_c maximizing validation AUROC.\n    \n    Important: train ONE Nyström model with shared (S, G); only γ varies per class.\n    This keeps everything additive and exact.\n    \"\"\"\n    var = VariantV1(len(classes), 1.0, nystrom_fn, n_anchors)\n    # Apply label smoothing if requested\n    if eps > 0:\n        var._psi = lambda Y: (1 - eps) * torch.as_tensor(Y, dtype=DTYPE_SERVER) + 0.5 * eps\n    var.add(train_features, train_labels)\n    # Now stats has S, G\n\n    F_val_lifted = var._phi(val_features)\n    d = var.stats.d\n    I_d = torch.eye(d, dtype=DTYPE_SERVER)\n\n    best_gammas = []\n    for j, cls in enumerate(classes):\n        if val_labels[:, j].sum() < 5:\n            best_gammas.append(1.0); continue\n        best_auc, best_g = -1, 1.0\n        for g in gamma_grid:\n            H = var.stats.S + g * I_d\n            try:\n                L = torch.linalg.cholesky(H)\n                W_j = torch.cholesky_solve(var.stats.G[:, j:j+1], L)\n            except Exception:\n                W_j = torch.linalg.solve(H, var.stats.G[:, j:j+1])\n            preds_j = (F_val_lifted @ W_j).cpu().numpy().flatten()\n            try:\n                auc = roc_auc_score(val_labels[:, j], preds_j)\n                if auc > best_auc:\n                    best_auc, best_g = auc, g\n            except Exception:\n                pass\n        best_gammas.append(best_g)\n    return best_gammas\n\n\ndef setup_nystrom_and_gamma(train_features, train_labels, val_features, val_labels,\n                            classes, config, seed):\n    \"\"\"One-time setup per (dataset, backbone, seed).\n    \n    NOTE: Anchors are sampled from val_features (held-out, non-revocable) so that\n    the Nyström lift ψ is independent of training data that may be deleted.\n    This ensures exactness in the privacy-preserving sense: deleting a training\n    sample fully removes its influence — its features never reside in ψ.\n    \"\"\"\n    rng = np.random.default_rng(seed)\n    # === CHANGED: Anchor pool from val instead of train ===\n    pool_size = min(5000, len(val_features))\n    pool_idx = rng.choice(len(val_features), pool_size, replace=False)\n    pool = val_features[pool_idx]\n\n    sample = pool[rng.choice(len(pool), min(500, len(pool)), replace=False)]\n    sqd = np.sum((sample[:, None, :] - sample[None, :, :])**2, axis=2)\n    median_sqd = np.median(sqd[sqd > 0])\n    rbf_gamma = 1.0 / max(median_sqd, 1e-6)\n\n    nystrom_fn, _ = make_nystrom(pool, config[\"nystrom_anchors\"], rbf_gamma, seed=seed)\n\n    # γ tuning still uses val_features (already non-revocable, no change needed)\n    gamma_vec = tune_per_class_gamma(\n        train_features, train_labels, val_features, val_labels,\n        classes, config[\"gamma_grid\"], nystrom_fn, config[\"nystrom_anchors\"],\n        eps=config[\"label_smoothing_eps\"]\n    )\n\n    class_freq = train_labels.sum(axis=0)\n\n    return {\n        \"nystrom_fn\": nystrom_fn,\n        \"rbf_gamma\": rbf_gamma,\n        \"gamma_vec_tuned\": gamma_vec,\n        \"class_freq\": class_freq,\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.731439Z","iopub.execute_input":"2026-05-10T04:40:02.731834Z","iopub.status.idle":"2026-05-10T04:40:02.75061Z","shell.execute_reply.started":"2026-05-10T04:40:02.73181Z","shell.execute_reply":"2026-05-10T04:40:02.749641Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_preds(preds, Y_test, classes):\n    out, aurocs, auprcs = {}, [], []\n    for j, cls in enumerate(classes):\n        if Y_test[:, j].sum() < 5:\n            out[f\"AUROC_{cls}\"] = float(\"nan\"); out[f\"AUPRC_{cls}\"] = float(\"nan\"); continue\n        try:\n            auc = roc_auc_score(Y_test[:, j], preds[:, j])\n            ap = average_precision_score(Y_test[:, j], preds[:, j])\n            aurocs.append(auc); auprcs.append(ap)\n            out[f\"AUROC_{cls}\"] = float(auc); out[f\"AUPRC_{cls}\"] = float(ap)\n        except Exception:\n            out[f\"AUROC_{cls}\"] = float(\"nan\"); out[f\"AUPRC_{cls}\"] = float(\"nan\")\n    out[\"macro_AUROC\"] = float(np.mean(aurocs)) if aurocs else float(\"nan\")\n    out[\"macro_AUPRC\"] = float(np.mean(auprcs)) if auprcs else float(\"nan\")\n    return out\n\n\ndef evaluate_variant(var, X_test, Y_test, classes):\n    return evaluate_preds(var.predict(X_test), Y_test, classes)\n\n\ndef make_variant(name, config, setup, classes):\n    c = len(classes)\n    if name == \"V0_plain\":\n        return VariantV0(config[\"feature_dim\"], c, config[\"ridge_gamma\"])\n    if name == \"V1_nystrom\":\n        return VariantV1(c, config[\"ridge_gamma\"], setup[\"nystrom_fn\"], config[\"nystrom_anchors\"])\n    if name == \"V2_freqgamma\":\n        return VariantV2_freq(c, config[\"ridge_gamma\"], setup[\"nystrom_fn\"],\n                              config[\"nystrom_anchors\"], config[\"label_smoothing_eps\"],\n                              setup[\"class_freq\"])\n    if name == \"V2_valgamma\":\n        return VariantV2(c, setup[\"gamma_vec_tuned\"], setup[\"nystrom_fn\"],\n                         config[\"nystrom_anchors\"], config[\"label_smoothing_eps\"])\n    raise ValueError(name)\n\n\ndef frob_deviation(W_a, W_b):\n    return float((torch.norm(W_a - W_b) / max(torch.norm(W_b).item(), 1e-30)).item())\n\n\ndef fed_train(name, config, setup, classes, train_features, train_labels, client_indices):\n    var = make_variant(name, config, setup, classes)\n    for idx in client_indices:\n        var.add(train_features[idx], train_labels[idx])\n    var.fit()\n    return var\n\n\ndef central_train(name, config, setup, classes, train_features, train_labels):\n    var = make_variant(name, config, setup, classes)\n    var.add(train_features, train_labels)\n    var.fit()\n    return var","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.751828Z","iopub.execute_input":"2026-05-10T04:40:02.752137Z","iopub.status.idle":"2026-05-10T04:40:02.770113Z","shell.execute_reply.started":"2026-05-10T04:40:02.752102Z","shell.execute_reply":"2026-05-10T04:40:02.769148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_central_ce(train_features, train_labels, test_features, test_labels,\n                     classes, num_epochs=50, lr=1e-3, wd=1e-4, bs=1024):\n    d, c = train_features.shape[1], train_labels.shape[1]\n    X = torch.tensor(train_features, dtype=torch.float32, device=DEVICE)\n    Y = torch.tensor(train_labels, dtype=torch.float32, device=DEVICE)\n    Xt = torch.tensor(test_features, dtype=torch.float32, device=DEVICE)\n    model = nn.Linear(d, c).to(DEVICE)\n    opt = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=wd)\n    for ep in range(num_epochs):\n        model.train()\n        perm = torch.randperm(len(X), device=DEVICE)\n        for i in range(0, len(X), bs):\n            idx = perm[i:i+bs]\n            loss = F.binary_cross_entropy_with_logits(model(X[idx]), Y[idx])\n            opt.zero_grad(); loss.backward(); opt.step()\n    model.eval()\n    with torch.no_grad():\n        preds = torch.sigmoid(model(Xt)).cpu().numpy()\n    return evaluate_preds(preds, test_labels, classes)\n\n\ndef train_fedavg_ce(train_features, train_labels, test_features, test_labels,\n                    classes, client_indices, num_rounds=20, local_epochs=2,\n                    lr=1e-3, bs=256):\n    d, c = train_features.shape[1], train_labels.shape[1]\n    K = len(client_indices)\n    Xt = torch.tensor(test_features, dtype=torch.float32, device=DEVICE)\n    global_model = nn.Linear(d, c).to(DEVICE)\n    Xs = [torch.tensor(train_features[idx], dtype=torch.float32, device=DEVICE)\n          for idx in client_indices]\n    Ys = [torch.tensor(train_labels[idx], dtype=torch.float32, device=DEVICE)\n          for idx in client_indices]\n    for rnd in range(num_rounds):\n        states, sizes = [], []\n        for k in range(K):\n            local = deepcopy(global_model)\n            opt = torch.optim.Adam(local.parameters(), lr=lr)\n            X_k, Y_k = Xs[k], Ys[k]\n            if len(X_k) == 0: continue\n            for _ in range(local_epochs):\n                perm = torch.randperm(len(X_k), device=DEVICE)\n                for i in range(0, len(X_k), bs):\n                    idx = perm[i:i+bs]\n                    loss = F.binary_cross_entropy_with_logits(local(X_k[idx]), Y_k[idx])\n                    opt.zero_grad(); loss.backward(); opt.step()\n            states.append({kk: vv.detach().clone() for kk, vv in local.state_dict().items()})\n            sizes.append(len(X_k))\n        total = sum(sizes)\n        new_state = {}\n        for key in states[0]:\n            new_state[key] = sum(s[key] * (sz/total) for s, sz in zip(states, sizes))\n        global_model.load_state_dict(new_state)\n    global_model.eval()\n    with torch.no_grad():\n        preds = torch.sigmoid(global_model(Xt)).cpu().numpy()\n    return evaluate_preds(preds, test_labels, classes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.773102Z","iopub.execute_input":"2026-05-10T04:40:02.773472Z","iopub.status.idle":"2026-05-10T04:40:02.789826Z","shell.execute_reply.started":"2026-05-10T04:40:02.773445Z","shell.execute_reply":"2026-05-10T04:40:02.788862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, shutil, pickle\n\ncache_dir = \"/kaggle/working/cache\"\n\n# Show what's in cache\nprint(\"=== Cache state BEFORE cleanup ===\")\nif os.path.exists(cache_dir):\n    for f in sorted(os.listdir(cache_dir)):\n        p = os.path.join(cache_dir, f)\n        size_mb = os.path.getsize(p) / 1e6\n        info = \"\"\n        if f.endswith('.pkl'):\n            try:\n                with open(p, 'rb') as fp:\n                    data = pickle.load(fp)\n                if isinstance(data, dict) and 'df' in data:\n                    info = f\" → {len(data['df'])} samples\"\n            except: pass\n        print(f\"  {f} ({size_mb:.1f} MB){info}\")\nelse:\n    print(\"  (empty)\")\n\n# DELETE ALL cached files\nshutil.rmtree(cache_dir, ignore_errors=True)\nos.makedirs(cache_dir, exist_ok=True)\nprint(\"\\n=== Cache cleared ===\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:40:02.791061Z","iopub.execute_input":"2026-05-10T04:40:02.791498Z","iopub.status.idle":"2026-05-10T04:40:02.80843Z","shell.execute_reply.started":"2026-05-10T04:40:02.79146Z","shell.execute_reply":"2026-05-10T04:40:02.807709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Patch: auto-detect feature_dim from train data\n# (V0_plain dùng raw features, các V1/V2 dùng Nyström n_anchors nên không bị)\n\n_fed_train_orig = fed_train\ndef fed_train(name, config, setup, classes, train_features, train_labels, client_indices):\n    config = {**config, \"feature_dim\": train_features.shape[1]}\n    return _fed_train_orig(name, config, setup, classes,\n                           train_features, train_labels, client_indices)\n\n_central_train_orig = central_train\ndef central_train(name, config, setup, classes, train_features, train_labels):\n    config = {**config, \"feature_dim\": train_features.shape[1]}\n    return _central_train_orig(name, config, setup, classes,\n                               train_features, train_labels)\n\nprint(\"✓ Patched fed_train and central_train to auto-detect feature_dim\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T04:58:42.606726Z","iopub.execute_input":"2026-05-10T04:58:42.607253Z","iopub.status.idle":"2026-05-10T04:58:42.614757Z","shell.execute_reply.started":"2026-05-10T04:58:42.607214Z","shell.execute_reply":"2026-05-10T04:58:42.614017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Patch setup_nystrom_and_gamma: use train_features for anchor pool\n_setup_orig = setup_nystrom_and_gamma\n\ndef setup_nystrom_and_gamma(train_features, train_labels, val_features, val_labels,\n                            classes, config, seed):\n    \"\"\"Anchors from train pool (always); val used only for γ tuning.\"\"\"\n    rng = np.random.default_rng(seed)\n    pool_size = min(5000, len(train_features))\n    pool_idx = rng.choice(len(train_features), pool_size, replace=False)\n    pool = train_features[pool_idx]\n\n    sample = pool[rng.choice(len(pool), min(500, len(pool)), replace=False)]\n    sqd = np.sum((sample[:, None, :] - sample[None, :, :])**2, axis=2)\n    median_sqd = np.median(sqd[sqd > 0])\n    rbf_gamma = 1.0 / max(median_sqd, 1e-6)\n\n    nystrom_fn, _ = make_nystrom(pool, config[\"nystrom_anchors\"], rbf_gamma, seed=seed)\n\n    gamma_vec = tune_per_class_gamma(\n        train_features, train_labels, val_features, val_labels,\n        classes, config[\"gamma_grid\"], nystrom_fn, config[\"nystrom_anchors\"],\n        eps=config[\"label_smoothing_eps\"]\n    )\n    class_freq = train_labels.sum(axis=0)\n    return {\n        \"nystrom_fn\": nystrom_fn, \"rbf_gamma\": rbf_gamma,\n        \"gamma_vec_tuned\": gamma_vec, \"class_freq\": class_freq,\n    }\n\nprint(\"✓ Patched setup_nystrom_and_gamma to use train_features for anchor pool\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:06:16.842947Z","iopub.execute_input":"2026-05-10T05:06:16.843837Z","iopub.status.idle":"2026-05-10T05:06:16.853749Z","shell.execute_reply.started":"2026-05-10T05:06:16.843797Z","shell.execute_reply":"2026-05-10T05:06:16.85306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport os\n\nbase = CONFIG[\"chexpert_data_dir\"]\nprint(f\"base: {base}\")\n\n# Find train.csv\ncsv_path = None\nfor cand in [os.path.join(base, \"train.csv\"),\n             os.path.join(base, \"CheXpert-v1.0/train.csv\"),\n             os.path.join(base, \"CheXpert-v1.0-small/train.csv\")]:\n    if os.path.exists(cand):\n        csv_path = cand\n        break\nprint(f\"CSV: {csv_path}\")\nprint(f\"CSV directory: {os.path.dirname(csv_path)}\")\n\ndf = pd.read_csv(csv_path)\nprint(f\"\\nFirst 3 paths in CSV:\")\nfor p in df[\"Path\"].head(3):\n    print(f\"  {p}\")\n\n# List directories near base\nprint(f\"\\nContents of base:\")\nfor f in sorted(os.listdir(base))[:10]:\n    print(f\"  {f}\")\n\n# Try common path resolution\nsample_csv_path = df[\"Path\"].iloc[0]\nprint(f\"\\nSample resolution attempts for: {sample_csv_path}\")\ncandidates = [\n    os.path.join(base, sample_csv_path),\n    os.path.join(os.path.dirname(base), sample_csv_path),\n    os.path.join(base, sample_csv_path.replace(\"CheXpert-v1.0-small/\", \"\")),\n    os.path.join(base, sample_csv_path.replace(\"CheXpert-v1.0/\", \"\")),\n]\nfor c in candidates:\n    print(f\"  {os.path.exists(c)}: {c}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:20:54.371666Z","iopub.execute_input":"2026-05-10T05:20:54.372034Z","iopub.status.idle":"2026-05-10T05:21:04.870551Z","shell.execute_reply.started":"2026-05-10T05:20:54.371998Z","shell.execute_reply":"2026-05-10T05:21:04.869648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Robust path resolver for CheXpert\ndef prepare_chexpert(config):\n    cache_path = Path(config[\"cache_dir\"]) / \"metadata_chexpert.pkl\"\n    if cache_path.exists():\n        with open(cache_path, \"rb\") as f:\n            cached = pickle.load(f)\n        if abs(len(cached[\"df\"]) - config[\"subset_size\"]) / config[\"subset_size\"] <= 0.2:\n            return cached\n\n    base = config[\"chexpert_data_dir\"]\n    csv_path = None\n    for cand in [os.path.join(base, \"train.csv\"),\n                 os.path.join(base, \"CheXpert-v1.0/train.csv\"),\n                 os.path.join(base, \"CheXpert-v1.0-small/train.csv\")]:\n        if os.path.exists(cand):\n            csv_path = cand; break\n    if csv_path is None:\n        for root, dirs, files in os.walk(base):\n            if \"train.csv\" in files:\n                csv_path = os.path.join(root, \"train.csv\"); break\n\n    df_raw = pd.read_csv(csv_path)\n    print(f\"[CheXpert] Raw rows: {len(df_raw)}\")\n\n    if \"Frontal/Lateral\" in df_raw.columns:\n        df_raw = df_raw[df_raw[\"Frontal/Lateral\"] == \"Frontal\"].copy()\n    print(f\"[CheXpert] After frontal filter: {len(df_raw)}\")\n\n    name_map = {\n        \"Atelectasis\": \"Atelectasis\",\n        \"Cardiomegaly\": \"Cardiomegaly\",\n        \"Consolidation\": \"Consolidation\",\n        \"Effusion\": \"Pleural Effusion\",\n        \"Pneumothorax\": \"Pneumothorax\",\n    }\n    selected = config[\"selected_classes\"]\n    label_matrix = np.zeros((len(df_raw), len(selected)), dtype=np.float32)\n    for j, c in enumerate(selected):\n        col = name_map[c]\n        if col in df_raw.columns:\n            vals = df_raw[col].fillna(0).values\n            label_matrix[:, j] = ((vals == 1) | (vals == -1)).astype(np.float32)\n\n    def extract_pid(path_str):\n        for p in str(path_str).split(\"/\"):\n            if p.startswith(\"patient\"):\n                return p\n        return str(path_str)\n\n    df_raw[\"patient_id\"] = df_raw[\"Path\"].apply(extract_pid)\n    df_raw = df_raw.reset_index(drop=True)\n\n    rng = np.random.default_rng(42)\n    pids = df_raw[\"patient_id\"].unique()\n    rng.shuffle(pids)\n    selected_p, n = [], 0\n    for pid in pids:\n        np_pid = (df_raw[\"patient_id\"] == pid).sum()\n        if n + np_pid > config[\"subset_size\"]:\n            if n == 0: selected_p.append(pid); n += np_pid\n            break\n        selected_p.append(pid); n += np_pid\n\n    mask = df_raw[\"patient_id\"].isin(selected_p).values\n    df_raw = df_raw.loc[mask].reset_index(drop=True)\n    label_matrix = label_matrix[mask]\n\n    # === Robust path resolver: try multiple bases + prefix strips ===\n    csv_dir = os.path.dirname(csv_path)\n    parent_dir = os.path.dirname(csv_dir)\n\n    paths = []\n    for p_str in df_raw[\"Path\"].values:\n        full = None\n        # Try bases\n        for base_cand in [csv_dir, base, parent_dir]:\n            # Try as-is and with prefix stripped\n            for variant in [p_str,\n                            p_str.replace(\"CheXpert-v1.0-small/\", \"\").replace(\"CheXpert-v1.0/\", \"\"),\n                            p_str.split(\"/\", 1)[-1] if \"/\" in p_str else p_str]:\n                cand = os.path.join(base_cand, variant)\n                if os.path.exists(cand):\n                    full = cand; break\n            if full: break\n        paths.append(full)\n\n    valid = [i for i, p in enumerate(paths) if p is not None]\n    print(f\"[CheXpert] Valid images: {len(valid)} / {len(paths)}\")\n    if len(valid) == 0:\n        # Show what was tried\n        print(f\"  csv_dir={csv_dir}\")\n        print(f\"  Sample p_str: {df_raw['Path'].iloc[0]}\")\n        raise FileNotFoundError(\"No CheXpert images found. Check structure.\")\n\n    df = pd.DataFrame({\n        \"Image Index\": [paths[i] for i in valid],\n        \"patient_id\": df_raw[\"patient_id\"].values[valid],\n        \"path\": [paths[i] for i in valid],\n    })\n    label_matrix = label_matrix[valid]\n\n    print(f\"[CheXpert] Final: {len(df)} samples, {df['patient_id'].nunique()} patients\")\n    for j, c in enumerate(selected):\n        pos = int(label_matrix[:, j].sum())\n        print(f\"  {c}: {pos} positives ({pos/len(df):.2%})\")\n\n    out = {\"dataset\": \"chexpert\", \"df\": df, \"labels\": label_matrix}\n    with open(cache_path, \"wb\") as f:\n        pickle.dump(out, f)\n    return out\n\nprint(\"✓ prepare_chexpert patched with robust path resolver\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:23:23.208973Z","iopub.execute_input":"2026-05-10T05:23:23.209676Z","iopub.status.idle":"2026-05-10T05:23:23.227342Z","shell.execute_reply.started":"2026-05-10T05:23:23.209635Z","shell.execute_reply":"2026-05-10T05:23:23.226483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nold = \"/kaggle/working/cache/metadata_chexpert.pkl\"\nif os.path.exists(old): os.remove(old)\nprint(\"Cleared old chexpert cache\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:23:26.422971Z","iopub.execute_input":"2026-05-10T05:23:26.423682Z","iopub.status.idle":"2026-05-10T05:23:26.428386Z","shell.execute_reply.started":"2026-05-10T05:23:26.423648Z","shell.execute_reply":"2026-05-10T05:23:26.427658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"RQ1 — Multi-seed × multi-config (clean configs only)\")\nprint(\"=\"*60)\n\nVARIANT_NAMES = [\"V0_plain\", \"V1_nystrom\", \"V2_freqgamma\", \"V2_valgamma\"]\nall_rq1 = {}\n\nfor dataset_name in CONFIG[\"datasets\"]:\n    meta = prepare_metadata(dataset_name, CONFIG)\n    for backbone_name in CONFIG[\"backbones\"]:\n        # Skip contaminated combinations\n        if not is_clean_config(dataset_name, backbone_name):\n            print(f\"\\n[SKIP] {dataset_name} × {backbone_name} — contaminated\")\n            continue\n\n        features = extract_features(meta, backbone_name, CONFIG)\n        config_key = f\"{dataset_name}__{backbone_name.replace('/','_')}\"\n        print(f\"\\n>>> {config_key}\")\n\n        per_seed_results = {name: [] for name in VARIANT_NAMES + [\"Centralized_CE\", \"FedAvg_CE\"]}\n        per_seed_frob = {name: [] for name in VARIANT_NAMES}\n\n        for seed in CONFIG[\"seeds\"]:\n            print(f\"\\n  --- seed {seed} ---\")\n            set_seed(seed)\n            split = split_train_val_test(meta, CONFIG, seed)\n            tr_F = features[split[\"train_idx\"]]\n            tr_Y = meta[\"labels\"][split[\"train_idx\"]]\n            va_F = features[split[\"val_idx\"]]\n            va_Y = meta[\"labels\"][split[\"val_idx\"]]\n            te_F = features[split[\"test_idx\"]]\n            te_Y = meta[\"labels\"][split[\"test_idx\"]]\n\n            client_indices = dirichlet_partition(\n                tr_Y, CONFIG[\"num_clients\"], CONFIG[\"non_iid_alpha\"], seed=seed\n            )\n            setup = setup_nystrom_and_gamma(\n                tr_F, tr_Y, va_F, va_Y, CONFIG[\"selected_classes\"], CONFIG, seed\n            )\n            print(f\"    tuned γ_c: {[f'{g:.2f}' for g in setup['gamma_vec_tuned']]}\")\n\n            for name in VARIANT_NAMES:\n                var_fed = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                    tr_F, tr_Y, client_indices)\n                var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                        tr_F, tr_Y)\n                metrics = evaluate_variant(var_fed, te_F, te_Y, CONFIG[\"selected_classes\"])\n                fdev = frob_deviation(var_fed.get_W(), var_cen.get_W())\n                per_seed_results[name].append(metrics)\n                per_seed_frob[name].append(fdev)\n                print(f\"    {name}: macro={metrics['macro_AUROC']:.4f}, frob={fdev:.2e}\")\n\n            ce_c = train_central_ce(tr_F, tr_Y, te_F, te_Y, CONFIG[\"selected_classes\"])\n            ce_f = train_fedavg_ce(tr_F, tr_Y, te_F, te_Y, CONFIG[\"selected_classes\"], client_indices)\n            per_seed_results[\"Centralized_CE\"].append(ce_c)\n            per_seed_results[\"FedAvg_CE\"].append(ce_f)\n            print(f\"    Cent_CE: macro={ce_c['macro_AUROC']:.4f}\")\n            print(f\"    FedAvg_CE: macro={ce_f['macro_AUROC']:.4f}\")\n\n        agg = {}\n        for name in per_seed_results:\n            seeds_metrics = per_seed_results[name]\n            if not seeds_metrics: continue\n            keys = seeds_metrics[0].keys()\n            agg[name] = {}\n            for k in keys:\n                vals = [m[k] for m in seeds_metrics if not np.isnan(m.get(k, np.nan))]\n                if vals:\n                    agg[name][k+\"_mean\"] = float(np.mean(vals))\n                    agg[name][k+\"_std\"] = float(np.std(vals))\n            if name in per_seed_frob:\n                f_vals = per_seed_frob[name]\n                agg[name][\"frob_mean\"] = float(np.mean(f_vals))\n                agg[name][\"frob_max\"] = float(np.max(f_vals))\n            agg[name][\"per_seed\"] = seeds_metrics\n        all_rq1[config_key] = agg\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"rq1_all.json\", \"w\") as f:\n    json.dump(all_rq1, f, indent=2, default=float)\nprint(f\"\\nRQ1 saved: {len(all_rq1)} clean configs × {len(CONFIG['seeds'])} seeds\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:23:29.492292Z","iopub.execute_input":"2026-05-10T05:23:29.493233Z","iopub.status.idle":"2026-05-10T05:47:31.288804Z","shell.execute_reply.started":"2026-05-10T05:23:29.4932Z","shell.execute_reply":"2026-05-10T05:47:31.287505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"RQ2 — primary config × multi-seed\")\nprint(\"=\"*60)\n\nprimary_ds = CONFIG[\"rq2_primary_dataset\"]\nprimary_bb = CONFIG[\"rq2_primary_backbone\"]\nprint(f\"Primary: {primary_ds} × {primary_bb}\")\n\nmeta = prepare_metadata(primary_ds, CONFIG)\nfeatures = extract_features(meta, primary_bb, CONFIG)\n\nRQ2_VARIANTS = [\"V0_plain\", \"V2_valgamma\"]\nrq2_results = {\"rq2a\": {}, \"rq2b\": {}, \"rq2c\": {}}\nrare_class_name = CONFIG[\"rare_class\"]\nrare_key = f\"AUROC_{rare_class_name}\"\nrare_idx = CONFIG[\"selected_classes\"].index(rare_class_name)\n\nfor seed in CONFIG[\"seeds\"]:\n    print(f\"\\n--- seed {seed} ---\")\n    set_seed(seed)\n    split = split_train_val_test(meta, CONFIG, seed)\n    tr_F, tr_Y = features[split[\"train_idx\"]], meta[\"labels\"][split[\"train_idx\"]]\n    va_F, va_Y = features[split[\"val_idx\"]], meta[\"labels\"][split[\"val_idx\"]]\n    te_F, te_Y = features[split[\"test_idx\"]], meta[\"labels\"][split[\"test_idx\"]]\n    tr_pids = meta[\"df\"][\"patient_id\"].values[split[\"train_idx\"]]\n\n    client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                          CONFIG[\"non_iid_alpha\"], seed=seed)\n    setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                     CONFIG[\"selected_classes\"], CONFIG, seed)\n\n    for name in RQ2_VARIANTS:\n        # ===== RQ2a: image-level withdrawal for VinDr primary config =====\n        var = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                        tr_F, tr_Y, client_indices)\n        rng = np.random.default_rng(seed)\n        unique_pids = np.unique(tr_pids); rng.shuffle(unique_pids)\n        pids_to_del = unique_pids[:CONFIG[\"n_patient_deletions\"]]\n        deletion_unit = \"image\" if primary_ds == \"vindr\" else \"patient\"\n        timeline = [{\"step\": 0, \"n_deleted_patients\": 0, \"n_deleted_samples\": 0,\n                     \"n_deleted_units\": 0, \"deletion_unit\": deletion_unit,\n                     \"cum_time\": 0.0,\n                     \"metrics\": evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])}]\n        cum, ns = 0.0, 0\n        for i, pid in enumerate(pids_to_del):\n            mask = tr_pids == pid\n            if not mask.any(): continue\n            ns += int(mask.sum())\n            t0 = time.time()\n            var.remove(tr_F[mask], tr_Y[mask]); var.fit()\n            cum += time.time() - t0\n            if (i+1) % 25 == 0 or (i+1) == len(pids_to_del):\n                timeline.append({\"step\": i+1, \"n_deleted_patients\": i+1,\n                                 \"n_deleted_samples\": ns,\n                                 \"n_deleted_units\": ns if deletion_unit == \"image\" else i+1,\n                                 \"deletion_unit\": deletion_unit, \"cum_time\": cum,\n                                 \"metrics\": evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])})\n        rem_mask = ~np.isin(tr_pids, pids_to_del)\n        var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                tr_F[rem_mask], tr_Y[rem_mask])\n        f2a = frob_deviation(var.get_W(), var_cen.get_W())\n        rq2_results[\"rq2a\"].setdefault(name, []).append({\"timeline\": timeline, \"frob\": f2a})\n        m_final = timeline[-1][\"metrics\"]\n        print(f\"  {name} RQ2a: macro={m_final['macro_AUROC']:.4f}, time={cum:.2f}s, frob={f2a:.2e}\")\n\n        # ===== RQ2b: site withdrawal =====\n        var = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                        tr_F, tr_Y, client_indices)\n        m_b = evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])\n        w_idx = client_indices[CONFIG[\"withdrawn_site_idx\"]]\n        t0 = time.time()\n        var.remove(tr_F[w_idx], tr_Y[w_idx]); var.fit()\n        elapsed = time.time() - t0\n        m_a = evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])\n        rem_idx = np.concatenate([client_indices[k] for k in range(len(client_indices))\n                                   if k != CONFIG[\"withdrawn_site_idx\"]])\n        var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                tr_F[rem_idx], tr_Y[rem_idx])\n        f2b = frob_deviation(var.get_W(), var_cen.get_W())\n        rq2_results[\"rq2b\"].setdefault(name, []).append({\n            \"before\": m_b, \"after\": m_a, \"n_removed\": int(len(w_idx)),\n            \"wall_time_s\": float(elapsed), \"frob\": f2b})\n        print(f\"  {name} RQ2b: macro {m_b['macro_AUROC']:.4f}→{m_a['macro_AUROC']:.4f}, frob={f2b:.2e}\")\n\n        # ===== RQ2c: rare disease deletion =====\n        var = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                        tr_F, tr_Y, client_indices)\n        m_b = evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])\n        pos_mask = tr_Y[:, rare_idx] > 0\n        t0 = time.time()\n        var.remove(tr_F[pos_mask], tr_Y[pos_mask]); var.fit()\n        elapsed = time.time() - t0\n        m_a = evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])\n        var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                tr_F[~pos_mask], tr_Y[~pos_mask])\n        f2c = frob_deviation(var.get_W(), var_cen.get_W())\n        rq2_results[\"rq2c\"].setdefault(name, []).append({\n            \"before\": m_b, \"after\": m_a, \"n_removed\": int(pos_mask.sum()),\n            \"wall_time_s\": float(elapsed), \"frob\": f2c})\n        print(f\"  {name} RQ2c: rare {m_b[rare_key]:.4f}→{m_a[rare_key]:.4f}, frob={f2c:.2e}\")\n\nrq2_results[\"primary_config\"] = f\"{primary_ds}__{primary_bb.replace('/','_')}\"\nwith open(Path(CONFIG[\"output_dir\"]) / \"rq2_all.json\", \"w\") as f:\n    json.dump(rq2_results, f, indent=2, default=float)\nprint(f\"\\nRQ2 saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:49:22.113958Z","iopub.execute_input":"2026-05-10T05:49:22.11431Z","iopub.status.idle":"2026-05-10T05:51:47.379986Z","shell.execute_reply.started":"2026-05-10T05:49:22.114276Z","shell.execute_reply":"2026-05-10T05:51:47.379139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"RQ2a — true patient-level withdrawal on NIH × DINOv2\")\nprint(\"=\"*60)\n\npatient_ds = \"nih\"\npatient_bb = \"facebook/dinov2-base\"\npatient_key = f\"{patient_ds}__{patient_bb.replace('/', '_')}\"\nprint(f\"Patient-level config: {patient_ds} × {patient_bb}\")\n\nmeta_patient = prepare_metadata(patient_ds, CONFIG)\nfeatures_patient = extract_features(meta_patient, patient_bb, CONFIG)\n\nrq2_patient_results = {\n    \"dataset\": patient_ds,\n    \"backbone\": patient_bb,\n    \"config_key\": patient_key,\n    \"deletion_unit\": \"patient\",\n    \"rq2a\": {},\n}\n\nfor seed in CONFIG[\"seeds\"]:\n    print(f\"\\n--- seed {seed} ---\")\n    set_seed(seed)\n    split = split_train_val_test(meta_patient, CONFIG, seed)\n    tr_F = features_patient[split[\"train_idx\"]]\n    tr_Y = meta_patient[\"labels\"][split[\"train_idx\"]]\n    va_F = features_patient[split[\"val_idx\"]]\n    va_Y = meta_patient[\"labels\"][split[\"val_idx\"]]\n    te_F = features_patient[split[\"test_idx\"]]\n    te_Y = meta_patient[\"labels\"][split[\"test_idx\"]]\n    tr_pids = meta_patient[\"df\"][\"patient_id\"].values[split[\"train_idx\"]]\n\n    client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                          CONFIG[\"non_iid_alpha\"], seed=seed)\n    setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                     CONFIG[\"selected_classes\"], CONFIG, seed)\n\n    rng = np.random.default_rng(seed)\n    unique_pids = np.unique(tr_pids)\n    rng.shuffle(unique_pids)\n    pids_to_del = unique_pids[:CONFIG[\"n_patient_deletions\"]]\n\n    for name in RQ2_VARIANTS:\n        var = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                        tr_F, tr_Y, client_indices)\n        timeline = [{\n            \"step\": 0,\n            \"n_deleted_patients\": 0,\n            \"n_deleted_samples\": 0,\n            \"cum_time\": 0.0,\n            \"metrics\": evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"]),\n        }]\n        cum, n_deleted_samples = 0.0, 0\n        for i, pid in enumerate(pids_to_del):\n            mask = tr_pids == pid\n            if not mask.any():\n                continue\n            n_deleted_samples += int(mask.sum())\n            t0 = time.time()\n            var.remove(tr_F[mask], tr_Y[mask])\n            var.fit()\n            cum += time.time() - t0\n            if (i + 1) % 25 == 0 or (i + 1) == len(pids_to_del):\n                timeline.append({\n                    \"step\": i + 1,\n                    \"n_deleted_patients\": i + 1,\n                    \"n_deleted_samples\": n_deleted_samples,\n                    \"cum_time\": cum,\n                    \"metrics\": evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"]),\n                })\n\n        rem_mask = ~np.isin(tr_pids, pids_to_del)\n        var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                tr_F[rem_mask], tr_Y[rem_mask])\n        fdev = frob_deviation(var.get_W(), var_cen.get_W())\n        rq2_patient_results[\"rq2a\"].setdefault(name, []).append({\n            \"seed\": seed,\n            \"deleted_patient_ids\": [int(x) for x in pids_to_del],\n            \"timeline\": timeline,\n            \"frob\": fdev,\n        })\n        final_metrics = timeline[-1][\"metrics\"]\n        print(\n            f\"  {name}: macro={final_metrics['macro_AUROC']:.4f}, \"\n            f\"patients={timeline[-1]['n_deleted_patients']}, \"\n            f\"images={timeline[-1]['n_deleted_samples']}, \"\n            f\"time={cum:.2f}s, frob={fdev:.2e}\"\n        )\n\n# Save raw JSON.\nout_json = Path(CONFIG[\"output_dir\"]) / \"rq2_patient_nih_dinov2.json\"\nwith open(out_json, \"w\") as f:\n    json.dump(rq2_patient_results, f, indent=2, default=float)\nprint(f\"Saved {out_json}\")\n\n# Save compact table for the manuscript/supplement.\nrows = []\nfor name in RQ2_VARIANTS:\n    per_seed = rq2_patient_results[\"rq2a\"][name]\n    before = [d[\"timeline\"][0][\"metrics\"][\"macro_AUROC\"] for d in per_seed]\n    after = [d[\"timeline\"][-1][\"metrics\"][\"macro_AUROC\"] for d in per_seed]\n    times = [d[\"timeline\"][-1][\"cum_time\"] for d in per_seed]\n    frobs = [d[\"frob\"] for d in per_seed]\n    n_patients = [d[\"timeline\"][-1][\"n_deleted_patients\"] for d in per_seed]\n    n_samples = [d[\"timeline\"][-1][\"n_deleted_samples\"] for d in per_seed]\n    rows.append({\n        \"Method\": name,\n        \"macro_AUROC_before\": f\"{np.mean(before):.4f} ± {np.std(before):.4f}\",\n        \"macro_AUROC_after\": f\"{np.mean(after):.4f} ± {np.std(after):.4f}\",\n        \"delta_pp\": f\"{(np.mean(after) - np.mean(before)) * 100:.2f}\",\n        \"n_patients_removed\": int(np.mean(n_patients)),\n        \"n_images_removed\": f\"{np.mean(n_samples):.1f} ± {np.std(n_samples):.1f}\",\n        \"time_s\": f\"{np.mean(times):.2f} ± {np.std(times):.2f}\",\n        \"Frob_max\": f\"{np.max(frobs):.2e}\",\n    })\npatient_table = pd.DataFrame(rows)\nout_csv = Path(CONFIG[\"tables_dir\"]) / \"table2b_patient_nih_dinov2.csv\"\npatient_table.to_csv(out_csv, index=False)\nprint(patient_table.to_string(index=False))\nprint(f\"Saved {out_csv}\")\n\n# Save patient-level timeline figure.\nfig, axes = plt.subplots(1, 2, figsize=(13, 4.5))\nfor name, color in zip(RQ2_VARIANTS, ['#888', '#7cb342']):\n    all_timelines = rq2_patient_results[\"rq2a\"][name]\n    steps = [t[\"n_deleted_patients\"] for t in all_timelines[0][\"timeline\"]]\n    macro_per_seed = np.array([[t[\"metrics\"][\"macro_AUROC\"] for t in tl[\"timeline\"]]\n                               for tl in all_timelines])\n    cum_per_seed = np.array([[t[\"cum_time\"] for t in tl[\"timeline\"]]\n                             for tl in all_timelines])\n    axes[0].errorbar(steps, macro_per_seed.mean(0), yerr=macro_per_seed.std(0),\n                     label=name, color=color, marker='o', capsize=3)\n    axes[1].errorbar(steps, cum_per_seed.mean(0), yerr=cum_per_seed.std(0),\n                     label=name, color=color, marker='o', capsize=3)\naxes[0].set_xlabel(\"# Patient deletions\")\naxes[0].set_ylabel(\"macro AUROC\")\naxes[0].set_title(\"True patient-level withdrawal on NIH×DINOv2\")\naxes[0].legend(); axes[0].grid(True, alpha=0.3)\naxes[1].set_xlabel(\"# Patient deletions\")\naxes[1].set_ylabel(\"Cumulative time (s)\")\naxes[1].set_title(\"Wall-clock cost\")\naxes[1].legend(); axes[1].grid(True, alpha=0.3)\nplt.tight_layout()\nout_fig = Path(CONFIG[\"figures_dir\"]) / \"fig3b_patient_nih_dinov2.png\"\nplt.savefig(out_fig, bbox_inches='tight')\nplt.show()\nprint(f\"Saved {out_fig}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nGenerating figures...\")\nplt.rcParams.update({\"font.size\": 10, \"figure.dpi\": 110})\n\nvariant_order = [\"V0_plain\", \"V1_nystrom\", \"V2_freqgamma\", \"V2_valgamma\"]\nclasses = CONFIG[\"selected_classes\"]\n\n# === Figure 1: macro AUROC across configs (mean ± std) ===\nn_configs = len(all_rq1)\nfig, axes = plt.subplots(1, n_configs, figsize=(5*n_configs, 5), squeeze=False)\naxes = axes.flatten()\nfor ax_idx, (config_key, agg) in enumerate(all_rq1.items()):\n    ax = axes[ax_idx]\n    means = [agg[v].get(\"macro_AUROC_mean\", np.nan) for v in variant_order]\n    stds = [agg[v].get(\"macro_AUROC_std\", 0) for v in variant_order]\n    bars = ax.bar(range(len(variant_order)), means, yerr=stds, capsize=5,\n                  color=['#888', '#4c9aaf', '#e89c5a', '#7cb342'], edgecolor='black')\n    # Reference lines\n    ce_mean = agg.get(\"Centralized_CE\", {}).get(\"macro_AUROC_mean\", None)\n    fa_mean = agg.get(\"FedAvg_CE\", {}).get(\"macro_AUROC_mean\", None)\n    if ce_mean: ax.axhline(ce_mean, color='red', linestyle='--', alpha=0.6, label=f'Cent CE ({ce_mean:.3f})')\n    if fa_mean: ax.axhline(fa_mean, color='blue', linestyle=':', alpha=0.6, label=f'FedAvg CE ({fa_mean:.3f})')\n    ax.set_xticks(range(len(variant_order)))\n    ax.set_xticklabels(variant_order, rotation=20, fontsize=8)\n    ax.set_ylabel(\"macro AUROC\")\n    ax.set_title(config_key, fontsize=10)\n    ax.legend(fontsize=8)\n    ax.grid(True, alpha=0.3, axis='y')\n    ax.set_ylim(min(means)-0.02 if means else 0.5, max(means)+0.04 if means else 1.0)\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig1_rq1_macro_auroc.png\", bbox_inches='tight')\nplt.show()\n\n# === Figure 2: Per-class AUROC for primary config ===\nprimary_key = f\"{CONFIG['rq2_primary_dataset']}__{CONFIG['rq2_primary_backbone'].replace('/','_')}\"\nif primary_key in all_rq1:\n    agg = all_rq1[primary_key]\n    fig, ax = plt.subplots(figsize=(11, 5))\n    x = np.arange(len(classes))\n    width = 0.18\n    colors = ['#888', '#4c9aaf', '#e89c5a', '#7cb342']\n    for i, v in enumerate(variant_order):\n        means = [agg[v].get(f\"AUROC_{c}_mean\", np.nan) for c in classes]\n        stds = [agg[v].get(f\"AUROC_{c}_std\", 0) for c in classes]\n        ax.bar(x + i*width, means, width, yerr=stds, capsize=3, label=v,\n               color=colors[i], edgecolor='black', linewidth=0.5)\n    ax.set_xticks(x + width*1.5)\n    ax.set_xticklabels(classes, rotation=15)\n    ax.set_ylabel(\"AUROC\")\n    ax.set_title(f\"Per-class AUROC ({primary_key}, {len(CONFIG['seeds'])} seeds)\")\n    ax.legend(loc='lower right')\n    ax.grid(True, alpha=0.3, axis='y')\n    plt.tight_layout()\n    plt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig2_per_class_primary.png\", bbox_inches='tight')\n    plt.show()\n\n# === Figure 3: RQ2a image-level withdrawal timeline (VinDr primary config) ===\nif \"rq2a\" in rq2_results:\n    fig, axes = plt.subplots(1, 2, figsize=(13, 4.5))\n    for name, color in zip(RQ2_VARIANTS, ['#888', '#7cb342']):\n        all_timelines = rq2_results[\"rq2a\"][name]\n        # Average across seeds at each step\n        steps = [t.get(\"n_deleted_units\", t.get(\"n_deleted_samples\", t.get(\"n_deleted_patients\", t[\"step\"])))\n                 for t in all_timelines[0][\"timeline\"]]\n        macro_per_seed = np.array([[t[\"metrics\"][\"macro_AUROC\"] for t in tl[\"timeline\"]] for tl in all_timelines])\n        cum_per_seed = np.array([[t[\"cum_time\"] for t in tl[\"timeline\"]] for tl in all_timelines])\n        macro_mean, macro_std = macro_per_seed.mean(0), macro_per_seed.std(0)\n        cum_mean, cum_std = cum_per_seed.mean(0), cum_per_seed.std(0)\n        axes[0].errorbar(steps, macro_mean, yerr=macro_std, label=name, color=color,\n                         marker='o', capsize=3)\n        axes[1].errorbar(steps, cum_mean, yerr=cum_std, label=name, color=color,\n                         marker='o', capsize=3)\n    axes[0].set_xlabel(\"# Image deletions\"); axes[0].set_ylabel(\"macro AUROC\")\n    axes[0].set_title(\"Utility during image-level withdrawal stream\"); axes[0].legend(); axes[0].grid(True, alpha=0.3)\n    axes[1].set_xlabel(\"# Image deletions\"); axes[1].set_ylabel(\"Cumulative time (s)\")\n    axes[1].set_title(\"Wall-clock cost\"); axes[1].legend(); axes[1].grid(True, alpha=0.3)\n    plt.tight_layout()\n    plt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig3_rq2a.png\", bbox_inches='tight')\n    plt.show()\n\n# === Figure 4: Exactness verification (log scale) ===\nfig, ax = plt.subplots(figsize=(11, 4))\nlabels_f, values_f = [], []\nfor config_key, agg in all_rq1.items():\n    for v in variant_order:\n        if v in agg and \"frob_max\" in agg[v]:\n            labels_f.append(f\"{config_key.split('__')[0]}\\n{v}\")\n            values_f.append(agg[v][\"frob_max\"])\nax.bar(range(len(values_f)), values_f, color='steelblue', edgecolor='black')\nax.set_yscale('log')\nax.set_xticks(range(len(values_f)))\nax.set_xticklabels(labels_f, rotation=45, ha='right', fontsize=7)\nax.set_ylabel(\"Max Frob deviation (log)\")\nax.set_title(\"Exactness verification across all configs and seeds\")\nax.axhline(1e-9, color='red', linestyle='--', label='10⁻⁹ target')\nax.legend(); ax.grid(True, alpha=0.3, axis='y')\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig4_exactness.png\", bbox_inches='tight')\nplt.show()\n\nprint(f\"Figures saved to {CONFIG['figures_dir']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:52:05.558945Z","iopub.execute_input":"2026-05-10T05:52:05.559916Z","iopub.status.idle":"2026-05-10T05:52:30.789301Z","shell.execute_reply.started":"2026-05-10T05:52:05.55988Z","shell.execute_reply":"2026-05-10T05:52:30.788323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nGenerating tables...\")\n\n# === Table 1: RQ1 with mean ± std ===\nrows = []\nfor config_key, agg in all_rq1.items():\n    for name in variant_order + [\"Centralized_CE\", \"FedAvg_CE\"]:\n        if name not in agg: continue\n        r = agg[name]\n        row = {\"Config\": config_key, \"Method\": name}\n        m = r.get(\"macro_AUROC_mean\", np.nan); s = r.get(\"macro_AUROC_std\", 0)\n        row[\"macro_AUROC\"] = f\"{m:.4f} ± {s:.4f}\"\n        for cls in classes:\n            mc = r.get(f\"AUROC_{cls}_mean\", np.nan); sc = r.get(f\"AUROC_{cls}_std\", 0)\n            row[cls] = f\"{mc:.4f} ± {sc:.4f}\"\n        if \"frob_max\" in r:\n            row[\"Frob_max\"] = f\"{r['frob_max']:.2e}\"\n        else:\n            row[\"Frob_max\"] = \"—\"\n        rows.append(row)\ntable1 = pd.DataFrame(rows)\ntable1.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table1_rq1.csv\", index=False)\nprint(\"\\n=== Table 1: RQ1 (mean ± std across seeds) ===\")\nprint(table1.to_string(index=False))\n\n# === Table 2: RQ2 summary ===\nrows = []\nfor scenario, key in [(\"Image-level withdrawal\", \"rq2a\"), (\"Site withdrawal\", \"rq2b\"), (\"Rare disease\", \"rq2c\")]:\n    if key not in rq2_results: continue\n    for name in RQ2_VARIANTS:\n        if name not in rq2_results[key]: continue\n        per_seed = rq2_results[key][name]\n        if scenario == \"Image-level withdrawal\":\n            macro_after = [tl[\"timeline\"][-1][\"metrics\"][\"macro_AUROC\"] for tl in per_seed]\n            time_total = [tl[\"timeline\"][-1][\"cum_time\"] for tl in per_seed]\n            frobs = [tl[\"frob\"] for tl in per_seed]\n            n_rem = per_seed[0][\"timeline\"][-1].get(\"n_deleted_units\", per_seed[0][\"timeline\"][-1][\"n_deleted_samples\"])\n        else:\n            macro_after = [d[\"after\"][\"macro_AUROC\"] for d in per_seed]\n            time_total = [d[\"wall_time_s\"] for d in per_seed]\n            frobs = [d[\"frob\"] for d in per_seed]\n            n_rem = per_seed[0][\"n_removed\"]\n        rows.append({\n            \"Method\": name, \"Scenario\": scenario,\n            \"macro_AUROC_after\": f\"{np.mean(macro_after):.4f} ± {np.std(macro_after):.4f}\",\n            \"n_removed\": n_rem,\n            \"time_s\": f\"{np.mean(time_total):.2f} ± {np.std(time_total):.2f}\",\n            \"Frob_max\": f\"{np.max(frobs):.2e}\",\n        })\ntable2 = pd.DataFrame(rows)\ntable2.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table2_rq2.csv\", index=False)\nprint(\"\\n=== Table 2: RQ2 ===\")\nprint(table2.to_string(index=False))\n\n# === Final consolidated dump ===\nfinal = {\n    \"config\": {k: (str(v) if isinstance(v, Path) else v) for k, v in CONFIG.items()},\n    \"rq1\": all_rq1,\n    \"rq2\": rq2_results,\n}\nwith open(Path(CONFIG[\"output_dir\"]) / \"all_results.json\", \"w\") as f:\n    json.dump(final, f, indent=2, default=float)\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"DONE.\")\nprint(f\"  All results: {CONFIG['output_dir']}/all_results.json\")\nprint(f\"  Tables:      {CONFIG['tables_dir']}/\")\nprint(f\"  Figures:     {CONFIG['figures_dir']}/\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:52:35.194489Z","iopub.execute_input":"2026-05-10T05:52:35.195452Z","iopub.status.idle":"2026-05-10T05:52:35.24783Z","shell.execute_reply.started":"2026-05-10T05:52:35.195413Z","shell.execute_reply":"2026-05-10T05:52:35.246908Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"Sensitivity: number of clients K\")\nprint(\"=\"*60)\n\nK_VALUES = [2, 5, 10, 20]\n\nprimary_ds = CONFIG[\"rq2_primary_dataset\"]\nprimary_bb = CONFIG[\"rq2_primary_backbone\"]\nmeta_p = prepare_metadata(primary_ds, CONFIG)\nfeatures_p = extract_features(meta_p, primary_bb, CONFIG)\n\nk_results = {}\nfor K in K_VALUES:\n    print(f\"\\n--- K = {K} ---\")\n    per_seed = {\"V0_plain\": [], \"V2_valgamma\": []}\n    per_seed_frob = {\"V0_plain\": [], \"V2_valgamma\": []}\n    for seed in CONFIG[\"seeds\"]:\n        set_seed(seed)\n        split = split_train_val_test(meta_p, CONFIG, seed)\n        tr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\n        va_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\n        te_F, te_Y = features_p[split[\"test_idx\"]], meta_p[\"labels\"][split[\"test_idx\"]]\n\n        client_indices = dirichlet_partition(tr_Y, K, CONFIG[\"non_iid_alpha\"], seed=seed)\n        setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                         CONFIG[\"selected_classes\"], CONFIG, seed)\n        for name in [\"V0_plain\", \"V2_valgamma\"]:\n            var_fed = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                tr_F, tr_Y, client_indices)\n            var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"], tr_F, tr_Y)\n            metrics = evaluate_variant(var_fed, te_F, te_Y, CONFIG[\"selected_classes\"])\n            fdev = frob_deviation(var_fed.get_W(), var_cen.get_W())\n            per_seed[name].append(metrics[\"macro_AUROC\"])\n            per_seed_frob[name].append(fdev)\n\n    k_results[K] = {}\n    for name in [\"V0_plain\", \"V2_valgamma\"]:\n        m = np.mean(per_seed[name]); s = np.std(per_seed[name])\n        fmax = np.max(per_seed_frob[name])\n        k_results[K][name] = {\n            \"macro_mean\": float(m), \"macro_std\": float(s), \"frob_max\": float(fmax),\n            \"per_seed\": per_seed[name],\n        }\n        print(f\"  {name}: macro={m:.4f}±{s:.4f}, frob_max={fmax:.2e}\")\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"k_sensitivity.json\", \"w\") as f:\n    json.dump(k_results, f, indent=2, default=float)\nprint(\"\\nSaved k_sensitivity.json\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:52:41.681185Z","iopub.execute_input":"2026-05-10T05:52:41.68216Z","iopub.status.idle":"2026-05-10T05:53:48.198324Z","shell.execute_reply.started":"2026-05-10T05:52:41.682117Z","shell.execute_reply":"2026-05-10T05:53:48.197462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"Sensitivity: non-IID degree α\")\nprint(\"=\"*60)\n\nALPHA_VALUES = [0.1, 0.5, 1.0, 5.0]\n\nalpha_results = {}\nfor alpha in ALPHA_VALUES:\n    print(f\"\\n--- α = {alpha} ---\")\n    per_seed = {\"V0_plain\": [], \"V2_valgamma\": []}\n    per_seed_frob = {\"V0_plain\": [], \"V2_valgamma\": []}\n    for seed in CONFIG[\"seeds\"]:\n        set_seed(seed)\n        split = split_train_val_test(meta_p, CONFIG, seed)\n        tr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\n        va_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\n        te_F, te_Y = features_p[split[\"test_idx\"]], meta_p[\"labels\"][split[\"test_idx\"]]\n\n        client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"], alpha, seed=seed)\n        setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                         CONFIG[\"selected_classes\"], CONFIG, seed)\n        for name in [\"V0_plain\", \"V2_valgamma\"]:\n            var_fed = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                                tr_F, tr_Y, client_indices)\n            var_cen = central_train(name, CONFIG, setup, CONFIG[\"selected_classes\"], tr_F, tr_Y)\n            metrics = evaluate_variant(var_fed, te_F, te_Y, CONFIG[\"selected_classes\"])\n            fdev = frob_deviation(var_fed.get_W(), var_cen.get_W())\n            per_seed[name].append(metrics[\"macro_AUROC\"])\n            per_seed_frob[name].append(fdev)\n\n    alpha_results[alpha] = {}\n    for name in [\"V0_plain\", \"V2_valgamma\"]:\n        m = np.mean(per_seed[name]); s = np.std(per_seed[name])\n        fmax = np.max(per_seed_frob[name])\n        alpha_results[alpha][name] = {\n            \"macro_mean\": float(m), \"macro_std\": float(s), \"frob_max\": float(fmax),\n            \"per_seed\": per_seed[name],\n        }\n        print(f\"  {name}: macro={m:.4f}±{s:.4f}, frob_max={fmax:.2e}\")\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"alpha_sensitivity.json\", \"w\") as f:\n    json.dump(alpha_results, f, indent=2, default=float)\nprint(\"\\nSaved alpha_sensitivity.json\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:53:48.199905Z","iopub.execute_input":"2026-05-10T05:53:48.201778Z","iopub.status.idle":"2026-05-10T05:54:53.600269Z","shell.execute_reply.started":"2026-05-10T05:53:48.201741Z","shell.execute_reply":"2026-05-10T05:54:53.599546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"Continual add-back invertibility test\")\nprint(\"=\"*60)\n\n# Helper to compare two W structures (handles V0 single tensor)\ndef W_norm(W):\n    if isinstance(W, list):\n        return float(np.sqrt(sum(torch.norm(w).item()**2 for w in W)))\n    return float(torch.norm(W).item())\n\ndef W_diff_norm(W_a, W_b):\n    if isinstance(W_a, list):\n        return float(np.sqrt(sum(torch.norm(a - b).item()**2 for a, b in zip(W_a, W_b))))\n    return float(torch.norm(W_a - W_b).item())\n\naddback_results = {}\nfor seed in CONFIG[\"seeds\"]:\n    print(f\"\\n--- seed {seed} ---\")\n    set_seed(seed)\n    split = split_train_val_test(meta_p, CONFIG, seed)\n    tr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\n    va_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\n    tr_pids = meta_p[\"df\"][\"patient_id\"].values[split[\"train_idx\"]]\n\n    client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                          CONFIG[\"non_iid_alpha\"], seed=seed)\n    setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                     CONFIG[\"selected_classes\"], CONFIG, seed)\n\n    for name in [\"V0_plain\", \"V2_valgamma\"]:\n        var = fed_train(name, CONFIG, setup, CONFIG[\"selected_classes\"],\n                        tr_F, tr_Y, client_indices)\n        # Snapshot initial W\n        W_init = var.get_W()\n        if not isinstance(W_init, list):\n            W_init = W_init.clone()\n        else:\n            W_init = [w.clone() for w in W_init]\n\n        # Pick 100 patients\n        rng = np.random.default_rng(seed)\n        unique_pids = np.unique(tr_pids); rng.shuffle(unique_pids)\n        pids = unique_pids[:100]\n\n        # Phase 1: delete all 100\n        for pid in pids:\n            mask = tr_pids == pid\n            if mask.any():\n                var.remove(tr_F[mask], tr_Y[mask])\n        var.fit()\n\n        # Phase 2: re-add in shuffled order\n        rng2 = np.random.default_rng(seed + 1000)\n        pids_shuffled = pids.copy(); rng2.shuffle(pids_shuffled)\n        for pid in pids_shuffled:\n            mask = tr_pids == pid\n            if mask.any():\n                var.add(tr_F[mask], tr_Y[mask])\n        var.fit()\n\n        W_final = var.get_W()\n        diff = W_diff_norm(W_final, W_init)\n        rel_diff = diff / max(W_norm(W_init), 1e-30)\n        addback_results.setdefault(name, []).append(rel_diff)\n        print(f\"  {name}: addback rel deviation = {rel_diff:.2e}\")\n\nprint(\"\\nSummary:\")\nfor name, devs in addback_results.items():\n    print(f\"  {name}: mean={np.mean(devs):.2e}, max={np.max(devs):.2e}\")\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"addback.json\", \"w\") as f:\n    json.dump({k: list(v) for k, v in addback_results.items()}, f, indent=2)\nprint(\"\\nSaved addback.json\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:54:53.601229Z","iopub.execute_input":"2026-05-10T05:54:53.601533Z","iopub.status.idle":"2026-05-10T05:55:12.978973Z","shell.execute_reply.started":"2026-05-10T05:54:53.6015Z","shell.execute_reply":"2026-05-10T05:55:12.978106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"Calibration analysis: raw vs Platt-scaled ECE\")\nprint(\"=\"*60)\n\nfrom sklearn.linear_model import LogisticRegression\n\ndef expected_calibration_error(probs, labels, n_bins=10):\n    ece_per_class = []\n    for j in range(probs.shape[1]):\n        if labels[:, j].sum() < 5: continue\n        p, l = probs[:, j], labels[:, j]\n        bins = np.linspace(0, 1, n_bins+1)\n        ece = 0.0\n        for i in range(n_bins):\n            mask = (p >= bins[i]) & (p < bins[i+1])\n            if mask.sum() > 0:\n                acc = l[mask].mean()\n                conf = p[mask].mean()\n                ece += (mask.sum() / len(p)) * abs(acc - conf)\n        ece_per_class.append(ece)\n    return float(np.mean(ece_per_class)) if ece_per_class else float(\"nan\")\n\n\ndef platt_scale_per_class(raw_val, labels_val, raw_test):\n    n_classes = raw_val.shape[1]\n    probs_test = np.zeros_like(raw_test, dtype=np.float64)\n    for j in range(n_classes):\n        if labels_val[:, j].sum() < 5:\n            probs_test[:, j] = 1 / (1 + np.exp(-raw_test[:, j]))\n            continue\n        try:\n            lr = LogisticRegression(max_iter=200)\n            lr.fit(raw_val[:, j:j+1], labels_val[:, j])\n            probs_test[:, j] = lr.predict_proba(raw_test[:, j:j+1])[:, 1]\n        except Exception:\n            probs_test[:, j] = 1 / (1 + np.exp(-raw_test[:, j]))\n    return probs_test\n\n\nprimary_ds = CONFIG[\"rq2_primary_dataset\"]\nprimary_bb = CONFIG[\"rq2_primary_backbone\"]\nmeta_p = prepare_metadata(primary_ds, CONFIG)\nfeatures_p = extract_features(meta_p, primary_bb, CONFIG)\n\nece_results = {}\nfor variant_name in [\"V0_plain\", \"V1_nystrom\", \"V2_valgamma\"]:\n    raw_eces, platt_eces = [], []\n    for seed in CONFIG[\"seeds\"]:\n        set_seed(seed)\n        split = split_train_val_test(meta_p, CONFIG, seed)\n        tr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\n        va_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\n        te_F, te_Y = features_p[split[\"test_idx\"]], meta_p[\"labels\"][split[\"test_idx\"]]\n\n        client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                              CONFIG[\"non_iid_alpha\"], seed=seed)\n        setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                         CONFIG[\"selected_classes\"], CONFIG, seed)\n        var = fed_train(variant_name, CONFIG, setup,\n                        CONFIG[\"selected_classes\"], tr_F, tr_Y, client_indices)\n\n        raw_val = var.predict(va_F)\n        raw_test = var.predict(te_F)\n        sigmoid_test = 1.0 / (1.0 + np.exp(-raw_test))\n        raw_eces.append(expected_calibration_error(sigmoid_test, te_Y))\n        platt_test = platt_scale_per_class(raw_val, va_Y, raw_test)\n        platt_eces.append(expected_calibration_error(platt_test, te_Y))\n\n    ece_results[variant_name] = {\"raw_ece\": raw_eces, \"platt_ece\": platt_eces}\n    print(f\"  {variant_name}: raw={np.mean(raw_eces):.4f}±{np.std(raw_eces):.4f}, \"\n          f\"Platt={np.mean(platt_eces):.4f}±{np.std(platt_eces):.4f}\")\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"ece_platt.json\", \"w\") as f:\n    json.dump(ece_results, f, indent=2, default=float)\nprint(\"Saved ece_platt.json\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:55:12.980713Z","iopub.execute_input":"2026-05-10T05:55:12.980967Z","iopub.status.idle":"2026-05-10T05:55:50.303459Z","shell.execute_reply.started":"2026-05-10T05:55:12.980943Z","shell.execute_reply":"2026-05-10T05:55:50.302616Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"Label smoothing ablation: V2_valgamma with eps={0.0, 0.1}\")\nprint(\"=\"*60)\n\nls_results = {0.0: [], 0.1: []}\n\nfor eps in [0.0, 0.1]:\n    print(f\"\\n--- eps = {eps} ---\")\n    config_eps = {**CONFIG, \"label_smoothing_eps\": eps}\n    for seed in CONFIG[\"seeds\"]:\n        set_seed(seed)\n        split = split_train_val_test(meta_p, CONFIG, seed)\n        tr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\n        va_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\n        te_F, te_Y = features_p[split[\"test_idx\"]], meta_p[\"labels\"][split[\"test_idx\"]]\n\n        client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                              CONFIG[\"non_iid_alpha\"], seed=seed)\n        setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                         CONFIG[\"selected_classes\"], config_eps, seed)\n        var = fed_train(\"V2_valgamma\", config_eps, setup, CONFIG[\"selected_classes\"],\n                        tr_F, tr_Y, client_indices)\n        metrics = evaluate_variant(var, te_F, te_Y, CONFIG[\"selected_classes\"])\n        # Also compute ECE\n        raw = var.predict(te_F)\n        probs = 1.0 / (1.0 + np.exp(-raw))\n        ece = expected_calibration_error(probs, te_Y)\n        ls_results[eps].append({\n            \"macro_AUROC\": metrics[\"macro_AUROC\"],\n            \"ECE\": ece,\n        })\n        print(f\"  seed {seed}: macro={metrics['macro_AUROC']:.4f}, ECE={ece:.4f}\")\n\nprint(\"\\nAggregated:\")\nfor eps in [0.0, 0.1]:\n    macros = [r[\"macro_AUROC\"] for r in ls_results[eps]]\n    eces = [r[\"ECE\"] for r in ls_results[eps]]\n    print(f\"  eps={eps}: macro={np.mean(macros):.4f}±{np.std(macros):.4f}, \"\n          f\"ECE={np.mean(eces):.4f}±{np.std(eces):.4f}\")\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"label_smoothing_ablation.json\", \"w\") as f:\n    json.dump({str(k): v for k, v in ls_results.items()}, f, indent=2, default=float)\nprint(\"\\nSaved label_smoothing_ablation.json\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:55:50.304514Z","iopub.execute_input":"2026-05-10T05:55:50.305281Z","iopub.status.idle":"2026-05-10T05:56:16.530454Z","shell.execute_reply.started":"2026-05-10T05:55:50.305245Z","shell.execute_reply":"2026-05-10T05:56:16.529663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"SISA baseline comparison on flagship config\")\nprint(\"=\"*60)\n\nclass SISAEnsemble:\n    \"\"\"SISA (Bourtoule et al. 2021): shard-and-ensemble exact unlearning.\"\"\"\n    def __init__(self, name, config, setup, classes, num_shards=5):\n        self.name = name; self.config = config; self.setup = setup\n        self.classes = classes; self.num_shards = num_shards\n        self.shard_models = []\n        self.shard_assignments = None\n        self.shard_data_indices = None\n\n    def fit(self, train_features, train_labels, seed=42):\n        rng = np.random.default_rng(seed)\n        n = len(train_features)\n        perm = rng.permutation(n)\n        self.shard_data_indices = np.array_split(perm, self.num_shards)\n        self.shard_assignments = np.zeros(n, dtype=int)\n        for s, idx in enumerate(self.shard_data_indices):\n            self.shard_assignments[idx] = s\n\n        self.shard_models = []\n        for s in range(self.num_shards):\n            var = make_variant(self.name, self.config, self.setup, self.classes)\n            sh_idx = self.shard_data_indices[s]\n            var.add(train_features[sh_idx], train_labels[sh_idx])\n            var.fit()\n            self.shard_models.append(var)\n        return self\n\n    def predict(self, features):\n        return sum(m.predict(features) for m in self.shard_models) / self.num_shards\n\n    def unlearn(self, deletion_indices, train_features, train_labels):\n        affected = set(int(self.shard_assignments[i]) for i in deletion_indices)\n        deleted_set = set(int(i) for i in deletion_indices)\n        for s in affected:\n            keep_idx = [int(i) for i in self.shard_data_indices[s] if int(i) not in deleted_set]\n            new_var = make_variant(self.name, self.config, self.setup, self.classes)\n            if len(keep_idx) > 0:\n                new_var.add(train_features[keep_idx], train_labels[keep_idx])\n                new_var.fit()\n            self.shard_models[s] = new_var\n            self.shard_data_indices[s] = np.array(keep_idx, dtype=int)\n        return list(affected)\n\n\nprimary_ds = CONFIG[\"rq2_primary_dataset\"]\nprimary_bb = CONFIG[\"rq2_primary_backbone\"]\nmeta_p = prepare_metadata(primary_ds, CONFIG)\nfeatures_p = extract_features(meta_p, primary_bb, CONFIG)\n\nsisa_v2_only = {\"utility\": [], \"delete_time\": []}\nv2_compare = {\"utility\": [], \"delete_time\": []}\n\nfor seed in CONFIG[\"seeds\"]:\n    print(f\"\\n--- seed {seed} ---\")\n    set_seed(seed)\n    split = split_train_val_test(meta_p, CONFIG, seed)\n    tr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\n    va_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\n    te_F, te_Y = features_p[split[\"test_idx\"]], meta_p[\"labels\"][split[\"test_idx\"]]\n\n    setup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                     CONFIG[\"selected_classes\"], CONFIG, seed)\n\n    # SISA-V2\n    sisa = SISAEnsemble(\"V2_valgamma\", CONFIG, setup,\n                        CONFIG[\"selected_classes\"], num_shards=CONFIG[\"sisa_num_shards\"])\n    sisa.fit(tr_F, tr_Y, seed=seed)\n    sisa_preds = sisa.predict(te_F)\n    sisa_metrics = evaluate_preds(sisa_preds, te_Y, CONFIG[\"selected_classes\"])\n    sisa_v2_only[\"utility\"].append(sisa_metrics[\"macro_AUROC\"])\n\n    rng = np.random.default_rng(seed + 100)\n    del_idx = rng.choice(len(tr_F), CONFIG[\"n_patient_deletions\"], replace=False).tolist()\n    t0 = time.time()\n    sisa.unlearn(del_idx, tr_F, tr_Y)\n    sisa_v2_only[\"delete_time\"].append(time.time() - t0)\n\n    # Single V2 for comparison\n    client_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                          CONFIG[\"non_iid_alpha\"], seed=seed)\n    var_v2 = fed_train(\"V2_valgamma\", CONFIG, setup,\n                       CONFIG[\"selected_classes\"], tr_F, tr_Y, client_indices)\n    v2_metrics = evaluate_variant(var_v2, te_F, te_Y, CONFIG[\"selected_classes\"])\n    v2_compare[\"utility\"].append(v2_metrics[\"macro_AUROC\"])\n\n    t0 = time.time()\n    for i in del_idx:\n        var_v2.remove(tr_F[i:i+1], tr_Y[i:i+1])\n    var_v2.fit()\n    v2_compare[\"delete_time\"].append(time.time() - t0)\n\n    print(f\"  SISA={sisa_v2_only['utility'][-1]:.4f}, V2={v2_compare['utility'][-1]:.4f}\")\n\nprint(\"\\n--- Summary ---\")\nprint(f\"SISA-V2 utility: {np.mean(sisa_v2_only['utility']):.4f} ± {np.std(sisa_v2_only['utility']):.4f}\")\nprint(f\"V2 (single)    : {np.mean(v2_compare['utility']):.4f} ± {np.std(v2_compare['utility']):.4f}\")\nprint(f\"SISA delete time: {np.mean(sisa_v2_only['delete_time']):.2f}s ± {np.std(sisa_v2_only['delete_time']):.2f}s\")\nprint(f\"V2 delete time  : {np.mean(v2_compare['delete_time']):.2f}s ± {np.std(v2_compare['delete_time']):.2f}s\")\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"sisa_comparison.json\", \"w\") as f:\n    json.dump({\"sisa_v2\": sisa_v2_only, \"v2_single\": v2_compare}, f, indent=2, default=float)\nprint(\"Saved sisa_comparison.json\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:58:20.524787Z","iopub.execute_input":"2026-05-10T05:58:20.525553Z","iopub.status.idle":"2026-05-10T05:58:49.849662Z","shell.execute_reply.started":"2026-05-10T05:58:20.525519Z","shell.execute_reply":"2026-05-10T05:58:49.848477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"Communication cost analysis\")\nprint(\"=\"*60)\n\nm_anchors = CONFIG[\"nystrom_anchors\"]\nn_classes = len(CONFIG[\"selected_classes\"])\nd_features = 768\n\nv0_round_bytes = (d_features * d_features + d_features * n_classes) * 8\nv2_round_bytes = (m_anchors * m_anchors + m_anchors * n_classes) * 8\nfedavg_round_bytes = (d_features * n_classes + n_classes) * 4\nfedavg_total = fedavg_round_bytes * 20 * CONFIG[\"num_clients\"]\nv2_per_delete = (m_anchors * m_anchors + m_anchors * n_classes) * 8\n\nprint(f\"\\n--- Per-round client-to-server (single client) ---\")\nprint(f\"V0_plain:    {v0_round_bytes/1e6:.2f} MB\")\nprint(f\"V2 (m=1024): {v2_round_bytes/1e6:.2f} MB\")\nprint(f\"FedAvg-CE:   {fedavg_round_bytes/1e3:.2f} KB per round, {20 * fedavg_round_bytes/1e6:.2f} MB total per client\")\n\nprint(f\"\\n--- Total full training (K={CONFIG['num_clients']} clients) ---\")\nprint(f\"V0:        {v0_round_bytes * CONFIG['num_clients']/1e6:.2f} MB\")\nprint(f\"V2:        {v2_round_bytes * CONFIG['num_clients']/1e6:.2f} MB\")\nprint(f\"FedAvg-CE: {fedavg_total/1e6:.2f} MB\")\n\nprint(f\"\\n--- Per single deletion ---\")\nprint(f\"V0 delete: {v0_round_bytes/1e6:.2f} MB\")\nprint(f\"V2 delete: {v2_per_delete/1e6:.2f} MB\")\nprint(f\"FedAvg retrain: {fedavg_total/1e6:.2f} MB\")\n\n# Empirical\nset_seed(42)\nsplit = split_train_val_test(meta_p, CONFIG, 42)\ntr_F, tr_Y = features_p[split[\"train_idx\"]], meta_p[\"labels\"][split[\"train_idx\"]]\nva_F, va_Y = features_p[split[\"val_idx\"]], meta_p[\"labels\"][split[\"val_idx\"]]\nclient_indices = dirichlet_partition(tr_Y, CONFIG[\"num_clients\"],\n                                      CONFIG[\"non_iid_alpha\"], seed=42)\nsetup = setup_nystrom_and_gamma(tr_F, tr_Y, va_F, va_Y,\n                                 CONFIG[\"selected_classes\"], CONFIG, 42)\nvar = fed_train(\"V2_valgamma\", CONFIG, setup,\n                CONFIG[\"selected_classes\"], tr_F, tr_Y, client_indices)\nS_bytes = var.stats.S.element_size() * var.stats.S.nelement()\nG_bytes = var.stats.G.element_size() * var.stats.G.nelement()\nprint(f\"\\n--- Empirical V2 statistics ---\")\nprint(f\"S={S_bytes/1e6:.2f} MB, G={G_bytes/1e6:.4f} MB, total={(S_bytes+G_bytes)/1e6:.2f} MB\")\n\ncomm_results = {\n    \"theoretical\": {\n        \"v0_per_round_mb\": v0_round_bytes/1e6,\n        \"v2_per_round_mb\": v2_round_bytes/1e6,\n        \"fedavg_per_round_kb\": fedavg_round_bytes/1e3,\n        \"fedavg_total_mb\": fedavg_total/1e6,\n        \"v2_per_delete_mb\": v2_per_delete/1e6,\n    },\n    \"empirical\": {\n        \"v2_S_mb\": S_bytes/1e6, \"v2_G_mb\": G_bytes/1e6,\n        \"v2_total_mb\": (S_bytes+G_bytes)/1e6,\n    }\n}\nwith open(Path(CONFIG[\"output_dir\"]) / \"comm_cost.json\", \"w\") as f:\n    json.dump(comm_results, f, indent=2)\nprint(\"Saved comm_cost.json\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:58:49.85154Z","iopub.execute_input":"2026-05-10T05:58:49.851983Z","iopub.status.idle":"2026-05-10T05:58:52.377446Z","shell.execute_reply.started":"2026-05-10T05:58:49.851953Z","shell.execute_reply":"2026-05-10T05:58:52.376637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nGenerating supplementary figures...\")\n\n# === Figure: K sensitivity ===\nfig, ax = plt.subplots(figsize=(8, 5))\nks = sorted(k_results.keys())\nv0_means = [k_results[k][\"V0_plain\"][\"macro_mean\"] for k in ks]\nv0_stds = [k_results[k][\"V0_plain\"][\"macro_std\"] for k in ks]\nv2_means = [k_results[k][\"V2_valgamma\"][\"macro_mean\"] for k in ks]\nv2_stds = [k_results[k][\"V2_valgamma\"][\"macro_std\"] for k in ks]\nax.errorbar(ks, v0_means, yerr=v0_stds, marker='o', label='V0_plain', color='#888', capsize=4)\nax.errorbar(ks, v2_means, yerr=v2_stds, marker='s', label='V2_valgamma', color='#7cb342', capsize=4)\nax.set_xlabel(\"Number of clients K\")\nax.set_ylabel(\"macro AUROC\")\nax.set_title(\"Sensitivity to K (primary config)\")\nax.set_xscale('log'); ax.set_xticks(ks); ax.set_xticklabels(ks)\nax.legend(); ax.grid(True, alpha=0.3)\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig5_k_sensitivity.png\", bbox_inches='tight')\nplt.show()\n\n# === Figure: α sensitivity ===\nfig, ax = plt.subplots(figsize=(8, 5))\nalphas = sorted(alpha_results.keys())\nv0_means = [alpha_results[a][\"V0_plain\"][\"macro_mean\"] for a in alphas]\nv0_stds = [alpha_results[a][\"V0_plain\"][\"macro_std\"] for a in alphas]\nv2_means = [alpha_results[a][\"V2_valgamma\"][\"macro_mean\"] for a in alphas]\nv2_stds = [alpha_results[a][\"V2_valgamma\"][\"macro_std\"] for a in alphas]\nax.errorbar(alphas, v0_means, yerr=v0_stds, marker='o', label='V0_plain', color='#888', capsize=4)\nax.errorbar(alphas, v2_means, yerr=v2_stds, marker='s', label='V2_valgamma', color='#7cb342', capsize=4)\nax.set_xlabel(\"Dirichlet α (lower = more non-IID)\")\nax.set_ylabel(\"macro AUROC\")\nax.set_title(\"Sensitivity to non-IID degree α (primary config)\")\nax.set_xscale('log')\nax.legend(); ax.grid(True, alpha=0.3)\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig6_alpha_sensitivity.png\", bbox_inches='tight')\nplt.show()\n\n# === Figure: Add-back invertibility ===\nfig, ax = plt.subplots(figsize=(7, 4))\nnames = list(addback_results.keys())\nmaxes = [np.max(addback_results[n]) for n in names]\nmeans = [np.mean(addback_results[n]) for n in names]\nx = np.arange(len(names))\nax.bar(x - 0.2, maxes, 0.4, label='max', color='#e89c5a', edgecolor='black')\nax.bar(x + 0.2, means, 0.4, label='mean', color='#4c9aaf', edgecolor='black')\nax.set_yscale('log')\nax.set_xticks(x); ax.set_xticklabels(names)\nax.set_ylabel(\"Relative Frob deviation (log)\")\nax.set_title(\"Add-back invertibility: ‖W_addback − W_init‖ / ‖W_init‖\")\nax.axhline(1e-9, color='red', linestyle='--', alpha=0.6, label='10⁻⁹ target')\nax.legend(); ax.grid(True, alpha=0.3, axis='y')\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig7_addback.png\", bbox_inches='tight')\nplt.show()\n\n# === Figure: ECE comparison (raw vs Platt-scaled) ===\nfig, ax = plt.subplots(figsize=(9, 4))\nece_names = list(ece_results.keys())\nraw_means = [np.mean(ece_results[n][\"raw_ece\"]) for n in ece_names]\nraw_stds = [np.std(ece_results[n][\"raw_ece\"]) for n in ece_names]\nplatt_means = [np.mean(ece_results[n][\"platt_ece\"]) for n in ece_names]\nplatt_stds = [np.std(ece_results[n][\"platt_ece\"]) for n in ece_names]\n\nx = np.arange(len(ece_names))\nwidth = 0.35\nax.bar(x - width/2, raw_means, width, yerr=raw_stds, label='Raw (sigmoid)',\n       color='#aaa', edgecolor='black', capsize=3)\nax.bar(x + width/2, platt_means, width, yerr=platt_stds, label='Platt-scaled',\n       color='#7cb342', edgecolor='black', capsize=3)\nax.set_xticks(x); ax.set_xticklabels(ece_names)\nax.set_ylabel(\"Expected Calibration Error (lower = better)\")\nax.set_title(\"Calibration: raw vs Platt-scaled ECE\")\nax.legend(); ax.grid(True, alpha=0.3, axis='y')\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig8_ece.png\", bbox_inches='tight')\nplt.show()\n\n# === Figure: Label smoothing ablation ===\nfig, axes = plt.subplots(1, 2, figsize=(11, 4))\neps_vals = [0.0, 0.1]\nmacros_e = [[r[\"macro_AUROC\"] for r in ls_results[e]] for e in eps_vals]\neces_e = [[r[\"ECE\"] for r in ls_results[e]] for e in eps_vals]\n\naxes[0].bar([f\"ε={e}\" for e in eps_vals],\n            [np.mean(m) for m in macros_e],\n            yerr=[np.std(m) for m in macros_e],\n            capsize=5, color=['#888', '#7cb342'], edgecolor='black')\naxes[0].set_ylabel(\"macro AUROC\"); axes[0].set_title(\"V2_valgamma utility\")\naxes[0].grid(True, alpha=0.3, axis='y')\n\naxes[1].bar([f\"ε={e}\" for e in eps_vals],\n            [np.mean(e) for e in eces_e],\n            yerr=[np.std(e) for e in eces_e],\n            capsize=5, color=['#888', '#7cb342'], edgecolor='black')\naxes[1].set_ylabel(\"ECE (lower = better)\"); axes[1].set_title(\"V2_valgamma calibration\")\naxes[1].grid(True, alpha=0.3, axis='y')\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig9_label_smoothing.png\", bbox_inches='tight')\nplt.show()\n\n# === Figure: SISA vs V2 single ===\nfig, axes = plt.subplots(1, 2, figsize=(11, 4))\nsisa_data = json.load(open(Path(CONFIG[\"output_dir\"]) / \"sisa_comparison.json\"))\nmethods = [\"SISA-V2\", \"V2 (single)\"]\nutils = [sisa_data[\"sisa_v2\"][\"utility\"], sisa_data[\"v2_single\"][\"utility\"]]\ntimes = [sisa_data[\"sisa_v2\"][\"delete_time\"], sisa_data[\"v2_single\"][\"delete_time\"]]\n\naxes[0].bar(methods, [np.mean(u) for u in utils], yerr=[np.std(u) for u in utils],\n            capsize=5, color=['#e89c5a', '#7cb342'], edgecolor='black')\naxes[0].set_ylabel(\"macro AUROC\"); axes[0].set_title(\"Utility\")\naxes[0].grid(True, alpha=0.3, axis='y')\n\naxes[1].bar(methods, [np.mean(t) for t in times], yerr=[np.std(t) for t in times],\n            capsize=5, color=['#e89c5a', '#7cb342'], edgecolor='black')\naxes[1].set_ylabel(f\"Wall-clock for {CONFIG['n_patient_deletions']} deletions (s)\")\naxes[1].set_title(\"Deletion cost\")\naxes[1].grid(True, alpha=0.3, axis='y')\nplt.tight_layout()\nplt.savefig(Path(CONFIG[\"figures_dir\"]) / \"fig10_sisa_comparison.png\", bbox_inches='tight')\nplt.show()\n\nprint(\"\\nSupplementary figures saved\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:58:59.327987Z","iopub.execute_input":"2026-05-10T05:58:59.328304Z","iopub.status.idle":"2026-05-10T05:59:02.462375Z","shell.execute_reply.started":"2026-05-10T05:58:59.328278Z","shell.execute_reply":"2026-05-10T05:59:02.461567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nGenerating supplementary tables...\")\n\n# === Table: K sensitivity ===\nrows = []\nfor K in sorted(k_results.keys()):\n    for name in [\"V0_plain\", \"V2_valgamma\"]:\n        r = k_results[K][name]\n        rows.append({\n            \"K\": K, \"Method\": name,\n            \"macro_AUROC\": f\"{r['macro_mean']:.4f} ± {r['macro_std']:.4f}\",\n            \"Frob_max\": f\"{r['frob_max']:.2e}\",\n        })\ntable_k = pd.DataFrame(rows)\ntable_k.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table3_k_sensitivity.csv\", index=False)\nprint(\"\\n=== Table 3: K sensitivity ===\")\nprint(table_k.to_string(index=False))\n\n# === Table: α sensitivity ===\nrows = []\nfor alpha in sorted(alpha_results.keys()):\n    for name in [\"V0_plain\", \"V2_valgamma\"]:\n        r = alpha_results[alpha][name]\n        rows.append({\n            \"alpha\": alpha, \"Method\": name,\n            \"macro_AUROC\": f\"{r['macro_mean']:.4f} ± {r['macro_std']:.4f}\",\n            \"Frob_max\": f\"{r['frob_max']:.2e}\",\n        })\ntable_a = pd.DataFrame(rows)\ntable_a.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table4_alpha_sensitivity.csv\", index=False)\nprint(\"\\n=== Table 4: α sensitivity ===\")\nprint(table_a.to_string(index=False))\n\n# === Table: Add-back ===\nrows = []\nfor name, devs in addback_results.items():\n    rows.append({\n        \"Method\": name,\n        \"rel_dev_mean\": f\"{np.mean(devs):.2e}\",\n        \"rel_dev_max\": f\"{np.max(devs):.2e}\",\n        \"rel_dev_min\": f\"{np.min(devs):.2e}\",\n    })\ntable_add = pd.DataFrame(rows)\ntable_add.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table5_addback.csv\", index=False)\nprint(\"\\n=== Table 5: Add-back invertibility ===\")\nprint(table_add.to_string(index=False))\n\n# === Table: ECE (raw vs Platt-scaled) ===\nrows = []\nfor name, eces in ece_results.items():\n    raw_arr = eces[\"raw_ece\"]\n    platt_arr = eces[\"platt_ece\"]\n    rows.append({\n        \"Method\": name,\n        \"Raw_ECE\": f\"{np.mean(raw_arr):.4f} ± {np.std(raw_arr):.4f}\",\n        \"Platt_ECE\": f\"{np.mean(platt_arr):.4f} ± {np.std(platt_arr):.4f}\",\n    })\ntable_ece = pd.DataFrame(rows)\ntable_ece.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table6_ece.csv\", index=False)\nprint(\"\\n=== Table 6: ECE (raw vs Platt-scaled) ===\")\nprint(table_ece.to_string(index=False))\n\n# === Table: Label smoothing ===\nrows = []\nfor eps in [0.0, 0.1]:\n    macros = [r[\"macro_AUROC\"] for r in ls_results[eps]]\n    eces = [r[\"ECE\"] for r in ls_results[eps]]\n    rows.append({\n        \"eps\": eps,\n        \"macro_AUROC\": f\"{np.mean(macros):.4f} ± {np.std(macros):.4f}\",\n        \"ECE\": f\"{np.mean(eces):.4f} ± {np.std(eces):.4f}\",\n    })\ntable_ls = pd.DataFrame(rows)\ntable_ls.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table7_label_smoothing.csv\", index=False)\nprint(\"\\n=== Table 7: Label smoothing ablation ===\")\nprint(table_ls.to_string(index=False))\n\n# === Table: SISA comparison ===\nsisa_data = json.load(open(Path(CONFIG[\"output_dir\"]) / \"sisa_comparison.json\"))\nrows = []\nfor label, data in [(\"SISA-V2 (5 shards)\", sisa_data[\"sisa_v2\"]),\n                    (\"V2-valgamma (single)\", sisa_data[\"v2_single\"])]:\n    rows.append({\n        \"Method\": label,\n        \"macro_AUROC\": f\"{np.mean(data['utility']):.4f} ± {np.std(data['utility']):.4f}\",\n        f\"Wall-clock {CONFIG['n_patient_deletions']} deletions (s)\":\n            f\"{np.mean(data['delete_time']):.2f} ± {np.std(data['delete_time']):.2f}\",\n    })\ntable_sisa = pd.DataFrame(rows)\ntable_sisa.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table8_sisa.csv\", index=False)\nprint(\"\\n=== Table 8: SISA comparison ===\")\nprint(table_sisa.to_string(index=False))\n\n# === Table: Communication cost ===\ncomm_data = json.load(open(Path(CONFIG[\"output_dir\"]) / \"comm_cost.json\"))\nrows = [\n    {\"Metric\": \"V0_plain per-round (MB)\",\n     \"Value\": f\"{comm_data['theoretical']['v0_per_round_mb']:.2f}\"},\n    {\"Metric\": \"V2 per-round (MB)\",\n     \"Value\": f\"{comm_data['theoretical']['v2_per_round_mb']:.2f}\"},\n    {\"Metric\": \"FedAvg-CE per-round (KB)\",\n     \"Value\": f\"{comm_data['theoretical']['fedavg_per_round_kb']:.2f}\"},\n    {\"Metric\": \"FedAvg-CE total full training (MB)\",\n     \"Value\": f\"{comm_data['theoretical']['fedavg_total_mb']:.2f}\"},\n    {\"Metric\": \"V2 per single deletion (MB)\",\n     \"Value\": f\"{comm_data['theoretical']['v2_per_delete_mb']:.2f}\"},\n]\ntable_comm = pd.DataFrame(rows)\ntable_comm.to_csv(Path(CONFIG[\"tables_dir\"]) / \"table9_comm_cost.csv\", index=False)\nprint(\"\\n=== Table 9: Communication cost ===\")\nprint(table_comm.to_string(index=False))\n\nprint(\"\\nAll supplementary tables saved\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:59:07.867001Z","iopub.execute_input":"2026-05-10T05:59:07.867691Z","iopub.status.idle":"2026-05-10T05:59:07.909022Z","shell.execute_reply.started":"2026-05-10T05:59:07.867663Z","shell.execute_reply":"2026-05-10T05:59:07.908184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\nConsolidating all results...\")\n\nimport json\nimport shutil\n\n# Update all_results.json with all supplementary experiments including new ones\nfinal = {\n    \"config\": {k: (str(v) if isinstance(v, Path) else v) for k, v in CONFIG.items()},\n    \"rq1\": all_rq1,\n    \"rq2\": rq2_results,\n    \"rq2_patient_nih_dinov2\": rq2_patient_results if \"rq2_patient_results\" in globals() else None,\n    \"supplementary\": {\n        \"k_sensitivity\": k_results,\n        \"alpha_sensitivity\": alpha_results,\n        \"addback\": {k: list(v) for k, v in addback_results.items()},\n        \"ece_platt\": ece_results,\n        \"label_smoothing_ablation\": {str(k): v for k, v in ls_results.items()},\n        \"sisa_comparison\": json.load(open(Path(CONFIG[\"output_dir\"]) / \"sisa_comparison.json\")),\n        \"comm_cost\": json.load(open(Path(CONFIG[\"output_dir\"]) / \"comm_cost.json\")),\n        \"rq2_patient_nih_dinov2\": rq2_patient_results if \"rq2_patient_results\" in globals() else None,\n    },\n}\n\nwith open(Path(CONFIG[\"output_dir\"]) / \"all_results.json\", \"w\") as f:\n    json.dump(final, f, indent=2, default=float)\nprint(\"Saved consolidated all_results.json\")\n\nzip_path = \"/kaggle/working/output_results_full.zip\"\nshutil.make_archive(\n    base_name=\"/kaggle/working/output_results_full\",\n    format=\"zip\",\n    root_dir=\"/kaggle/working/output\",\n)\nprint(f\"\\nZipped: {zip_path}\")\nimport os\nprint(f\"Size: {os.path.getsize(zip_path)/1e6:.1f} MB\")\nprint(\"\\nAll done — download output_results_full.zip from /kaggle/working/\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T05:59:12.430715Z","iopub.execute_input":"2026-05-10T05:59:12.43139Z","iopub.status.idle":"2026-05-10T05:59:12.496472Z","shell.execute_reply.started":"2026-05-10T05:59:12.431356Z","shell.execute_reply":"2026-05-10T05:59:12.495737Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"zip_path = \"/kaggle/working/cache.zip\"\nshutil.make_archive(\n    base_name=\"/kaggle/working/cache\",\n    format=\"zip\",\n    root_dir=\"/kaggle/working/output\",\n)\nprint(f\"\\nZipped: {zip_path}\")\nimport os\nprint(f\"Size: {os.path.getsize(zip_path)/1e6:.1f} MB\")\nprint(\"\\nAll done — download output_results_full.zip from /kaggle/working/\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-10T06:06:54.002434Z","iopub.execute_input":"2026-05-10T06:06:54.003061Z","iopub.status.idle":"2026-05-10T06:06:54.05121Z","shell.execute_reply.started":"2026-05-10T06:06:54.003026Z","shell.execute_reply":"2026-05-10T06:06:54.05046Z"}},"outputs":[],"execution_count":null}]}