{"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"},{"sourceId":1101206,"sourceType":"datasetVersion","datasetId":615046}],"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 (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":"import os\nimport pandas as pd\nfrom PIL import Image\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom sklearn.model_selection import train_test_split\nimport torch.nn as nn\nimport torch.optim as optim\nfrom sklearn.metrics import accuracy_score, f1_score, cohen_kappa_score, roc_auc_score, confusion_matrix, ConfusionMatrixDisplay\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T11:38:29.886212Z","iopub.execute_input":"2025-05-07T11:38:29.886573Z","iopub.status.idle":"2025-05-07T11:38:40.985345Z","shell.execute_reply.started":"2025-05-07T11:38:29.886545Z","shell.execute_reply":"2025-05-07T11:38:40.984226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MILDataset(Dataset):\n    def __init__(self, csv_file, transform=None):\n        self.data = pd.read_csv(csv_file)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        image = Image.open(row[\"image_path\"]).convert(\"RGB\")\n        label = int(row[\"isup_grade\"])\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image.unsqueeze(0), label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T11:38:45.861097Z","iopub.execute_input":"2025-05-07T11:38:45.861582Z","iopub.status.idle":"2025-05-07T11:38:45.868662Z","shell.execute_reply.started":"2025-05-07T11:38:45.861556Z","shell.execute_reply":"2025-05-07T11:38:45.86745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LABEL_CSV = \"/kaggle/input/prostate-cancer-grade-assessment/train.csv\"\nRESIZED_IMG_DIR = \"/kaggle/input/panda-resized-train-data-512x512/train_images/train_images\"\n\ndf = pd.read_csv(LABEL_CSV)[[\"image_id\", \"isup_grade\"]]\ndf[\"image_path\"] = df[\"image_id\"].apply(lambda x: os.path.join(RESIZED_IMG_DIR, f\"{x}.png\"))\ndf = df[df[\"image_path\"].apply(os.path.exists)].reset_index(drop=True)\n\ntrain_df, val_df = train_test_split(df, test_size=0.2, stratify=df[\"isup_grade\"], random_state=42)\ntrain_df.to_csv(\"/kaggle/working/train.csv\", index=False)\nval_df.to_csv(\"/kaggle/working/val.csv\", index=False)\n\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.5]*3, [0.5]*3)\n])\n\ntrain_dataset = MILDataset(\"/kaggle/working/train.csv\", transform=transform)\nval_dataset = MILDataset(\"/kaggle/working/val.csv\", transform=transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T11:38:49.701102Z","iopub.execute_input":"2025-05-07T11:38:49.701456Z","iopub.status.idle":"2025-05-07T11:39:32.035878Z","shell.execute_reply.started":"2025-05-07T11:38:49.701395Z","shell.execute_reply":"2025-05-07T11:39:32.034756Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.models import resnet18, ResNet18_Weights\n\nclass SimpleMILClassifier(nn.Module):\n    def __init__(self, num_classes=6):\n        super(SimpleMILClassifier, self).__init__()\n        self.feature_extractor = resnet18(weights=ResNet18_Weights.DEFAULT)\n        self.feature_extractor.fc = nn.Identity()\n\n        self.attention = nn.Sequential(\n            nn.Linear(512, 128),\n            nn.Tanh(),\n            nn.Linear(128, 1)\n        )\n\n        self.classifier = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        B = x.size(0)\n        x = x.squeeze(1)\n        feats = self.feature_extractor(x)\n        attn_weights = torch.softmax(self.attention(feats), dim=0)\n        bag_rep = torch.sum(attn_weights * feats, dim=0, keepdim=True)\n        out = self.classifier(bag_rep)\n        return out\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T11:39:38.918039Z","iopub.execute_input":"2025-05-07T11:39:38.918366Z","iopub.status.idle":"2025-05-07T11:39:38.927109Z","shell.execute_reply.started":"2025-05-07T11:39:38.91834Z","shell.execute_reply":"2025-05-07T11:39:38.925193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = SimpleMILClassifier().to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=3e-4)  # Slightly higher LR for faster convergence","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T11:39:43.911194Z","iopub.execute_input":"2025-05-07T11:39:43.911535Z","iopub.status.idle":"2025-05-07T11:39:44.998703Z","shell.execute_reply.started":"2025-05-07T11:39:43.911509Z","shell.execute_reply":"2025-05-07T11:39:44.997439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader):\n    model.train()\n    total_loss = 0\n    for x, y in loader:\n        x, y = x.to(device), y.to(device)\n        out = model(x)\n        loss = criterion(out, y)\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    print(f\"Training Loss: {total_loss:.4f}\")\n\ndef evaluate_full(model, loader):\n    model.eval()\n    y_true, y_pred, y_prob = [], [], []\n\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(device), y.to(device)\n            out = model(x)\n            prob = torch.softmax(out, dim=1)\n            pred = torch.argmax(prob, dim=1)\n            y_true.append(y.item())\n            y_pred.append(pred.item())\n            y_prob.append(prob.cpu().numpy()[0])\n\n    acc = accuracy_score(y_true, y_pred)\n    f1 = f1_score(y_true, y_pred, average='weighted')\n    kappa = cohen_kappa_score(y_true, y_pred)\n    roc = roc_auc_score(y_true, y_prob, multi_class='ovr')\n    cm = confusion_matrix(y_true, y_pred)\n\n    print(f\"Accuracy: {acc:.4f}\")\n    print(f\"F1 Score: {f1:.4f}\")\n    print(f\"Kappa Score: {kappa:.4f}\")\n    print(f\"ROC AUC: {roc:.4f}\")\n\n    disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=list(range(6)))\n    disp.plot(cmap='Blues')\n    plt.title(\"Validation Confusion Matrix\")\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T11:39:51.900986Z","iopub.execute_input":"2025-05-07T11:39:51.901404Z","iopub.status.idle":"2025-05-07T11:39:51.9271Z","shell.execute_reply.started":"2025-05-07T11:39:51.901374Z","shell.execute_reply":"2025-05-07T11:39:51.925022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(7):  \n    print(f\"Epoch {epoch+1}\")\n    train_one_epoch(model, train_loader)\n    evaluate_full(model, val_loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-07T12:03:50.269362Z","iopub.execute_input":"2025-05-07T12:03:50.269741Z","iopub.status.idle":"2025-05-07T16:16:20.807529Z","shell.execute_reply.started":"2025-05-07T12:03:50.26971Z","shell.execute_reply":"2025-05-07T16:16:20.806491Z"}},"outputs":[],"execution_count":null}]}