{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"none","dataSources":[{"sourceType":"competition","sourceId":29653,"databundleVersionId":2420395}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# HYBRID GRAPH-RAG MRI FRAMEWORK\n# CNN + GNN + RETRIEVAL AUGMENTATION\n# ============================================================\n\n# =========================\n# INSTALL\n# =========================\n!pip install -q pydicom torch-geometric sentence-transformers faiss-cpu\n\n# =========================\n# IMPORTS\n# =========================\nimport os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.models as models\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\nfrom skimage.transform import resize\nfrom skimage.segmentation import slic\n\nfrom torch_geometric.data import Data, Batch\nfrom torch_geometric.nn import GATConv, global_mean_pool\n\nfrom sentence_transformers import SentenceTransformer\nimport faiss\n\n# =========================\n# CONFIG\n# =========================\nIMG_SIZE = 224\nNUM_SUPERPIXELS = 100\nEMBED_DIM = 128\nBATCH_SIZE = 8\nEPOCHS = 3\n\n# =========================\n# MRI DATASET\n# =========================\nclass MRIDataset(Dataset):\n    def __init__(self, df, data_dir):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        pid = str(row[\"BraTS21ID\"]).zfill(5)\n        path = os.path.join(self.data_dir, pid)\n\n        img = np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n\n        try:\n            flair_path = os.path.join(path, \"FLAIR\")\n            files = sorted([\n                f for f in os.listdir(flair_path)\n                if f.endswith(\".dcm\")\n            ])\n\n            if len(files) > 0:\n                f = files[len(files) // 2]\n                dcm = pydicom.dcmread(os.path.join(flair_path, f))\n\n                img = dcm.pixel_array.astype(np.float32)\n                img = resize(\n                    img,\n                    (IMG_SIZE, IMG_SIZE),\n                    preserve_range=True,\n                    anti_aliasing=True\n                ).astype(np.float32)\n\n                img = (img - np.mean(img)) / (np.std(img) + 1e-8)\n\n        except Exception:\n            pass\n\n        img3 = np.stack([img] * 3, axis=0).astype(np.float32)\n        label = np.float32(row[\"MGMT_value\"])\n\n        return (\n            torch.FloatTensor(img3),\n            torch.FloatTensor([label]),\n            torch.FloatTensor(img)\n        )\n\n# =========================\n# CNN FEATURE EXTRACTOR\n# =========================\nclass CNNEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        backbone = models.densenet121(weights=None)\n        self.features = backbone.features\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(1024, EMBED_DIM)\n\n    def forward(self, x):\n        x = self.features(x)\n        x = torch.relu(x)\n        x = self.pool(x)\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n        return x\n\n# =========================\n# GRAPH CONSTRUCTION\n# =========================\ndef build_graph(image):\n    image = np.asarray(image, dtype=np.float32)\n    image = np.squeeze(image)\n\n    if image.ndim != 2:\n        raise ValueError(f\"Expected 2D grayscale image, got shape {image.shape}\")\n\n    image = np.nan_to_num(image)\n\n    img_min = image.min()\n    img_max = image.max()\n\n    if img_max > img_min:\n        slic_image = (image - img_min) / (img_max - img_min)\n    else:\n        slic_image = np.zeros_like(image)\n\n    segments = slic(\n        slic_image,\n        n_segments=NUM_SUPERPIXELS,\n        compactness=10,\n        start_label=0,\n        channel_axis=None\n    )\n\n    node_features = []\n    centers = []\n\n    for seg_id in np.unique(segments):\n        mask = segments == seg_id\n        pixels = image[mask]\n\n        node_features.append([\n            float(pixels.mean()),\n            float(pixels.std()),\n            float(pixels.max())\n        ])\n\n        coords = np.argwhere(mask)\n        centers.append(coords.mean(axis=0))\n\n    edge_index = []\n\n    for i in range(len(centers)):\n        for j in range(i + 1, len(centers)):\n            dist = np.linalg.norm(centers[i] - centers[j])\n\n            if dist < 30:\n                edge_index.append([i, j])\n                edge_index.append([j, i])\n\n    if len(edge_index) == 0:\n        edge_index = torch.empty((2, 0), dtype=torch.long)\n    else:\n        edge_index = torch.LongTensor(edge_index).t().contiguous()\n\n    x = torch.FloatTensor(node_features)\n\n    return Data(x=x, edge_index=edge_index)\n\n# =========================\n# GRAPH NEURAL NETWORK\n# =========================\nclass GraphEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.gat1 = GATConv(3, 32)\n        self.gat2 = GATConv(32, EMBED_DIM)\n\n    def forward(self, data):\n        x = data.x\n        edge_index = data.edge_index\n\n        x = self.gat1(x, edge_index)\n        x = torch.relu(x)\n        x = self.gat2(x, edge_index)\n\n        if hasattr(data, \"batch\") and data.batch is not None:\n            batch = data.batch\n        else:\n            batch = torch.zeros(\n                x.size(0),\n                dtype=torch.long,\n                device=x.device\n            )\n\n        x = global_mean_pool(x, batch)\n        return x\n\n# =========================\n# MEDICAL RETRIEVAL\n# =========================\nclass MedicalRetriever:\n    def __init__(self):\n        self.embedder = SentenceTransformer(\"all-MiniLM-L6-v2\")\n\n        self.documents = [\n            \"MGMT methylated glioma often responds better to temozolomide.\",\n            \"FLAIR hyperintensity is associated with glioma edema.\",\n            \"Tumor heterogeneity is important in radiogenomics.\",\n            \"MRI texture patterns can indicate genomic mutations.\"\n        ]\n\n        embeddings = self.embedder.encode(\n            self.documents,\n            convert_to_numpy=True\n        ).astype(\"float32\")\n\n        self.index = faiss.IndexFlatL2(embeddings.shape[1])\n        self.index.add(embeddings)\n\n    def retrieve(self, query_text):\n        q = self.embedder.encode(\n            [query_text],\n            convert_to_numpy=True\n        ).astype(\"float32\")\n\n        _, I = self.index.search(q, 1)\n        return self.documents[I[0][0]]\n\n# =========================\n# FUSION MODEL\n# =========================\nclass HybridGraphRAG(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.cnn = CNNEncoder()\n        self.gnn = GraphEncoder()\n\n        self.text_embedder = SentenceTransformer(\"all-MiniLM-L6-v2\")\n        self.text_fc = nn.Linear(384, EMBED_DIM)\n\n        self.classifier = nn.Sequential(\n            nn.Linear(EMBED_DIM * 3, 128),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(128, 1)\n        )\n\n    def forward(self, img_tensor, graph_data, text):\n        cnn_feat = self.cnn(img_tensor)\n        gnn_feat = self.gnn(graph_data)\n\n        if isinstance(text, str):\n            text = [text] * img_tensor.size(0)\n\n        # Important fix:\n        # use NumPy output, then create a normal Torch tensor.\n        text_emb_np = self.text_embedder.encode(\n            text,\n            convert_to_numpy=True\n        )\n\n        text_emb = torch.tensor(\n            text_emb_np,\n            dtype=torch.float32,\n            device=img_tensor.device\n        )\n\n        text_feat = self.text_fc(text_emb)\n\n        fused = torch.cat(\n            [cnn_feat, gnn_feat, text_feat],\n            dim=1\n        )\n\n        out = self.classifier(fused)\n        return out\n\n# =========================\n# FIND DATASET PATH\n# =========================\ntrain_csv = glob.glob(\n    \"/kaggle/input/**/train_labels.csv\",\n    recursive=True\n)[0]\n\ntrain_dir = train_csv.replace(\"train_labels.csv\", \"train\")\n\nprint(train_csv)\nprint(train_dir)\n\n# =========================\n# LOAD DATA\n# =========================\ndf = pd.read_csv(train_csv)\n\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df[\"MGMT_value\"],\n    random_state=42\n)\n\ntrain_dataset = MRIDataset(train_df, train_dir)\nval_dataset = MRIDataset(val_df, train_dir)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False\n)\n\n# =========================\n# DEVICE\n# =========================\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"Device:\", device)\n\n# =========================\n# MODEL\n# =========================\nmodel = HybridGraphRAG().to(device)\n\noptimizer = optim.Adam(\n    model.parameters(),\n    lr=1e-4\n)\n\ncriterion = nn.BCEWithLogitsLoss()\n\nretriever = MedicalRetriever()\nretrieved_text = retriever.retrieve(\"glioma MRI MGMT\")\n\n# =========================\n# TRAINING\n# =========================\nfor epoch in range(EPOCHS):\n    print(f\"\\nEpoch {epoch + 1}\")\n\n    model.train()\n    running_loss = 0.0\n\n    for imgs, labels, raw_imgs in train_loader:\n        imgs = imgs.to(device)\n        labels = labels.to(device)\n\n        graphs = [\n            build_graph(raw_img.numpy())\n            for raw_img in raw_imgs\n        ]\n\n        graph_batch = Batch.from_data_list(graphs).to(device)\n        batch_texts = [retrieved_text] * imgs.size(0)\n\n        optimizer.zero_grad()\n\n        outputs = model(\n            imgs,\n            graph_batch,\n            batch_texts\n        )\n\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n\n    avg_loss = running_loss / len(train_loader)\n    print(\"Training loss:\", avg_loss)\n\n# =========================\n# VALIDATION\n# =========================\nmodel.eval()\n\npreds = []\nlabels_all = []\n\nwith torch.no_grad():\n    for imgs, labels, raw_imgs in val_loader:\n        imgs = imgs.to(device)\n\n        graphs = [\n            build_graph(raw_img.numpy())\n            for raw_img in raw_imgs\n        ]\n\n        graph_batch = Batch.from_data_list(graphs).to(device)\n        batch_texts = [retrieved_text] * imgs.size(0)\n\n        outputs = torch.sigmoid(\n            model(\n                imgs,\n                graph_batch,\n                batch_texts\n            )\n        )\n\n        preds.extend(outputs.cpu().numpy().flatten())\n        labels_all.extend(labels.numpy().flatten())\n\nauc = roc_auc_score(labels_all, preds)\nprint(\"\\nValidation AUC:\", auc)\n\n# =========================\n# SAVE MODEL\n# =========================\ntorch.save(\n    model.state_dict(),\n    \"hybrid_graph_rag_model.pth\"\n)\n\nprint(\"\\nDONE!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T21:46:01.450184Z","iopub.execute_input":"2026-05-08T21:46:01.450512Z","iopub.status.idle":"2026-05-08T21:56:58.853549Z","shell.execute_reply.started":"2026-05-08T21:46:01.450477Z","shell.execute_reply":"2026-05-08T21:56:58.852481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CARE-RULE HYBRID FRAMEWORK\n# ERROR-FREE KAGGLE VERSION\n# MRI + RULES + RETRIEVAL AUGMENTATION\n# ============================================================\n\n# =========================\n# INSTALL\n# =========================\n!pip install -q pydicom sentence-transformers faiss-cpu\n\n# =========================\n# IMPORTS\n# =========================\nimport os\nimport glob\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\n\nfrom skimage.transform import resize\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, roc_auc_score, f1_score\n\nimport torchvision.models as models\n\nfrom sentence_transformers import SentenceTransformer\nimport faiss\n\n# =========================\n# CONFIG\n# =========================\nIMG_SIZE = 224\nBATCH_SIZE = 4\nEPOCHS = 1\nEMBED_DIM = 256\n\n# =========================\n# FIND DATASET AUTOMATICALLY\n# =========================\ncsv_files = glob.glob(\n    \"/kaggle/input/**/train_labels.csv\",\n    recursive=True\n)\n\nif len(csv_files) == 0:\n    raise FileNotFoundError(\n        \"Dataset not found. Add RSNA-MICCAI dataset using Kaggle -> Add Input\"\n    )\n\ntrain_csv = csv_files[0]\ntrain_dir = train_csv.replace(\"train_labels.csv\", \"train\")\n\nprint(\"CSV FILE:\", train_csv)\nprint(\"TRAIN DIR:\", train_dir)\n\n# =========================\n# DATASET\n# =========================\nclass MRIDataset(Dataset):\n    def __init__(self, df, data_dir):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        pid = str(row[\"BraTS21ID\"]).zfill(5)\n        patient_path = os.path.join(self.data_dir, pid)\n\n        img = np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n\n        try:\n            flair_path = os.path.join(patient_path, \"FLAIR\")\n\n            files = sorted([\n                f for f in os.listdir(flair_path)\n                if f.endswith(\".dcm\")\n            ])\n\n            if len(files) > 0:\n                middle = files[len(files) // 2]\n\n                dcm = pydicom.dcmread(\n                    os.path.join(flair_path, middle)\n                )\n\n                img = dcm.pixel_array.astype(np.float32)\n\n                img = resize(\n                    img,\n                    (IMG_SIZE, IMG_SIZE),\n                    preserve_range=True,\n                    anti_aliasing=True\n                ).astype(np.float32)\n\n                img = (img - np.mean(img)) / (np.std(img) + 1e-8)\n\n        except Exception:\n            pass\n\n        img3 = np.stack([img] * 3, axis=0).astype(np.float32)\n        label = np.float32(row[\"MGMT_value\"])\n\n        return (\n            torch.FloatTensor(img3),\n            torch.FloatTensor([label])\n        )\n\n# =========================\n# IMAGE ENCODER\n# =========================\nclass CAREEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        backbone = models.efficientnet_b0(weights=None)\n\n        self.features = backbone.features\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(1280, EMBED_DIM)\n\n    def forward(self, x):\n        x = self.features(x)\n        x = self.pool(x)\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n        return x\n\n# =========================\n# TEXT ENCODER\n# =========================\nclass TextEncoder(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.net = nn.Sequential(\n            nn.Linear(384, EMBED_DIM),\n            nn.ReLU(),\n            nn.Dropout(0.2)\n        )\n\n    def forward(self, x):\n        return self.net(x)\n\n# =========================\n# RULE ENGINE\n# =========================\nclass RuleReasoner:\n    def apply_rules(self, batch_imgs):\n        scores = []\n\n        for i in range(batch_imgs.shape[0]):\n            img = batch_imgs[i][0].detach().cpu().numpy()\n\n            mean_intensity = img.mean()\n            std_intensity = img.std()\n\n            score = 0\n\n            if mean_intensity > 0.5:\n                score += 1\n\n            if std_intensity > 0.5:\n                score += 1\n\n            if mean_intensity + std_intensity > 1.0:\n                score += 1\n\n            scores.append([score / 3.0])\n\n        return torch.FloatTensor(scores)\n\n# =========================\n# RETRIEVAL MODULE\n# =========================\nclass MedicalRetriever:\n    def __init__(self):\n        self.embedder = SentenceTransformer(\n            \"all-MiniLM-L6-v2\",\n            device=\"cpu\"\n        )\n\n        self.knowledge = [\n            \"MGMT methylation improves treatment response.\",\n            \"FLAIR hyperintensity is linked with glioma edema.\",\n            \"MRI heterogeneity indicates tumor aggressiveness.\",\n            \"Radiogenomics predicts genomic mutations from MRI.\",\n            \"Tumor texture is important in glioma classification.\"\n        ]\n\n        embeddings = self.embedder.encode(\n            self.knowledge,\n            convert_to_numpy=True\n        ).astype(\"float32\")\n\n        self.index = faiss.IndexFlatL2(embeddings.shape[1])\n        self.index.add(embeddings)\n\n    def retrieve(self, query):\n        q = self.embedder.encode(\n            [query],\n            convert_to_numpy=True\n        ).astype(\"float32\")\n\n        _, I = self.index.search(q, 1)\n\n        return self.knowledge[I[0][0]]\n\n# =========================\n# HYBRID MODEL\n# =========================\nclass CARERULEModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.image_encoder = CAREEncoder()\n        self.text_encoder = TextEncoder()\n\n        self.fusion = nn.Sequential(\n            nn.Linear(EMBED_DIM * 2 + 1, 256),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, 64),\n            nn.ReLU(),\n            nn.Linear(64, 1)\n        )\n\n    def forward(self, images, text_embeddings, rule_scores):\n        image_feat = self.image_encoder(images)\n        text_feat = self.text_encoder(text_embeddings)\n\n        rule_scores = rule_scores.view(-1, 1)\n\n        fused = torch.cat(\n            [\n                image_feat,\n                text_feat,\n                rule_scores\n            ],\n            dim=1\n        )\n\n        output = self.fusion(fused)\n        return output\n\n# =========================\n# LOAD DATA\n# =========================\ndf = pd.read_csv(train_csv)\n\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df[\"MGMT_value\"],\n    random_state=42\n)\n\ntrain_dataset = MRIDataset(train_df, train_dir)\nval_dataset = MRIDataset(val_df, train_dir)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False\n)\n\n# =========================\n# DEVICE\n# =========================\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"Using Device:\", device)\n\n# =========================\n# INITIALIZE\n# =========================\nmodel = CARERULEModel().to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\n\noptimizer = optim.Adam(\n    model.parameters(),\n    lr=1e-4\n)\n\nrule_engine = RuleReasoner()\nretriever = MedicalRetriever()\n\nembedder = SentenceTransformer(\n    \"all-MiniLM-L6-v2\",\n    device=\"cpu\"\n)\n\n# =========================\n# RETRIEVE TEXT ONCE\n# =========================\nretrieved_text = retriever.retrieve(\n    \"MRI glioma MGMT FLAIR\"\n)\n\nprint(\"\\nRetrieved Medical Knowledge:\")\nprint(retrieved_text)\n\n# =========================\n# FIXED TEXT EMBEDDING\n# =========================\nretrieved_embedding_np = embedder.encode(\n    [retrieved_text],\n    convert_to_numpy=True\n).astype(\"float32\")\n\nretrieved_embedding = torch.from_numpy(\n    np.array(retrieved_embedding_np, copy=True)\n).clone().float().to(device)\n\n# =========================\n# TRAINING\n# =========================\nfor epoch in range(EPOCHS):\n    print(f\"\\nEpoch {epoch + 1}\")\n\n    model.train()\n    losses = []\n\n    for imgs, labels in train_loader:\n        imgs = imgs.to(device)\n        labels = labels.to(device)\n\n        text_embeddings = retrieved_embedding.repeat(\n            imgs.shape[0],\n            1\n        ).clone()\n\n        rule_scores = rule_engine.apply_rules(\n            imgs\n        ).to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(\n            imgs,\n            text_embeddings,\n            rule_scores\n        )\n\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        losses.append(loss.item())\n\n    print(\"Train Loss:\", np.mean(losses))\n\n# =========================\n# VALIDATION\n# =========================\nmodel.eval()\n\npreds = []\nlabels_all = []\n\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        imgs = imgs.to(device)\n\n        text_embeddings = retrieved_embedding.repeat(\n            imgs.shape[0],\n            1\n        ).clone()\n\n        rule_scores = rule_engine.apply_rules(\n            imgs\n        ).to(device)\n\n        outputs = torch.sigmoid(\n            model(\n                imgs,\n                text_embeddings,\n                rule_scores\n            )\n        )\n\n        preds.extend(\n            outputs.cpu().numpy().flatten()\n        )\n\n        labels_all.extend(\n            labels.numpy().flatten()\n        )\n\n# =========================\n# METRICS\n# =========================\npreds = np.array(preds)\nlabels_all = np.array(labels_all)\n\npreds_bin = (preds > 0.5).astype(int)\n\nauc = roc_auc_score(labels_all, preds)\nacc = accuracy_score(labels_all, preds_bin)\nf1 = f1_score(labels_all, preds_bin)\n\n# =========================\n# RESULTS\n# =========================\nprint(\"\\n======================\")\nprint(\"FINAL RESULTS\")\nprint(\"======================\")\n\nprint(\"AUC :\", round(auc, 4))\nprint(\"ACC :\", round(acc, 4))\nprint(\"F1  :\", round(f1, 4))\n\n# =========================\n# SAVE MODEL\n# =========================\ntorch.save(\n    model.state_dict(),\n    \"/kaggle/working/care_rule_model.pth\"\n)\n\nprint(\"\\nMODEL SAVED:\")\nprint(\"/kaggle/working/care_rule_model.pth\")\n\nprint(\"\\nDONE SUCCESSFULLY!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T22:15:55.991088Z","iopub.execute_input":"2026-05-08T22:15:55.991401Z","iopub.status.idle":"2026-05-08T22:25:36.979207Z","shell.execute_reply.started":"2026-05-08T22:15:55.991373Z","shell.execute_reply":"2026-05-08T22:25:36.978167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# HYBRID MRI FRAMEWORK\n# CNN + BrainGNN Hybrid\n# Generates Images + Metrics + Saved Model\n# FULL KAGGLE SAFE VERSION\n# ============================================================\n\n# =========================\n# INSTALL\n# =========================\n!pip install -q pydicom torch-geometric\n\n# =========================\n# IMPORTS\n# =========================\nimport os\nimport glob\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\n\nimport pydicom\n\nfrom skimage.transform import resize\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    roc_auc_score,\n    roc_curve,\n    confusion_matrix,\n    ConfusionMatrixDisplay\n)\n\nimport torchvision.models as models\n\nfrom torch_geometric.data import Data\nfrom torch_geometric.nn import GCNConv\nfrom torch_geometric.nn import global_mean_pool\n\n# =========================\n# CONFIG\n# =========================\nIMG_SIZE = 224\nBATCH_SIZE = 4\nEPOCHS = 1\n\nDEVICE = torch.device(\n    'cuda' if torch.cuda.is_available() else 'cpu'\n)\n\nprint(\"DEVICE:\", DEVICE)\n\n# =========================\n# FIND DATASET\n# =========================\ncsv_files = glob.glob(\n    '/kaggle/input/**/train_labels.csv',\n    recursive=True\n)\n\nif len(csv_files) == 0:\n    raise FileNotFoundError(\n        'Dataset not found.\\n'\n        'Add RSNA dataset in Kaggle.'\n    )\n\ntrain_csv = csv_files[0]\n\ntrain_dir = train_csv.replace(\n    'train_labels.csv',\n    'train'\n)\n\nprint(\"CSV:\", train_csv)\nprint(\"TRAIN DIR:\", train_dir)\n\n# =========================\n# DATASET\n# =========================\nclass MRIDataset(Dataset):\n\n    def __init__(self, df, data_dir):\n\n        self.df = df.reset_index(drop=True)\n\n        self.data_dir = data_dir\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        pid = str(row['BraTS21ID']).zfill(5)\n\n        patient_path = os.path.join(\n            self.data_dir,\n            pid\n        )\n\n        img = np.zeros(\n            (IMG_SIZE, IMG_SIZE),\n            dtype=np.float32\n        )\n\n        try:\n\n            flair_path = os.path.join(\n                patient_path,\n                'FLAIR'\n            )\n\n            files = sorted([\n                f for f in os.listdir(flair_path)\n                if f.endswith('.dcm')\n            ])\n\n            if len(files) > 0:\n\n                mid = files[len(files)//2]\n\n                dcm = pydicom.dcmread(\n                    os.path.join(flair_path, mid)\n                )\n\n                img = dcm.pixel_array.astype(np.float32)\n\n                img = resize(\n                    img,\n                    (IMG_SIZE, IMG_SIZE)\n                )\n\n                img = (\n                    img - np.mean(img)\n                ) / (\n                    np.std(img) + 1e-8\n                )\n\n        except:\n            pass\n\n        img3 = np.stack([img]*3, axis=0)\n\n        label = row['MGMT_value']\n\n        return (\n            torch.FloatTensor(img3),\n            torch.FloatTensor([label]),\n            torch.FloatTensor(img)\n        )\n\n# =========================\n# GRAPH CREATION\n# =========================\ndef create_graph(image_tensor):\n\n    image = image_tensor.numpy()\n\n    patches = []\n\n    step = 32\n\n    for i in range(0, IMG_SIZE, step):\n\n        for j in range(0, IMG_SIZE, step):\n\n            patch = image[\n                i:i+step,\n                j:j+step\n            ]\n\n            feat = [\n                patch.mean(),\n                patch.std()\n            ]\n\n            patches.append(feat)\n\n    x = torch.tensor(\n        patches,\n        dtype=torch.float\n    )\n\n    edge_index = []\n\n    num_nodes = len(patches)\n\n    for i in range(num_nodes-1):\n\n        edge_index.append([i, i+1])\n\n        edge_index.append([i+1, i])\n\n    # SAFE EDGE HANDLING\n    if len(edge_index) == 0:\n\n        edge_index = torch.tensor(\n            [[0],[0]],\n            dtype=torch.long\n        )\n\n    else:\n\n        edge_index = torch.tensor(\n            edge_index,\n            dtype=torch.long\n        ).t().contiguous()\n\n    graph = Data(\n        x=x,\n        edge_index=edge_index\n    )\n\n    return graph\n\n# =========================\n# CNN ENCODER\n# =========================\nclass CNNEncoder(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        backbone = models.resnet18(\n            weights=None\n        )\n\n        self.features = nn.Sequential(\n            *list(backbone.children())[:-1]\n        )\n\n        self.fc = nn.Linear(\n            512,\n            128\n        )\n\n    def forward(self, x):\n\n        x = self.features(x)\n\n        x = torch.flatten(x, 1)\n\n        x = self.fc(x)\n\n        return x\n\n# =========================\n# GRAPH NETWORK\n# =========================\nclass BrainGNN(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        self.gcn1 = GCNConv(2, 32)\n\n        self.gcn2 = GCNConv(32, 64)\n\n        self.fc = nn.Linear(64, 128)\n\n    def forward(self, graph):\n\n        x = graph.x.to(DEVICE)\n\n        edge_index = graph.edge_index.to(DEVICE)\n\n        x = self.gcn1(x, edge_index)\n\n        x = torch.relu(x)\n\n        x = self.gcn2(x, edge_index)\n\n        x = torch.relu(x)\n\n        batch = torch.zeros(\n            x.shape[0],\n            dtype=torch.long\n        ).to(DEVICE)\n\n        x = global_mean_pool(x, batch)\n\n        x = self.fc(x)\n\n        return x\n\n# =========================\n# HYBRID MODEL\n# =========================\nclass HybridModel(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        self.cnn = CNNEncoder()\n\n        self.gnn = BrainGNN()\n\n        self.classifier = nn.Sequential(\n\n            nn.Linear(256, 128),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.3),\n\n            nn.Linear(128, 1)\n\n        )\n\n    def forward(self, imgs, graphs):\n\n        cnn_feat = self.cnn(imgs)\n\n        graph_feats = []\n\n        for g in graphs:\n\n            gf = self.gnn(g)\n\n            graph_feats.append(gf)\n\n        graph_feat = torch.cat(\n            graph_feats,\n            dim=0\n        )\n\n        fused = torch.cat([\n            cnn_feat,\n            graph_feat\n        ], dim=1)\n\n        out = self.classifier(fused)\n\n        return out\n\n# =========================\n# LOAD DATA\n# =========================\ndf = pd.read_csv(train_csv)\n\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df['MGMT_value'],\n    random_state=42\n)\n\ntrain_dataset = MRIDataset(\n    train_df,\n    train_dir\n)\n\nval_dataset = MRIDataset(\n    val_df,\n    train_dir\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE\n)\n\n# =========================\n# SHOW SAMPLE MRI\n# =========================\nsample_img = train_dataset[0][2]\n\nplt.figure(figsize=(5,5))\n\nplt.imshow(sample_img, cmap='gray')\n\nplt.title('Sample MRI')\n\nplt.axis('off')\n\nplt.show()\n\n# =========================\n# MODEL\n# =========================\nmodel = HybridModel().to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss()\n\noptimizer = optim.Adam(\n    model.parameters(),\n    lr=1e-4\n)\n\n# =========================\n# TRAINING\n# =========================\ntrain_losses = []\n\nfor epoch in range(EPOCHS):\n\n    print(f'\\nEpoch {epoch+1}')\n\n    model.train()\n\n    losses = []\n\n    for imgs, labels, raw_imgs in train_loader:\n\n        imgs = imgs.to(DEVICE)\n\n        labels = labels.to(DEVICE)\n\n        graphs = []\n\n        for i in range(raw_imgs.shape[0]):\n\n            g = create_graph(\n                raw_imgs[i]\n            )\n\n            graphs.append(g)\n\n        optimizer.zero_grad()\n\n        outputs = model(\n            imgs,\n            graphs\n        )\n\n        outputs = outputs.view(-1,1)\n\n        labels = labels.view(-1,1)\n\n        loss = criterion(\n            outputs,\n            labels\n        )\n\n        loss.backward()\n\n        optimizer.step()\n\n        losses.append(loss.item())\n\n    epoch_loss = np.mean(losses)\n\n    train_losses.append(epoch_loss)\n\n    print(\"Train Loss:\", epoch_loss)\n\n# =========================\n# TRAINING CURVE\n# =========================\nplt.figure(figsize=(6,4))\n\nplt.plot(train_losses, marker='o')\n\nplt.title('Training Loss')\n\nplt.xlabel('Epoch')\n\nplt.ylabel('Loss')\n\nplt.grid(True)\n\nplt.show()\n\n# =========================\n# VALIDATION\n# =========================\nmodel.eval()\n\npreds = []\n\nlabels_all = []\n\nwith torch.no_grad():\n\n    for imgs, labels, raw_imgs in val_loader:\n\n        imgs = imgs.to(DEVICE)\n\n        graphs = []\n\n        for i in range(raw_imgs.shape[0]):\n\n            g = create_graph(\n                raw_imgs[i]\n            )\n\n            graphs.append(g)\n\n        outputs = torch.sigmoid(\n            model(imgs, graphs)\n        ).view(-1,1)\n\n        preds.extend(\n            outputs.cpu().numpy().flatten()\n        )\n\n        labels_all.extend(\n            labels.numpy().flatten()\n        )\n\n# =========================\n# METRICS\n# =========================\npreds_bin = (\n    np.array(preds) > 0.5\n).astype(int)\n\ntry:\n\n    auc = roc_auc_score(\n        labels_all,\n        preds\n    )\n\nexcept:\n\n    auc = 0.0\n\nacc = accuracy_score(\n    labels_all,\n    preds_bin\n)\n\nf1 = f1_score(\n    labels_all,\n    preds_bin,\n    zero_division=0\n)\n\nprint(\"\\n====================\")\nprint(\"FINAL RESULTS\")\nprint(\"====================\")\n\nprint(\"AUC :\", round(auc,4))\nprint(\"ACC :\", round(acc,4))\nprint(\"F1  :\", round(f1,4))\n\n# =========================\n# ROC CURVE\n# =========================\ntry:\n\n    fpr, tpr, _ = roc_curve(\n        labels_all,\n        preds\n    )\n\n    plt.figure(figsize=(5,5))\n\n    plt.plot(fpr, tpr)\n\n    plt.plot([0,1],[0,1],'--')\n\n    plt.title('ROC Curve')\n\n    plt.xlabel('False Positive Rate')\n\n    plt.ylabel('True Positive Rate')\n\n    plt.grid(True)\n\n    plt.show()\n\nexcept:\n\n    print(\"ROC curve skipped\")\n\n# =========================\n# CONFUSION MATRIX\n# =========================\ncm = confusion_matrix(\n    labels_all,\n    preds_bin\n)\n\ndisp = ConfusionMatrixDisplay(\n    confusion_matrix=cm\n)\n\ndisp.plot()\n\nplt.title('Confusion Matrix')\n\nplt.show()\n\n# =========================\n# GRAPH VISUALIZATION\n# =========================\nsample_graph = create_graph(\n    train_dataset[0][2]\n)\n\nedge_index = sample_graph.edge_index.numpy()\n\nplt.figure(figsize=(6,6))\n\nfor i in range(edge_index.shape[1]):\n\n    x1 = edge_index[0][i]\n\n    x2 = edge_index[1][i]\n\n    plt.plot(\n        [x1, x2],\n        [x1, x2]\n    )\n\nplt.title('Graph Structure')\n\nplt.show()\n\n# =========================\n# SAVE MODEL\n# =========================\ntorch.save(\n    model.state_dict(),\n    '/kaggle/working/hybrid_model.pth'\n)\n\nprint(\"\\nMODEL SAVED!\")\n\nprint(\"/kaggle/working/hybrid_model.pth\")\n\nprint(\"\\nDONE SUCCESSFULLY!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T22:32:33.849653Z","iopub.execute_input":"2026-05-08T22:32:33.849938Z","iopub.status.idle":"2026-05-08T22:41:43.418732Z","shell.execute_reply.started":"2026-05-08T22:32:33.849915Z","shell.execute_reply":"2026-05-08T22:41:43.417910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# HYBRID MRI FRAMEWORK\n# HI-GCN + A-GCL Inspired Hybrid Model\n# Generates Images + Metrics + Saved Model\n# FULL KAGGLE SAFE VERSION\n# ============================================================\n\n# =========================\n# INSTALL\n# =========================\n!pip install -q pydicom torch-geometric\n\n# =========================\n# IMPORTS\n# =========================\nimport os\nimport glob\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\n\nimport pydicom\n\nfrom skimage.transform import resize\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    f1_score,\n    roc_auc_score,\n    roc_curve,\n    confusion_matrix,\n    ConfusionMatrixDisplay\n)\n\nimport torchvision.models as models\n\nfrom torch_geometric.data import Data\nfrom torch_geometric.nn import GCNConv\nfrom torch_geometric.nn import global_mean_pool\n\n# =========================\n# CONFIG\n# =========================\nIMG_SIZE = 224\nBATCH_SIZE = 4\nEPOCHS = 1\n\nDEVICE = torch.device(\n    'cuda' if torch.cuda.is_available() else 'cpu'\n)\n\nprint(\"DEVICE:\", DEVICE)\n\n# =========================\n# FIND DATASET\n# =========================\ncsv_files = glob.glob(\n    '/kaggle/input/**/train_labels.csv',\n    recursive=True\n)\n\nif len(csv_files) == 0:\n    raise FileNotFoundError(\n        'Dataset not found.\\n'\n        'Add RSNA dataset in Kaggle.'\n    )\n\ntrain_csv = csv_files[0]\n\ntrain_dir = train_csv.replace(\n    'train_labels.csv',\n    'train'\n)\n\nprint(\"CSV:\", train_csv)\nprint(\"TRAIN DIR:\", train_dir)\n\n# =========================\n# DATASET\n# =========================\nclass MRIDataset(Dataset):\n\n    def __init__(self, df, data_dir):\n\n        self.df = df.reset_index(drop=True)\n\n        self.data_dir = data_dir\n\n    def __len__(self):\n\n        return len(self.df)\n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n\n        pid = str(row['BraTS21ID']).zfill(5)\n\n        patient_path = os.path.join(\n            self.data_dir,\n            pid\n        )\n\n        img = np.zeros(\n            (IMG_SIZE, IMG_SIZE),\n            dtype=np.float32\n        )\n\n        try:\n\n            flair_path = os.path.join(\n                patient_path,\n                'FLAIR'\n            )\n\n            files = sorted([\n                f for f in os.listdir(flair_path)\n                if f.endswith('.dcm')\n            ])\n\n            if len(files) > 0:\n\n                mid = files[len(files)//2]\n\n                dcm = pydicom.dcmread(\n                    os.path.join(flair_path, mid)\n                )\n\n                img = dcm.pixel_array.astype(np.float32)\n\n                img = resize(\n                    img,\n                    (IMG_SIZE, IMG_SIZE)\n                )\n\n                img = (\n                    img - np.mean(img)\n                ) / (\n                    np.std(img) + 1e-8\n                )\n\n        except:\n            pass\n\n        img3 = np.stack([img]*3, axis=0)\n\n        label = row['MGMT_value']\n\n        return (\n            torch.FloatTensor(img3),\n            torch.FloatTensor([label]),\n            torch.FloatTensor(img)\n        )\n\n# =========================\n# GRAPH CREATION\n# =========================\ndef create_graph(image_tensor):\n\n    image = image_tensor.numpy()\n\n    patches = []\n\n    step = 32\n\n    for i in range(0, IMG_SIZE, step):\n\n        for j in range(0, IMG_SIZE, step):\n\n            patch = image[\n                i:i+step,\n                j:j+step\n            ]\n\n            feat = [\n                patch.mean(),\n                patch.std()\n            ]\n\n            patches.append(feat)\n\n    x = torch.tensor(\n        patches,\n        dtype=torch.float\n    )\n\n    edge_index = []\n\n    num_nodes = len(patches)\n\n    for i in range(num_nodes-1):\n\n        edge_index.append([i, i+1])\n\n        edge_index.append([i+1, i])\n\n    if len(edge_index) == 0:\n\n        edge_index = torch.tensor(\n            [[0],[0]],\n            dtype=torch.long\n        )\n\n    else:\n\n        edge_index = torch.tensor(\n            edge_index,\n            dtype=torch.long\n        ).t().contiguous()\n\n    graph = Data(\n        x=x,\n        edge_index=edge_index\n    )\n\n    return graph\n\n# =========================\n# CNN ENCODER\n# =========================\nclass CNNEncoder(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        backbone = models.resnet18(\n            weights=None\n        )\n\n        self.features = nn.Sequential(\n            *list(backbone.children())[:-1]\n        )\n\n        self.fc = nn.Linear(\n            512,\n            128\n        )\n\n    def forward(self, x):\n\n        x = self.features(x)\n\n        x = torch.flatten(x, 1)\n\n        x = self.fc(x)\n\n        return x\n\n# =========================\n# HI-GCN INSPIRED GRAPH MODEL\n# =========================\nclass HIGCN(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        self.gcn1 = GCNConv(2, 32)\n\n        self.gcn2 = GCNConv(32, 64)\n\n        self.fc = nn.Linear(64, 128)\n\n    def forward(self, graph):\n\n        x = graph.x.to(DEVICE)\n\n        edge_index = graph.edge_index.to(DEVICE)\n\n        x = self.gcn1(x, edge_index)\n\n        x = torch.relu(x)\n\n        x = self.gcn2(x, edge_index)\n\n        x = torch.relu(x)\n\n        batch = torch.zeros(\n            x.shape[0],\n            dtype=torch.long\n        ).to(DEVICE)\n\n        x = global_mean_pool(x, batch)\n\n        x = self.fc(x)\n\n        return x\n\n# =========================\n# A-GCL INSPIRED CONTRASTIVE MODULE\n# =========================\nclass AGCLModule(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        self.projector = nn.Sequential(\n\n            nn.Linear(128, 128),\n\n            nn.ReLU(),\n\n            nn.Linear(128, 64)\n\n        )\n\n    def forward(self, x):\n\n        x = self.projector(x)\n\n        return x\n\n# =========================\n# HYBRID MODEL\n# =========================\nclass HybridModel(nn.Module):\n\n    def __init__(self):\n\n        super().__init__()\n\n        self.cnn = CNNEncoder()\n\n        self.higcn = HIGCN()\n\n        self.agcl = AGCLModule()\n\n        self.classifier = nn.Sequential(\n\n            nn.Linear(128 + 64, 128),\n\n            nn.ReLU(),\n\n            nn.Dropout(0.3),\n\n            nn.Linear(128, 1)\n\n        )\n\n    def forward(self, imgs, graphs):\n\n        cnn_feat = self.cnn(imgs)\n\n        graph_feats = []\n\n        for g in graphs:\n\n            gf = self.higcn(g)\n\n            graph_feats.append(gf)\n\n        graph_feat = torch.cat(\n            graph_feats,\n            dim=0\n        )\n\n        contrast_feat = self.agcl(graph_feat)\n\n        fused = torch.cat([\n            cnn_feat,\n            contrast_feat\n        ], dim=1)\n\n        out = self.classifier(fused)\n\n        return out\n\n# =========================\n# LOAD DATA\n# =========================\ndf = pd.read_csv(train_csv)\n\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df['MGMT_value'],\n    random_state=42\n)\n\ntrain_dataset = MRIDataset(\n    train_df,\n    train_dir\n)\n\nval_dataset = MRIDataset(\n    val_df,\n    train_dir\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE\n)\n\n# =========================\n# SHOW SAMPLE MRI\n# =========================\nsample_img = train_dataset[0][2]\n\nplt.figure(figsize=(5,5))\n\nplt.imshow(sample_img, cmap='gray')\n\nplt.title('Sample MRI')\n\nplt.axis('off')\n\nplt.show()\n\n# =========================\n# MODEL\n# =========================\nmodel = HybridModel().to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss()\n\noptimizer = optim.Adam(\n    model.parameters(),\n    lr=1e-4\n)\n\n# =========================\n# TRAINING\n# =========================\ntrain_losses = []\n\nfor epoch in range(EPOCHS):\n\n    print(f'\\nEpoch {epoch+1}')\n\n    model.train()\n\n    losses = []\n\n    for imgs, labels, raw_imgs in train_loader:\n\n        imgs = imgs.to(DEVICE)\n\n        labels = labels.to(DEVICE)\n\n        graphs = []\n\n        for i in range(raw_imgs.shape[0]):\n\n            g = create_graph(\n                raw_imgs[i]\n            )\n\n            graphs.append(g)\n\n        optimizer.zero_grad()\n\n        outputs = model(\n            imgs,\n            graphs\n        )\n\n        outputs = outputs.view(-1,1)\n\n        labels = labels.view(-1,1)\n\n        loss = criterion(\n            outputs,\n            labels\n        )\n\n        loss.backward()\n\n        optimizer.step()\n\n        losses.append(loss.item())\n\n    epoch_loss = np.mean(losses)\n\n    train_losses.append(epoch_loss)\n\n    print(\"Train Loss:\", epoch_loss)\n\n# =========================\n# TRAINING CURVE\n# =========================\nplt.figure(figsize=(6,4))\n\nplt.plot(train_losses, marker='o')\n\nplt.title('Training Loss')\n\nplt.xlabel('Epoch')\n\nplt.ylabel('Loss')\n\nplt.grid(True)\n\nplt.show()\n\n# =========================\n# VALIDATION\n# =========================\nmodel.eval()\n\npreds = []\n\nlabels_all = []\n\nwith torch.no_grad():\n\n    for imgs, labels, raw_imgs in val_loader:\n\n        imgs = imgs.to(DEVICE)\n\n        graphs = []\n\n        for i in range(raw_imgs.shape[0]):\n\n            g = create_graph(\n                raw_imgs[i]\n            )\n\n            graphs.append(g)\n\n        outputs = torch.sigmoid(\n            model(imgs, graphs)\n        ).view(-1,1)\n\n        preds.extend(\n            outputs.cpu().numpy().flatten()\n        )\n\n        labels_all.extend(\n            labels.numpy().flatten()\n        )\n\n# =========================\n# METRICS\n# =========================\npreds_bin = (\n    np.array(preds) > 0.5\n).astype(int)\n\ntry:\n\n    auc = roc_auc_score(\n        labels_all,\n        preds\n    )\n\nexcept:\n\n    auc = 0.0\n\nacc = accuracy_score(\n    labels_all,\n    preds_bin\n)\n\nf1 = f1_score(\n    labels_all,\n    preds_bin,\n    zero_division=0\n)\n\nprint(\"\\n====================\")\nprint(\"FINAL RESULTS\")\nprint(\"====================\")\n\nprint(\"AUC :\", round(auc,4))\nprint(\"ACC :\", round(acc,4))\nprint(\"F1  :\", round(f1,4))\n\n# =========================\n# ROC CURVE\n# =========================\ntry:\n\n    fpr, tpr, _ = roc_curve(\n        labels_all,\n        preds\n    )\n\n    plt.figure(figsize=(5,5))\n\n    plt.plot(fpr, tpr)\n\n    plt.plot([0,1],[0,1],'--')\n\n    plt.title('ROC Curve')\n\n    plt.xlabel('False Positive Rate')\n\n    plt.ylabel('True Positive Rate')\n\n    plt.grid(True)\n\n    plt.show()\n\nexcept:\n\n    print(\"ROC curve skipped\")\n\n# =========================\n# CONFUSION MATRIX\n# =========================\ncm = confusion_matrix(\n    labels_all,\n    preds_bin\n)\n\ndisp = ConfusionMatrixDisplay(\n    confusion_matrix=cm\n)\n\ndisp.plot()\n\nplt.title('Confusion Matrix')\n\nplt.show()\n\n# =========================\n# GRAPH VISUALIZATION\n# =========================\nsample_graph = create_graph(\n    train_dataset[0][2]\n)\n\nedge_index = sample_graph.edge_index.numpy()\n\nplt.figure(figsize=(6,6))\n\nfor i in range(edge_index.shape[1]):\n\n    x1 = edge_index[0][i]\n\n    x2 = edge_index[1][i]\n\n    plt.plot(\n        [x1, x2],\n        [x1, x2]\n    )\n\nplt.title('HI-GCN Graph Structure')\n\nplt.show()\n\n# =========================\n# SAVE MODEL\n# =========================\ntorch.save(\n    model.state_dict(),\n    '/kaggle/working/higcn_agcl_hybrid_model.pth'\n)\n\nprint(\"\\nMODEL SAVED!\")\n\nprint(\"/kaggle/working/higcn_agcl_hybrid_model.pth\")\n\nprint(\"\\nDONE SUCCESSFULLY!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-08T22:46:57.918919Z","iopub.execute_input":"2026-05-08T22:46:57.919228Z","iopub.status.idle":"2026-05-08T22:57:17.734942Z","shell.execute_reply.started":"2026-05-08T22:46:57.919196Z","shell.execute_reply":"2026-05-08T22:57:17.734009Z"}},"outputs":[],"execution_count":null}]}