{"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# RSNA ICH – DATNet-style Graph + Knowledge Graph\n# - Triple-window CT → 3-channel EfficientNet-B0\n# - Stage 1: Fine-tune EfficientNet on balanced subset\n# - Stage 2: Extract embeddings → PCA → build KG graph\n# - Stage 3: DATNet encoder (graph propagation + residual MLP)\n# - Train/Val/Test metrics: Acc, Prec, Rec, F1, AUC, Kappa\n# ==========================================================\nimport os, time, random\nimport numpy as np, pandas as pd\nfrom tqdm import tqdm\nimport pydicom, cv2\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    accuracy_score, precision_score, recall_score,\n    f1_score, cohen_kappa_score, roc_auc_score\n)\nfrom sklearn.metrics.pairwise import cosine_similarity\nfrom sklearn.decomposition import PCA\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\n\n# ---------------- CONFIG ----------------\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\nBASE = \"/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection\"\nCSV_PATH = os.path.join(BASE, \"stage_2_train.csv\")\nTRAIN_DIR = os.path.join(BASE, \"stage_2_train\")\n\nSAMPLES_PER_CLASS = 1200\nIMG_SIZE = 256\nBATCH = 32\nEMB_BATCH = 32\n\nEPOCHS_EFF = 5\nEPOCHS_DATNET = 30\nWARMUP_EPOCHS = 5\n\nLR_EFF = 3e-4\nLR_DATNET = 3e-4\nWEIGHT_DECAY = 1e-5\n\nK_NEIGH = 15\nPCA_DIM = 256\nCONCEPT_EMB_DIM = 256\nLABEL_SMOOTH = 0.05\n\nALPHA_SUP = 1.0\nALPHA_CCA = 0.1\n\n# ---------------- Triple-window CT → 3-ch ----------------\ndef window_image(img, wl, ww):\n    minv = wl - ww/2.0\n    maxv = wl + ww/2.0\n    out = (img - minv) / (maxv - minv + 1e-9)\n    return np.clip(out, 0.0, 1.0)\n\ndef make_3ch_from_dcm(path):\n    dcm = pydicom.dcmread(path)\n    img = dcm.pixel_array.astype(np.float32)\n    ch1 = window_image(img, 40, 80)\n    ch2 = window_image(img, 80, 200)\n    ch3 = window_image(img, 600, 2800)\n    ch1 = cv2.resize(ch1, (IMG_SIZE, IMG_SIZE))\n    ch2 = cv2.resize(ch2, (IMG_SIZE, IMG_SIZE))\n    ch3 = cv2.resize(ch3, (IMG_SIZE, IMG_SIZE))\n    return np.stack([ch1, ch2, ch3], axis=0)\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])\n\ndf_group = df.groupby([\"Image\",\"Subtype\"], as_index=False)[\"Label\"].max()\ndf_pivot = df_group.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\nsamples = []\nfor lbl in [0,1]:\n    sub = df_pivot[df_pivot[\"Label_binary\"]==lbl]\n    n = min(len(sub), SAMPLES_PER_CLASS)\n    samples.append(sub.sample(n, random_state=SEED))\ndf_bal = pd.concat(samples).reset_index(drop=True)\n\nsubtype_cols = [c for c in df_bal.columns if c not in [\"Image\",\"Label_binary\"]]\nC = len(subtype_cols)\ndf_indexed = df_bal.set_index(\"Image\")\nB_full = df_bal[subtype_cols].values.astype(float)\n\n# ---------------- Dataset ----------------\nclass RSNAEffDataset(Dataset):\n    def __init__(self, df, img_dir, img_size=IMG_SIZE, augment=False):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.img_size = img_size\n        self.augment = augment\n        \n        self.aug_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.RandomResizedCrop(img_size, scale=(0.85,1.0)),\n            transforms.RandomHorizontalFlip(),\n            transforms.RandomRotation(10),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n        ])\n        self.eval_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((img_size, img_size)),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n        ])\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        img_id = self.df.loc[idx,\"Image\"]\n        label = int(self.df.loc[idx,\"Label_binary\"])\n        path = os.path.join(self.img_dir, f\"ID_{img_id}.dcm\")\n        img3 = make_3ch_from_dcm(path)\n        img3 = (img3*255).astype(np.uint8)\n        tf = self.aug_tf if self.augment else self.eval_tf\n        img_t = tf(np.transpose(img3, (1,2,0)))\n        return img_t, label, img_id\n\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 = RSNAEffDataset(train_df, TRAIN_DIR, augment=True)\nval_ds   = RSNAEffDataset(val_df, TRAIN_DIR, augment=False)\ntest_ds  = RSNAEffDataset(test_df, TRAIN_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# ---------------- EfficientNet Embeddings ----------------\ndef make_efficientnet_model(num_classes=2):\n    eff = models.efficientnet_b0(pretrained=True)\n    in_feat = eff.classifier[1].in_features\n    eff.classifier[1] = nn.Linear(in_feat, num_classes)\n    return eff\n\neffnet = make_efficientnet_model(2).to(DEVICE)\nfor _, p in effnet.named_parameters(): p.requires_grad=True\ncrit_eff = nn.CrossEntropyLoss(label_smoothing=LABEL_SMOOTH)\nopt_eff = torch.optim.AdamW(effnet.parameters(), lr=LR_EFF, weight_decay=WEIGHT_DECAY)\nsched_eff = torch.optim.lr_scheduler.ReduceLROnPlateau(opt_eff, mode='max', factor=0.5, patience=1, verbose=True)\n\nclass EffnetEmbedder(nn.Module):\n    def __init__(self, eff):\n        super().__init__()\n        self.features = eff.features\n        self.avgpool = eff.avgpool\n        self.dropout = eff.classifier[0]\n        self.feat_dim = eff.classifier[1].in_features\n    def forward(self, x):\n        x = self.features(x)\n        x = self.avgpool(x)\n        x = torch.flatten(x,1)\n        x = self.dropout(x)\n        return x\n\nembedder = EffnetEmbedder(effnet).to(DEVICE)\nembedder.eval()\n\ndef extract_embeddings_for_df(df_subset, batch_size=EMB_BATCH):\n    ds = RSNAEffDataset(df_subset, TRAIN_DIR, augment=False)\n    dl = DataLoader(ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True)\n    embs,labs,ids=[],[],[]\n    with torch.no_grad():\n        for imgs, labels, id_batch in tqdm(dl, desc=\"Extract embeddings\"):\n            imgs = imgs.to(DEVICE)\n            feat = embedder(imgs)\n            embs.append(feat.cpu().numpy())\n            labs.extend(labels.numpy().tolist())\n            ids.extend(id_batch)\n    X = np.vstack(embs); y = np.array(labs)\n    return X,y,ids\n\nX_tr, y_tr, ids_tr = extract_embeddings_for_df(train_df)\nX_val, y_val, ids_val = extract_embeddings_for_df(val_df)\nX_te,  y_te, ids_te  = extract_embeddings_for_df(test_df)\n\nif PCA_DIM is not None and PCA_DIM < X_tr.shape[1]:\n    from sklearn.decomposition import PCA\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# ---------------- Graph Construction ----------------\ndef construct_combined_graph(X_images, ids_images, k_neighbors=K_NEIGH):\n    N = X_images.shape[0]\n    sim = cosine_similarity(X_images)\n    np.fill_diagonal(sim,0)\n    sim = np.power(sim,3)\n    A_ii = np.zeros_like(sim)\n    for i in range(N):\n        nbrs = np.argsort(sim[i])[-k_neighbors:]\n        A_ii[i,nbrs] = sim[i,nbrs]\n    A_ii = np.minimum(A_ii, A_ii.T)\n    B_img = df_indexed.loc[ids_images, subtype_cols].values.astype(float)\n    row_sum = B_img.sum(axis=1,keepdims=True)\n    B_norm = B_img/(row_sum+1e-6); B_norm[row_sum.squeeze()==0]=0.0\n    A_ic = B_norm; A_ci = A_ic.T.copy()\n    p = B_full.mean(axis=0)+1e-9; co = B_full.T @ B_full\n    pmi = np.log((co+1e-6)/(p[:,None]*p[None,:])); pmi[pmi<0]=0\n    A_cc = pmi\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(axis=1)\n    D_inv = np.diag(1/np.sqrt(deg+1e-9))\n    A_norm = D_inv @ A @ D_inv\n    X_all = np.vstack([X_images, np.zeros((C,X_images.shape[1]),dtype=float)])\n    X_all_t = torch.tensor(X_all,dtype=torch.float32,device=DEVICE)\n    A_norm_t = torch.tensor(A_norm,dtype=torch.float32,device=DEVICE)\n    img_idx = np.arange(N)\n    concept_idx = np.arange(N,N+C)\n    labels_img = np.array([df_indexed.loc[iid,\"Label_binary\"] for iid in ids_images],dtype=int)\n    return X_all_t, A_norm_t, img_idx, concept_idx, labels_img\n\nX_all_tr, A_tr, img_idx_tr, concept_idx_tr, labels_tr = construct_combined_graph(X_tr, ids_tr)\nX_all_val, A_val, img_idx_val, concept_idx_val, labels_val = construct_combined_graph(X_val, ids_val)\nX_all_te,  A_te, img_idx_te, concept_idx_te, labels_te  = construct_combined_graph(X_te, ids_te)\n\n# ---------------- DATNet Encoder ----------------\nclass DATNet(nn.Module):\n    def __init__(self, in_feats, hidden=256, out_feats=2, num_concepts=C, concept_emb_dim=CONCEPT_EMB_DIM):\n        super().__init__()\n        self.num_concepts = num_concepts\n        self.concept_embed = nn.Embedding(num_concepts, concept_emb_dim)\n        nn.init.xavier_uniform_(self.concept_embed.weight)\n        self.concept_proj = nn.Linear(concept_emb_dim, in_feats)\n        self.fc1 = nn.Linear(in_feats, hidden)\n        self.fc2 = nn.Linear(hidden, hidden)\n        self.dropout = nn.Dropout(0.4)\n        self.classifier = nn.Linear(hidden, out_feats)\n\n    def encode(self,X_all,A_norm,concept_idx):\n        if concept_idx is not None and len(concept_idx)>0:\n            c_ids = torch.arange(len(concept_idx), device=DEVICE)\n            c_emb = self.concept_proj(self.concept_embed(c_ids))\n            X_all = X_all.clone()\n            X_all[concept_idx,:] = c_emb\n        h = F.relu(self.fc1(X_all))\n        h = A_norm @ h + h  # residual propagation\n        h = F.relu(self.fc2(h))\n        h = self.dropout(h)\n        return h\n\n    def forward(self,X_all,A_norm,concept_idx,img_idx):\n        H_all = self.encode(X_all,A_norm,concept_idx)\n        H_img = H_all[img_idx]\n        logits = self.classifier(H_img)\n        return logits,H_img,H_all\n\nmodel = DATNet(X_all_tr.shape[1]).to(DEVICE)\nopt_datnet = torch.optim.AdamW(model.parameters(), lr=LR_DATNET, weight_decay=WEIGHT_DECAY)\ncrit_sup = nn.CrossEntropyLoss(label_smoothing=LABEL_SMOOTH)\nsched_datnet = torch.optim.lr_scheduler.ReduceLROnPlateau(opt_datnet, mode='max', factor=0.5, patience=2, verbose=True)\n\nlabels_tr_t = torch.tensor(labels_tr,dtype=torch.long,device=DEVICE)\n\ndef cross_correlation_loss(Z1,Z2,lamb=5e-3):\n    N,d = Z1.shape\n    z1 = (Z1-Z1.mean(0))/(Z1.std(0)+1e-9)\n    z2 = (Z2-Z2.mean(0))/(Z2.std(0)+1e-9)\n    c = (z1.T @ z2)/N\n    on_diag = torch.diag(c).add_(-1).pow(2).sum()\n    off_diag = (c - torch.diag(torch.diag(c))).pow(2).sum()\n    return on_diag + lamb*off_diag\n\ndef eval_datnet(model,X_all,A_norm,img_idx,concept_idx,labels_img):\n    model.eval()\n    with torch.no_grad():\n        logits,H_img,_ = model(X_all,A_norm,concept_idx,img_idx)\n        preds = torch.argmax(logits,1).cpu().numpy()\n        probs = F.softmax(logits,dim=1)[:,1].cpu().numpy()\n        labels = labels_img\n        acc = accuracy_score(labels,preds)\n        prec = precision_score(labels,preds,zero_division=0)\n        rec = recall_score(labels,preds,zero_division=0)\n        f1 = f1_score(labels,preds,zero_division=0)\n        try: auc=roc_auc_score(labels,probs)\n        except: auc=float(\"nan\")\n        kappa = cohen_kappa_score(labels,preds)\n    return acc,prec,rec,f1,auc,kappa\n\n# ---------------- Stage 3: Train DATNet ----------------\nprint(\"\\n=== Stage 3: Training DATNet ===\")\nbest_val_acc = -1.0; best_state=None; pat=0\n\nfor ep in range(1,EPOCHS_DATNET+1):\n    t0=time.time()\n    model.train()\n    opt_datnet.zero_grad()\n    if ep<=WARMUP_EPOCHS:\n        X_view1 = X_all_tr\n        logits1,Z1_img,_ = model(X_view1,A_tr,concept_idx_tr,img_idx_tr)\n        loss_sup = crit_sup(logits1,labels_tr_t)\n        loss = loss_sup\n    else:\n        noise1 = 0.02*torch.randn_like(X_all_tr)\n        noise2 = 0.02*torch.randn_like(X_all_tr)\n        mask1 = (torch.rand_like(X_all_tr)>0.1).float()\n        mask2 = (torch.rand_like(X_all_tr)>0.1).float()\n        X_view1 = X_all_tr*mask1 + noise1\n        X_view2 = X_all_tr*mask2 + noise2\n        logits1,Z1_img,_ = model(X_view1,A_tr,concept_idx_tr,img_idx_tr)\n        _,Z2_img,_ = model(X_view2,A_tr,concept_idx_tr,img_idx_tr)\n        loss_sup = crit_sup(logits1,labels_tr_t)\n        loss_cca = cross_correlation_loss(Z1_img,Z2_img)\n        loss = ALPHA_SUP*loss_sup + ALPHA_CCA*loss_cca\n\n    loss.backward()\n    torch.nn.utils.clip_grad_norm_(model.parameters(),max_norm=5.0)\n    opt_datnet.step()\n\n    train_acc,train_prec,train_rec,train_f1,train_auc,train_kappa = eval_datnet(\n        model,X_all_tr,A_tr,img_idx_tr,concept_idx_tr,labels_tr\n    )\n    val_acc,val_prec,val_rec,val_f1,val_auc,val_kappa = eval_datnet(\n        model,X_all_val,A_val,img_idx_val,concept_idx_val,labels_val\n    )\n    test_acc,test_prec,test_rec,test_f1,test_auc,test_kappa = eval_datnet(\n        model,X_all_te,A_te,img_idx_te,concept_idx_te,labels_te\n    )\n\n    sched_datnet.step(val_acc)\n\n    print(f\"Ep {ep}/{EPOCHS_DATNET} | Loss:{loss.item():.5f} | \"\n          f\"TrAcc:{train_acc:.4f} ValAcc:{val_acc:.4f} TeAcc:{test_acc:.4f} | \"\n          f\"ValF1:{val_f1:.4f} ValAUC:{val_auc:.4f} | time:{time.time()-t0:.1f}s\")\n\n    if val_acc>best_val_acc:\n        best_val_acc=val_acc\n        best_state={\"model\":model.state_dict(),\"epoch\":ep}\n        pat=0\n    else:\n        pat+=1\n        if pat>=6:\n            print(\"Early stopping DATNet.\")\n            break\n\nif best_state is not None:\n    model.load_state_dict(best_state[\"model\"])\n    print(f\"Loaded best DATNet model from epoch {best_state['epoch']} (ValAcc={best_val_acc:.4f})\")\n\n# ---------------- Final Metrics ----------------\ntrain_acc,train_prec,train_rec,train_f1,train_auc,train_kappa = eval_datnet(\n    model,X_all_tr,A_tr,img_idx_tr,concept_idx_tr,labels_tr\n)\nval_acc,val_prec,val_rec,val_f1,val_auc,val_kappa = eval_datnet(\n    model,X_all_val,A_val,img_idx_val,concept_idx_val,labels_val\n)\ntest_acc,test_prec,test_rec,test_f1,test_auc,test_kappa = eval_datnet(\n    model,X_all_te,A_te,img_idx_te,concept_idx_te,labels_te\n)\n\nprint(\"\\n=== FINAL METRICS (DATNet + KG) ===\")\nprint(f\"Train Acc:{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   Acc:{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  Acc:{test_acc:.4f} | Prec:{test_prec:.4f} | Rec:{test_rec:.4f} | F1:{test_f1:.4f} | AUC:{test_auc:.4f} | Kappa:{test_kappa:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-11T05:20:26.603264Z","iopub.execute_input":"2025-12-11T05:20:26.604236Z","iopub.status.idle":"2025-12-11T05:24:14.888231Z","shell.execute_reply.started":"2025-12-11T05:20:26.604195Z","shell.execute_reply":"2025-12-11T05:24:14.886949Z"}},"outputs":[],"execution_count":null}]}