{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":11848,"databundleVersionId":862157},{"sourceType":"competition","sourceId":45867,"databundleVersionId":6924515},{"sourceType":"kernelVersion","sourceId":41716214}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"edb3aa5b-cf8a-4c0a-9232-6757b8e38dc3","cell_type":"markdown","source":"# Multi-View Meta-Fusion Autoencoder for Histopathology\n## Complete Working Notebook — UBC-OCEAN Dataset\n\n**Dataset required:** Add `UBC Ovarian Cancer Subtype Classification` competition dataset\n- Path will be: `/kaggle/input/competitions/UBC-OCEAN/`\n- Enable GPU: Settings → Accelerator → GPU T4 x2\n\n| Cell | Content |\n|------|--------|\n| 1 | Install dependencies |\n| 2 | Imports + device setup |\n| 3 | Global config + label maps |\n| 4 | Dataset explorer + file check |\n| 5 | Multi-view generator (RGB/Saliency/H&E) |\n| 6 | Dataset class (real + synthetic fallback) |\n| 7 | Build DataLoaders |\n| 8 | Model architecture (Encoder/Fusion/Decoder) |\n| 9 | Loss functions |\n| 10 | Trainer (AE pretrain → joint finetune) |\n| 11 | **Train the model** |\n| 12 | Extract latents + compute all metrics |\n| 13 | Training curves plot |\n| 14 | t-SNE visualisation |\n| 15 | Confusion matrix |\n| 16 | Fusion gate weights |\n| 17 | Reconstruction samples |\n| 18 | Fusion strategy comparison |\n| 19 | Ablation study |\n| 20 | Per-class metrics + summary table |\n| 21 | Save checkpoint + final summary |\n","metadata":{}},{"id":"4bed8c0a-b732-4772-9a62-2e98c6f694ba","cell_type":"markdown","source":"## Cell 1: Install Dependencies\n","metadata":{}},{"id":"8bfb60a4-6269-461a-88e9-ae7b9a6944df","cell_type":"code","source":"!pip install -q opencv-python-headless scikit-image shap\nimport subprocess, sys\npkgs = [\"shap\", \"scikit-image\"]\nfor p in pkgs:\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", p])\nprint(\"Dependencies ready ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:09.149275Z","iopub.execute_input":"2026-03-07T07:05:09.149501Z","iopub.status.idle":"2026-03-07T07:05:20.068975Z","shell.execute_reply.started":"2026-03-07T07:05:09.149478Z","shell.execute_reply":"2026-03-07T07:05:20.068274Z"}},"outputs":[],"execution_count":null},{"id":"f96179ad-8446-493d-8f67-87ba8c388c94","cell_type":"markdown","source":"## Cell 2: Imports + Device Setup\n","metadata":{}},{"id":"a7ae5976-3812-4023-aea4-96c02da9e1e6","cell_type":"code","source":"import os, glob, json, random, warnings, time\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib\nmatplotlib.rcParams['figure.dpi'] = 100\nimport seaborn as sns\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\nfrom PIL import Image\nfrom collections import defaultdict\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport torchvision.transforms as T\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport cv2\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.manifold import TSNE\nfrom sklearn.cluster import KMeans\nfrom sklearn.metrics import (\n    accuracy_score, precision_score, recall_score, f1_score,\n    roc_auc_score, confusion_matrix,\n    mean_squared_error, mean_absolute_error, silhouette_score\n)\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.discriminant_analysis import LinearDiscriminantAnalysis\n\n# ── Reproducibility ──────────────────────────────────────────\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\nif torch.cuda.is_available(): torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device  : {DEVICE}\")\nif torch.cuda.is_available():\n    print(f\"GPU     : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM    : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\nelse:\n    print(\"No GPU found — running on CPU (will be slow)\")\nprint(\"Imports complete ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:20.070766Z","iopub.execute_input":"2026-03-07T07:05:20.071303Z","iopub.status.idle":"2026-03-07T07:05:29.202263Z","shell.execute_reply.started":"2026-03-07T07:05:20.071273Z","shell.execute_reply":"2026-03-07T07:05:29.201507Z"}},"outputs":[],"execution_count":null},{"id":"dd88f2cd-3d79-41f0-b48e-6e23f8706199","cell_type":"markdown","source":"## Cell 3: Global Config\n","metadata":{}},{"id":"8d4e7f05-1fa6-4e39-9858-5b6e8024881e","cell_type":"code","source":"CFG = {\n    # Paths  ← update DATASET_ROOT if your path differs\n    \"dataset_root\":    \"/kaggle/input/competitions/UBC-OCEAN\",\n    \"output_dir\":      \"/kaggle/working/outputs\",\n\n    # Data\n    \"patch_size\":      224,     # resize tiles/thumbnails to this\n\n    # Model\n    \"latent_dim\":      256,\n    \"n_views\":         4,\n    \"num_classes\":     5,\n\n    # Training\n    \"batch_size\":      16,      # lower to 8 if OOM\n    \"lr\":              1e-3,\n    \"weight_decay\":    1e-4,\n    \"ae_epochs\":       15,\n    \"joint_epochs\":    25,\n    \"lambda_rec\":      1.0,\n    \"lambda_cons\":     0.5,\n    \"lambda_kl\":       0.1,\n}\n\nos.makedirs(CFG[\"output_dir\"], exist_ok=True)\n\n# ── Label maps ────────────────────────────────────────────────\nLABEL_MAP   = {\"HGSC\": 0, \"EC\": 1, \"CC\": 2, \"MC\": 3, \"LGSC\": 4}\nCLASS_NAMES = [\"HGSC\", \"EC\", \"CC\", \"MC\", \"LGSC\"]\nCOLORS      = [\"#E63946\",\"#457B9D\",\"#2A9D8F\",\"#E9C46A\",\"#F4A261\"]\nVIEW_NAMES  = [\"RGB Morphology\",\"Saliency Map\",\"Haematoxylin\",\"Eosin\"]\n\nprint(\"Config loaded ✓\")\nprint(f\"Dataset root : {CFG['dataset_root']}\")\nprint(f\"Exists       : {os.path.exists(CFG['dataset_root'])}\")\nif os.path.exists(CFG['dataset_root']):\n    print(f\"Contents     : {os.listdir(CFG['dataset_root'])[:8]}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:29.203322Z","iopub.execute_input":"2026-03-07T07:05:29.203803Z","iopub.status.idle":"2026-03-07T07:05:29.211222Z","shell.execute_reply.started":"2026-03-07T07:05:29.20378Z","shell.execute_reply":"2026-03-07T07:05:29.21042Z"}},"outputs":[],"execution_count":null},{"id":"563c08d9-d6f5-4259-84a7-50825e98c0a3","cell_type":"markdown","source":"## Cell 4: Dataset Explorer + File Verification\n","metadata":{}},{"id":"6e471312-028d-475c-8e94-3392e3bd5f87","cell_type":"code","source":"root       = CFG[\"dataset_root\"]\ntrain_csv  = os.path.join(root, \"train.csv\")\nthumb_dir  = os.path.join(root, \"train_thumbnails\")\nwsi_dir    = os.path.join(root, \"train_images\")\n\nprint(\"=\" * 55)\nprint(\"UBC-OCEAN FILE CHECK\")\nprint(\"=\" * 55)\nfor name, path in [(\"train.csv\", train_csv),\n                   (\"train_thumbnails/\", thumb_dir),\n                   (\"train_images/\", wsi_dir)]:\n    exists = os.path.exists(path)\n    count  = \"\"\n    if exists and os.path.isdir(path):\n        count = f\"  ({len(os.listdir(path))} files)\"\n    print(f\"  {'✓' if exists else '✗'} {name}{count}\")\n\nif not os.path.exists(train_csv):\n    raise FileNotFoundError(\n        f\"train.csv not found at {train_csv}\\n\"\n        \"Make sure the UBC-OCEAN competition dataset is added as input.\"\n    )\n\ndf = pd.read_csv(train_csv)\ndf = df[df[\"label\"].isin(LABEL_MAP)].reset_index(drop=True)\n\nprint(f\"\\ntrain.csv shape  : {df.shape}\")\nprint(f\"Columns          : {df.columns.tolist()}\")\nprint(f\"\\nClass distribution:\")\nprint(df[\"label\"].value_counts().to_string())\n\nif \"is_tma\" in df.columns:\n    print(f\"\\nis_tma distribution:\")\n    print(df[\"is_tma\"].value_counts().to_string())\n\n# ── Verify thumbnail filenames ────────────────────────────────\nif os.path.exists(thumb_dir):\n    sample_files = os.listdir(thumb_dir)[:5]\n    print(f\"\\nSample thumbnail names: {sample_files}\")\n    # Build id→path cache\n    thumb_cache = {}\n    for fname in os.listdir(thumb_dir):\n        stem = os.path.splitext(fname)[0].replace(\"_thumbnail\",\"\")\n        thumb_cache[str(stem)] = os.path.join(thumb_dir, fname)\n    covered = set(df[\"image_id\"].astype(str)) & set(thumb_cache.keys())\n    print(f\"Thumbnails available  : {len(covered)} / {len(df)}\")\n\n# ── Plot class distribution ───────────────────────────────────\nfig, axes = plt.subplots(1, 2, figsize=(13, 5))\nfig.suptitle(\"UBC-OCEAN — Dataset Overview\", fontsize=14, fontweight=\"bold\")\n\nlc = df[\"label\"].value_counts().reindex(CLASS_NAMES, fill_value=0)\naxes[0].bar(lc.index, lc.values, color=COLORS, edgecolor=\"black\")\naxes[0].set_title(\"Class Distribution\"); axes[0].set_ylabel(\"Count\")\nfor i, v in enumerate(lc.values):\n    axes[0].text(i, v+1, str(v), ha=\"center\", fontweight=\"bold\")\naxes[0].grid(True, alpha=0.3, axis=\"y\")\n\naxes[1].pie(lc.values, labels=lc.index, autopct=\"%1.1f%%\",\n            colors=COLORS, startangle=90)\naxes[1].set_title(\"Class Proportions\")\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/dataset_overview.png\",\n            dpi=120, bbox_inches=\"tight\")\nplt.show()\nprint(\"\\n✓ Dataset verified\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:29.212226Z","iopub.execute_input":"2026-03-07T07:05:29.212506Z","iopub.status.idle":"2026-03-07T07:05:29.736027Z","shell.execute_reply.started":"2026-03-07T07:05:29.212476Z","shell.execute_reply":"2026-03-07T07:05:29.735368Z"}},"outputs":[],"execution_count":null},{"id":"5a34f86d-1dbb-4210-ab25-41e3157277c5","cell_type":"markdown","source":"## Cell 5: Multi-View Generator\n","metadata":{}},{"id":"52c2d0c9-659d-4074-ae5d-743e02b749d8","cell_type":"code","source":"def generate_views(img: np.ndarray, size: int = 224) -> dict:\n    '''\n    Input : (H,W,3) uint8 RGB numpy array\n    Output: dict  v1..v4  each (3,H,W) float32 tensor [0,1]\n      v1 — RGB morphology\n      v2 — Saliency (Sobel gradient)\n      v3 — Haematoxylin channel\n      v4 — Eosin channel\n    '''\n    img = cv2.resize(img, (size, size), interpolation=cv2.INTER_AREA)\n\n    # V1: RGB\n    v1 = img.astype(np.float32) / 255.0\n\n    # V2: Saliency\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY).astype(np.float32)\n    gx   = cv2.Sobel(gray, cv2.CV_32F, 1, 0, ksize=3)\n    gy   = cv2.Sobel(gray, cv2.CV_32F, 0, 1, ksize=3)\n    mag  = np.sqrt(gx**2 + gy**2)\n    mag  = cv2.GaussianBlur(mag, (5,5), 0)\n    mag -= mag.min(); mag /= (mag.max() + 1e-6)\n    v2   = np.stack([mag]*3, axis=-1)\n\n    # V3 & V4: H&E colour deconvolution\n    HE    = np.array([[0.6442,0.7166,0.2668],\n                      [0.0928,0.9541,0.2834],\n                      [0.6340,0.3960,0.6640]])\n    rgb_f  = img.astype(np.float32)/255.0 + 1e-6\n    OD     = -np.log(rgb_f)\n    stains = (OD.reshape(-1,3) @ np.linalg.pinv(HE).T).reshape(size,size,3)\n    stains = np.clip(stains, 0, None)\n\n    def nc(x):\n        mn, mx = x.min(), x.max()\n        return ((x-mn)/(mx-mn+1e-6)).astype(np.float32)\n\n    v3 = np.stack([nc(stains[:,:,0])]*3, axis=-1)\n    v4 = np.stack([nc(stains[:,:,1])]*3, axis=-1)\n\n    def to_t(a): return torch.from_numpy(a).permute(2,0,1).float()\n    return {\"v1\":to_t(v1), \"v2\":to_t(v2), \"v3\":to_t(v3), \"v4\":to_t(v4)}\n\nprint(\"Multi-view generator ready ✓\")\n# Quick test\n_test = np.random.randint(0,255,(256,256,3),dtype=np.uint8)\n_v    = generate_views(_test, CFG[\"patch_size\"])\nfor k,t in _v.items():\n    print(f\"  {k}: {tuple(t.shape)}  min={t.min():.3f}  max={t.max():.3f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:29.736812Z","iopub.execute_input":"2026-03-07T07:05:29.737086Z","iopub.status.idle":"2026-03-07T07:05:29.816962Z","shell.execute_reply.started":"2026-03-07T07:05:29.737065Z","shell.execute_reply":"2026-03-07T07:05:29.816343Z"}},"outputs":[],"execution_count":null},{"id":"1bea3726-403c-41c3-9809-679268d58f9d","cell_type":"markdown","source":"## Cell 6: Dataset Classes (Real + Synthetic Fallback)\n","metadata":{}},{"id":"e7167a14-beb5-4777-9265-e6ea95650cf9","cell_type":"code","source":"class UBCOceanDataset(Dataset):\n    TRAIN_AUG = T.Compose([\n        T.RandomHorizontalFlip(0.5),\n        T.RandomVerticalFlip(0.5),\n        T.RandomRotation(90),\n        T.ColorJitter(brightness=0.2, contrast=0.2,\n                      saturation=0.2, hue=0.05),\n    ])\n\n    def __init__(self, df, thumb_dir, patch_size=224, split=\"train\"):\n        self.patch_size = patch_size\n        self.augment    = (split == \"train\")\n        self.samples    = []   # list of (img_path, label_id)\n\n        # Build filename → path lookup once\n        cache = {}\n        if os.path.exists(thumb_dir):\n            for fname in os.listdir(thumb_dir):\n                stem = os.path.splitext(fname)[0].replace(\"_thumbnail\",\"\")\n                cache[str(stem)] = os.path.join(thumb_dir, fname)\n\n        missing = 0\n        for _, row in df.iterrows():\n            img_id   = str(row[\"image_id\"])\n            label_id = LABEL_MAP[row[\"label\"]]\n            if img_id in cache:\n                self.samples.append((cache[img_id], label_id))\n            else:\n                missing += 1\n\n        # Class-balanced weights\n        labs    = np.array([s[1] for s in self.samples])\n        counts  = np.bincount(labs, minlength=len(LABEL_MAP))\n        weights = 1.0 / (counts[labs] + 1e-6)\n        self.sample_weights = weights.tolist()\n\n        print(f\"  [{split:5s}] {len(df)} rows → \"\n              f\"{len(self.samples)} samples  (missing: {missing})\")\n\n    def __len__(self): return len(self.samples)\n\n    def __getitem__(self, idx):\n        path, label = self.samples[idx]\n        img = cv2.imread(path, cv2.IMREAD_COLOR)\n        if img is None:\n            img = np.full((self.patch_size,self.patch_size,3),128,dtype=np.uint8)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.augment:\n            img = np.array(self.TRAIN_AUG(Image.fromarray(img)))\n\n        return generate_views(img, self.patch_size), label\n\n\n# ── Synthetic fallback (used only if real data unavailable) ───\nclass SyntheticDataset(Dataset):\n    CLASS_COLORS = [(180,100,160),(120,180,120),(200,160,80),(100,140,200),(200,100,100)]\n\n    def __init__(self, n=1500, size=64, n_classes=5):\n        self.n=n; self.size=size; self.n_classes=n_classes\n        self.labels=np.random.randint(0,n_classes,n)\n        self.sample_weights=[1.0]*n\n\n    def __len__(self): return self.n\n\n    def __getitem__(self, idx):\n        c  = int(self.labels[idx])\n        cr,cg,cb = self.CLASS_COLORS[c]\n        img = np.zeros((self.size,self.size,3),dtype=np.float32)\n        img[:,:,0]=cr+np.random.randn(self.size,self.size)*25\n        img[:,:,1]=cg+np.random.randn(self.size,self.size)*25\n        img[:,:,2]=cb+np.random.randn(self.size,self.size)*25\n        for _ in range(random.randint(3,12)):\n            x,y=random.randint(4,self.size-4),random.randint(4,self.size-4)\n            cv2.circle(img,(x,y),random.randint(2,5),(40,20,70),-1)\n        img = np.clip(img,0,255).astype(np.uint8)\n        return generate_views(img, self.size), c\n\nprint(\"Dataset classes defined ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:29.817985Z","iopub.execute_input":"2026-03-07T07:05:29.818357Z","iopub.status.idle":"2026-03-07T07:05:30.014176Z","shell.execute_reply.started":"2026-03-07T07:05:29.818334Z","shell.execute_reply":"2026-03-07T07:05:30.01346Z"}},"outputs":[],"execution_count":null},{"id":"90c031c3-5729-49d6-b905-915123defe4c","cell_type":"markdown","source":"## Cell 7: Build DataLoaders\n","metadata":{}},{"id":"dea53640-d10e-4c46-bb21-5bfaf24f96ab","cell_type":"code","source":"def collate_fn(batch):\n    vl, ll = zip(*batch)\n    return {k: torch.stack([v[k] for v in vl]) for k in vl[0]}, \\\n           torch.tensor(ll, dtype=torch.long)\n\ndef build_dataloaders():\n    root      = CFG[\"dataset_root\"]\n    train_csv = os.path.join(root, \"train.csv\")\n    thumb_dir = os.path.join(root, \"train_thumbnails\")\n\n    USE_REAL = os.path.exists(train_csv) and os.path.exists(thumb_dir)\n\n    if USE_REAL:\n        print(\"✅  Using REAL UBC-OCEAN data\")\n        df = pd.read_csv(train_csv)\n        df = df[df[\"label\"].isin(LABEL_MAP)].reset_index(drop=True)\n\n        train_df, tmp = train_test_split(df, test_size=0.2,\n                                          stratify=df[\"label\"], random_state=SEED)\n        val_df, test_df = train_test_split(tmp, test_size=0.5,\n                                            stratify=tmp[\"label\"], random_state=SEED)\n\n        PATCH = CFG[\"patch_size\"]\n        train_ds = UBCOceanDataset(train_df, thumb_dir, PATCH, \"train\")\n        val_ds   = UBCOceanDataset(val_df,   thumb_dir, PATCH, \"val\")\n        test_ds  = UBCOceanDataset(test_df,  thumb_dir, PATCH, \"test\")\n\n        if len(train_ds) == 0:\n            print(\"⚠  No thumbnails matched — switching to synthetic data\")\n            USE_REAL = False\n\n    if not USE_REAL:\n        print(\"⚠  Using SYNTHETIC data (real dataset not found or empty)\")\n        PATCH      = 64\n        CFG[\"patch_size\"] = PATCH\n        train_ds = SyntheticDataset(1500, PATCH)\n        val_ds   = SyntheticDataset(300,  PATCH)\n        test_ds  = SyntheticDataset(300,  PATCH)\n\n    sampler = WeightedRandomSampler(\n        torch.tensor(train_ds.sample_weights, dtype=torch.float32),\n        num_samples=len(train_ds), replacement=True)\n\n    train_dl = DataLoader(train_ds, batch_size=CFG[\"batch_size\"],\n                          sampler=sampler, collate_fn=collate_fn,\n                          num_workers=2, pin_memory=True)\n    val_dl   = DataLoader(val_ds, batch_size=CFG[\"batch_size\"],\n                          shuffle=False, collate_fn=collate_fn,\n                          num_workers=2, pin_memory=True)\n    test_dl  = DataLoader(test_ds, batch_size=CFG[\"batch_size\"],\n                          shuffle=False, collate_fn=collate_fn,\n                          num_workers=2, pin_memory=True)\n\n    # Verify one batch\n    views, labels = next(iter(train_dl))\n    print(f\"\\nBatch shapes:\")\n    for k,v in views.items():\n        print(f\"  {k}: {tuple(v.shape)}\")\n    print(f\"  labels: {labels[:6].tolist()} → \"\n          f\"{[CLASS_NAMES[l] for l in labels[:6].tolist()]}\")\n\n    return train_dl, val_dl, test_dl\n\ntrain_dl, val_dl, test_dl = build_dataloaders()\nprint(\"\\nDataLoaders ready ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:30.016176Z","iopub.execute_input":"2026-03-07T07:05:30.016726Z","iopub.status.idle":"2026-03-07T07:05:47.551362Z","shell.execute_reply.started":"2026-03-07T07:05:30.016702Z","shell.execute_reply":"2026-03-07T07:05:47.550117Z"}},"outputs":[],"execution_count":null},{"id":"1b7db486-3b05-4212-ad91-b2f1c2c76c5c","cell_type":"markdown","source":"## Cell 8: Model Architecture\n","metadata":{}},{"id":"ba6f7034-07c9-4d96-b9f9-1763cd5d9e49","cell_type":"code","source":"# ── View Encoder (VAE) ────────────────────────────────────────\nclass ViewEncoder(nn.Module):\n    def __init__(self, latent_dim=256, variational=True):\n        super().__init__()\n        self.variational = variational\n        self.conv = nn.Sequential(\n            nn.Conv2d(3,32,3,stride=2,padding=1),  nn.BatchNorm2d(32),  nn.ReLU(True),\n            nn.Conv2d(32,64,3,stride=2,padding=1), nn.BatchNorm2d(64),  nn.ReLU(True),\n            nn.Conv2d(64,128,3,stride=2,padding=1),nn.BatchNorm2d(128), nn.ReLU(True),\n            nn.Conv2d(128,256,3,stride=2,padding=1),nn.BatchNorm2d(256),nn.ReLU(True),\n            nn.AdaptiveAvgPool2d((4,4)),\n        )\n        self.fc = nn.Sequential(\n            nn.Flatten(), nn.Linear(256*4*4, 512),\n            nn.ReLU(True), nn.Dropout(0.3),\n        )\n        if variational:\n            self.fc_mu = nn.Linear(512, latent_dim)\n            self.fc_lv = nn.Linear(512, latent_dim)\n        else:\n            self.fc_out = nn.Linear(512, latent_dim)\n\n    def forward(self, x):\n        h = self.fc(self.conv(x))\n        if self.variational:\n            mu, lv = self.fc_mu(h), self.fc_lv(h)\n            z = mu + torch.randn_like(mu)*torch.exp(0.5*lv) if self.training else mu\n            return z, mu, lv\n        return self.fc_out(h), None, None\n\n\n# ── Meta-Fusion Gate ──────────────────────────────────────────\nclass MetaFusionGate(nn.Module):\n    def __init__(self, n_views=4, latent_dim=256):\n        super().__init__()\n        d = n_views * latent_dim\n        self.net = nn.Sequential(\n            nn.Linear(d, d//2),  nn.ReLU(True), nn.Dropout(0.2),\n            nn.Linear(d//2, d//4), nn.ReLU(True),\n            nn.Linear(d//4, n_views),\n        )\n    def forward(self, zs):\n        return F.softmax(self.net(torch.cat(zs, dim=-1)), dim=-1)\n\n\n# ── View Decoder ──────────────────────────────────────────────\nclass ViewDecoder(nn.Module):\n    def __init__(self, latent_dim=256, out_size=224):\n        super().__init__()\n        self.out_size = out_size\n        self.fc = nn.Sequential(\n            nn.Linear(latent_dim, 512), nn.ReLU(True),\n            nn.Linear(512, 256*4*4),   nn.ReLU(True),\n        )\n        self.deconv = nn.Sequential(\n            nn.ConvTranspose2d(256,128,4,stride=2,padding=1),nn.BatchNorm2d(128),nn.ReLU(True),\n            nn.ConvTranspose2d(128,64, 4,stride=2,padding=1),nn.BatchNorm2d(64), nn.ReLU(True),\n            nn.ConvTranspose2d(64, 32, 4,stride=2,padding=1),nn.BatchNorm2d(32), nn.ReLU(True),\n            nn.ConvTranspose2d(32, 3,  4,stride=2,padding=1),nn.Sigmoid(),\n        )\n    def forward(self, z):\n        x = self.deconv(self.fc(z).view(-1,256,4,4))\n        if x.shape[-1] != self.out_size:\n            x = F.interpolate(x, self.out_size, mode=\"bilinear\", align_corners=False)\n        return x\n\n\n# ── Full Multi-View Autoencoder ───────────────────────────────\nclass MultiViewAE(nn.Module):\n    def __init__(self, latent_dim=256, n_views=4, patch_size=224):\n        super().__init__()\n        self.n_views    = n_views\n        self.latent_dim = latent_dim\n        self.encoders   = nn.ModuleList([ViewEncoder(latent_dim) for _ in range(n_views)])\n        self.fusion     = MetaFusionGate(n_views, latent_dim)\n        self.decoders   = nn.ModuleList([ViewDecoder(latent_dim, patch_size) for _ in range(n_views)])\n\n    def forward(self, views):\n        keys = [\"v1\",\"v2\",\"v3\",\"v4\"]\n        zs, mus, lvs = [], [], []\n        for i,k in enumerate(keys):\n            z, mu, lv = self.encoders[i](views[k])\n            zs.append(z); mus.append(mu); lvs.append(lv)\n        w   = self.fusion(zs)\n        z_f = sum(w[:,i:i+1]*zs[i] for i in range(self.n_views))\n        return {\"z_f\":z_f, \"zs\":zs, \"mus\":mus, \"logvars\":lvs,\n                \"weights\":w, \"recons\":[dec(z_f) for dec in self.decoders]}\n\n    # Alternative fusion strategies for comparison\n    def fuse_mean(self, zs):      return torch.stack(zs,dim=1).mean(1)\n    def fuse_concat(self, zs):    return torch.cat(zs, dim=-1)\n    def fuse_attention(self, zs):\n        z = torch.stack(zs,dim=1); d = z.size(-1)**0.5\n        a = torch.softmax((z@z.transpose(-1,-2))/d, dim=-1)\n        return (a@z).mean(1)\n\n\n# ── MLP Classifier ────────────────────────────────────────────\nclass MLPClassifier(nn.Module):\n    def __init__(self, latent_dim=256, num_classes=5, dropout=0.4):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(latent_dim,256), nn.BatchNorm1d(256), nn.ReLU(True), nn.Dropout(dropout),\n            nn.Linear(256,128),        nn.BatchNorm1d(128), nn.ReLU(True), nn.Dropout(dropout),\n            nn.Linear(128, num_classes),\n        )\n    def forward(self, z): return self.net(z)\n\n\n# ── Full System ───────────────────────────────────────────────\nclass FullSystem(nn.Module):\n    def __init__(self, latent_dim=256, n_views=4, patch_size=224, num_classes=5):\n        super().__init__()\n        self.ae  = MultiViewAE(latent_dim, n_views, patch_size)\n        self.clf = MLPClassifier(latent_dim, num_classes)\n    def forward(self, views):\n        out = self.ae(views)\n        out[\"logits\"] = self.clf(out[\"z_f\"])\n        return out\n\nprint(\"Model architecture defined ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:47.553774Z","iopub.execute_input":"2026-03-07T07:05:47.554108Z","iopub.status.idle":"2026-03-07T07:05:47.580316Z","shell.execute_reply.started":"2026-03-07T07:05:47.554065Z","shell.execute_reply":"2026-03-07T07:05:47.579664Z"}},"outputs":[],"execution_count":null},{"id":"8d0955dc-74de-4790-b613-d99be7f2b474","cell_type":"markdown","source":"## Cell 9: Loss Functions\n","metadata":{}},{"id":"6a6699a9-1815-4416-893c-6007e864b974","cell_type":"code","source":"class MultiViewLoss(nn.Module):\n    def __init__(self, l_rec=1.0, l_cons=0.5, l_kl=0.1):\n        super().__init__()\n        self.l_rec, self.l_cons, self.l_kl = l_rec, l_cons, l_kl\n\n    def forward(self, out, views):\n        keys = [\"v1\",\"v2\",\"v3\",\"v4\"]\n\n        # Reconstruction loss\n        rec = sum(F.mse_loss(out[\"recons\"][i], views[keys[i]])\n                  for i in range(4)) / 4\n\n        # Cross-view consistency\n        zs  = out[\"zs\"]; n = len(zs)\n        cons = sum((1 - F.cosine_similarity(zs[i],zs[j],dim=-1)).mean()\n                   for i in range(n) for j in range(i+1,n)) / max(n*(n-1)//2,1)\n\n        # KL divergence (VAE regularisation)\n        kl = torch.tensor(0., device=zs[0].device)\n        cnt = 0\n        for mu, lv in zip(out[\"mus\"], out[\"logvars\"]):\n            if mu is None: continue\n            kl  += -0.5 * torch.mean(1 + lv - mu.pow(2) - lv.exp())\n            cnt += 1\n        if cnt > 0: kl /= cnt\n\n        total = self.l_rec*rec + self.l_cons*cons + self.l_kl*kl\n        return total, {\"rec\":rec.item(), \"cons\":cons.item(), \"kl\":kl.item()}\n\nprint(\"Loss functions defined ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:47.582029Z","iopub.execute_input":"2026-03-07T07:05:47.582322Z","iopub.status.idle":"2026-03-07T07:05:47.596191Z","shell.execute_reply.started":"2026-03-07T07:05:47.582301Z","shell.execute_reply":"2026-03-07T07:05:47.595411Z"}},"outputs":[],"execution_count":null},{"id":"6a1c21a3-4119-4d80-b807-7c2ecbbda53f","cell_type":"markdown","source":"## Cell 10: Trainer\n","metadata":{}},{"id":"8f6d1cda-2bc4-418e-bd59-3eceb91d8567","cell_type":"code","source":"class Trainer:\n    def __init__(self, model, device, cfg):\n        self.model    = model.to(device)\n        self.device   = device\n        self.ae_crit  = MultiViewLoss(cfg[\"lambda_rec\"],cfg[\"lambda_cons\"],cfg[\"lambda_kl\"])\n        self.clf_crit = nn.CrossEntropyLoss(label_smoothing=0.1)\n        self.ae_opt   = optim.AdamW(model.ae.parameters(),\n                                    lr=cfg[\"lr\"], weight_decay=cfg[\"weight_decay\"])\n        self.clf_opt  = optim.AdamW(model.clf.parameters(),\n                                    lr=cfg[\"lr\"]*0.5, weight_decay=cfg[\"weight_decay\"])\n        self.ae_sch   = optim.lr_scheduler.CosineAnnealingLR(self.ae_opt, cfg[\"ae_epochs\"])\n        self.clf_sch  = optim.lr_scheduler.CosineAnnealingLR(self.clf_opt, cfg[\"joint_epochs\"])\n        self.scaler   = GradScaler()\n        self.history  = defaultdict(list)\n\n    def _to(self, v): return {k:t.to(self.device) for k,t in v.items()}\n\n    # ── AE pretraining step ───────────────────────────────────\n    def _ae_epoch(self, dl):\n        self.model.ae.train(); m = defaultdict(float)\n        for views, _ in tqdm(dl, desc=\"  AE\", leave=False):\n            views = self._to(views); self.ae_opt.zero_grad()\n            with autocast():\n                out = self.model.ae(views)\n                loss, bd = self.ae_crit(out, views)\n            self.scaler.scale(loss).backward()\n            self.scaler.unscale_(self.ae_opt)\n            nn.utils.clip_grad_norm_(self.model.ae.parameters(), 1.0)\n            self.scaler.step(self.ae_opt); self.scaler.update()\n            m[\"loss\"] += loss.item()\n            for k,v in bd.items(): m[k] += v\n        n = max(len(dl),1); return {k:v/n for k,v in m.items()}\n\n    # ── Joint fine-tuning step ────────────────────────────────\n    def _joint_epoch(self, dl):\n        self.model.train(); m = defaultdict(float); c = tot = 0\n        for views, labels in tqdm(dl, desc=\"  Joint\", leave=False):\n            views = self._to(views); labels = labels.to(self.device)\n            self.ae_opt.zero_grad(); self.clf_opt.zero_grad()\n            with autocast():\n                out  = self.model(views)\n                ae_l, bd = self.ae_crit(out, views)\n                clf_l    = self.clf_crit(out[\"logits\"], labels)\n                loss     = ae_l + clf_l\n            self.scaler.scale(loss).backward()\n            self.scaler.unscale_(self.ae_opt); self.scaler.unscale_(self.clf_opt)\n            nn.utils.clip_grad_norm_(self.model.parameters(), 1.0)\n            self.scaler.step(self.ae_opt); self.scaler.step(self.clf_opt)\n            self.scaler.update()\n            m[\"loss\"] += loss.item(); m[\"clf\"] += clf_l.item()\n            for k,v in bd.items(): m[k] += v\n            c   += (out[\"logits\"].argmax(1)==labels).sum().item()\n            tot += labels.size(0)\n        n = max(len(dl),1); res = {k:v/n for k,v in m.items()}\n        res[\"acc\"] = c/max(tot,1); return res\n\n    # ── Evaluation ────────────────────────────────────────────\n    @torch.no_grad()\n    def evaluate(self, dl, phase=\"val\"):\n        self.model.eval(); pa, la, lga = [], [], []\n        for views, labels in tqdm(dl, desc=f\"  {phase}\", leave=False):\n            out = self.model(self._to(views))\n            pa.append(out[\"logits\"].argmax(1).cpu())\n            la.append(labels)\n            lga.append(out[\"logits\"].softmax(-1).cpu())\n        p  = torch.cat(pa).numpy()\n        l  = torch.cat(la).numpy()\n        lg = torch.cat(lga).numpy()\n        f1  = f1_score(l, p, average=\"macro\", zero_division=0)\n        acc = accuracy_score(l, p)\n        try:    auc = roc_auc_score(l, lg, multi_class=\"ovr\", average=\"macro\")\n        except: auc = 0.0\n        return {\"acc\":acc, \"f1\":f1, \"auc\":auc, \"preds\":p, \"labels\":l, \"logits\":lg}\n\n    # ── Master fit ────────────────────────────────────────────\n    def fit(self, train_dl, val_dl, ae_epochs=15, joint_epochs=25):\n        best_f1 = 0; best_w = None\n\n        print(\"\\n\" + \"=\"*50)\n        print(\"PHASE 1 — AE PRETRAINING\")\n        print(\"=\"*50)\n        for ep in range(ae_epochs):\n            tr = self._ae_epoch(train_dl); self.ae_sch.step()\n            print(f\"  Ep {ep+1:3d}/{ae_epochs}  \"\n                  f\"loss={tr['loss']:.4f}  rec={tr['rec']:.4f}  \"\n                  f\"cons={tr['cons']:.4f}  kl={tr['kl']:.4f}\")\n            for k,v in tr.items(): self.history[f\"ae_{k}\"].append(v)\n\n        print(\"\\n\" + \"=\"*50)\n        print(\"PHASE 2 — JOINT FINE-TUNING\")\n        print(\"=\"*50)\n        for ep in range(joint_epochs):\n            tr  = self._joint_epoch(train_dl)\n            val = self.evaluate(val_dl, \"val\")\n            self.clf_sch.step()\n            print(f\"  Ep {ep+1:3d}/{joint_epochs}  \"\n                  f\"loss={tr['loss']:.4f}  acc={tr['acc']:.3f}  \"\n                  f\"val_acc={val['acc']:.3f}  val_f1={val['f1']:.3f}  \"\n                  f\"val_auc={val['auc']:.3f}\")\n            for k,v in tr.items():\n                if not isinstance(v, np.ndarray): self.history[f\"tr_{k}\"].append(v)\n            for k,v in val.items():\n                if not isinstance(v, np.ndarray): self.history[f\"val_{k}\"].append(v)\n            if val[\"f1\"] > best_f1:\n                best_f1 = val[\"f1\"]\n                best_w  = {k:v.clone() for k,v in self.model.state_dict().items()}\n                print(f\"    ✓ Best saved  (val F1={best_f1:.4f})\")\n\n        if best_w: self.model.load_state_dict(best_w)\n        print(f\"\\n✓ Training complete  |  Best val F1 = {best_f1:.4f}\")\n        return self.history\n\nprint(\"Trainer defined ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:47.597323Z","iopub.execute_input":"2026-03-07T07:05:47.597572Z","iopub.status.idle":"2026-03-07T07:05:47.620314Z","shell.execute_reply.started":"2026-03-07T07:05:47.59755Z","shell.execute_reply":"2026-03-07T07:05:47.619581Z"}},"outputs":[],"execution_count":null},{"id":"018d4078-2453-4dc4-9703-b3b217101b12","cell_type":"markdown","source":"## Cell 11: Initialise Model + Train\n","metadata":{}},{"id":"e2cd1703-91b5-4319-aa9a-ffa3e2810739","cell_type":"code","source":"system = FullSystem(\n    latent_dim  = CFG[\"latent_dim\"],\n    n_views     = CFG[\"n_views\"],\n    patch_size  = CFG[\"patch_size\"],\n    num_classes = CFG[\"num_classes\"],\n).to(DEVICE)\n\nn_params = sum(p.numel() for p in system.parameters() if p.requires_grad)\nprint(f\"Model parameters : {n_params:,}\")\n\ntrainer = Trainer(system, DEVICE, CFG)\nhistory = trainer.fit(\n    train_dl, val_dl,\n    ae_epochs    = CFG[\"ae_epochs\"],\n    joint_epochs = CFG[\"joint_epochs\"],\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T07:05:47.621351Z","iopub.execute_input":"2026-03-07T07:05:47.621812Z","iopub.status.idle":"2026-03-07T08:48:38.136924Z","shell.execute_reply.started":"2026-03-07T07:05:47.621786Z","shell.execute_reply":"2026-03-07T08:48:38.136091Z"}},"outputs":[],"execution_count":null},{"id":"40ae9417-d76c-4bc9-be3e-810a02d0f22b","cell_type":"markdown","source":"## Cell 12: Extract Latents + Compute All Metrics\n","metadata":{}},{"id":"9b52d7a7-542c-4966-a32e-00f241e38850","cell_type":"code","source":"@torch.no_grad()\ndef extract_latents(model, dl):\n    model.eval(); zfs, ws, labs, lgs = [], [], [], []\n    for views, labels in tqdm(dl, desc=\"Extracting latents\"):\n        out = model({k:v.to(DEVICE) for k,v in views.items()})\n        zfs.append(out[\"z_f\"].cpu());  ws.append(out[\"weights\"].cpu())\n        labs.append(labels);           lgs.append(out[\"logits\"].softmax(-1).cpu())\n    return {\n        \"z_f\":     torch.cat(zfs).numpy(),\n        \"weights\": torch.cat(ws).numpy(),\n        \"labels\":  torch.cat(labs).numpy(),\n        \"logits\":  torch.cat(lgs).numpy(),\n    }\n\ntest_lat = extract_latents(system, test_dl)\npreds    = test_lat[\"logits\"].argmax(1)\nlabs     = test_lat[\"labels\"].astype(int)\n\n# ── Classification metrics ────────────────────────────────────\nacc  = accuracy_score(labs, preds)\nprec = precision_score(labs, preds, average=\"macro\", zero_division=0)\nrec  = recall_score(labs, preds,    average=\"macro\", zero_division=0)\nf1   = f1_score(labs, preds,        average=\"macro\", zero_division=0)\ntry:    auc = roc_auc_score(labs, test_lat[\"logits\"], multi_class=\"ovr\", average=\"macro\")\nexcept: auc = 0.0\n\n# ── Regression metrics (on softmax vs one-hot) ────────────────\noh   = np.eye(CFG[\"num_classes\"])[labs]\nmse  = mean_squared_error(oh.ravel(), test_lat[\"logits\"].ravel())\nmae  = mean_absolute_error(oh.ravel(), test_lat[\"logits\"].ravel())\nmape = (np.abs(oh.ravel()-test_lat[\"logits\"].ravel()) /\n        (np.abs(oh.ravel())+1e-6)).mean()*100\n\n# ── Latent quality metrics ────────────────────────────────────\ntry:    sil_gt = silhouette_score(test_lat[\"z_f\"], labs,\n                                   sample_size=min(len(labs),500))\nexcept: sil_gt = 0.0\n\nkm = KMeans(n_clusters=CFG[\"num_classes\"], random_state=SEED, n_init=10)\nkm_labels = km.fit_predict(test_lat[\"z_f\"])\ntry:    sil_km = silhouette_score(test_lat[\"z_f\"], km_labels,\n                                   sample_size=min(len(labs),500))\nexcept: sil_km = 0.0\n\n# ── Fisher Discriminant Ratio ─────────────────────────────────\ndef fisher_dr(Z, y):\n    classes = np.unique(y); total = 0; cnt = 0\n    for c1 in classes:\n        for c2 in classes:\n            if c2 <= c1: continue\n            z1,z2 = Z[y==c1], Z[y==c2]\n            if len(z1)<2 or len(z2)<2: continue\n            num   = ((z1.mean(0)-z2.mean(0))**2).mean()\n            denom = (z1.var(0)+z2.var(0)).mean()+1e-8\n            total += num/denom; cnt += 1\n    return float(total/max(cnt,1))\n\nfdr = fisher_dr(test_lat[\"z_f\"], labs)\n\nMETRICS = {\n    \"Accuracy\":              acc,\n    \"Precision (macro)\":     prec,\n    \"Recall (macro)\":        rec,\n    \"F1 (macro)\":            f1,\n    \"AUC (OvR)\":             auc,\n    \"MSE\":                   mse,\n    \"RMSE\":                  np.sqrt(mse),\n    \"MAE\":                   mae,\n    \"MAPE (%)\":              mape,\n    \"Fisher Disc. Ratio\":    fdr,\n    \"Silhouette (GT)\":       sil_gt,\n    \"Silhouette (KMeans)\":   sil_km,\n}\n\nprint(\"\\n\" + \"=\"*50)\nprint(\"FINAL TEST SET METRICS\")\nprint(\"=\"*50)\nfor k,v in METRICS.items():\n    print(f\"  {k:<28} {v:.4f}\")\n\nwith open(f\"{CFG['output_dir']}/metrics.json\",\"w\") as f:\n    json.dump({k:float(v) for k,v in METRICS.items()}, f, indent=2)\nprint(\"\\nMetrics saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:48:38.138527Z","iopub.execute_input":"2026-03-07T08:48:38.138903Z","iopub.status.idle":"2026-03-07T08:48:46.965376Z","shell.execute_reply.started":"2026-03-07T08:48:38.138876Z","shell.execute_reply":"2026-03-07T08:48:46.961992Z"}},"outputs":[],"execution_count":null},{"id":"17ebde6f-74cd-4b13-b395-60f39a98c8fe","cell_type":"markdown","source":"## Cell 13: Training Curves\n","metadata":{}},{"id":"5b233d8c-a000-439a-bdb4-a6d3ba8161f8","cell_type":"code","source":"fig, axes = plt.subplots(2, 3, figsize=(18,10))\nfig.suptitle(\"Training History — Multi-View Meta-Fusion AE\",\n             fontsize=15, fontweight=\"bold\")\n\nplots = [(\"ae_loss\",\"AE Pretraining Loss\"),\n         (\"tr_loss\",\"Joint Training Loss\"),\n         (\"tr_acc\", \"Training Accuracy\"),\n         (\"val_acc\",\"Validation Accuracy\"),\n         (\"val_f1\", \"Validation F1 (macro)\"),\n         (\"val_auc\",\"Validation AUC\")]\n\nfor ax, (key, title) in zip(axes.flatten(), plots):\n    if key in history and history[key]:\n        ax.plot(history[key], lw=2, color=\"#457B9D\", marker=\"o\",\n                markersize=3, markevery=max(1,len(history[key])//10))\n        ax.set_title(title, fontweight=\"bold\")\n        ax.set_xlabel(\"Epoch\"); ax.grid(True, alpha=0.3)\n        ax.set_ylim(bottom=0)\n    else:\n        ax.set_visible(False)\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/training_curves.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"Training curves saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:48:46.966735Z","iopub.execute_input":"2026-03-07T08:48:46.970147Z","iopub.status.idle":"2026-03-07T08:48:48.636596Z","shell.execute_reply.started":"2026-03-07T08:48:46.970112Z","shell.execute_reply":"2026-03-07T08:48:48.635783Z"}},"outputs":[],"execution_count":null},{"id":"6f9ff899-91c2-4555-bb5a-491364fd186c","cell_type":"markdown","source":"## Cell 14: t-SNE Visualisation\n","metadata":{}},{"id":"ed1e8b99-3c26-497a-8583-7cbbd219f75b","cell_type":"code","source":"print(\"Computing t-SNE (may take ~30s) …\")\nperp = min(30, len(test_lat[\"z_f\"])//5, len(test_lat[\"z_f\"])-1)\nperp = max(5, perp)\nZ2d  = TSNE(2, perplexity=perp, random_state=SEED,\n             n_iter=1000, learning_rate=\"auto\", init=\"pca\"\n             ).fit_transform(test_lat[\"z_f\"])\n\nfig, axes = plt.subplots(1, 2, figsize=(16,7))\nfig.suptitle(\"t-SNE — Fused Latent Space\", fontsize=14, fontweight=\"bold\")\n\nfor ci,(cn,col) in enumerate(zip(CLASS_NAMES,COLORS)):\n    m = labs==ci\n    if m.sum()==0: continue\n    axes[0].scatter(Z2d[m,0],Z2d[m,1], c=col, label=cn,\n                    alpha=0.65, s=20, edgecolors=\"none\")\naxes[0].set_title(\"Ground Truth Labels\")\naxes[0].legend(markerscale=2, fontsize=9)\naxes[0].grid(True, alpha=0.3)\naxes[0].set_xlabel(\"t-SNE 1\"); axes[0].set_ylabel(\"t-SNE 2\")\n\nsc = axes[1].scatter(Z2d[:,0],Z2d[:,1], c=km_labels,\n                      cmap=\"tab10\", alpha=0.65, s=20, edgecolors=\"none\")\nplt.colorbar(sc, ax=axes[1], label=\"Cluster ID\")\naxes[1].set_title(f\"K-Means Clusters (k={CFG['num_classes']})\")\naxes[1].grid(True, alpha=0.3)\naxes[1].set_xlabel(\"t-SNE 1\"); axes[1].set_ylabel(\"t-SNE 2\")\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/tsne.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"t-SNE saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:50:36.970661Z","iopub.execute_input":"2026-03-07T08:50:36.970947Z","iopub.status.idle":"2026-03-07T08:50:37.947259Z","shell.execute_reply.started":"2026-03-07T08:50:36.970906Z","shell.execute_reply":"2026-03-07T08:50:37.946552Z"}},"outputs":[],"execution_count":null},{"id":"e23c951e-1d2e-4f6a-a290-275a52d065f0","cell_type":"markdown","source":"## Cell 15: Confusion Matrix\n","metadata":{}},{"id":"0c078616-6f2b-4198-9750-cb1b19f887aa","cell_type":"code","source":"cm  = confusion_matrix(labs, preds, labels=list(range(CFG[\"num_classes\"])))\ncmp = cm.astype(float) / (cm.sum(1, keepdims=True) + 1e-8)\n\nfig, axes = plt.subplots(1, 2, figsize=(14,6))\nfig.suptitle(\"Confusion Matrix — Test Set\", fontsize=14, fontweight=\"bold\")\n\nfor ax, data, fmt, title in zip(axes,\n                                  [cm, cmp],\n                                  [\"d\", \".2f\"],\n                                  [\"Counts\",\"Row-Normalised\"]):\n    sns.heatmap(data, annot=True, fmt=fmt, cmap=\"Blues\",\n                xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES,\n                ax=ax, linewidths=0.5, cbar=True)\n    ax.set_title(title); ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\")\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/confusion_matrix.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"Confusion matrix saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:50:48.280038Z","iopub.execute_input":"2026-03-07T08:50:48.28086Z","iopub.status.idle":"2026-03-07T08:50:49.10844Z","shell.execute_reply.started":"2026-03-07T08:50:48.280825Z","shell.execute_reply":"2026-03-07T08:50:49.107706Z"}},"outputs":[],"execution_count":null},{"id":"5c0f7dd2-90e8-4c3d-910b-4a1b90ce48ae","cell_type":"markdown","source":"## Cell 16: Meta-Fusion Gate Weights\n","metadata":{}},{"id":"843bf2a2-ff8c-47e3-ab7b-e3092d8387dd","cell_type":"code","source":"ws = test_lat[\"weights\"]   # (N, 4)\n\nfig, axes = plt.subplots(1, 2, figsize=(14,5))\nfig.suptitle(\"Meta-Fusion Gate Weights\", fontsize=14, fontweight=\"bold\")\n\n# Per-class mean weights\nmeans = np.array([ws[labs==c].mean(0) if (labs==c).sum()>0\n                   else np.zeros(4) for c in range(CFG[\"num_classes\"])])\nx = np.arange(CFG[\"num_classes\"]); bw = 0.18\nfor i,(vn,col) in enumerate(zip(VIEW_NAMES,COLORS)):\n    axes[0].bar(x+i*bw, means[:,i], bw, label=vn, color=col)\naxes[0].set_xticks(x+1.5*bw); axes[0].set_xticklabels(CLASS_NAMES)\naxes[0].set_ylabel(\"Mean Weight\"); axes[0].set_title(\"Per-Class Mean Weights\")\naxes[0].legend(fontsize=8); axes[0].grid(True, alpha=0.3, axis=\"y\")\n\n# Distribution box-plots\naxes[1].boxplot(ws, labels=VIEW_NAMES, patch_artist=True,\n                boxprops=dict(facecolor=\"#457B9D\", alpha=0.6),\n                medianprops=dict(color=\"black\",lw=2))\naxes[1].set_ylabel(\"Attention Weight\"); axes[1].set_title(\"Weight Distributions\")\naxes[1].grid(True, alpha=0.3, axis=\"y\")\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/fusion_weights.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"Fusion weights saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:50:54.283111Z","iopub.execute_input":"2026-03-07T08:50:54.283886Z","iopub.status.idle":"2026-03-07T08:50:54.905329Z","shell.execute_reply.started":"2026-03-07T08:50:54.283842Z","shell.execute_reply":"2026-03-07T08:50:54.904464Z"}},"outputs":[],"execution_count":null},{"id":"cc8b64ef-edcb-46ba-8256-036f5084400e","cell_type":"markdown","source":"## Cell 17: Reconstruction Samples\n","metadata":{}},{"id":"76136086-fd6a-4009-a530-567128c2f8d7","cell_type":"code","source":"system.eval()\nviews_b, labs_b = next(iter(test_dl))\nviews_b_gpu = {k:v.to(DEVICE) for k,v in views_b.items()}\n\nwith torch.no_grad():\n    out_b = system(views_b_gpu)\n\nrecons = out_b[\"recons\"]\nn_show = min(3, views_b[\"v1\"].size(0))\n\nfig, axes = plt.subplots(n_show, 8, figsize=(20, 3*n_show))\nif n_show == 1: axes = axes[np.newaxis, :]\nfig.suptitle(\"Original vs Reconstructed Views\", fontsize=13, fontweight=\"bold\")\n\nfor row in range(n_show):\n    for col, (key, vname) in enumerate(zip([\"v1\",\"v2\",\"v3\",\"v4\"], VIEW_NAMES)):\n        # Original\n        orig = views_b[key][row].permute(1,2,0).numpy()\n        axes[row, col*2].imshow(np.clip(orig,0,1))\n        axes[row, col*2].set_title(vname if row==0 else \"\", fontsize=8)\n        axes[row, col*2].axis(\"off\")\n        # Reconstructed\n        rec = recons[col][row].cpu().permute(1,2,0).numpy()\n        axes[row, col*2+1].imshow(np.clip(rec,0,1))\n        axes[row, col*2+1].set_title(\"Recon\" if row==0 else \"\", fontsize=8)\n        axes[row, col*2+1].axis(\"off\")\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/reconstructions.png\", dpi=120, bbox_inches=\"tight\")\nplt.show()\nprint(\"Reconstructions saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:51:02.603213Z","iopub.execute_input":"2026-03-07T08:51:02.603478Z","iopub.status.idle":"2026-03-07T08:51:12.686536Z","shell.execute_reply.started":"2026-03-07T08:51:02.603455Z","shell.execute_reply":"2026-03-07T08:51:12.685494Z"}},"outputs":[],"execution_count":null},{"id":"1bd0420e-7b64-4b15-80b6-35c03e4f51a7","cell_type":"markdown","source":"## Cell 18: Fusion Strategy Comparison\n","metadata":{}},{"id":"84552cd6-95b7-40e8-bd10-15b2a98b7ce5","cell_type":"code","source":"print(\"Comparing fusion strategies …\")\nsystem.eval()\n\nreps = defaultdict(list); all_labs = []\nwith torch.no_grad():\n    for views, labels in tqdm(test_dl, desc=\"  Collecting\", leave=False):\n        views_gpu = {k:v.to(DEVICE) for k,v in views.items()}\n        out = system.ae(views_gpu); zs = out[\"zs\"]\n        reps[\"Meta-Fusion\"].append(out[\"z_f\"].cpu())\n        reps[\"Mean\"].append(system.ae.fuse_mean(zs).cpu())\n        reps[\"Attention\"].append(system.ae.fuse_attention(zs).cpu())\n        # concat is 4x bigger — use PCA to reduce\n        concat_z = system.ae.fuse_concat(zs).cpu()\n        reps[\"Concatenation\"].append(concat_z)\n        all_labs.append(labels)\n\nall_labs = torch.cat(all_labs).numpy()\nresults  = {}\n\nfor name, rep_list in reps.items():\n    Z = torch.cat(rep_list).numpy()\n    Z = StandardScaler().fit_transform(Z)\n    n = len(Z); idx = np.random.permutation(n)\n    tr_idx, te_idx = idx[:int(0.8*n)], idx[int(0.8*n):]\n    try:\n        nc  = len(np.unique(all_labs))\n        lda = LinearDiscriminantAnalysis(n_components=min(nc-1, Z.shape[1]-1))\n        lda.fit(Z[tr_idx], all_labs[tr_idx])\n        p   = lda.predict(Z[te_idx])\n        f1_s = f1_score(all_labs[te_idx], p, average=\"macro\", zero_division=0)\n        acc_s = accuracy_score(all_labs[te_idx], p)\n    except Exception as e:\n        f1_s = 0.0; acc_s = 0.0\n    try:    sil_s = silhouette_score(Z, all_labs, sample_size=min(n,500))\n    except: sil_s = 0.0\n    results[name] = {\"F1\":f1_s, \"Accuracy\":acc_s, \"Silhouette\":sil_s}\n    print(f\"  {name:<18}  F1={f1_s:.3f}  Acc={acc_s:.3f}  Sil={sil_s:.3f}\")\n\ndf_res = pd.DataFrame(results).T.reset_index()\ndf_res.columns = [\"Strategy\",\"F1\",\"Accuracy\",\"Silhouette\"]\n\nfig, axes = plt.subplots(1,3, figsize=(15,5))\nfig.suptitle(\"Fusion Strategy Comparison\", fontsize=14, fontweight=\"bold\")\nfor ax, metric in zip(axes, [\"F1\",\"Accuracy\",\"Silhouette\"]):\n    bars = ax.bar(df_res[\"Strategy\"], df_res[metric],\n                  color=COLORS[:4], edgecolor=\"black\")\n    ax.set_title(f\"Linear Probe — {metric}\")\n    ax.set_ylabel(metric); ax.set_ylim(0,1.05)\n    ax.grid(True, alpha=0.3, axis=\"y\")\n    ax.tick_params(axis=\"x\", rotation=15)\n    for bar,v in zip(bars, df_res[metric]):\n        ax.text(bar.get_x()+bar.get_width()/2, bar.get_height()+0.01,\n                f\"{v:.3f}\", ha=\"center\", va=\"bottom\", fontsize=9)\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/fusion_comparison.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"Fusion comparison saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:51:31.358299Z","iopub.execute_input":"2026-03-07T08:51:31.358859Z","iopub.status.idle":"2026-03-07T08:51:39.644687Z","shell.execute_reply.started":"2026-03-07T08:51:31.358829Z","shell.execute_reply":"2026-03-07T08:51:39.643943Z"}},"outputs":[],"execution_count":null},{"id":"940dfe22-cba8-4bfc-bdef-a9d7f33af97d","cell_type":"markdown","source":"## Cell 19: Ablation Study\n","metadata":{}},{"id":"3d8846a9-efaf-49be-ae8d-40d472a4ced0","cell_type":"code","source":"print(\"Running ablation study …\")\nsystem.eval()\n\nconditions = [\"Clean\"] + [f\"Noisy {n}\" for n in VIEW_NAMES]\nview_keys  = [\"v1\",\"v2\",\"v3\",\"v4\"]\nablation   = {}\n\nfor cond_idx, cond_name in enumerate(conditions):\n    corrupt_view = cond_idx - 1   # -1 = no corruption\n    all_w, all_p, all_l = [], [], []\n\n    with torch.no_grad():\n        for views, labels in test_dl:\n            views_gpu = {k:v.to(DEVICE) for k,v in views.items()}\n            if corrupt_view >= 0:\n                key = view_keys[corrupt_view]\n                views_gpu = dict(views_gpu)\n                views_gpu[key] = torch.clamp(\n                    views_gpu[key] + torch.randn_like(views_gpu[key])*1.5, 0, 1)\n            out = system(views_gpu)\n            all_w.append(out[\"weights\"].cpu())\n            all_p.append(out[\"logits\"].argmax(1).cpu())\n            all_l.append(labels)\n\n    w  = torch.cat(all_w).numpy().mean(0)\n    p  = torch.cat(all_p).numpy()\n    l  = torch.cat(all_l).numpy()\n    f1_v = f1_score(l, p, average=\"macro\", zero_division=0)\n    ablation[cond_name] = {\"weights\":w, \"F1\":f1_v}\n    print(f\"  {cond_name:<22}  \"\n          f\"w=[{', '.join(f'{x:.3f}' for x in w)}]  F1={f1_v:.3f}\")\n\n# Plot\nfig, axes = plt.subplots(1, 2, figsize=(14,5))\nfig.suptitle(\"Ablation Study — Adaptive Fusion under View Corruption\",\n             fontsize=13, fontweight=\"bold\")\n\nwmat = np.array([ablation[c][\"weights\"] for c in conditions])\nim   = axes[0].imshow(wmat, aspect=\"auto\", cmap=\"YlOrRd\", vmin=0, vmax=0.6)\naxes[0].set_xticks(range(4)); axes[0].set_xticklabels(VIEW_NAMES, rotation=15, fontsize=8)\naxes[0].set_yticks(range(len(conditions))); axes[0].set_yticklabels(conditions, fontsize=8)\naxes[0].set_title(\"Mean Fusion Weights\")\nplt.colorbar(im, ax=axes[0])\nfor i in range(len(conditions)):\n    for j in range(4):\n        axes[0].text(j, i, f\"{wmat[i,j]:.2f}\", ha=\"center\", va=\"center\",\n                     fontsize=8, color=\"black\" if wmat[i,j]<0.4 else \"white\")\n\nf1s  = [ablation[c][\"F1\"] for c in conditions]\ncols = [\"#2A9D8F\"] + COLORS[:4]\nbars = axes[1].bar(conditions, f1s, color=cols, edgecolor=\"black\")\naxes[1].set_ylabel(\"F1 Score (macro)\"); axes[1].set_ylim(0,1.05)\naxes[1].set_title(\"Robustness to View Corruption\")\naxes[1].grid(True, alpha=0.3, axis=\"y\"); axes[1].tick_params(axis=\"x\", rotation=15)\nfor bar,v in zip(bars,f1s):\n    axes[1].text(bar.get_x()+bar.get_width()/2, bar.get_height()+0.01,\n                 f\"{v:.3f}\", ha=\"center\", va=\"bottom\", fontsize=9)\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/ablation.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"Ablation study saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:51:46.640062Z","iopub.execute_input":"2026-03-07T08:51:46.640393Z","iopub.status.idle":"2026-03-07T08:52:22.86347Z","shell.execute_reply.started":"2026-03-07T08:51:46.640366Z","shell.execute_reply":"2026-03-07T08:52:22.862749Z"}},"outputs":[],"execution_count":null},{"id":"8ee24ba0-4198-4aab-836e-0f4d862f21e6","cell_type":"markdown","source":"## Cell 20: Per-Class Metrics + Summary Table\n","metadata":{}},{"id":"51b831f6-f786-4c73-97bc-695227f004c3","cell_type":"code","source":"# Per-class metrics\nprec_pc = precision_score(labs, preds, average=None, zero_division=0)\nrec_pc  = recall_score(labs,  preds, average=None, zero_division=0)\nf1_pc   = f1_score(labs,      preds, average=None, zero_division=0)\n\ndf_pc = pd.DataFrame({\n    \"Class\":     CLASS_NAMES[:len(prec_pc)],\n    \"Precision\": prec_pc,\n    \"Recall\":    rec_pc,\n    \"F1\":        f1_pc,\n})\nprint(\"Per-Class Metrics:\")\nprint(df_pc.to_string(index=False))\n\nfig, axes = plt.subplots(1, 2, figsize=(16,5))\nfig.suptitle(\"Per-Class Metrics + Summary Table\", fontsize=14, fontweight=\"bold\")\n\nx = np.arange(len(df_pc)); bw = 0.25\naxes[0].bar(x-bw, df_pc[\"Precision\"], bw, label=\"Precision\", color=\"#457B9D\")\naxes[0].bar(x,    df_pc[\"Recall\"],    bw, label=\"Recall\",    color=\"#2A9D8F\")\naxes[0].bar(x+bw, df_pc[\"F1\"],        bw, label=\"F1\",        color=\"#E9C46A\")\naxes[0].set_xticks(x); axes[0].set_xticklabels(df_pc[\"Class\"])\naxes[0].set_ylabel(\"Score\"); axes[0].set_ylim(0,1.05)\naxes[0].set_title(\"Per-Class Classification Metrics\")\naxes[0].legend(); axes[0].grid(True, alpha=0.3, axis=\"y\")\n\n# Summary table\nrows = [(k, f\"{v:.4f}\") for k,v in METRICS.items()]\naxes[1].axis(\"off\")\ntbl = axes[1].table(cellText=rows, colLabels=[\"Metric\",\"Value\"],\n                     cellLoc=\"center\", loc=\"center\")\ntbl.auto_set_font_size(False); tbl.set_fontsize(9); tbl.scale(1.1, 1.4)\nfor i,key in enumerate(tbl._cells):\n    if key[0] % 2 == 0 and key[0] > 0:\n        tbl._cells[key].set_facecolor(\"#EAF4FB\")\naxes[1].set_title(\"All Evaluation Metrics\", fontweight=\"bold\")\n\nplt.tight_layout()\nplt.savefig(f\"{CFG['output_dir']}/per_class_metrics.png\", dpi=130, bbox_inches=\"tight\")\nplt.show()\nprint(\"Per-class metrics saved ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:54:42.087464Z","iopub.execute_input":"2026-03-07T08:54:42.087817Z","iopub.status.idle":"2026-03-07T08:54:42.781048Z","shell.execute_reply.started":"2026-03-07T08:54:42.087785Z","shell.execute_reply":"2026-03-07T08:54:42.780271Z"}},"outputs":[],"execution_count":null},{"id":"7c765d18-d493-4791-bf5b-1613399516d8","cell_type":"markdown","source":"## Cell 21: Save Checkpoint + Final Summary\n","metadata":{}},{"id":"62c8f8ed-fb29-4a60-8aea-8805ba18ef3b","cell_type":"code","source":"# Save model\ntorch.save({\n    \"model_state_dict\": system.state_dict(),\n    \"history\":          dict(history),\n    \"metrics\":          {k:float(v) for k,v in METRICS.items()},\n    \"cfg\":              CFG,\n}, f\"{CFG['output_dir']}/model_checkpoint.pth\")\n\n# Print summary\nOUT = CFG[\"output_dir\"]\nprint(\"\\n\" + \"=\"*55)\nprint(\"COMPLETE — ALL OUTPUTS SAVED\")\nprint(\"=\"*55)\nprint(f\"\\nOutput directory: {OUT}\")\nfor fname in sorted(os.listdir(OUT)):\n    fsize = os.path.getsize(os.path.join(OUT,fname))\n    print(f\"  {fname:<40} {fsize/1024:.1f} KB\")\n\nprint(\"\\n\" + \"=\"*55)\nprint(\"FINAL RESULTS SUMMARY\")\nprint(\"=\"*55)\nprint(f\"  Accuracy        : {METRICS['Accuracy']:.4f}\")\nprint(f\"  F1 (macro)      : {METRICS['F1 (macro)']:.4f}\")\nprint(f\"  AUC (OvR)       : {METRICS['AUC (OvR)']:.4f}\")\nprint(f\"  Precision       : {METRICS['Precision (macro)']:.4f}\")\nprint(f\"  Recall          : {METRICS['Recall (macro)']:.4f}\")\nprint(f\"  MSE             : {METRICS['MSE']:.6f}\")\nprint(f\"  RMSE            : {METRICS['RMSE']:.6f}\")\nprint(f\"  MAE             : {METRICS['MAE']:.6f}\")\nprint(f\"  MAPE (%)        : {METRICS['MAPE (%)']:.2f}\")\nprint(f\"  FDR             : {METRICS['Fisher Disc. Ratio']:.4f}\")\nprint(f\"  Silhouette (GT) : {METRICS['Silhouette (GT)']:.4f}\")\nprint(f\"  Silhouette (KM) : {METRICS['Silhouette (KMeans)']:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T08:54:52.937057Z","iopub.execute_input":"2026-03-07T08:54:52.937789Z","iopub.status.idle":"2026-03-07T08:54:53.079323Z","shell.execute_reply.started":"2026-03-07T08:54:52.93776Z","shell.execute_reply":"2026-03-07T08:54:53.078697Z"}},"outputs":[],"execution_count":null}]}