{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":13451,"datasetId":654585,"databundleVersionId":1188070}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===============================================================\n# AGCL + Knowledge Graph (single-cell script) for RSNA Intracranial Hemorrhage\n# - ResNet18 encoder (probe fine-tune)\n# - Combined graph: image nodes + concept nodes (KG integrated)\n# - Adaptive adjacency (learnable gates)\n# - Contrastive loss (NT-Xent) + supervised CE loss\n# - Prints Train/Val/Test metrics each epoch\n# ===============================================================\nimport os, time, random\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport pydicom, cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, precision_score, f1_score, roc_auc_score, cohen_kappa_score\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom sklearn.decomposition import PCA\n\n# -------------------- CONFIG --------------------\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\nBASE_PATH = \"/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection\"\nCSV_PATH = os.path.join(BASE_PATH, \"stage_2_train.csv\")\nIMG_DIR = os.path.join(BASE_PATH, \"stage_2_train\")\n\nSAMPLES_PER_CLASS = 2000   # balanced small sample for fast run\nIMG_SIZE = 160\nBATCH = 64\nFT_EPOCHS = 2\nEPOCHS = 20\nLR_PROBE = 2e-4\nLR_AGCL = 1e-3\nK_NEIGH = 16\nPCA_DIM = 256\nKG_EMB_DIM = 16\nREBUILD_EVERY = 5\nTEMP = 0.2\nALPHA_CON = 1.0\nALPHA_SUP = 1.0\nSAVE_EMB = \"agcl_embs.npz\"\nKG_TRIPLES_CSV = \"kg_triples.csv\"\nUSE_KG = os.path.exists(KG_TRIPLES_CSV)\n\nSUBTYPE_COLS = ['any','epidural','intraparenchymal','intraventricular','subarachnoid','subdural']\n\n# -------------------- HELPERS --------------------\ndef set_seed(s=SEED):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s)\nset_seed()\n\ndef apply_brain_window(img, level=40, width=80):\n    low = level - width/2.0\n    high = level + width/2.0\n    img_w = np.clip(img, low, high)\n    img_w = (img_w - low) / (high - low + 1e-6)\n    return img_w.astype(np.float32)\n\ndef load_dicom(path):\n    d = pydicom.dcmread(path)\n    arr = d.pixel_array.astype(np.float32)\n    arr = apply_brain_window(arr)\n    return arr\n\ndef make_3ch(path):\n    arr = load_dicom(path)\n    arr = cv2.resize(arr, (IMG_SIZE, IMG_SIZE))\n    return np.stack([arr,arr,arr], axis=-1)  # pseudo-RGB\n\n# -------------------- LOAD & BALANCE LABELS --------------------\ndf = pd.read_csv(CSV_PATH)\ndf[\"Image\"] = df[\"ID\"].apply(lambda x: x.split(\"_\")[1])\ndf[\"Subtype\"] = df[\"ID\"].apply(lambda x: x.split(\"_\")[2])\ndf2 = df.groupby([\"Image\",\"Subtype\"], as_index=False)[\"Label\"].max()\ndf_pivot = df2.pivot(index=\"Image\", columns=\"Subtype\", values=\"Label\").reset_index().fillna(0)\ndf_pivot[\"Label_binary\"] = df_pivot.iloc[:,1:].max(axis=1).astype(int)\n\n# balanced sampling\npos = df_pivot[df_pivot[\"Label_binary\"]==1].sample(n=SAMPLES_PER_CLASS, random_state=SEED)\nneg = df_pivot[df_pivot[\"Label_binary\"]==0].sample(n=SAMPLES_PER_CLASS, random_state=SEED)\ndf_bal = pd.concat([pos,neg]).sample(frac=1, random_state=SEED).reset_index(drop=True)\n\n# -------------------- DATASET --------------------\nclass RSNADataset(Dataset):\n    def __init__(self, df, img_dir, augment=False):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.augment = augment\n        self.aug_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((IMG_SIZE,IMG_SIZE)),\n            transforms.RandomResizedCrop(IMG_SIZE, scale=(0.85,1.0)),\n            transforms.RandomHorizontalFlip(),\n            transforms.RandomRotation(10),\n            transforms.ToTensor()\n        ])\n        self.eval_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((IMG_SIZE,IMG_SIZE)),\n            transforms.ToTensor()\n        ])\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.img_dir, f\"ID_{row.Image}.dcm\")\n        img = make_3ch(img_path)\n        tf = self.aug_tf if self.augment else self.eval_tf\n        img_t = tf(img)\n        label = int(row.Label_binary)\n        return img_t.float(), label, row.Image\n\n# -------------------- SPLIT --------------------\ntrain_df, temp_df = train_test_split(df_bal, test_size=0.3, stratify=df_bal[\"Label_binary\"], random_state=SEED)\nval_df, test_df = train_test_split(temp_df, test_size=0.5, stratify=temp_df[\"Label_binary\"], random_state=SEED)\n\ntrain_ds = RSNADataset(train_df, IMG_DIR, augment=True)\nval_ds = RSNADataset(val_df, IMG_DIR, augment=False)\ntest_ds = RSNADataset(test_df, IMG_DIR, augment=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH, shuffle=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH, shuffle=False)\ntest_loader = DataLoader(test_ds, batch_size=BATCH, shuffle=False)\n\n# -------------------- ResNet18 backbone + probe --------------------\nresnet = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\nembedding_dim = resnet.fc.in_features\nresnet.fc = nn.Identity()\nresnet = resnet.to(DEVICE)\n\nprobe = nn.Linear(embedding_dim, 2).to(DEVICE)\n\n# freeze except layer4\nfor name,p in resnet.named_parameters():\n    p.requires_grad = False\n    if \"layer4\" in name:\n        p.requires_grad = True\nparams_ft = list(filter(lambda p: p.requires_grad, resnet.parameters())) + list(probe.parameters())\nopt_ft = torch.optim.AdamW(params_ft, lr=LR_PROBE)\ncrit_ce = nn.CrossEntropyLoss()\n\n# quick probe fine-tune\nbest_val = -1; best_state=None\nfor ep in range(FT_EPOCHS):\n    resnet.train(); probe.train()\n    preds=[]; labs=[]; running_loss=0\n    for imgs, lab, _ in train_loader:\n        imgs = imgs.to(DEVICE); lab=lab.to(DEVICE)\n        opt_ft.zero_grad()\n        feats = resnet(imgs)\n        logits = probe(feats)\n        loss = crit_ce(logits, lab)\n        loss.backward(); opt_ft.step()\n        running_loss += loss.item()*imgs.size(0)\n        preds.extend(torch.argmax(logits,1).cpu().numpy()); labs.extend(lab.cpu().numpy())\n    train_acc = accuracy_score(labs, preds)\n    # val\n    resnet.eval(); probe.eval()\n    v_preds=[]; v_labs=[]\n    with torch.no_grad():\n        for imgs, lab, _ in val_loader:\n            imgs = imgs.to(DEVICE); lab=lab.to(DEVICE)\n            logits = probe(resnet(imgs))\n            v_preds.extend(torch.argmax(logits,1).cpu().numpy())\n            v_labs.extend(lab.cpu().numpy())\n    val_acc = accuracy_score(v_labs, v_preds)\n    if val_acc>best_val:\n        best_val=val_acc; best_state=(resnet.state_dict(), probe.state_dict())\n    print(f\"FT Epoch {ep+1}/{FT_EPOCHS} Loss:{running_loss/len(train_ds):.4f} TrainAcc:{train_acc:.4f} ValAcc:{val_acc:.4f}\")\nif best_state:\n    resnet.load_state_dict(best_state[0]); probe.load_state_dict(best_state[1])\n\n# -------------------- Extract embeddings --------------------\ndef extract_embeddings(loader):\n    resnet.eval()\n    embs=[]; labs=[]; ids=[]\n    with torch.no_grad():\n        for imgs, lab, idlist in tqdm(loader):\n            imgs = imgs.to(DEVICE)\n            feat = resnet(imgs)\n            embs.append(feat.cpu().numpy())\n            labs.extend(lab.numpy())\n            ids.extend(idlist)\n    return np.vstack(embs), np.array(labs), ids\n\nX_tr, y_tr, ids_tr = extract_embeddings(train_loader)\nX_val, y_val, ids_val = extract_embeddings(val_loader)\nX_te, y_te, ids_te = extract_embeddings(test_loader)\n\n# PCA to reduce dim\nif PCA_DIM is not None and PCA_DIM < X_tr.shape[1]:\n    pca = PCA(n_components=PCA_DIM, random_state=SEED)\n    X_tr = pca.fit_transform(X_tr)\n    X_val = pca.transform(X_val)\n    X_te = pca.transform(X_te)\n\n# -------------------- Knowledge graph concept feats --------------------\ndf_indexed = df_pivot.set_index(\"Image\")\nconcept_feats = np.random.normal(0,0.01,(len(SUBTYPE_COLS), KG_EMB_DIM))\n\n# -------------------- Build combined image-concept graph --------------------\ndef build_img_knn_weighted(X, k=K_NEIGH):\n    sim = cosine_similarity(X); np.fill_diagonal(sim,0)\n    N = sim.shape[0]; A=np.zeros_like(sim)\n    for i in range(N):\n        idx = np.argsort(sim[i])[-k:]\n        A[i, idx] = sim[i, idx]\n    A = np.maximum(A, A.T) + np.eye(N)*1e-6\n    deg = A.sum(1); deg_inv_sqrt=1.0/np.sqrt(deg)\n    return (deg_inv_sqrt[:,None]*A)*deg_inv_sqrt[None,:]\n\ndef build_combined(X_img, img_ids, concept_feats_local):\n    A_ii = build_img_knn_weighted(X_img)\n    N=C=0\n    N=X_img.shape[0]; C=concept_feats_local.shape[0]\n    A_ic=np.zeros((N,C))\n    for i,iid in enumerate(img_ids):\n        if iid in df_indexed.index:\n            A_ic[i,:] = df_indexed.loc[iid,SUBTYPE_COLS].values\n    A_ci = A_ic.T\n    simc = cosine_similarity(concept_feats_local); np.fill_diagonal(simc,0)\n    A_cc = simc\n    top=np.concatenate([A_ii,A_ic],1); bottom=np.concatenate([A_ci,A_cc],1)\n    A = np.concatenate([top,bottom],0) + np.eye(N+C)*1e-6\n    deg = A.sum(1); deg_inv_sqrt=1.0/np.sqrt(deg)\n    A_norm = (deg_inv_sqrt[:,None]*A)*deg_inv_sqrt[None,:]\n    D_img=X_img.shape[1]; D_con=concept_feats_local.shape[1]\n    if D_con<D_img:\n        con_padded = np.concatenate([concept_feats_local,np.zeros((C,D_img-D_con))],1)\n    else: con_padded = concept_feats_local[:,:D_img]\n    X_all = np.vstack([X_img,con_padded])\n    labels_img = np.array([df_indexed.loc[iid,\"Label_binary\"] if iid in df_indexed.index else 0 for iid in img_ids])\n    return X_all, A_norm, labels_img\n\nX_all_tr, A_tr, labels_tr = build_combined(X_tr, ids_tr, concept_feats)\nX_all_val, A_val, labels_val = build_combined(X_val, ids_val, concept_feats)\nX_all_te,  A_te,  labels_te  = build_combined(X_te, ids_te, concept_feats)\n\n# torch tensors\nX_tr_t = torch.tensor(X_all_tr, dtype=torch.float32, device=DEVICE)\nA_tr_t = torch.tensor(A_tr, dtype=torch.float32, device=DEVICE)\ny_tr_img = torch.tensor(labels_tr, dtype=torch.long, device=DEVICE)\nX_val_t = torch.tensor(X_all_val, dtype=torch.float32, device=DEVICE)\nA_val_t = torch.tensor(A_val, dtype=torch.float32, device=DEVICE)\ny_val_img = torch.tensor(labels_val, dtype=torch.long, device=DEVICE)\nX_te_t = torch.tensor(X_all_te, dtype=torch.float32, device=DEVICE)\nA_te_t = torch.tensor(A_te, dtype=torch.float32, device=DEVICE)\nN_img_tr = X_tr.shape[0]; N_img_val = X_val.shape[0]; N_img_te = X_te.shape[0]\n\n# -------------------- Model --------------------\nclass SimpleGCNBlock(nn.Module):\n    def __init__(self, in_dim, hid_dim):\n        super().__init__()\n        self.lin1=nn.Linear(in_dim,hid_dim)\n        self.lin2=nn.Linear(hid_dim,hid_dim)\n        self.dropout=nn.Dropout(0.4)\n        self.bn=nn.LayerNorm(hid_dim)\n    def forward(self,X,A):\n        h=F.relu(self.bn(self.lin1(X)))\n        h=A@h\n        h=self.dropout(h)\n        h=F.relu(self.lin2(h))\n        h=A@h\n        return h\n\nclass AGCL_KG_Model(nn.Module):\n    def __init__(self, feat_dim,hid=512,n_classes=2):\n        super().__init__()\n        self.encoder=nn.Linear(feat_dim,hid)\n        self.gate_proj=nn.Linear(feat_dim,1)\n        self.gcn=SimpleGCNBlock(hid,hid)\n        self.classifier=nn.Linear(hid,n_classes)\n    def forward(self,X_all,A_all):\n        H = F.relu(self.encoder(X_all))\n        gate = torch.sigmoid(self.gate_proj(X_all)).squeeze(-1)\n        Hg = self.gcn(H,A_all)\n        logits = self.classifier(Hg)\n        return logits,H,gate\n\n# -------------------- Contrastive loss --------------------\ndef nt_xent_loss(Z1,Z2,temperature=TEMP):\n    N = Z1.shape[0]\n    z = torch.cat([Z1,Z2],0)\n    sim = (z@z.T)/temperature\n    sim_exp = torch.exp(sim - torch.max(sim,1,keepdim=True)[0])\n    mask = (~torch.eye(2*N, dtype=bool, device=Z1.device)).float()\n    denom = (sim_exp*mask).sum(1)\n    positives = torch.exp(torch.sum(Z1*Z2,1)/temperature)\n    positives = torch.cat([positives,positives],0)\n    loss = -torch.log(positives/denom)\n    return loss.mean()\n\n# -------------------- Training --------------------\nmodel = AGCL_KG_Model(X_all_tr.shape[1], hid=256, n_classes=2).to(DEVICE)\nopt = torch.optim.AdamW(model.parameters(), lr=LR_AGCL)\ncrit = nn.CrossEntropyLoss()\n\nbase_sim_tr = cosine_similarity(X_tr)\nbase_sim_val = cosine_similarity(X_val)\nbase_sim_te  = cosine_similarity(X_te)\n\nbest_val=-1; patience_ct=0\nA_tr_curr = A_tr.copy(); A_val_curr=A_val.copy(); A_te_curr=A_te.copy()\nA_tr_t = torch.tensor(A_tr_curr, dtype=torch.float32, device=DEVICE)\nA_val_t = torch.tensor(A_val_curr, dtype=torch.float32, device=DEVICE)\nA_te_t = torch.tensor(A_te_curr, dtype=torch.float32, device=DEVICE)\n\nfor ep in range(1,EPOCHS+1):\n    t0=time.time(); model.train(); opt.zero_grad()\n    logits_all,H_all,gate_all = model(X_tr_t,A_tr_t)\n    logits_img = logits_all[:N_img_tr]\n    loss_sup = crit(logits_img,y_tr_img)\n\n    # contrastive views\n    X_view1 = X_tr_t + 0.01*torch.randn_like(X_tr_t)\n    mask = (torch.rand_like(X_tr_t)>0.1).float()\n    X_view1 = X_view1*mask\n    X_view2 = X_tr_t*(1+0.02*torch.randn_like(X_tr_t))\n    model.eval()\n    with torch.no_grad():\n        Z1 = F.normalize(F.relu(model.encoder(X_view1)),dim=1)\n        Z2 = F.normalize(F.relu(model.encoder(X_view2)),dim=1)\n    model.train()\n    Z1_img = Z1[:N_img_tr]; Z2_img = Z2[:N_img_tr]\n    loss_con = nt_xent_loss(Z1_img,Z2_img)\n    loss = ALPHA_SUP*loss_sup + ALPHA_CON*loss_con\n    loss.backward(); opt.step()\n\n    # evaluation\n# ================================\n# FINAL BEST-EPOCH EVALUATION\n# ================================\n\n\n# Load the best model\nmodel.load_state_dict(best_state[\"model\"])\nmodel.eval()\n\nwith torch.no_grad():\n    logits_te_all, _, _ = model(X_te_t, A_te_t)\n\n# Only image nodes\npreds = torch.argmax(logits_te_all[:N_img_te], dim=1).cpu().numpy()\nprobs = F.softmax(logits_te_all[:N_img_te], dim=1)[:, 1].cpu().numpy()\n\n# Compute test metrics\nacc   = accuracy_score(labels_te, preds)\nprec  = precision_score(labels_te, preds, zero_division=0)\nrec   = recall_score(labels_te, preds, zero_division=0)\nf1s   = f1_score(labels_te, preds, zero_division=0)\nauc   = roc_auc_score(labels_te, probs)\nkappa = cohen_kappa_score(labels_te, preds)\n\n# Pretty print\nprint(\"=========== FINAL RESULTS ===========\")\nprint(f\" Best Epoch        : {best_state['epoch']}\")\nprint(f\" Val Acc (best)    : {best_state['val_acc']:.4f}\")\nprint(f\" Val F1  (best)    : {best_state['val_f1']:.4f}\")\nprint(f\" Val AUC (best)    : {best_state['val_auc']:.4f}\")\nprint(\"-------------------------------------\")\nprint(f\" Test Accuracy     : {acc:.4f}\")\nprint(f\" Test Precision    : {prec:.4f}\")\nprint(f\" Test Recall       : {rec:.4f}\")\nprint(f\" Test F1 Score     : {f1s:.4f}\")\nprint(f\" Test ROC-AUC      : {auc:.4f}\")\nprint(f\" Cohen Kappa       : {kappa:.4f}\")\nprint(\"=====================================\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T20:41:25.678279Z","iopub.execute_input":"2025-12-12T20:41:25.678639Z","iopub.status.idle":"2025-12-12T20:49:47.150887Z","shell.execute_reply.started":"2025-12-12T20:41:25.678608Z","shell.execute_reply":"2025-12-12T20:49:47.149433Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===============================================================\n# AGCL + Knowledge Graph (single-cell script) for RSNA Intracranial Hemorrhage\n# - ResNet18 encoder (probe fine-tune)\n# - Combined graph: image nodes + concept nodes (KG integrated)\n# - Adaptive adjacency (learnable gates)\n# - Contrastive loss (NT-Xent) + supervised CE loss\n# - Prints Train/Val/Test metrics each epoch\n# ===============================================================\nimport os, time, random\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport pydicom, cv2\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    accuracy_score, precision_score, recall_score, f1_score,\n    roc_auc_score, cohen_kappa_score\n)\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom sklearn.decomposition import PCA\n\n# -------------------- CONFIG --------------------\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\nBASE_PATH = \"/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection\"\nCSV_PATH = os.path.join(BASE_PATH, \"stage_2_train.csv\")\nIMG_DIR = os.path.join(BASE_PATH, \"stage_2_train\")\n\nSAMPLES_PER_CLASS = 2000   # balanced small sample for fast runs (adjust)\nIMG_SIZE = 160\nBATCH = 64\nFT_EPOCHS = 2\nEPOCHS = 20\nLR_PROBE = 2e-4\nLR_AGCL = 1e-3\nK_NEIGH = 16\nPCA_DIM = 256\nKG_EMB_DIM = 16\nREBUILD_EVERY = 5\nTEMP = 0.2\nALPHA_CON = 1.0\nALPHA_SUP = 1.0\n\nSUBTYPE_COLS = ['any','epidural','intraparenchymal','intraventricular','subarachnoid','subdural']\n\n# -------------------- HELPERS --------------------\ndef set_seed(s=SEED):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s)\nset_seed()\n\ndef apply_brain_window(img, level=40, width=80):\n    low = level - width/2.0\n    high = level + width/2.0\n    img_w = np.clip(img, low, high)\n    img_w = (img_w - low) / (high - low + 1e-6)\n    return img_w.astype(np.float32)\n\ndef load_dicom(path):\n    d = pydicom.dcmread(path)\n    arr = d.pixel_array.astype(np.float32)\n    arr = apply_brain_window(arr)\n    return arr\n\ndef make_3ch(path):\n    arr = load_dicom(path)\n    arr = cv2.resize(arr, (IMG_SIZE, IMG_SIZE))\n    # return pseudo-RGB 3-channel\n    return np.stack([arr, arr, arr], axis=-1)\n\n# -------------------- LOAD & BALANCE LABELS --------------------\ndf = pd.read_csv(CSV_PATH)\ndf[\"Image\"] = df[\"ID\"].apply(lambda x: x.split(\"_\")[1])\ndf[\"Subtype\"] = df[\"ID\"].apply(lambda x: x.split(\"_\")[2])\ndf2 = df.groupby([\"Image\",\"Subtype\"], as_index=False)[\"Label\"].max()\ndf_pivot = df2.pivot(index=\"Image\", columns=\"Subtype\", values=\"Label\").reset_index().fillna(0)\ndf_pivot[\"Label_binary\"] = df_pivot.iloc[:,1:].max(axis=1).astype(int)\n\n# safe check\nn_pos = df_pivot[\"Label_binary\"].sum()\nn_neg = (df_pivot[\"Label_binary\"]==0).sum()\nprint(\"Total images available:\", len(df_pivot), \"Pos:\", n_pos, \"Neg:\", n_neg)\n\n# balanced sampling (if not enough, reduce SAMPLES_PER_CLASS)\nn_pos_take = min(int(SAMPLES_PER_CLASS), int(n_pos))\nn_neg_take = min(int(SAMPLES_PER_CLASS), int(n_neg))\nif n_pos_take == 0 or n_neg_take == 0:\n    raise RuntimeError(\"Not enough positive or negative samples available — reduce SAMPLES_PER_CLASS or check CSV.\")\n\npos = df_pivot[df_pivot[\"Label_binary\"]==1].sample(n=n_pos_take, random_state=SEED)\nneg = df_pivot[df_pivot[\"Label_binary\"]==0].sample(n=n_neg_take, random_state=SEED)\ndf_bal = pd.concat([pos,neg]).sample(frac=1, random_state=SEED).reset_index(drop=True)\nprint(\"Balanced samples:\", df_bal[\"Label_binary\"].value_counts().to_dict())\n\n# -------------------- DATASET --------------------\nclass RSNADataset(Dataset):\n    def __init__(self, df, img_dir, augment=False):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.augment = augment\n        self.aug_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((IMG_SIZE,IMG_SIZE)),\n            transforms.RandomResizedCrop(IMG_SIZE, scale=(0.85,1.0)),\n            transforms.RandomHorizontalFlip(),\n            transforms.RandomRotation(10),\n            transforms.ToTensor()\n        ])\n        self.eval_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((IMG_SIZE,IMG_SIZE)),\n            transforms.ToTensor()\n        ])\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.img_dir, f\"ID_{row.Image}.dcm\")\n        img = make_3ch(img_path)\n        tf = self.aug_tf if self.augment else self.eval_tf\n        img_t = tf(img)\n        label = int(row.Label_binary)\n        return img_t.float(), label, row.Image\n\n# -------------------- SPLIT --------------------\ntrain_df, temp_df = train_test_split(df_bal, test_size=0.3, stratify=df_bal[\"Label_binary\"], random_state=SEED)\nval_df, test_df = train_test_split(temp_df, test_size=0.5, stratify=temp_df[\"Label_binary\"], random_state=SEED)\n\ntrain_ds = RSNADataset(train_df, IMG_DIR, augment=True)\nval_ds = RSNADataset(val_df, IMG_DIR, augment=False)\ntest_ds = RSNADataset(test_df, IMG_DIR, augment=False)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH, shuffle=True, num_workers=2, pin_memory=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH, shuffle=False, num_workers=2, pin_memory=True)\ntest_loader = DataLoader(test_ds, batch_size=BATCH, shuffle=False, num_workers=2, pin_memory=True)\n\n# -------------------- ResNet18 backbone + probe --------------------\nresnet = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\nembedding_dim = resnet.fc.in_features\nresnet.fc = nn.Identity()\nresnet = resnet.to(DEVICE)\n\nprobe = nn.Linear(embedding_dim, 2).to(DEVICE)\n\n# freeze except layer4\nfor name,p in resnet.named_parameters():\n    p.requires_grad = False\n    if \"layer4\" in name:\n        p.requires_grad = True\nparams_ft = list(filter(lambda p: p.requires_grad, resnet.parameters())) + list(probe.parameters())\nopt_ft = torch.optim.AdamW(params_ft, lr=LR_PROBE)\ncrit_ce = nn.CrossEntropyLoss()\n\n# quick probe fine-tune\nbest_val = -1.0\nbest_probe_state = None\nprint(\"\\n=== Probe fine-tune ===\")\nfor ep in range(FT_EPOCHS):\n    resnet.train(); probe.train()\n    preds=[]; labs=[]; running_loss=0.0\n    for imgs, lab, _ in train_loader:\n        imgs = imgs.to(DEVICE); lab = lab.to(DEVICE)\n        opt_ft.zero_grad()\n        feats = resnet(imgs)\n        logits = probe(feats)\n        loss = crit_ce(logits, lab)\n        loss.backward(); opt_ft.step()\n        running_loss += loss.item()*imgs.size(0)\n        preds.extend(torch.argmax(logits,1).cpu().numpy()); labs.extend(lab.cpu().numpy())\n    train_acc = accuracy_score(labs, preds)\n    # val\n    resnet.eval(); probe.eval()\n    v_preds=[]; v_labs=[]\n    with torch.no_grad():\n        for imgs, lab, _ in val_loader:\n            imgs = imgs.to(DEVICE); lab = lab.to(DEVICE)\n            logits = probe(resnet(imgs))\n            v_preds.extend(torch.argmax(logits,1).cpu().numpy())\n            v_labs.extend(lab.cpu().numpy())\n    val_acc = accuracy_score(v_labs, v_preds)\n    if val_acc > best_val:\n        best_val = val_acc\n        best_probe_state = {\"resnet\": resnet.state_dict(), \"probe\": probe.state_dict(), \"epoch\": ep+1, \"val_acc\": val_acc}\n    print(f\"FT Epoch {ep+1}/{FT_EPOCHS} Loss:{running_loss/len(train_ds):.4f} TrainAcc:{train_acc:.4f} ValAcc:{val_acc:.4f}\")\nif best_probe_state is not None:\n    resnet.load_state_dict(best_probe_state[\"resnet\"])\n    probe.load_state_dict(best_probe_state[\"probe\"])\n    print(\"Loaded best probe checkpoint (epoch\", best_probe_state[\"epoch\"], \"val_acc\", best_probe_state[\"val_acc\"], \")\")\n\n# -------------------- Extract embeddings --------------------\ndef extract_embeddings(loader):\n    resnet.eval()\n    embs=[]; labs=[]; ids=[]\n    with torch.no_grad():\n        for imgs, lab, idlist in tqdm(loader, desc=\"Extracting embeddings\"):\n            imgs = imgs.to(DEVICE)\n            feat = resnet(imgs)\n            embs.append(feat.cpu().numpy())\n            labs.extend(lab.numpy())\n            ids.extend(idlist)\n    return np.vstack(embs), np.array(labs), ids\n\nprint(\"\\n=== Extract embeddings ===\")\nX_tr, y_tr, ids_tr = extract_embeddings(train_loader)\nX_val, y_val, ids_val = extract_embeddings(val_loader)\nX_te, y_te, ids_te = extract_embeddings(test_loader)\n\n# PCA to reduce dim (optional)\nif PCA_DIM is not None and PCA_DIM < X_tr.shape[1]:\n    pca = PCA(n_components=PCA_DIM, random_state=SEED)\n    X_tr = pca.fit_transform(X_tr)\n    X_val = pca.transform(X_val)\n    X_te = pca.transform(X_te)\n\n# -------------------- Knowledge graph concept feats --------------------\ndf_indexed = df_pivot.set_index(\"Image\")\nconcept_feats = np.random.normal(0,0.01,(len(SUBTYPE_COLS), KG_EMB_DIM))\n\n# -------------------- Build combined image-concept graph --------------------\ndef build_img_knn_weighted(X, k=K_NEIGH):\n    sim = cosine_similarity(X); np.fill_diagonal(sim,0)\n    N = sim.shape[0]; A=np.zeros_like(sim)\n    for i in range(N):\n        idx = np.argsort(sim[i])[-k:]\n        A[i, idx] = sim[i, idx]\n    A = np.maximum(A, A.T) + np.eye(N)*1e-6\n    deg = A.sum(1); deg_inv_sqrt=1.0/np.sqrt(deg + 1e-12)\n    return (deg_inv_sqrt[:,None]*A)*deg_inv_sqrt[None,:]\n\ndef build_combined(X_img, img_ids, concept_feats_local):\n    A_ii = build_img_knn_weighted(X_img)\n    N = X_img.shape[0]; C = concept_feats_local.shape[0]\n    A_ic = np.zeros((N,C))\n    for i,iid in enumerate(img_ids):\n        if iid in df_indexed.index:\n            A_ic[i,:] = df_indexed.loc[iid, SUBTYPE_COLS].values\n    A_ci = A_ic.T\n    simc = cosine_similarity(concept_feats_local); np.fill_diagonal(simc,0)\n    A_cc = simc\n    top = np.concatenate([A_ii, A_ic], axis=1)\n    bottom = np.concatenate([A_ci, A_cc], axis=1)\n    A = np.concatenate([top, bottom], axis=0)\n    A = A + np.eye(A.shape[0]) * 1e-6\n    deg = A.sum(1); deg_inv_sqrt = 1.0/np.sqrt(deg + 1e-12)\n    A_norm = (deg_inv_sqrt[:,None] * A) * deg_inv_sqrt[None,:]\n    D_img = X_img.shape[1]; D_con = concept_feats_local.shape[1]\n    if D_con < D_img:\n        con_padded = np.concatenate([concept_feats_local, np.zeros((C, D_img - D_con))], axis=1)\n    else:\n        con_padded = concept_feats_local[:, :D_img]\n    X_all = np.vstack([X_img, con_padded])\n    labels_img = np.array([df_indexed.loc[iid, \"Label_binary\"] if iid in df_indexed.index else 0 for iid in img_ids], dtype=int)\n    return X_all, A_norm, labels_img\n\nX_all_tr, A_tr, labels_tr = build_combined(X_tr, ids_tr, concept_feats)\nX_all_val, A_val, labels_val = build_combined(X_val, ids_val, concept_feats)\nX_all_te,  A_te,  labels_te  = build_combined(X_te, ids_te, concept_feats)\n\n# torch tensors for AGCL training (full combined graphs)\nX_tr_t = torch.tensor(X_all_tr, dtype=torch.float32, device=DEVICE)\nA_tr_t = torch.tensor(A_tr, dtype=torch.float32, device=DEVICE)\ny_tr_img = torch.tensor(labels_tr, dtype=torch.long, device=DEVICE)\nX_val_t = torch.tensor(X_all_val, dtype=torch.float32, device=DEVICE)\nA_val_t = torch.tensor(A_val, dtype=torch.float32, device=DEVICE)\ny_val_img = torch.tensor(labels_val, dtype=torch.long, device=DEVICE)\nX_te_t = torch.tensor(X_all_te, dtype=torch.float32, device=DEVICE)\nA_te_t = torch.tensor(A_te, dtype=torch.float32, device=DEVICE)\nlabels_te_np = labels_te  # keep numpy for sklearn metrics\nN_img_tr = X_tr.shape[0]; N_img_val = X_val.shape[0]; N_img_te = X_te.shape[0]\n\n# -------------------- Model --------------------\nclass SimpleGCNBlock(nn.Module):\n    def __init__(self, in_dim, hid_dim):\n        super().__init__()\n        self.lin1=nn.Linear(in_dim,hid_dim)\n        self.lin2=nn.Linear(hid_dim,hid_dim)\n        self.dropout=nn.Dropout(0.4)\n        self.bn=nn.LayerNorm(hid_dim)\n    def forward(self,X,A):\n        h=F.relu(self.bn(self.lin1(X)))\n        h=A@h\n        h=self.dropout(h)\n        h=F.relu(self.lin2(h))\n        h=A@h\n        return h\n\nclass AGCL_KG_Model(nn.Module):\n    def __init__(self, feat_dim,hid=512,n_classes=2):\n        super().__init__()\n        self.encoder=nn.Linear(feat_dim,hid)\n        self.gate_proj=nn.Linear(feat_dim,1)\n        self.gcn=SimpleGCNBlock(hid,hid)\n        self.classifier=nn.Linear(hid,n_classes)\n    def forward(self,X_all,A_all):\n        H = F.relu(self.encoder(X_all))\n        gate = torch.sigmoid(self.gate_proj(X_all)).squeeze(-1)\n        Hg = self.gcn(H,A_all)\n        logits = self.classifier(Hg)\n        return logits,H,gate\n\n# -------------------- Contrastive loss (NT-Xent, simple) --------------------\ndef nt_xent_loss(Z1,Z2,temperature=TEMP):\n    # Z1, Z2: (N, d) assumed normalized already\n    N = Z1.shape[0]\n    z = torch.cat([Z1, Z2], dim=0)  # (2N, d)\n    sim = (z @ z.T) / temperature\n    # numerical stability\n    sim_max, _ = torch.max(sim, dim=1, keepdim=True)\n    sim_exp = torch.exp(sim - sim_max.detach())\n    mask = (~torch.eye(2*N, dtype=bool, device=Z1.device)).float()\n    denom = (sim_exp * mask).sum(dim=1)\n    positives = torch.exp((torch.sum(Z1*Z2, dim=1)) / temperature)\n    positives = torch.cat([positives, positives], dim=0)\n    loss = -torch.log(positives / denom)\n    return loss.mean()\n\n# -------------------- Eval helper --------------------\ndef eval_model(model, X_all_t, A_t, N_img, labels_np):\n    model.eval()\n    with torch.no_grad():\n        logits_all, _, _ = model(X_all_t, A_t)\n        logits_img = logits_all[:N_img]\n        probs = F.softmax(logits_img, dim=1)[:,1].cpu().numpy()\n        preds = torch.argmax(logits_img, dim=1).cpu().numpy()\n    acc = accuracy_score(labels_np, preds)\n    prec = precision_score(labels_np, preds, zero_division=0)\n    rec = recall_score(labels_np, preds, zero_division=0)\n    f1s = f1_score(labels_np, preds, zero_division=0)\n    try:\n        auc = roc_auc_score(labels_np, probs)\n    except Exception:\n        auc = float(\"nan\")\n    kappa = cohen_kappa_score(labels_np, preds)\n    return acc, prec, rec, f1s, auc, kappa\n\n# -------------------- Training --------------------\nmodel = AGCL_KG_Model(X_all_tr.shape[1], hid=256, n_classes=2).to(DEVICE)\nopt = torch.optim.AdamW(model.parameters(), lr=LR_AGCL, weight_decay=1e-5)\ncrit = nn.CrossEntropyLoss()\n\nbest_state = None\nbest_val_acc = -1.0\npatience = 0\nPATIENCE_MAX = 6\n\nprint(\"\\n=== AGCL training ===\")\nfor ep in range(1, EPOCHS+1):\n    t0 = time.time()\n    model.train(); opt.zero_grad()\n    logits_all, H_all, gate_all = model(X_tr_t, A_tr_t)  # full combined graph\n    logits_img = logits_all[:N_img_tr]\n    loss_sup = crit(logits_img, y_tr_img)\n\n    # contrastive views (image-node embeddings only)\n    # create two noisy views of X_tr_t (image node rows only, then pad concept rows unchanged)\n    # but simpler: apply noise to entire X_tr_t; we then take encoded img rows\n    X_view1 = X_tr_t + 0.01 * torch.randn_like(X_tr_t)\n    mask = (torch.rand_like(X_tr_t) > 0.1).float()\n    X_view1 = X_view1 * mask\n    X_view2 = X_tr_t * (1 + 0.02 * torch.randn_like(X_tr_t))\n\n    model.eval()\n    with torch.no_grad():\n        Z1 = F.normalize(F.relu(model.encoder(X_view1)), dim=1)\n        Z2 = F.normalize(F.relu(model.encoder(X_view2)), dim=1)\n    model.train()\n    Z1_img = Z1[:N_img_tr]\n    Z2_img = Z2[:N_img_tr]\n    loss_con = nt_xent_loss(Z1_img, Z2_img)\n\n    loss = ALPHA_SUP * loss_sup + ALPHA_CON * loss_con\n    opt.zero_grad(); loss.backward(); opt.step()\n\n    # validation & test metrics\n    train_acc, train_prec, train_rec, train_f1, train_auc, train_kappa = eval_model(model, X_tr_t, A_tr_t, N_img_tr, labels_tr)\n    val_acc, val_prec, val_rec, val_f1, val_auc, val_kappa = eval_model(model, X_val_t, A_val_t, N_img_val, labels_val)\n    test_acc, test_prec, test_rec, test_f1, test_auc, test_kappa = eval_model(model, X_te_t, A_te_t, N_img_te, labels_te)\n\n    # update scheduler/patience simple\n    if val_acc > best_val_acc:\n        best_val_acc = val_acc\n        best_state = {\n            \"model\": model.state_dict(),\n            \"epoch\": ep,\n            \"val_acc\": val_acc,\n            \"val_f1\": val_f1,\n            \"val_auc\": val_auc\n        }\n        patience = 0\n    else:\n        patience += 1\n\n    print(f\"Ep {ep}/{EPOCHS} | Loss:{loss.item():.5f} | TrAcc:{train_acc:.4f} ValAcc:{val_acc:.4f} TeAcc:{test_acc:.4f} | ValF1:{val_f1:.4f} ValAUC:{val_auc:.4f} | time:{time.time()-t0:.1f}s\")\n\n    if patience >= PATIENCE_MAX:\n        print(\"Early stopping AGCL (patience reached).\")\n        break\n\n# ================================\n# FINAL BEST-EPOCH EVALUATION\n# ================================\nif best_state is None:\n    # fallback to current model\n    best_state = {\"model\": model.state_dict(), \"epoch\": ep, \"val_acc\": val_acc, \"val_f1\": val_f1, \"val_auc\": val_auc}\n\nmodel.load_state_dict(best_state[\"model\"])\nmodel.eval()\n\ntrain_acc, train_prec, train_rec, train_f1, train_auc, train_kappa = eval_model(model, X_tr_t, A_tr_t, N_img_tr, labels_tr)\nval_acc, val_prec, val_rec, val_f1, val_auc, val_kappa = eval_model(model, X_val_t, A_val_t, N_img_val, labels_val)\ntest_acc, test_prec, test_rec, test_f1, test_auc, test_kappa = eval_model(model, X_te_t, A_te_t, N_img_te, labels_te)\n\nprint(\"\\n=========== FINAL RESULTS ===========\")\nprint(f\" Best Epoch        : {best_state['epoch']}\")\nprint(f\" Val Acc (best)    : {best_state['val_acc']:.4f}\")\nprint(f\" Val F1  (best)    : {best_state['val_f1']:.4f}\")\nprint(f\" Val AUC (best)    : {best_state['val_auc']:.4f}\")\nprint(\"-------------------------------------\")\nprint(f\" Train Accuracy     : {train_acc:.4f} | Prec: {train_prec:.4f} | Rec: {train_rec:.4f} | F1: {train_f1:.4f} | AUC: {train_auc:.4f} | Kappa: {train_kappa:.4f}\")\nprint(f\" Val   Accuracy     : {val_acc:.4f} | Prec: {val_prec:.4f} | Rec: {val_rec:.4f} | F1: {val_f1:.4f} | AUC: {val_auc:.4f} | Kappa: {val_kappa:.4f}\")\nprint(f\" Test  Accuracy     : {test_acc:.4f} | Prec: {test_prec:.4f} | Rec: {test_rec:.4f} | F1: {test_f1:.4f} | AUC: {test_auc:.4f} | Kappa: {test_kappa:.4f}\")\nprint(\"=====================================\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T02:49:44.631872Z","iopub.execute_input":"2025-12-13T02:49:44.632186Z","iopub.status.idle":"2025-12-13T02:56:30.999344Z","shell.execute_reply.started":"2025-12-13T02:49:44.63216Z","shell.execute_reply":"2025-12-13T02:56:30.998283Z"}},"outputs":[],"execution_count":null}]}