{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":18647,"databundleVersionId":1126921,"sourceType":"competition"}],"dockerImageVersionId":31012,"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\n\n# Input data files are available in the read-only \"../input/\" directory\n# Running this code 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# Ignore warnings for cleaner output\nimport warnings\nwarnings.filterwarnings('ignore')\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:53:57.051576Z","iopub.execute_input":"2025-05-14T19:53:57.051951Z","iopub.status.idle":"2025-05-14T19:54:04.459306Z","shell.execute_reply.started":"2025-05-14T19:53:57.051929Z","shell.execute_reply":"2025-05-14T19:54:04.457929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport openslide\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score, roc_auc_score, classification_report, confusion_matrix\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:54:04.46159Z","iopub.execute_input":"2025-05-14T19:54:04.461972Z","iopub.status.idle":"2025-05-14T19:54:04.469176Z","shell.execute_reply.started":"2025-05-14T19:54:04.461944Z","shell.execute_reply":"2025-05-14T19:54:04.468178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CSV_PATH = \"/kaggle/input/prostate-cancer-grade-assessment/train.csv\"\nIMG_DIR = \"/kaggle/input/prostate-cancer-grade-assessment/train_images\"\n\ndf = pd.read_csv(CSV_PATH)\ndf = df[['image_id', 'isup_grade']]\ndf = df.sample(500, random_state=42).reset_index(drop=True)\ndf['image_path'] = df['image_id'].apply(lambda x: os.path.join(IMG_DIR, f\"{x}.tiff\"))\n\nprint(df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:54:04.470114Z","iopub.execute_input":"2025-05-14T19:54:04.470411Z","iopub.status.idle":"2025-05-14T19:54:04.506122Z","shell.execute_reply.started":"2025-05-14T19:54:04.470388Z","shell.execute_reply":"2025-05-14T19:54:04.505331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 224\n\ntransform = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406],\n                         [0.229, 0.224, 0.225])\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:54:04.506733Z","iopub.execute_input":"2025-05-14T19:54:04.507025Z","iopub.status.idle":"2025-05-14T19:54:04.51233Z","shell.execute_reply.started":"2025-05-14T19:54:04.507007Z","shell.execute_reply":"2025-05-14T19:54:04.511411Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ProstateFeatureDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.df = dataframe\n        self.transform = transform\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = row['image_path']\n        slide = openslide.OpenSlide(img_path)\n        w, h = slide.dimensions\n        patch = slide.read_region((w//2 - 256, h//2 - 256), 0, (512, 512)).convert(\"RGB\")\n        if self.transform:\n            patch = self.transform(patch)\n        label = row['isup_grade']  # multiclass 0 to 5\n        return patch, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:54:04.514693Z","iopub.execute_input":"2025-05-14T19:54:04.51511Z","iopub.status.idle":"2025-05-14T19:54:04.531223Z","shell.execute_reply.started":"2025-05-14T19:54:04.515088Z","shell.execute_reply":"2025-05-14T19:54:04.530289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load pretrained ResNet50 and MobileNetV2 models\n\nresnet = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)\nresnet_features = nn.Sequential(*list(resnet.children())[:-1])\nresnet_features.eval()\n\nmobilenet = models.mobilenet_v2(weights=models.MobileNet_V2_Weights.IMAGENET1K_V1)\nmobilenet_features = nn.Sequential(*list(mobilenet.children())[:-1])\nmobilenet_features.eval()\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nresnet_features = resnet_features.to(device)\nmobilenet_features = mobilenet_features.to(device)\n\n# Dataset and DataLoader\nfeature_dataset = ProstateFeatureDataset(df, transform=transform)\nfeature_loader = DataLoader(feature_dataset, batch_size=8, shuffle=False)\n\nresnet_feats = []\nmobilenet_feats = []\nlabels_list = []\n\nwith torch.no_grad():\n    for images, labels in feature_loader:\n        images = images.to(device)\n\n        # --- ResNet50 features ---\n        r_features = resnet_features(images)  # [batch, 2048, 1, 1]\n        r_features = r_features.view(r_features.size(0), -1)  # flatten to [batch, 2048]\n\n        # --- MobileNetV2 features ---\n        m_features = mobilenet_features(images)  # [batch, 1280, 7, 7]\n        m_features = torch.nn.functional.adaptive_avg_pool2d(m_features, (1, 1))\n        m_features = m_features.view(m_features.size(0), -1)  # flatten to [batch, 1280]\n\n        # --- Concatenate features ---\n        fused = torch.cat((r_features, m_features), dim=1)  # [batch, 3328]\n\n        resnet_feats.append(r_features.cpu().numpy())\n        mobilenet_feats.append(m_features.cpu().numpy())\n        labels_list.extend(labels.cpu().numpy())\n\n# Combine all batches\nresnet_array = np.vstack(resnet_feats)     # [500, 2048]\nmobilenet_array = np.vstack(mobilenet_feats)  # [500, 1280]\nfused_features = np.hstack((resnet_array, mobilenet_array))  # [500, 3328]\nlabels_array = np.array(labels_list)\n\nprint(\"Fused feature shape:\", fused_features.shape)\nprint(\"Labels shape:\", labels_array.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:54:04.531984Z","iopub.execute_input":"2025-05-14T19:54:04.532245Z","iopub.status.idle":"2025-05-14T19:55:29.986671Z","shell.execute_reply.started":"2025-05-14T19:54:04.532226Z","shell.execute_reply":"2025-05-14T19:55:29.985733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X_train, X_temp, y_train, y_temp = train_test_split(fused_features, labels_array,\n                                                    test_size=0.4, stratify=labels_array,\n                                                    random_state=42)\n\nX_val, X_test, y_val, y_test = train_test_split(X_temp, y_temp,\n                                                test_size=0.5, stratify=y_temp,\n                                                random_state=42)\n\nprint(f\"Training set size: {X_train.shape[0]}\")\nprint(f\"Validation set size: {X_val.shape[0]}\")\nprint(f\"Test set size: {X_test.shape[0]}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:29.987583Z","iopub.execute_input":"2025-05-14T19:55:29.987878Z","iopub.status.idle":"2025-05-14T19:55:30.000017Z","shell.execute_reply.started":"2025-05-14T19:55:29.987828Z","shell.execute_reply":"2025-05-14T19:55:29.999021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HiFuseClassifier(nn.Module):\n    def __init__(self, input_dim, hidden_dim=512, num_classes=6):\n        super(HiFuseClassifier, self).__init__()\n        self.fc1 = nn.Linear(input_dim, hidden_dim)\n        self.relu = nn.ReLU()\n        self.fc2 = nn.Linear(hidden_dim, num_classes)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.relu(x)\n        x = self.fc2(x)\n        return x\n\ninput_dim = X_train.shape[1]  # 3328\nmodel_hifuse = HiFuseClassifier(input_dim).to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:30.001144Z","iopub.execute_input":"2025-05-14T19:55:30.001574Z","iopub.status.idle":"2025-05-14T19:55:30.033659Z","shell.execute_reply.started":"2025-05-14T19:55:30.001543Z","shell.execute_reply":"2025-05-14T19:55:30.032895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import TensorDataset, DataLoader\n\n# Convert to tensors\nX_train_tensor = torch.tensor(X_train, dtype=torch.float32)\ny_train_tensor = torch.tensor(y_train, dtype=torch.long)\n\nX_val_tensor = torch.tensor(X_val, dtype=torch.float32)\ny_val_tensor = torch.tensor(y_val, dtype=torch.long)\n\nX_test_tensor = torch.tensor(X_test, dtype=torch.float32)\ny_test_tensor = torch.tensor(y_test, dtype=torch.long)\n\n# Datasets and DataLoaders\ntrain_dataset_hf = TensorDataset(X_train_tensor, y_train_tensor)\nval_dataset_hf = TensorDataset(X_val_tensor, y_val_tensor)\ntest_dataset_hf = TensorDataset(X_test_tensor, y_test_tensor)\n\ntrain_loader_hf = DataLoader(train_dataset_hf, batch_size=16, shuffle=True)\nval_loader_hf = DataLoader(val_dataset_hf, batch_size=16)\ntest_loader_hf = DataLoader(test_dataset_hf, batch_size=16)\n\n# Confirm shape\nsample_features, _ = next(iter(train_loader_hf))\nprint(\"Sample shape from HiFuse loader:\", sample_features.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:30.034721Z","iopub.execute_input":"2025-05-14T19:55:30.035138Z","iopub.status.idle":"2025-05-14T19:55:30.045928Z","shell.execute_reply.started":"2025-05-14T19:55:30.035108Z","shell.execute_reply":"2025-05-14T19:55:30.045173Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model_hifuse.parameters(), lr=0.001)\n\nepochs = 10\nfor epoch in range(epochs):\n    model_hifuse.train()\n    total_loss = 0\n    correct = 0\n    total = 0\n\n    for features, labels in train_loader_hf:\n        features, labels = features.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model_hifuse(features)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item() * labels.size(0)\n        preds = outputs.argmax(dim=1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n    train_acc = correct / total\n    avg_loss = total_loss / total\n    print(f\"Epoch {epoch+1}: Loss={avg_loss:.4f}, Train Accuracy={train_acc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:30.047542Z","iopub.execute_input":"2025-05-14T19:55:30.047891Z","iopub.status.idle":"2025-05-14T19:55:31.720541Z","shell.execute_reply.started":"2025-05-14T19:55:30.047863Z","shell.execute_reply":"2025-05-14T19:55:31.719574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_hifuse.eval()\ny_true_train = []\ny_pred_train = []\n\nwith torch.no_grad():\n    for features, labels in train_loader_hf:\n        features = features.to(device)\n        outputs = model_hifuse(features)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n\n        y_true_train.extend(labels.cpu().numpy())\n        y_pred_train.extend(preds)\n\ntrain_accuracy = accuracy_score(y_true_train, y_pred_train)\nprint(f\"Training Accuracy: {train_accuracy:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:31.721624Z","iopub.execute_input":"2025-05-14T19:55:31.722374Z","iopub.status.idle":"2025-05-14T19:55:31.774429Z","shell.execute_reply.started":"2025-05-14T19:55:31.722326Z","shell.execute_reply":"2025-05-14T19:55:31.77346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_true_val = []\ny_pred_val = []\n\nwith torch.no_grad():\n    for features, labels in val_loader_hf:\n        features = features.to(device)\n        outputs = model_hifuse(features)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n\n        y_true_val.extend(labels.cpu().numpy())\n        y_pred_val.extend(preds)\n\nval_accuracy = accuracy_score(y_true_val, y_pred_val)\nval_f1 = f1_score(y_true_val, y_pred_val, average='weighted')\nval_kappa = cohen_kappa_score(y_true_val, y_pred_val)\n\ny_true_val_bin = [1 if y > 0 else 0 for y in y_true_val]\ny_pred_val_bin = [1 if y > 0 else 0 for y in y_pred_val]\nval_roc = roc_auc_score(y_true_val_bin, y_pred_val_bin)\n\nprint(f\"Validation Accuracy: {val_accuracy:.4f}\")\nprint(f\"F1 Score: {val_f1:.4f}\")\nprint(f\"Kappa: {val_kappa:.4f}\")\nprint(f\"ROC AUC: {val_roc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:31.775563Z","iopub.execute_input":"2025-05-14T19:55:31.775932Z","iopub.status.idle":"2025-05-14T19:55:31.806469Z","shell.execute_reply.started":"2025-05-14T19:55:31.775896Z","shell.execute_reply":"2025-05-14T19:55:31.805709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_true_test = []\ny_pred_test = []\n\nwith torch.no_grad():\n    for features, labels in test_loader_hf:\n        features = features.to(device)\n        outputs = model_hifuse(features)\n        preds = outputs.argmax(dim=1).cpu().numpy()\n\n        y_true_test.extend(labels.cpu().numpy())\n        y_pred_test.extend(preds)\n\ntest_accuracy = accuracy_score(y_true_test, y_pred_test)\ntest_f1 = f1_score(y_true_test, y_pred_test, average='weighted')\ntest_kappa = cohen_kappa_score(y_true_test, y_pred_test)\n\ny_true_test_bin = [1 if y > 0 else 0 for y in y_true_test]\ny_pred_test_bin = [1 if y > 0 else 0 for y in y_pred_test]\ntest_roc = roc_auc_score(y_true_test_bin, y_pred_test_bin)\n\nprint(f\"Test Accuracy: {test_accuracy:.4f}\")\nprint(f\"F1 Score: {test_f1:.4f}\")\nprint(f\"Kappa: {test_kappa:.4f}\")\nprint(f\"ROC AUC: {test_roc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:31.807392Z","iopub.execute_input":"2025-05-14T19:55:31.807721Z","iopub.status.idle":"2025-05-14T19:55:31.834867Z","shell.execute_reply.started":"2025-05-14T19:55:31.8077Z","shell.execute_reply":"2025-05-14T19:55:31.83398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"results = pd.DataFrame({\n    \"Set\": [\"Training\", \"Validation\", \"Test\"],\n    \"Accuracy\": [train_accuracy, val_accuracy, test_accuracy],\n    \"F1 Measure\": [\"-\", f\"{val_f1:.4f}\", f\"{test_f1:.4f}\"],\n    \"Kappa\": [\"-\", f\"{val_kappa:.4f}\", f\"{test_kappa:.4f}\"],\n    \"ROC Area\": [\"-\", f\"{val_roc:.4f}\", f\"{test_roc:.4f}\"]\n})\n\nprint(results.to_string(index=False))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:31.837679Z","iopub.execute_input":"2025-05-14T19:55:31.838044Z","iopub.status.idle":"2025-05-14T19:55:31.847072Z","shell.execute_reply.started":"2025-05-14T19:55:31.838022Z","shell.execute_reply":"2025-05-14T19:55:31.846142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Test Set Accuracy:\", test_accuracy)\nprint(\"\\n📊 Classification Report:\")\nprint(classification_report(y_true_test, y_pred_test))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:31.848052Z","iopub.execute_input":"2025-05-14T19:55:31.848346Z","iopub.status.idle":"2025-05-14T19:55:31.871802Z","shell.execute_reply.started":"2025-05-14T19:55:31.848326Z","shell.execute_reply":"2025-05-14T19:55:31.870724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(y_true_test, y_pred_test)\n\nplt.figure(figsize=(6,5))\nsns.heatmap(cm, annot=True, fmt='d', cmap=\"Blues\",\n            xticklabels=[0, 1, 2, 3, 4, 5],\n            yticklabels=[0, 1, 2, 3, 4, 5])\nplt.title(\"🧩 HiFuse Confusion Matrix\")\nplt.xlabel(\"Predicted Label\")\nplt.ylabel(\"True Label\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T19:55:31.87307Z","iopub.execute_input":"2025-05-14T19:55:31.873383Z","iopub.status.idle":"2025-05-14T19:55:32.212705Z","shell.execute_reply.started":"2025-05-14T19:55:31.873362Z","shell.execute_reply":"2025-05-14T19:55:32.211666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}