{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":1512919,"sourceType":"datasetVersion","datasetId":611716},{"sourceId":11132125,"sourceType":"datasetVersion","datasetId":6942900}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%env CUDA_LAUNCH_BLOCKING=1\n\n####################################\n# 1. Import Libraries and Set Seed\n####################################\nimport os\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom torchvision import transforms, models\nimport numpy as np\nfrom PIL import Image\nimport random\nimport matplotlib.pyplot as plt\n\n# For reproducibility\nseed = 42\ntorch.manual_seed(seed)\nnp.random.seed(seed)\nrandom.seed(seed)\n\n####################################\n# 2. Data Paths and Transforms\n####################################\n# Update these paths based on your dataset structure on Kaggle.\nTRAIN_CSV = \"/kaggle/input/aptos2019-blindness-detection/train.csv\"\nTRAIN_IMAGES_DIR = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\n\n# Enhanced training augmentation: random resized crop, horizontal flip, rotation, and color jitter.\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),  \n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n# Validation transform: simply resize and normalize.\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\n####################################\n# 3. Custom Dataset for APTOS\n####################################\nclass APTOSDataset(Dataset):\n    def __init__(self, csv_file, images_dir, transform=None):\n        \"\"\"\n        csv_file: Path to train.csv (columns: id_code and diagnosis)\n        images_dir: Folder containing images (files named <id_code>.png)\n        transform: torchvision transforms to apply.\n        \"\"\"\n        self.data = pd.read_csv(csv_file)\n        self.images_dir = images_dir\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        img_id = row['id_code']\n        label = int(row['diagnosis'])  # Labels: 0 to 4\n        img_path = os.path.join(self.images_dir, img_id + \".png\")\n        image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if image is None:\n            # Fallback to a black image if reading fails.\n            image = np.zeros((224, 224, 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\n####################################\n# 4. Create DataLoaders\n####################################\nfull_dataset = APTOSDataset(TRAIN_CSV, TRAIN_IMAGES_DIR, transform=train_transform)\nprint(\"Total samples in dataset:\", len(full_dataset))\n\n# (Optional) Visualize a few samples.\nfor i in range(3):\n    img, label = full_dataset[i]\n    img_np = img.permute(1, 2, 0).numpy()\n    plt.imshow(img_np)\n    plt.title(f\"Sample {i} - Label: {label}\")\n    plt.axis(\"off\")\n    plt.show()\n\n# Split the dataset: 80% training, 20% validation.\ntrain_size = int(0.8 * len(full_dataset))\nval_size = len(full_dataset) - train_size\ntrain_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])\n# For validation, override the transform.\nval_dataset.dataset.transform = val_transform\n\nbatch_size = 32\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=2)\n\n####################################\n# 5. Define the Model (Pretrained ResNet18 with Dropout)\n####################################\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\nmodel = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\nnum_ftrs = model.fc.in_features\n# Replace the final layer with a dropout layer followed by a linear layer.\nmodel.fc = nn.Sequential(\n    nn.Dropout(0.6),  # Increased dropout rate\n    nn.Linear(num_ftrs, 5)  # APTOS has 5 classes.\n)\nmodel = model.to(device)\n\n####################################\n# 6. Loss, Optimizer, Scheduler, and Early Stopping Setup\n####################################\ncriterion = nn.CrossEntropyLoss()\n# Use Adam with weight decay for regularization.\noptimizer = optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n\n# Early stopping parameters.\nbest_val_loss = float('inf')\npatience = 5\npatience_counter = 0\n\n####################################\n# 7. Training Loop with Early Stopping\n####################################\nnum_epochs = 20\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item() * images.size(0)\n        _, preds = torch.max(outputs, 1)\n        total_train += labels.size(0)\n        correct_train += (preds == labels).sum().item()\n    scheduler.step()\n    train_loss = running_loss / total_train\n    train_acc = 100.0 * correct_train / total_train\n    \n    model.eval()\n    running_val_loss = 0.0\n    correct_val = 0\n    total_val = 0\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            running_val_loss += loss.item() * images.size(0)\n            _, preds = torch.max(outputs, 1)\n            total_val += labels.size(0)\n            correct_val += (preds == labels).sum().item()\n    val_loss = running_val_loss / total_val\n    val_acc = 100.0 * correct_val / total_val\n    \n    print(f\"Epoch [{epoch+1}/{num_epochs}] - \"\n          f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% | \"\n          f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n    \n    # Early Stopping Check\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        patience_counter = 0\n        torch.save(model.state_dict(), \"best_model.pt\")\n    else:\n        patience_counter += 1\n        if patience_counter >= patience:\n            print(\"Early stopping triggered.\")\n            break\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%env CUDA_LAUNCH_BLOCKING=1\n\n####################################\n# 1. Import Libraries and Set Seed\n####################################\nimport os\nimport pandas as pd\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom torchvision import transforms, models\nimport numpy as np\nfrom PIL import Image\nimport random\nimport matplotlib.pyplot as plt\n\nfrom sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score\n\n# For reproducibility\nseed = 42\ntorch.manual_seed(seed)\nnp.random.seed(seed)\nrandom.seed(seed)\n\n####################################\n# 2. Data Paths and Transforms\n####################################\nTRAIN_CSV = \"/kaggle/input/aptos2019-blindness-detection/train.csv\"\nTRAIN_IMAGES_DIR = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\n\n# Training: enhanced augmentation\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n# Validation: simple resize and normalize\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\n####################################\n# 3. Custom Dataset for APTOS\n####################################\nclass APTOSDataset(Dataset):\n    def __init__(self, csv_file, images_dir, transform=None):\n        \"\"\"\n        csv_file: Path to train.csv (columns: id_code and diagnosis)\n        images_dir: Folder containing images (named <id_code>.png)\n        transform: torchvision transforms to apply.\n        \"\"\"\n        self.data = pd.read_csv(csv_file)\n        self.images_dir = images_dir\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        img_id = row['id_code']\n        label = int(row['diagnosis'])  # Labels: 0 to 4\n        img_path = os.path.join(self.images_dir, img_id + \".png\")\n        image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if image is None:\n            # Fallback to a black image if reading fails.\n            image = np.zeros((224, 224, 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\n####################################\n# 4. Create DataLoaders\n####################################\nfull_dataset = APTOSDataset(TRAIN_CSV, TRAIN_IMAGES_DIR, transform=train_transform)\nprint(\"Total samples in dataset:\", len(full_dataset))\n\n# (Optional) Visualize a few samples\nfor i in range(3):\n    img, label = full_dataset[i]\n    img_np = img.permute(1, 2, 0).numpy()\n    plt.imshow(img_np)\n    plt.title(f\"Sample {i} - Label: {label}\")\n    plt.axis(\"off\")\n    plt.show()\n\n# Split: 80% training, 20% validation\ntrain_size = int(0.8 * len(full_dataset))\nval_size = len(full_dataset) - train_size\ntrain_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])\n# Override transform for validation\nval_dataset.dataset.transform = val_transform\n\nbatch_size = 32\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=2)\n\n####################################\n# 5. Define the Model (Pretrained ResNet18 with Dropout)\n####################################\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\nmodel = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\nnum_ftrs = model.fc.in_features\nmodel.fc = nn.Sequential(\n    nn.Dropout(0.6),  # Increased dropout rate\n    nn.Linear(num_ftrs, 5)  # 5 classes for APTOS\n)\nmodel = model.to(device)\n\n####################################\n# 6. Loss, Optimizer, and Scheduler Setup\n####################################\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n\n####################################\n# 7. Training Loop for 100 Epochs (No Early Stopping)\n####################################\nnum_epochs = 100\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n\n    # Training phase\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item() * images.size(0)\n        _, preds = torch.max(outputs, 1)\n        total_train += labels.size(0)\n        correct_train += (preds == labels).sum().item()\n    \n    scheduler.step()\n    train_loss = running_loss / total_train\n    train_acc = 100.0 * correct_train / total_train\n\n    # Validation phase\n    model.eval()\n    running_val_loss = 0.0\n    correct_val = 0\n    total_val = 0\n    all_preds = []\n    all_labels = []\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            running_val_loss += loss.item() * images.size(0)\n            _, preds = torch.max(outputs, 1)\n            total_val += labels.size(0)\n            correct_val += (preds == labels).sum().item()\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n    val_loss = running_val_loss / total_val\n    val_acc = 100.0 * correct_val / total_val\n\n    # Compute confusion matrix and additional metrics (weighted averages)\n    cm = confusion_matrix(all_labels, all_preds)\n    precision = precision_score(all_labels, all_preds, average='weighted', zero_division=0)\n    recall = recall_score(all_labels, all_preds, average='weighted', zero_division=0)\n    f1 = f1_score(all_labels, all_preds, average='weighted', zero_division=0)\n    accuracy = np.mean(np.array(all_preds) == np.array(all_labels)) * 100\n\n    # Detailed epoch logging (format similar to your example)\n    print(f\"Epoch [{epoch+1}/{num_epochs}] - \"\n          f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% | \"\n          f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n    print(\"Confusion Matrix:\")\n    print(cm)\n    print(f\"Metrics — Precision: {precision:.4f}, Recall: {recall:.4f}, F1: {f1:.4f}, Accuracy: {accuracy:.2f}%\\n\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%env CUDA_LAUNCH_BLOCKING=1\nimport os\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport pandas as pd\nimport random\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom torchvision import transforms, models\nfrom torch.utils.data import Dataset, random_split\n\n# Torch Geometric Imports\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader\nfrom torch_geometric.nn import GATConv, global_mean_pool\n\nfrom sklearn.metrics import confusion_matrix, precision_score, recall_score, f1_score\n\n# For reproducibility\nseed = 42\ntorch.manual_seed(seed)\nnp.random.seed(seed)\nrandom.seed(seed)\n\n####################################\n# 1. Data Paths and Transforms\n####################################\nTRAIN_CSV = \"/kaggle/input/aptos2019-blindness-detection/train.csv\"\nTRAIN_IMAGES_DIR = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\n\n# Standard image transforms (same as before)\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\n####################################\n# 2. Graph Dataset for CNN-GAT\n####################################\nclass APTOSGraphDataset(Dataset):\n    def __init__(self, csv_file, images_dir, transform=None):\n        \"\"\"\n        csv_file: CSV file with columns \"id_code\" and \"diagnosis\"\n        images_dir: Directory with images named <id_code>.png\n        transform: Image transforms to apply.\n        \"\"\"\n        self.data = pd.read_csv(csv_file)\n        self.images_dir = images_dir\n        self.transform = transform\n\n        # Load a pretrained ResNet18 up to a certain convolutional layer to extract spatial features.\n        # We remove the average pool and fc layers.\n        resnet = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n        modules = list(resnet.children())[:-2]  # keep up to last conv layer\n        self.feature_extractor = nn.Sequential(*modules)\n        self.feature_extractor.eval()  # set to eval mode\n        self.feature_extractor.to(device)\n\n    def __len__(self):\n        return len(self.data)\n    \n    def _build_grid_edges(self, height, width):\n        # Create a list of edges for a regular grid\n        # We connect each node to its 4 neighbors (if they exist)\n        edges = []\n        def node_idx(i, j):\n            return i * width + j\n        for i in range(height):\n            for j in range(width):\n                idx = node_idx(i, j)\n                # Right neighbor\n                if j < width - 1:\n                    edges.append([idx, node_idx(i, j+1)])\n                    edges.append([node_idx(i, j+1), idx])\n                # Down neighbor\n                if i < height - 1:\n                    edges.append([idx, node_idx(i+1, j)])\n                    edges.append([node_idx(i+1, j), idx])\n        return torch.tensor(edges, dtype=torch.long).t().contiguous()\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        img_id = row['id_code']\n        label = int(row['diagnosis'])  # integer label 0 to 4\n\n        img_path = os.path.join(self.images_dir, img_id + \".png\")\n        image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if image is None:\n            # Fallback: a black image\n            image = np.zeros((224,224,3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n\n        # Extract spatial features using the CNN extractor.\n        # image shape: (3,224,224). Add batch dimension.\n        with torch.no_grad():\n            x = image.unsqueeze(0).to(device)  # shape [1,3,224,224]\n            feature_map = self.feature_extractor(x)  # shape: [1, C, H, W]\n        feature_map = feature_map.squeeze(0).cpu()  # [C, H, W]\n        C, H, W = feature_map.shape\n        node_features = feature_map.view(C, H*W).t()  # shape: [H*W, C]\n\n        # Build graph edges based on grid connectivity.\n        edge_index = self._build_grid_edges(H, W)\n\n        # Create a PyG Data object.\n        data = Data(x=node_features, edge_index=edge_index)\n        # Attach the label (for whole–graph classification).\n        data.y = torch.tensor([label], dtype=torch.long)\n        return data\n\n####################################\n# 3. Create Dataset and DataLoaders (Graph Version)\n####################################\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# Use the same CSV and images directory (APTOS)\nfull_graph_dataset = APTOSGraphDataset(TRAIN_CSV, TRAIN_IMAGES_DIR, transform=train_transform)\nprint(\"Total samples in APTOS dataset:\", len(full_graph_dataset))\n\n# Optionally, you may visualize one sample's graph properties:\nsample = full_graph_dataset[0]\nprint(\"Graph sample: number of nodes =\", sample.x.size(0), \", number of edges =\", sample.edge_index.size(1))\nprint(\"Graph label:\", sample.y.item())\n\n# Split into training and validation (80/20)\ntrain_size = int(0.8 * len(full_graph_dataset))\nval_size = len(full_graph_dataset) - train_size\ntrain_dataset, val_dataset = torch.utils.data.random_split(full_graph_dataset, [train_size, val_size])\n\nbatch_size = 16  # You may adjust batch size (each sample is a graph)\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\n####################################\n# 4. Define the CNN-GAT Model\n####################################\nclass CNNGAT(nn.Module):\n    def __init__(self, in_channels, hidden_channels, num_classes, num_heads=4):\n        \"\"\"\n        in_channels: Dimension of node features (from CNN extractor)\n        hidden_channels: Hidden dimension in GAT layers.\n        num_classes: Number of output classes (5 for APTOS)\n        \"\"\"\n        super(CNNGAT, self).__init__()\n        self.gat1 = GATConv(in_channels, hidden_channels, heads=num_heads, dropout=0.3)\n        # Here, we combine the heads (by concatenation) so next layer input is hidden_channels*num_heads.\n        self.gat2 = GATConv(hidden_channels * num_heads, hidden_channels, heads=1, concat=False, dropout=0.3)\n        # Global pooling (mean) to get graph-level representation.\n        self.lin = nn.Linear(hidden_channels, num_classes)\n    \n    def forward(self, data):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        x = self.gat1(x, edge_index)\n        x = torch.relu(x)\n        x = self.gat2(x, edge_index)\n        x = torch.relu(x)\n        # Global mean pooling over nodes in each graph\n        x = global_mean_pool(x, batch)\n        x = self.lin(x)\n        return x\n\n# Instantiate model. The node feature dimension is the output of ResNet18 last conv.\n# You can check sample.x.shape[1] for that.\nin_channels = sample.x.size(1)\nmodel = CNNGAT(in_channels=in_channels, hidden_channels=64, num_classes=5, num_heads=4)\nmodel = model.to(device)\nprint(\"CNN-GAT model instantiated.\")\n\n####################################\n# 5. Loss, Optimizer, and Scheduler\n####################################\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n\n####################################\n# 6. Training Loop for 100 Epochs (Graph Model)\n####################################\nnum_epochs = 100\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n\n    for data in train_loader:\n        data = data.to(device)\n        optimizer.zero_grad()\n        outputs = model(data)\n        loss = criterion(outputs, data.y)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item() * data.num_graphs\n        _, preds = torch.max(outputs, 1)\n        total_train += data.num_graphs\n        correct_train += (preds == data.y).sum().item()\n    scheduler.step()\n    train_loss = running_loss / total_train\n    train_acc = 100.0 * correct_train / total_train\n\n    # Validation\n    model.eval()\n    running_val_loss = 0.0\n    correct_val = 0\n    total_val = 0\n    all_preds = []\n    all_labels = []\n    with torch.no_grad():\n        for data in val_loader:\n            data = data.to(device)\n            outputs = model(data)\n            loss = criterion(outputs, data.y)\n            running_val_loss += loss.item() * data.num_graphs\n            _, preds = torch.max(outputs, 1)\n            total_val += data.num_graphs\n            correct_val += (preds == data.y).sum().item()\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(data.y.cpu().numpy())\n    val_loss = running_val_loss / total_val\n    val_acc = 100.0 * correct_val / total_val\n\n    # Compute confusion matrix and metrics\n    cm = confusion_matrix(all_labels, all_preds)\n    precision = precision_score(all_labels, all_preds, average='weighted', zero_division=0)\n    recall = recall_score(all_labels, all_preds, average='weighted', zero_division=0)\n    f1 = f1_score(all_labels, all_preds, average='weighted', zero_division=0)\n    overall_acc = np.mean(np.array(all_preds) == np.array(all_labels)) * 100\n\n    # Detailed logging\n    print(f\"Epoch [{epoch+1}/{num_epochs}] - \"\n          f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% | \"\n          f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n    print(\"Confusion Matrix:\")\n    print(cm)\n    print(f\"Metrics — Precision: {precision:.4f}, Recall: {recall:.4f}, F1: {f1:.4f}, Accuracy: {overall_acc:.2f}%\\n\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Uncomment the next line if you're running in a Jupyter/Colab environment:\n%env CUDA_LAUNCH_BLOCKING=1\n\n####################################\n# 1. Import Libraries and Set Seed\n####################################\nimport os\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport pandas as pd\nimport random\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom torchvision import transforms, models\nfrom torch.utils.data import Dataset, random_split\n\n# Torch Geometric Imports\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader  # Use PyG DataLoader\nfrom torch_geometric.nn import GATConv, global_mean_pool\n\n# For reproducibility\nseed = 42\ntorch.manual_seed(seed)\nnp.random.seed(seed)\nrandom.seed(seed)\n\n####################################\n# 2. Define Device\n####################################\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n####################################\n# 3. Data Paths, Transforms, and Dataset\n####################################\n# Update these paths as per your dataset structure on Kaggle\nTRAIN_CSV = \"/kaggle/input/ocular-disease-recognition-odir5k/full_df.csv\"\nTRAIN_IMAGES_DIR = \"/kaggle/input/ocular-disease-recognition-odir5k/preprocessed_images\"\n\n# Define image transforms\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\n# For your CSV, targets are lists with 8 elements (thus, 8 classes: 0-7)\nnum_classes = 8\n\nclass OcularGraphDataset(Dataset):\n    def __init__(self, csv_file, images_dir, transform=None):\n        \"\"\"\n        csv_file: CSV with ocular data. Expected columns include \"ID\" and \"target\".\n                  The \"target\" column should be a string like \"[1, 0, 0, 0, 0, 0, 0, 0]\".\n                  np.argmax converts it to an integer label.\n        images_dir: Directory with images named as <ID>.png.\n        transform: Image transforms to apply.\n        \"\"\"\n        self.data = pd.read_csv(csv_file)\n        self.images_dir = images_dir\n        self.transform = transform\n\n        # Use a pretrained ResNet18 and remove its avgpool and fc layers for feature extraction.\n        resnet = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n        modules = list(resnet.children())[:-2]\n        self.feature_extractor = nn.Sequential(*modules)\n        self.feature_extractor.eval()\n        self.feature_extractor.to(device)\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def _build_grid_edges(self, height, width):\n        edges = []\n        def node_idx(i, j):\n            return i * width + j\n        for i in range(height):\n            for j in range(width):\n                idx = node_idx(i, j)\n                if j < width - 1:\n                    edges.append([idx, node_idx(i, j+1)])\n                    edges.append([node_idx(i, j+1), idx])\n                if i < height - 1:\n                    edges.append([idx, node_idx(i+1, j)])\n                    edges.append([node_idx(i+1, j), idx])\n        return torch.tensor(edges, dtype=torch.long).t().contiguous()\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        img_id = str(row['ID'])\n        target_str = row['target']\n        try:\n            target_list = eval(target_str) if isinstance(target_str, str) else target_str\n        except Exception as e:\n            raise ValueError(f\"Error parsing target for index {idx}: {target_str}\") from e\n        target_arr = np.array(target_list)\n        label = int(np.argmax(target_arr))\n        # Ensure label is in the expected range [0, num_classes-1]\n        assert 0 <= label < num_classes, f\"Label {label} out-of-bound at index {idx}\"\n        \n        img_path = os.path.join(self.images_dir, img_id + \".png\")\n        image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if image is None:\n            image = np.zeros((224,224,3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n        \n        # Extract spatial features using the CNN feature extractor.\n        with torch.no_grad():\n            x = image.unsqueeze(0).to(device)  # [1, 3, 224, 224]\n            feature_map = self.feature_extractor(x)  # [1, C, H, W]\n        feature_map = feature_map.squeeze(0).cpu()  # [C, H, W]\n        C, H, W = feature_map.shape\n        node_features = feature_map.view(C, H * W).t()  # [H*W, C]\n        edge_index = self._build_grid_edges(H, W)\n        data_obj = Data(x=node_features, edge_index=edge_index)\n        data_obj.y = torch.tensor([label], dtype=torch.long)\n        return data_obj\n\nprint(\"Building graph dataset for Ocular data...\")\nfull_graph_dataset = OcularGraphDataset(TRAIN_CSV, TRAIN_IMAGES_DIR, transform=train_transform)\nprint(\"Total samples in dataset:\", len(full_graph_dataset))\nsample = full_graph_dataset[0]\nprint(\"Graph sample: nodes =\", sample.x.size(0), \", edges =\", sample.edge_index.size(1))\nprint(\"Graph label:\", sample.y.item())\n\n# Split dataset: 80% training, 20% validation\ntrain_size = int(0.8 * len(full_graph_dataset))\nval_size = len(full_graph_dataset) - train_size\ntrain_dataset, val_dataset = random_split(full_graph_dataset, [train_size, val_size])\nbatch_size = 16\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\n####################################\n# 4. Define the CNN-GAT Model\n####################################\nclass CNNGAT(nn.Module):\n    def __init__(self, in_channels, hidden_channels, num_classes, num_heads=4):\n        super(CNNGAT, self).__init__()\n        self.gat1 = GATConv(in_channels, hidden_channels, heads=num_heads, dropout=0.3)\n        # Output of gat1 has dimension hidden_channels * num_heads\n        self.gat2 = GATConv(hidden_channels * num_heads, hidden_channels, heads=1, concat=False, dropout=0.3)\n        self.lin = nn.Linear(hidden_channels, num_classes)\n    \n    def forward(self, data):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        x = self.gat1(x, edge_index)\n        x = torch.relu(x)\n        x = self.gat2(x, edge_index)\n        x = torch.relu(x)\n        x = global_mean_pool(x, batch)\n        x = self.lin(x)\n        return x\n\nin_channels = sample.x.size(1)\ncnn_gat_model = CNNGAT(in_channels, hidden_channels=64, num_classes=num_classes, num_heads=4).to(device)\nprint(\"CNN-GAT model instantiated.\")\n\n####################################\n# 5. Check Target Labels and Model Outputs\n####################################\n# --- 1. Check target labels in one batch ---\ndata = next(iter(train_loader))\ndata = data.to(device)\nprint(\"\\n--- Checking target labels in one batch ---\")\nprint(\"Target tensor shape:\", data.y.shape)  # Expected shape: [batch_size]\ntargets = data.y.squeeze()\nprint(\"Unique target labels in this batch:\", torch.unique(targets))\n\n# --- 2. Check model output dimensions ---\noutputs = cnn_gat_model(data)\nprint(\"\\n--- Checking model output dimensions ---\")\nprint(\"Model output shape:\", outputs.shape)  # Expected: [batch_size, 8]\n\n# --- 3. Inspect a few training iterations ---\nprint(\"\\n--- Inspecting a few training iterations ---\")\nfor i, data in enumerate(train_loader):\n    data = data.to(device)\n    outputs = cnn_gat_model(data)\n    if i < 3:\n        print(f\"\\nBatch {i}:\")\n        print(\"  Output shape:\", outputs.shape)\n        print(\"  Target shape:\", data.y.shape)\n        print(\"  Unique targets:\", torch.unique(data.y.squeeze()))\n    if i == 3:\n        break\n\n####################################\n# 6. Fine-Tuning Training Loop\n####################################\nprint(\"\\n--- Starting supervised fine-tuning ---\")\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(cnn_gat_model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n\nnum_epochs = 2  # Adjust as needed\n\nfor epoch in range(num_epochs):\n    cnn_gat_model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n    \n    # Training loop\n    for data in train_loader:\n        data = data.to(device)\n        optimizer.zero_grad()\n        outputs = cnn_gat_model(data)\n        loss = criterion(outputs, data.y)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * data.num_graphs\n        _, preds = torch.max(outputs, 1)\n        total_train += data.num_graphs\n        correct_train += (preds == data.y).sum().item()\n    \n    scheduler.step()\n    train_loss = running_loss / total_train\n    train_acc = 100.0 * correct_train / total_train\n\n    # Validation loop\n    cnn_gat_model.eval()\n    running_val_loss = 0.0\n    correct_val = 0\n    total_val = 0\n    all_preds = []\n    all_labels = []\n    with torch.no_grad():\n        for data in val_loader:\n            data = data.to(device)\n            outputs = cnn_gat_model(data)\n            loss = criterion(outputs, data.y)\n            running_val_loss += loss.item() * data.num_graphs\n            _, preds = torch.max(outputs, 1)\n            total_val += data.num_graphs\n            correct_val += (preds == data.y).sum().item()\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(data.y.cpu().numpy())\n    \n    val_loss = running_val_loss / total_val\n    val_acc = 100.0 * correct_val / total_val\n\n    print(f\"Epoch [{epoch+1}/{num_epochs}] - Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% | \" +\n          f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n\nprint(\"Supervised fine-tuning complete!\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Uncomment the next line if running in a notebook to enforce synchronous CUDA error checking.\n%env CUDA_LAUNCH_BLOCKING=1\n\n####################################\n# 1. Import Libraries and Set Seed\n####################################\nimport os\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport pandas as pd\nimport random\nfrom PIL import Image\nfrom torchvision import transforms, models\nfrom torch.utils.data import Dataset, random_split\n\n# Torch Geometric Imports\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader  # PyG DataLoader\nfrom torch_geometric.nn import GATConv, global_mean_pool\n\n# Sklearn metrics for evaluation\nfrom sklearn.metrics import precision_score, recall_score, f1_score\n\n# For reproducibility\nseed = 42\ntorch.manual_seed(seed)\nnp.random.seed(seed)\nrandom.seed(seed)\n\n####################################\n# 2. Define Device\n####################################\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n####################################\n# 3. Data Paths, Transforms, and Dataset\n####################################\n# Update these paths according to your dataset structure on Kaggle\nTRAIN_CSV = \"/kaggle/input/ocular-disease-recognition-odir5k/full_df.csv\"\nTRAIN_IMAGES_DIR = \"/kaggle/input/ocular-disease-recognition-odir5k/preprocessed_images\"\n\n# Define image transforms for training (with augmentation) and validation\ntrain_transform = transforms.Compose([\n    transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\n# For this dataset, targets are lists with 8 elements (i.e. 8 classes: indices 0-7)\nnum_classes = 8\n\nclass OcularGraphDataset(Dataset):\n    def __init__(self, csv_file, images_dir, transform=None):\n        \"\"\"\n        csv_file: CSV with ocular data. Expected columns include \"ID\" and \"target\".\n                  The \"target\" column should be a string like \"[1, 0, 0, 0, 0, 0, 0, 0]\".\n                  We convert it to an integer label via np.argmax.\n        images_dir: Directory with images named as <ID>.png.\n        transform: Image transforms to apply.\n        \"\"\"\n        self.data = pd.read_csv(csv_file)\n        self.images_dir = images_dir\n        self.transform = transform\n\n        # Load a pretrained ResNet18, remove its avgpool and fc layers, for feature extraction.\n        resnet = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n        modules = list(resnet.children())[:-2]\n        self.feature_extractor = nn.Sequential(*modules)\n        self.feature_extractor.eval()\n        self.feature_extractor.to(device)\n    \n    def __len__(self):\n        return len(self.data)\n    \n    def _build_grid_edges(self, height, width):\n        edges = []\n        def node_idx(i, j):\n            return i * width + j\n        for i in range(height):\n            for j in range(width):\n                idx = node_idx(i, j)\n                if j < width - 1:\n                    edges.append([idx, node_idx(i, j+1)])\n                    edges.append([node_idx(i, j+1), idx])\n                if i < height - 1:\n                    edges.append([idx, node_idx(i+1, j)])\n                    edges.append([node_idx(i+1, j), idx])\n        return torch.tensor(edges, dtype=torch.long).t().contiguous()\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        img_id = str(row['ID'])\n        target_str = row['target']\n        try:\n            target_list = eval(target_str) if isinstance(target_str, str) else target_str\n        except Exception as e:\n            raise ValueError(f\"Error parsing target for index {idx}: {target_str}\") from e\n        target_arr = np.array(target_list)\n        label = int(np.argmax(target_arr))\n        assert 0 <= label < num_classes, f\"Label {label} out-of-bound at index {idx}\"\n        \n        img_path = os.path.join(self.images_dir, img_id + \".png\")\n        image = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        if image is None:\n            image = np.zeros((224,224,3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n        \n        # Extract spatial features using the CNN feature extractor.\n        with torch.no_grad():\n            x = image.unsqueeze(0).to(device)  # [1, 3, 224, 224]\n            feature_map = self.feature_extractor(x)  # [1, C, H, W]\n        feature_map = feature_map.squeeze(0).cpu()  # [C, H, W]\n        C, H, W = feature_map.shape\n        node_features = feature_map.view(C, H * W).t()  # [H*W, C]\n        edge_index = self._build_grid_edges(H, W)\n        data_obj = Data(x=node_features, edge_index=edge_index)\n        data_obj.y = torch.tensor([label], dtype=torch.long)\n        return data_obj\n\nprint(\"Building graph dataset for Ocular data...\")\nfull_graph_dataset = OcularGraphDataset(TRAIN_CSV, TRAIN_IMAGES_DIR, transform=train_transform)\nprint(\"Total samples in dataset:\", len(full_graph_dataset))\nsample = full_graph_dataset[0]\nprint(\"Graph sample: nodes =\", sample.x.size(0), \", edges =\", sample.edge_index.size(1))\nprint(\"Graph label:\", sample.y.item())\n\n# Split dataset: 80% training, 20% validation\ntrain_size = int(0.8 * len(full_graph_dataset))\nval_size = len(full_graph_dataset) - train_size\ntrain_dataset, val_dataset = random_split(full_graph_dataset, [train_size, val_size])\nbatch_size = 16\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\n####################################\n# 4. Define the DGI Module for Self-Supervised Pretraining\n####################################\nclass DGI(nn.Module):\n    def __init__(self, in_channels, hidden_channels, num_heads=4):\n        \"\"\"\n        in_channels: Dimension of node features.\n        hidden_channels: Hidden dimension for the encoder.\n        \"\"\"\n        super(DGI, self).__init__()\n        # GATConv returns concatenated output: dimension = hidden_channels * num_heads\n        self.encoder = GATConv(in_channels, hidden_channels, heads=num_heads, dropout=0.3)\n        self.readout = global_mean_pool\n        # Discriminator adjusted to accept dimension hidden_channels * num_heads\n        self.disc = nn.Bilinear(hidden_channels * num_heads, hidden_channels * num_heads, 1)\n    \n    def forward(self, x, edge_index, batch):\n        h_pos = self.encoder(x, edge_index)  # [num_nodes, hidden_channels * num_heads]\n        h_pos = torch.relu(h_pos)\n        summary = self.readout(h_pos, batch)  # [num_graphs, hidden_channels * num_heads]\n        summary = torch.sigmoid(summary)\n        # Negative sample: shuffle node features\n        perm = torch.randperm(x.size(0))\n        x_neg = x[perm]\n        h_neg = self.encoder(x_neg, edge_index)\n        h_neg = torch.relu(h_neg)\n        pos_score = self.disc(h_pos, summary[batch])\n        neg_score = self.disc(h_neg, summary[batch])\n        return pos_score, neg_score\n\ndef dgi_loss(pos_score, neg_score):\n    pos_loss = -torch.mean(torch.log(torch.sigmoid(pos_score) + 1e-15))\n    neg_loss = -torch.mean(torch.log(1 - torch.sigmoid(neg_score) + 1e-15))\n    return pos_loss + neg_loss\n\n####################################\n# 5. Self-Supervised Pretraining (DGI)\n####################################\nprint(\"\\n--- Starting Self-Supervised Pretraining (DGI) ---\")\nin_channels = sample.x.size(1)\nhidden_channels = 64\ndgi_model = DGI(in_channels, hidden_channels, num_heads=4).to(device)\ndgi_optimizer = optim.Adam(dgi_model.parameters(), lr=1e-4, weight_decay=1e-4)\nnum_dgi_epochs = 100  # Adjust as needed (here we use 20 epochs for self-supervised pretraining)\n\nfor epoch in range(num_dgi_epochs):\n    dgi_model.train()\n    total_loss = 0.0\n    total_batches = 0\n    for data in train_loader:\n        data = data.to(device)\n        dgi_optimizer.zero_grad()\n        pos_score, neg_score = dgi_model(data.x, data.edge_index, data.batch)\n        loss = dgi_loss(pos_score, neg_score)\n        loss.backward()\n        dgi_optimizer.step()\n        total_loss += loss.item()\n        total_batches += 1\n    avg_loss = total_loss / total_batches\n    print(f\"DGI Epoch [{epoch+1}/{num_dgi_epochs}] - Loss: {avg_loss:.4f}\")\nprint(\"Self-supervised pretraining complete!\\n\")\n\n# Transfer pretrained encoder weights to CNN-GAT's first layer (if available)\n# This transfer assumes that the internal structure of the GATConv layer is identical.\n####################################\n# 6. Define the CNN-GAT Model for Supervised Fine-Tuning\n####################################\nclass CNNGAT(nn.Module):\n    def __init__(self, in_channels, hidden_channels, num_classes, num_heads=4):\n        \"\"\"\n        in_channels: Dimension of node features from CNN extractor.\n        hidden_channels: Hidden dimension in GAT layers.\n        num_classes: Number of output classes.\n        \"\"\"\n        super(CNNGAT, self).__init__()\n        self.gat1 = GATConv(in_channels, hidden_channels, heads=num_heads, dropout=0.3)\n        self.gat2 = GATConv(hidden_channels * num_heads, hidden_channels, heads=1, concat=False, dropout=0.3)\n        self.lin = nn.Linear(hidden_channels, num_classes)\n    \n    def forward(self, data):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        x = self.gat1(x, edge_index)\n        x = torch.relu(x)\n        x = self.gat2(x, edge_index)\n        x = torch.relu(x)\n        x = global_mean_pool(x, batch)\n        x = self.lin(x)\n        return x\n\ncnn_gat_model = CNNGAT(in_channels, hidden_channels=64, num_classes=num_classes, num_heads=4).to(device)\nprint(\"CNN-GAT model instantiated.\")\n\n# Transfer weights from DGI encoder to CNN-GAT's first layer\ntry:\n    cnn_gat_model.gat1.lin.weight.data = dgi_model.encoder.lin.weight.data.clone()\n    if dgi_model.encoder.lin.bias is not None:\n        cnn_gat_model.gat1.lin.bias.data = dgi_model.encoder.lin.bias.data.clone()\n    print(\"Transferred weights from DGI encoder to CNN-GAT.\")\nexcept Exception as e:\n    print(\"Weight transfer failed:\", e)\n\n####################################\n# 7. Supervised Fine-Tuning Training Loop (100 Epochs)\n####################################\nprint(\"\\n--- Starting Supervised Fine-Tuning ---\")\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(cnn_gat_model.parameters(), lr=1e-4, weight_decay=1e-4)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n\nnum_supervised_epochs = 100  # Set to 100 epochs\n\nfor epoch in range(num_supervised_epochs):\n    cnn_gat_model.train()\n    running_loss = 0.0\n    correct_train = 0\n    total_train = 0\n    \n    for data in train_loader:\n        data = data.to(device)\n        optimizer.zero_grad()\n        outputs = cnn_gat_model(data)\n        loss = criterion(outputs, data.y)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * data.num_graphs\n        _, preds = torch.max(outputs, 1)\n        total_train += data.num_graphs\n        correct_train += (preds == data.y).sum().item()\n    \n    scheduler.step()\n    train_loss = running_loss / total_train\n    train_acc = 100.0 * correct_train / total_train\n\n    # Validation\n    cnn_gat_model.eval()\n    running_val_loss = 0.0\n    correct_val = 0\n    total_val = 0\n    all_preds = []\n    all_labels = []\n    with torch.no_grad():\n        for data in val_loader:\n            data = data.to(device)\n            outputs = cnn_gat_model(data)\n            loss = criterion(outputs, data.y)\n            running_val_loss += loss.item() * data.num_graphs\n            _, preds = torch.max(outputs, 1)\n            total_val += data.num_graphs\n            correct_val += (preds == data.y).sum().item()\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(data.y.cpu().numpy())\n    \n    val_loss = running_val_loss / total_val\n    val_acc = 100.0 * correct_val / total_val\n    \n    # Compute additional metrics without printing the raw confusion matrix\n    precision = precision_score(all_labels, all_preds, average='weighted', zero_division=0)\n    recall = recall_score(all_labels, all_preds, average='weighted', zero_division=0)\n    f1 = f1_score(all_labels, all_preds, average='weighted', zero_division=0)\n    overall_acc = 100.0 * np.mean(np.array(all_preds) == np.array(all_labels))\n    \n    print(f\"Epoch [{epoch+1}/{num_supervised_epochs}] - \"\n          f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}% | \"\n          f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}% | \"\n          f\"Precision: {precision:.4f}, Recall: {recall:.4f}, F1: {f1:.4f}, Overall Acc: {overall_acc:.2f}%\")\n\nprint(\"Supervised fine-tuning complete!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-24T19:25:06.102235Z","iopub.execute_input":"2025-03-24T19:25:06.102536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}