{"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 – DGAT-v3 (MULTI-HEAD, BEST) + Knowledge Graph\n# SINGLE CELL – Stable (float32), Self-loops, Residual, LayerNorm,\n# Balanced loss, Early stop on Val F1, SAFE AUC\n# ================================================================\nimport os, random\nimport numpy as np\nimport pandas as pd\nimport cv2, pydicom\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    accuracy_score, precision_score, recall_score,\n    f1_score, roc_auc_score, cohen_kappa_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\n\nfrom torchvision import transforms\nfrom torchvision.models import efficientnet_b0\nfrom torchvision.models import EfficientNet_B0_Weights\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\")\nIMG_DIR  = os.path.join(BASE, \"stage_2_train\")\n\nIMG_SIZE = 256\nSAMPLES_PER_CLASS = 500         # increase if you have GPU\nBATCH = 16\nEMB_BATCH = 16\n\nPCA_DIM = 256\nK_NEIGH = 12\n\nEPOCHS_EFF = 2                  # short feature warmup\nEPOCHS_DGAT = 50\nPATIENCE = 8                    # early stop\n\nHEADS = 4                       # multi-head attention\nDROP = 0.25\nATT_TEMP = 0.7                  # smaller = sharper attention\n\n# ================================================================\n# 1) CT windowing → 3-channel\n# ================================================================\ndef window(img, wl, ww):\n    lo, hi = wl - ww/2, wl + ww/2\n    return np.clip((img - lo) / (hi - lo + 1e-9), 0, 1)\n\ndef load_dcm_3ch(path):\n    dcm = pydicom.dcmread(path)\n    img = dcm.pixel_array.astype(np.float32)\n    ch1 = cv2.resize(window(img, 40, 80), (IMG_SIZE, IMG_SIZE))\n    ch2 = cv2.resize(window(img, 80, 200), (IMG_SIZE, IMG_SIZE))\n    ch3 = cv2.resize(window(img, 600, 2800), (IMG_SIZE, IMG_SIZE))\n    return np.stack([ch1, ch2, ch3], 0)  # (3,H,W)\n\n# ================================================================\n# 2) Load CSV → pivot multi-label → binary → balanced subset\n# ================================================================\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_g = df.groupby([\"Image\",\"Subtype\"], as_index=False)[\"Label\"].max()\ndf_p = df_g.pivot(index=\"Image\", columns=\"Subtype\", values=\"Label\").fillna(0.0)\ndf_p[\"Label_binary\"] = df_p.max(axis=1).astype(int)\n\nsubtype_cols = list(df_p.columns[:-1])\nC = len(subtype_cols)\n\nparts = []\nfor lbl in [0, 1]:\n    sub = df_p[df_p[\"Label_binary\"] == lbl]\n    parts.append(sub.sample(min(SAMPLES_PER_CLASS, len(sub)), random_state=SEED))\ndf_bal = pd.concat(parts).reset_index()            # has Image column now\ndf_indexed = df_bal.set_index(\"Image\")\nB_full = df_bal[subtype_cols].values.astype(np.float32)\n\nprint(\"Balanced counts:\", df_bal[\"Label_binary\"].value_counts().to_dict())\nprint(\"Concepts (subtypes):\", subtype_cols)\n\n# ================================================================\n# 3) Dataset\n# ================================================================\nclass RSNADataset(Dataset):\n    def __init__(self, df_subset):\n        self.df = df_subset.reset_index(drop=True)\n        self.tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((IMG_SIZE, IMG_SIZE)),\n            transforms.ToTensor(),   # float32 [0..1]\n        ])\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        r = self.df.iloc[idx]\n        img_id = r[\"Image\"]\n        y = int(r[\"Label_binary\"])\n        path = os.path.join(IMG_DIR, f\"ID_{img_id}.dcm\")\n        img3 = load_dcm_3ch(path)                 # (3,H,W) float32 [0..1]\n        img3 = (img3 * 255).astype(np.uint8).transpose(1,2,0)  # (H,W,3) uint8\n        x = self.tf(img3)\n        return x, y, img_id\n\n# ================================================================\n# 4) Splits\n# ================================================================\ntrain_df, temp = train_test_split(\n    df_bal, test_size=0.3, stratify=df_bal[\"Label_binary\"], random_state=SEED\n)\nval_df, test_df = train_test_split(\n    temp, test_size=0.5, stratify=temp[\"Label_binary\"], random_state=SEED\n)\n\ntrain_loader = DataLoader(RSNADataset(train_df), batch_size=BATCH, shuffle=True, num_workers=2)\n\n# ================================================================\n# 5) EfficientNet (short) + Embeddings\n# ================================================================\neff = efficientnet_b0(weights=EfficientNet_B0_Weights.DEFAULT)\neff.classifier[1] = nn.Linear(eff.classifier[1].in_features, 2)\neff = eff.to(DEVICE)\n\nopt_eff = torch.optim.AdamW(eff.parameters(), lr=3e-4, weight_decay=1e-5)\ncrit_eff = nn.CrossEntropyLoss()\n\nprint(\"\\n=== Stage 1: EfficientNet short warmup ===\")\nfor ep in range(1, EPOCHS_EFF + 1):\n    eff.train()\n    for x, y, _ in train_loader:\n        x, y = x.to(DEVICE), y.to(DEVICE)\n        opt_eff.zero_grad()\n        loss = crit_eff(eff(x), y)\n        loss.backward()\n        opt_eff.step()\n    print(f\"Eff Ep {ep}/{EPOCHS_EFF} done\")\n\nclass Embed(nn.Module):\n    def __init__(self, m):\n        super().__init__()\n        self.f = m.features\n        self.a = m.avgpool\n        self.drop = m.classifier[0]\n    def forward(self, x):\n        x = self.f(x)\n        x = self.a(x)\n        x = torch.flatten(x, 1)\n        x = self.drop(x)\n        return x\n\nembedder = Embed(eff).to(DEVICE).eval()\n\ndef extract_embeddings(df_subset):\n    dl = DataLoader(RSNADataset(df_subset), batch_size=EMB_BATCH, shuffle=False, num_workers=2)\n    X, ids = [], []\n    with torch.no_grad():\n        for x, _, ids_b in dl:\n            z = embedder(x.to(DEVICE)).cpu().numpy()\n            X.append(z)\n            ids += list(ids_b)\n    return np.vstack(X).astype(np.float32), ids\n\nprint(\"\\n=== Stage 2: Extracting embeddings ===\")\nX_tr, ids_tr   = extract_embeddings(train_df)\nX_val, ids_val = extract_embeddings(val_df)\nX_te, ids_te   = extract_embeddings(test_df)\n\npca = PCA(PCA_DIM, random_state=SEED)\nX_tr = pca.fit_transform(X_tr).astype(np.float32)\nX_val = pca.transform(X_val).astype(np.float32)\nX_te  = pca.transform(X_te).astype(np.float32)\n\n# ================================================================\n# 6) Graph construction (Image-Image kNN + Image-Concept + PMI KG)\n#    + self-loops + symmetric norm  (ALL float32)\n# ================================================================\ndef build_graph(X, ids):\n    N = X.shape[0]\n\n    # image-image\n    sim = cosine_similarity(X).astype(np.float32)\n    np.fill_diagonal(sim, 0.0)\n    sim = sim ** 3\n\n    A_ii = np.zeros((N, N), dtype=np.float32)\n    for i in range(N):\n        k = min(K_NEIGH, N - 1)\n        if k <= 0: continue\n        nbrs = np.argsort(sim[i])[-k:]\n        A_ii[i, nbrs] = sim[i, nbrs]\n    A_ii = np.minimum(A_ii, A_ii.T)\n\n    # image-concept\n    B = df_indexed.loc[ids, subtype_cols].values.astype(np.float32)  # (N,C)\n    B = B / (B.sum(1, keepdims=True) + 1e-6)\n\n    # concept-concept PMI\n    p = B_full.mean(0) + 1e-9\n    co = (B_full.T @ B_full).astype(np.float32)\n    pmi = np.log((co + 1e-6) / (p[:, None] * p[None, :])).astype(np.float32)\n    pmi[pmi < 0] = 0.0\n\n    # combined adjacency\n    A = np.block([[A_ii, B],\n                  [B.T, pmi]]).astype(np.float32)\n\n    # self-loops (prevents isolated nodes -> prevents NaNs)\n    A = A + np.eye(A.shape[0], dtype=np.float32)\n\n    # symmetric normalization\n    deg = A.sum(1).astype(np.float32)\n    inv_sqrt = 1.0 / (np.sqrt(deg) + 1e-9)\n    A = (inv_sqrt[:, None] * A) * inv_sqrt[None, :]\n\n    # features: images have X, concepts start zeros (will be replaced by concept embeddings)\n    X_all = np.vstack([X, np.zeros((C, X.shape[1]), dtype=np.float32)]).astype(np.float32)\n\n    y = np.array([df_indexed.loc[i, \"Label_binary\"] for i in ids], dtype=int)\n    return (\n        torch.tensor(X_all, dtype=torch.float32, device=DEVICE),\n        torch.tensor(A, dtype=torch.float32, device=DEVICE),\n        np.arange(N),\n        np.arange(N, N + C),\n        y\n    )\n\nX_all_tr, A_tr, img_tr, con_tr, y_tr = build_graph(X_tr, ids_tr)\nX_all_val, A_val, img_val, con_val, y_val = build_graph(X_val, ids_val)\nX_all_te, A_te, img_te, con_te, y_te = build_graph(X_te, ids_te)\n\n# ================================================================\n# 7) DGAT-v3 Multi-Head Dense Attention Block\n#    - masked attention by adjacency\n#    - residual + layernorm\n# ================================================================\nclass MH_DGAT_Block(nn.Module):\n    def __init__(self, d, heads=4, dropout=0.25, att_temp=0.7):\n        super().__init__()\n        assert d % heads == 0, \"d must be divisible by heads\"\n        self.d = d\n        self.h = heads\n        self.dk = d // heads\n        self.att_temp = att_temp\n\n        self.Wq = nn.Linear(d, d, bias=False)\n        self.Wk = nn.Linear(d, d, bias=False)\n        self.Wv = nn.Linear(d, d, bias=False)\n        self.Wo = nn.Linear(d, d, bias=False)\n\n        self.ln = nn.LayerNorm(d)\n        self.drop = nn.Dropout(dropout)\n\n        # small FFN improves accuracy\n        self.ffn = nn.Sequential(\n            nn.Linear(d, 2*d),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(2*d, d),\n            nn.Dropout(dropout),\n        )\n        self.ln2 = nn.LayerNorm(d)\n\n    def forward(self, X, A_mask):\n        # X: (N,d) float32\n        N = X.size(0)\n\n        Q = self.Wq(X).view(N, self.h, self.dk).transpose(0, 1)  # (h,N,dk)\n        K = self.Wk(X).view(N, self.h, self.dk).transpose(0, 1)  # (h,N,dk)\n        V = self.Wv(X).view(N, self.h, self.dk).transpose(0, 1)  # (h,N,dk)\n\n        # scores: (h,N,N)\n        scores = (Q @ K.transpose(1, 2)) / (self.dk ** 0.5)\n        scores = scores / self.att_temp\n\n        # mask non-edges using adjacency (A_mask: (N,N), bool/float)\n        scores = scores.masked_fill(A_mask[None, :, :] == 0, float(\"-inf\"))\n\n        att = torch.softmax(scores, dim=2)  # along neighbors\n        att = self.drop(att)\n\n        H = att @ V  # (h,N,dk)\n        H = H.transpose(0, 1).contiguous().view(N, self.d)  # (N,d)\n        H = self.Wo(H)\n        H = self.drop(H)\n\n        # Residual + LN\n        X = self.ln(X + H)\n\n        # FFN + Residual + LN\n        X2 = self.ffn(X)\n        X = self.ln2(X + X2)\n        return X\n\nclass DGATv3(nn.Module):\n    def __init__(self, d, heads=4, dropout=0.25, att_temp=0.7):\n        super().__init__()\n        self.d = d\n\n        # learned concept embeddings (KG nodes)\n        self.ce = nn.Embedding(C, d)\n        nn.init.xavier_uniform_(self.ce.weight)\n\n        self.block1 = MH_DGAT_Block(d, heads=heads, dropout=dropout, att_temp=att_temp)\n        self.block2 = MH_DGAT_Block(d, heads=heads, dropout=dropout, att_temp=att_temp)\n\n        self.cls = nn.Linear(d, 2)\n\n    def forward(self, X_all, A, img_idx, con_idx):\n        X = X_all.clone().float()\n        A_mask = (A > 0).float()\n\n        # inject concepts\n        c_ids = torch.arange(C, device=X.device)\n        X[torch.tensor(con_idx, device=X.device)] = self.ce(c_ids)\n\n        X = self.block1(X, A_mask)\n        X = self.block2(X, A_mask)\n\n        img_t = torch.tensor(img_idx, device=X.device)\n        return self.cls(X[img_t])\n\nmodel = DGATv3(PCA_DIM, heads=HEADS, dropout=DROP, att_temp=ATT_TEMP).to(DEVICE)\n\n# balanced CE (helps recall & F1)\npos_w = (len(y_tr) - int(np.sum(y_tr))) / (int(np.sum(y_tr)) + 1e-6)\nw = torch.tensor([1.0, float(pos_w)], device=DEVICE, dtype=torch.float32)\ncrit = nn.CrossEntropyLoss(weight=w)\n\nopt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-5)\ny_tr_t = torch.tensor(y_tr, dtype=torch.long, device=DEVICE)\n\n# ================================================================\n# 8) Train (early stop on Val F1)\n# ================================================================\ndef val_f1():\n    model.eval()\n    with torch.no_grad():\n        p = model(X_all_val, A_val, img_val, con_val)\n        pr = torch.argmax(p, 1).cpu().numpy()\n    return f1_score(y_val, pr, zero_division=0)\n\nbest_f1 = -1.0\nbest_state = None\npat = 0\n\nprint(\"\\n=== Stage 3: Training DGAT-v3 (multi-head) ===\")\nfor ep in range(1, EPOCHS_DGAT + 1):\n    model.train()\n    opt.zero_grad()\n    logits = model(X_all_tr, A_tr, img_tr, con_tr)\n    loss = crit(logits, y_tr_t)\n    loss.backward()\n    torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n    opt.step()\n\n    f1v = val_f1()\n    print(f\"Ep {ep}/{EPOCHS_DGAT} | Loss:{loss.item():.4f} | ValF1:{f1v:.4f}\")\n\n    if f1v > best_f1:\n        best_f1 = f1v\n        best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n        pat = 0\n    else:\n        pat += 1\n        if pat >= PATIENCE:\n            print(\"Early stopping.\")\n            break\n\nif best_state is not None:\n    model.load_state_dict(best_state)\n    model.to(DEVICE)\n    print(f\"Loaded best DGAT-v3 (ValF1={best_f1:.4f})\")\n\n# ================================================================\n# 9) Final metrics (SAFE AUC)\n# ================================================================\ndef metrics(X_all, A, img_idx, con_idx, y):\n    model.eval()\n    with torch.no_grad():\n        p = model(X_all, A, img_idx, con_idx)\n        pr = torch.argmax(p, 1).cpu().numpy()\n        pb = torch.softmax(p, 1)[:, 1].detach().cpu().numpy()\n        pb = np.nan_to_num(pb, nan=0.0, posinf=1.0, neginf=0.0)\n\n    acc = accuracy_score(y, pr)\n    prec = precision_score(y, pr, zero_division=0)\n    rec = recall_score(y, pr, zero_division=0)\n    f1 = f1_score(y, pr, zero_division=0)\n    kappa = cohen_kappa_score(y, pr)\n\n    if len(np.unique(y)) < 2:\n        auc = float(\"nan\")\n    else:\n        try:\n            auc = roc_auc_score(y, pb)\n        except Exception:\n            auc = float(\"nan\")\n\n    return acc, prec, rec, f1, auc, kappa\n\nprint(\"\\n=== FINAL METRICS (DGAT-v3 Multi-Head + KG) ===\")\nprint(\"Train (Acc,Prec,Rec,F1,AUC,Kappa):\", metrics(X_all_tr, A_tr, img_tr, con_tr, y_tr))\nprint(\"Val   (Acc,Prec,Rec,F1,AUC,Kappa):\", metrics(X_all_val, A_val, img_val, con_val, y_val))\nprint(\"Test  (Acc,Prec,Rec,F1,AUC,Kappa):\", metrics(X_all_te, A_te, img_te, con_te, y_te))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-12T17:02:29.448121Z","iopub.execute_input":"2025-12-12T17:02:29.448984Z","iopub.status.idle":"2025-12-12T17:08:14.901788Z","shell.execute_reply.started":"2025-12-12T17:02:29.448949Z","shell.execute_reply":"2025-12-12T17:08:14.900698Z"}},"outputs":[],"execution_count":null}]}