{"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":[{"sourceId":13451,"databundleVersionId":1188070,"sourceType":"competition"}],"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# GATE Model + Knowledge Graph (RSNA Intracranial Hemorrhage)\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, recall_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\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\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 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\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 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# -------------------- PROBE FINE-TUNE --------------------\nbest_val=-1; best_state=None\nfor ep in range(FT_EPOCHS):\n    resnet.train(); probe.train()\n    running_loss=0; preds=[]; labs=[]\n    for imgs, lab, _ in train_loader:\n        imgs, lab = imgs.to(DEVICE), 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\n    # validation\n    resnet.eval(); probe.eval()\n    v_preds=[]; v_labs=[]\n    with torch.no_grad():\n        for imgs, lab, _ in val_loader:\n            imgs, lab = imgs.to(DEVICE), 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: 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}\")\n\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\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 --------------------\ndf_indexed = df_pivot.set_index(\"Image\")\nconcept_feats = np.random.normal(0,0.01,(len(SUBTYPE_COLS), KG_EMB_DIM))\n\nclass ConstructGraph:\n    def __init__(self, X_img, img_ids, concept_feats):\n        self.X_img = X_img\n        self.img_ids = img_ids\n        self.concept_feats = concept_feats\n        self.N = X_img.shape[0]; self.C = concept_feats.shape[0]\n\n    def build_img_knn_weighted(self, k=K_NEIGH):\n        sim = cosine_similarity(self.X_img); np.fill_diagonal(sim,0)\n        A=np.zeros_like(sim)\n        for i in range(self.N):\n            idx = np.argsort(sim[i])[-k:]\n            A[i, idx] = sim[i, idx]\n        A = np.maximum(A, A.T) + np.eye(self.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\n    def build_combined_graph(self):\n        A_ii = self.build_img_knn_weighted()\n        A_ic = np.zeros((self.N, self.C))\n        for i,iid in enumerate(self.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(self.concept_feats); 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(self.N+self.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\n        # combine features\n        D_img = self.X_img.shape[1]; D_con = self.concept_feats.shape[1]\n        if D_con<D_img: con_padded = np.concatenate([self.concept_feats, np.zeros((self.C,D_img-D_con))],1)\n        else: con_padded = self.concept_feats[:,:D_img]\n        X_all = np.vstack([self.X_img, con_padded])\n\n        labels_img = np.array([df_indexed.loc[iid,\"Label_binary\"] if iid in df_indexed.index else 0 for iid in self.img_ids])\n        return X_all, A_norm, labels_img\n\n# -------------------- Build Graphs --------------------\ngraph_tr = ConstructGraph(X_tr, ids_tr, concept_feats)\nX_all_tr, A_tr, labels_tr = graph_tr.build_combined_graph()\n\ngraph_val = ConstructGraph(X_val, ids_val, concept_feats)\nX_all_val, A_val, labels_val = graph_val.build_combined_graph()\n\ngraph_te = ConstructGraph(X_te, ids_te, concept_feats)\nX_all_te, A_te, labels_te = graph_te.build_combined_graph()\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# -------------------- GATE 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 GATE_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# -------------------- TRAIN --------------------\nmodel = GATE_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\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    print(f\"Epoch {ep}/{EPOCHS} Loss:{loss.item():.4f} Time:{time.time()-t0:.1f}s\")\n\n# -------------------- FINAL EVALUATION --------------------\nmodel.eval()\nwith torch.no_grad():\n    t0 = time.time()\n    logits_te_all, _, _ = model(X_te_t, A_te_t)\n    preds = torch.argmax(logits_te_all[:N_img_te], dim=1).cpu().numpy()\n    probs = F.softmax(logits_te_all[:N_img_te], dim=1)[:,1].cpu().numpy()\n    acc = accuracy_score(labels_te, preds)\n    prec = precision_score(labels_te, preds, zero_division=0)\n    rec = recall_score(labels_te, preds, zero_division=0)\n    f1s = f1_score(labels_te, preds, zero_division=0)\n    auc = roc_auc_score(labels_te, probs)\n    kappa = cohen_kappa_score(labels_te, preds)\n\nprint(\"\\n=== FINAL METRICS ===\")\nprint(f\"Accuracy : {acc:.4f}\")\nprint(f\"ROC-AUC  : {auc:.4f}\")\nprint(f\"Precision: {prec:.4f}\")\nprint(f\"Recall   : {rec:.4f}\")\nprint(f\"F1 Score : {f1s:.4f}\")\nprint(f\"Cohen Kappa : {kappa:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-19T17:36:05.687289Z","iopub.execute_input":"2025-11-19T17:36:05.68764Z","iopub.status.idle":"2025-11-19T17:43:17.115339Z","shell.execute_reply.started":"2025-11-19T17:36:05.687618Z","shell.execute_reply":"2025-11-19T17:43:17.114356Z"}},"outputs":[],"execution_count":null}]}