{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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\n# import numpy as np # linear algebra\n# import 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\n# import os\n# for 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,"execution":{"iopub.status.busy":"2025-08-24T21:26:49.265653Z","iopub.execute_input":"2025-08-24T21:26:49.265941Z","iopub.status.idle":"2025-08-24T21:26:49.281592Z","shell.execute_reply.started":"2025-08-24T21:26:49.265917Z","shell.execute_reply":"2025-08-24T21:26:49.280946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================\n# All Required Imports\n# ============================\n\nimport os\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torchvision import transforms\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\n# ML utilities\nfrom sklearn.cluster import KMeans\nfrom sklearn.metrics import silhouette_score\nfrom sklearn.decomposition import PCA\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:59:55.726309Z","iopub.execute_input":"2025-08-25T08:59:55.726557Z","iopub.status.idle":"2025-08-25T09:00:03.691548Z","shell.execute_reply.started":"2025-08-25T08:59:55.726533Z","shell.execute_reply":"2025-08-25T09:00:03.690942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================\n# Check Imports & Versions\n# ============================\n\nprint(\"Imports successful!\")\n\nprint(\"\\nLibrary versions:\")\nprint(\"os (builtin)\")\nprint(\"torch:\", torch.__version__)\nprint(\"torchvision:\", __import__('torchvision').__version__)\nprint(\"PIL:\", __import__('PIL').__version__)\nprint(\"matplotlib:\", plt.matplotlib.__version__)\nprint(\"scikit-learn:\", __import__('sklearn').__version__)\n\nprint(\"\\nCUDA available:\", torch.cuda.is_available())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T09:00:34.407249Z","iopub.execute_input":"2025-08-25T09:00:34.407923Z","iopub.status.idle":"2025-08-25T09:00:34.47833Z","shell.execute_reply.started":"2025-08-25T09:00:34.4079Z","shell.execute_reply":"2025-08-25T09:00:34.477597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 0. Safe handling for very large images\n# =========================================\nImage.MAX_IMAGE_PIXELS = None  # bypass DecompressionBombError\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T09:00:39.540943Z","iopub.execute_input":"2025-08-25T09:00:39.541217Z","iopub.status.idle":"2025-08-25T09:00:39.544836Z","shell.execute_reply.started":"2025-08-25T09:00:39.541198Z","shell.execute_reply":"2025-08-25T09:00:39.543975Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1. Custom Dataset\n# ======================\nclass UBCOceanDataset(Dataset):\n    def __init__(self, image_dir, transform=None):\n        self.image_dir = image_dir\n        self.transform = transform\n        self.images = [f for f in os.listdir(image_dir) if f.lower().endswith(('.png','.jpg','.jpeg'))]\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.image_dir, self.images[idx])\n        image = Image.open(img_path).convert(\"L\")  # grayscale\n        if self.transform:\n            image = self.transform(image)\n        return image, 0  # dummy label for unsupervised\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T09:00:45.139192Z","iopub.execute_input":"2025-08-25T09:00:45.139545Z","iopub.status.idle":"2025-08-25T09:00:45.144579Z","shell.execute_reply.started":"2025-08-25T09:00:45.139524Z","shell.execute_reply":"2025-08-25T09:00:45.143853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ======================\n# # Test Custom Dataset\n# # ======================\n\n# # Define transform\n# transform = transforms.Compose([\n#     transforms.Resize((32, 32)),\n#     transforms.ToTensor()\n# ])\n\n# # Point to your dataset directory\n# image_dir = \"/kaggle/input/UBC-OCEAN/train_images\"\n\n# # Create dataset\n# dataset = UBCOceanDataset(image_dir, transform=transform)\n\n# # Check length\n# print(\"Number of images found:\", len(dataset))\n\n# # Try loading one sample\n# if len(dataset) > 0:\n#     img, label = dataset[0]\n#     print(\"Image shape:\", img.shape)\n#     print(\"Label:\", label)\n\n#     # Plot sample\n#     plt.imshow(img[0], cmap=\"gray\")\n#     plt.title(\"Sample Image (grayscale)\")\n#     plt.axis(\"off\")\n#     plt.show()\n# else:\n#     print(\"No images found in directory:\", image_dir)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.262481Z","iopub.status.idle":"2025-08-24T21:26:56.262725Z","shell.execute_reply.started":"2025-08-24T21:26:56.262605Z","shell.execute_reply":"2025-08-24T21:26:56.262615Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_dir = '/kaggle/input/UBC-OCEAN/train_images'\nprint(\"Exists?\", os.path.exists(image_dir))\nif os.path.exists(image_dir):\n    print(\"Number of images:\", len(os.listdir(image_dir)))\n    print(\"First 5 images:\", os.listdir(image_dir)[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T09:00:52.031712Z","iopub.execute_input":"2025-08-25T09:00:52.032662Z","iopub.status.idle":"2025-08-25T09:00:52.086995Z","shell.execute_reply.started":"2025-08-25T09:00:52.032637Z","shell.execute_reply":"2025-08-25T09:00:52.086262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((32,32)),\n    transforms.ToTensor()\n])\n\ndataset = UBCOceanDataset('/kaggle/input/UBC-OCEAN/train_images', transform=transform)\nprint(\"Dataset length:\", len(dataset))\n\nfor i in range(min(5, len(dataset))):\n    try:\n        img, label = dataset[i]\n        print(f\"Sample {i}: img shape={img.shape}, label={label}\")\n    except Exception as e:\n        print(f\"❌ Error at sample {i}: {e}\")\n        break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T09:02:57.903517Z","iopub.execute_input":"2025-08-25T09:02:57.903795Z","iopub.status.idle":"2025-08-25T09:06:08.857414Z","shell.execute_reply.started":"2025-08-25T09:02:57.903775Z","shell.execute_reply":"2025-08-25T09:06:08.856653Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UBCOceanDataset(Dataset):\n    def __init__(self, root, transform=None):\n        self.root = root\n        self.transform = transform\n        self.img_paths = [os.path.join(root, f) for f in os.listdir(root) if f.endswith(\".jpg\")]\n        self.labels = [0]*len(self.img_paths)  # <-- replace with real labels\n        print(f\"Found {len(self.img_paths)} images at {root}\")\n\n    def __getitem__(self, idx):\n        img_path = self.img_paths[idx]\n        try:\n            img = Image.open(img_path).convert(\"RGB\")\n        except Exception as e:\n            print(f\"❌ Corrupted image: {img_path}, error={e}\")\n            return torch.zeros(3,32,32), -1  # dummy return\n\n        if self.transform:\n            img = self.transform(img)\n\n        label = self.labels[idx]\n        return img, label\n\n    def __len__(self):\n        return len(self.img_paths)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T09:06:50.830239Z","iopub.execute_input":"2025-08-25T09:06:50.830529Z","iopub.status.idle":"2025-08-25T09:06:50.836174Z","shell.execute_reply.started":"2025-08-25T09:06:50.830511Z","shell.execute_reply":"2025-08-25T09:06:50.835574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2. DataLoader\n# ======================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)\n\n# transform = transforms.Compose([\n#     transforms.Resize((32,32)),\n#     transforms.ToTensor()\n# ])\n\n# dataset = UBCOceanDataset('/kaggle/input/UBC-OCEAN/train_images', transform=transform)\n# Use subset for testing speed\n# dataset = Subset(dataset, range(min(200, len(dataset))))  # first 200 images\nloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=0, pin_memory=False)\n\n\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:43:33.846535Z","iopub.execute_input":"2025-08-25T08:43:33.846853Z","iopub.status.idle":"2025-08-25T08:43:33.857866Z","shell.execute_reply.started":"2025-08-25T08:43:33.846833Z","shell.execute_reply":"2025-08-25T08:43:33.857286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Test \n# print(\"Testing Dataset\")\n# print(\"Dataset length:\", len(dataset))\n# sample = dataset[0]\n# print(type(sample), len(sample))\n# img, label = sample\n# print(\"Image shape:\", img.shape)\n# print(\"Label:\", label)\n\n# print(\"Testing Loader\")\n# # 2. Check dataloader length\n# print(\"Number of batches:\", len(loader))\n\n# # 3. Get one batch\n# images, labels = next(iter(loader))\n# print(\"Batch image shape:\", images.shape)   # Expect [B, C, H, W]\n# print(\"Batch label shape:\", labels.shape)   # Expect [B] or [B, ...]\n\n# # 4. Visualize a few samples\n# grid = torchvision.utils.make_grid(images[:8], nrow=4)  # show first 8\n# plt.imshow(grid.permute(1, 2, 0))  # convert C,H,W → H,W,C\n# plt.title(f\"Labels: {labels[:8].tolist()}\")\n# plt.axis(\"off\")\n# plt.show()\n\n\n\n# # 5. Manually inspect one sample\n# img, label = dataset[0]\n# print(\"Single image shape:\", img.shape)\n# print(\"Single label:\", label)\n# plt.imshow(img.permute(1, 2, 0))  # if image is Tensor\n# plt.show()\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3. Patch Splitting\n# ======================\ndef split_patches(x, patch_size=8):\n    B, C, H, W = x.shape\n    patches = x.unfold(2, patch_size, patch_size).unfold(3, patch_size, patch_size)\n    patches = patches.contiguous().view(B, C, -1, patch_size, patch_size)\n    patches = patches.permute(0,2,1,3,4)  # B, num_patches, C, patch_size, patch_size\n    return patches\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:44:25.538984Z","iopub.execute_input":"2025-08-25T08:44:25.539279Z","iopub.status.idle":"2025-08-25T08:44:25.543924Z","shell.execute_reply.started":"2025-08-25T08:44:25.539257Z","shell.execute_reply":"2025-08-25T08:44:25.543145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4. Patch Encoder\n# ======================\nclass PatchEncoder(nn.Module):\n    def __init__(self, embed_dim=64):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(1, 32, 3, padding=1), nn.ReLU(),\n            nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(),\n            nn.AdaptiveAvgPool2d(1)\n        )\n        self.fc = nn.Linear(64, embed_dim)\n        \n    def forward(self, patches):\n        B, N, C, H, W = patches.shape\n        patches = patches.view(B*N, C, H, W)\n        out = self.conv(patches).view(B*N, -1)\n        out = self.fc(out)\n        return out.view(B, N, -1)  # B, num_patches, embed_dim\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:44:35.462092Z","iopub.execute_input":"2025-08-25T08:44:35.462387Z","iopub.status.idle":"2025-08-25T08:44:35.468264Z","shell.execute_reply.started":"2025-08-25T08:44:35.462365Z","shell.execute_reply":"2025-08-25T08:44:35.467486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ======================\n# # Test Patch Encoder\n# # ======================\n\n\n\n# # Create encoder\n# encoder = PatchEncoder(embed_dim=64).to(device)\n\n# # Take a batch of images and split into patches\n# imgs, _ = next(iter(loader))                # [B, 1, 32, 32]\n# imgs = imgs.to(device)\n# patches = split_patches(imgs, patch_size=8).to(device) # [B, 16, 1, 8, 8]\n\n# # Pass through encoder\n# patch_embeddings = encoder(patches)\n\n# print(\"Patches input shape:\", patches.shape)             # [B, 16, 1, 8, 8]\n# print(\"Patch embeddings shape:\", patch_embeddings.shape) # [B, 16, 64]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.269189Z","iopub.status.idle":"2025-08-24T21:26:56.269505Z","shell.execute_reply.started":"2025-08-24T21:26:56.269345Z","shell.execute_reply":"2025-08-24T21:26:56.26936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5. Graph Encoder\n# ======================\nclass GraphEncoder(nn.Module):\n    def __init__(self, embed_dim=64):\n        super().__init__()\n        self.fc1 = nn.Linear(embed_dim, embed_dim)\n        self.fc2 = nn.Linear(embed_dim, embed_dim)\n        \n    def forward(self, x):\n        h = F.relu(self.fc1(x))\n        h = F.relu(self.fc2(h))\n        return h.mean(dim=1)  # global embedding per image\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:44:41.176297Z","iopub.execute_input":"2025-08-25T08:44:41.176576Z","iopub.status.idle":"2025-08-25T08:44:41.181798Z","shell.execute_reply.started":"2025-08-25T08:44:41.176557Z","shell.execute_reply":"2025-08-25T08:44:41.1809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # ======================\n# # Test Graph Encoder\n# # ======================\n\n# # Create graph encoder\n# graph_encoder = GraphEncoder(embed_dim=64)\n\n# # Assume patch embeddings already computed from PatchEncoder\n# # (shape: [B, N, 64])\n# global_embeddings = graph_encoder(patch_embeddings)\n\n# print(\"Patch embeddings shape:\", patch_embeddings.shape)   # [B, N, 64]\n# print(\"Global embeddings shape:\", global_embeddings.shape) # [B, 64]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.271501Z","iopub.status.idle":"2025-08-24T21:26:56.271789Z","shell.execute_reply.started":"2025-08-24T21:26:56.271627Z","shell.execute_reply":"2025-08-24T21:26:56.271642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 6. Edge-Contrastive Loss (Vectorized)\n# ======================\ndef edge_contrastive_loss(embeddings):\n    B, N, D = embeddings.shape\n    normed = embeddings / (embeddings.norm(dim=2, keepdim=True) + 1e-8)\n    sim_matrix = torch.bmm(normed, normed.transpose(1,2))  # B x N x N\n    pos_mask = torch.eye(N, device=sim_matrix.device).bool()\n    neg_mask = ~pos_mask\n    pos_loss = -torch.log(sim_matrix[:,pos_mask] + 1e-8).mean()\n    neg_loss = -torch.log(1 - sim_matrix[:,neg_mask] + 1e-8).mean()\n    return pos_loss + neg_loss\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:44:46.898264Z","iopub.execute_input":"2025-08-25T08:44:46.898935Z","iopub.status.idle":"2025-08-25T08:44:46.90338Z","shell.execute_reply.started":"2025-08-25T08:44:46.898911Z","shell.execute_reply":"2025-08-25T08:44:46.902796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Dummy embeddings: batch=4, patches=8, dim=64\ndummy_embeddings = torch.randn(4, 8, 64)\n\nloss = edge_contrastive_loss(dummy_embeddings)\nprint(\"Loss value:\", loss.item())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:44:52.441314Z","iopub.execute_input":"2025-08-25T08:44:52.441584Z","iopub.status.idle":"2025-08-25T08:44:52.492303Z","shell.execute_reply.started":"2025-08-25T08:44:52.441564Z","shell.execute_reply":"2025-08-25T08:44:52.491486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 7. PG-ECA Model\n# ======================\nclass PGECA(nn.Module):\n    def __init__(self, embed_dim=64, patch_size=8):\n        super().__init__()\n        self.encoder = PatchEncoder(embed_dim)\n        self.graph_encoder = GraphEncoder(embed_dim)\n        self.decoder = nn.Sequential(\n            nn.Linear(embed_dim, 64),\n            nn.ReLU(),\n            nn.Linear(64, patch_size*patch_size),\n            nn.Sigmoid()\n        )\n        self.patch_size = patch_size\n        \n    def forward(self, x):\n        patches = split_patches(x, self.patch_size)\n        patch_embed = self.encoder(patches)\n        global_embed = self.graph_encoder(patch_embed)\n        recon = self.decoder(global_embed).view(-1,1,self.patch_size,self.patch_size)\n        return patch_embed, global_embed, recon\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:45:04.986285Z","iopub.execute_input":"2025-08-25T08:45:04.987048Z","iopub.status.idle":"2025-08-25T08:45:04.992002Z","shell.execute_reply.started":"2025-08-25T08:45:04.987021Z","shell.execute_reply":"2025-08-25T08:45:04.99132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test for 7. PG-ECA Model with decoder\nmodel = PGECA(embed_dim=64, patch_size=8)\n\n# Dummy grayscale input: batch=2, channel=1, 32x32 images\ndummy = torch.randn(2, 1, 32, 32)\n\n# Forward pass\npatch_embed, global_embed, recon = model(dummy)\n\nprint(\"Patch embeddings:\", patch_embed.shape)   # [B, N, 64]\nprint(\"Global embeddings:\", global_embed.shape) # [B, 64]\nprint(\"Reconstruction:\", recon.shape)           # [B, 1, 8, 8]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:45:12.861276Z","iopub.execute_input":"2025-08-25T08:45:12.861543Z","iopub.status.idle":"2025-08-25T08:45:12.949354Z","shell.execute_reply.started":"2025-08-25T08:45:12.861523Z","shell.execute_reply":"2025-08-25T08:45:12.948716Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # 8. Training Loop\n# # ======================\n# device = 'cuda' if torch.cuda.is_available() else 'cpu'\n# model = PGECA().to(device)\n# optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n# epochs = 5\n\n# for epoch in range(epochs):\n#     print(\"Epoch-\",epoch)\n#     total_loss = 0\n#     for imgs,_ in loader:\n#         print(\"here!\")\n#         imgs = imgs.to(device)\n#         patch_embed, global_embed, recon = model(imgs)\n        \n#         loss_edge = edge_contrastive_loss(patch_embed)\n#         target = split_patches(imgs)[:,:,0,:,:].mean(dim=1, keepdim=True)\n#         loss_recon = F.mse_loss(recon, target)\n#         loss = loss_edge + 0.1*loss_recon\n        \n#         optimizer.zero_grad()\n#         loss.backward()\n#         optimizer.step()\n#         total_loss += loss.item()\n        \n#         # Print every 10 batches\n#         if (batch_idx + 1) % 10 == 0:\n#             print(f\"Epoch [{epoch+1}/{epochs}], Batch [{batch_idx+1}/{len(loader)}], Loss: {loss.item():.4f}\")\n        \n#     print(f\"Epoch {epoch+1}/{epochs}, Loss: {total_loss/len(loader):.4f}\")\n\n# # ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.278404Z","iopub.status.idle":"2025-08-24T21:26:56.278623Z","shell.execute_reply.started":"2025-08-24T21:26:56.278523Z","shell.execute_reply":"2025-08-24T21:26:56.278532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import time\n# import torch.nn.functional as F\n\n# device = 'cuda' if torch.cuda.is_available() else 'cpu'\n# print(\"Using device:\", device)\n\n# model = PGECA().to(device)\n# optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n# epochs = 1\n\n# for epoch in range(epochs):\n#     print(f\"\\nEpoch {epoch+1}/{epochs}\")\n#     total_loss = 0\n\n#     for batch_idx, batch in enumerate(loader):\n#         start_time = time.time()\n#         try:\n#             imgs, _ = batch\n#             print(f\"✅ Batch {batch_idx+1} fetched in {time.time()-start_time:.2f}s\")\n\n#             imgs = imgs.to(device)\n#             print(\"   → Images moved to device\")\n\n#             # Forward pass\n#             patch_embed, global_embed, recon = model(imgs)\n#             print(\"   → Model forward done\")\n\n#             # Target\n#             target = split_patches(imgs).to(device)[:,:,0,:,:].mean(dim=1, keepdim=True)\n#             print(\"   → Target created\")\n\n#             # Loss\n#             loss_edge = edge_contrastive_loss(patch_embed)\n#             loss_recon = F.mse_loss(recon, target)\n#             loss = loss_edge + 0.1*loss_recon\n#             print(f\"   → Loss computed: {loss.item():.4f}\")\n\n#             # Backprop\n#             optimizer.zero_grad()\n#             loss.backward()\n#             optimizer.step()\n#             print(\"   → Backprop done\")\n\n#             total_loss += loss.item()\n\n#             if (batch_idx + 1) % 10 == 0:\n#                 print(f\"📊 Batch [{batch_idx+1}/{len(loader)}], Loss: {loss.item():.4f}\")\n\n#         except Exception as e:\n#             print(f\"❌ Error at batch {batch_idx+1}: {e}\")\n#             break\n\n#         # Timeout check (if a batch takes too long)\n#         elapsed = time.time() - start_time\n#         if elapsed > 30:  # adjust threshold as needed\n#             print(f\"⚠️ Batch {batch_idx+1} took too long ({elapsed:.2f}s). Stopping debug.\")\n#             break\n\n#     avg_loss = total_loss / max(1, (batch_idx+1))\n#     print(f\"Epoch [{epoch+1}/{epochs}], Average Loss: {avg_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:55:43.21371Z","iopub.execute_input":"2025-08-25T08:55:43.214036Z","execution_failed":"2025-08-25T08:56:49.267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# device = 'cuda' if torch.cuda.is_available() else 'cpu'\n# model = PGECA().to(device)\n\n# optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n# epochs = 1\n\n# for epoch in range(epochs):\n#     print(f\"Epoch {epoch+1}/{epochs}\")\n#     total_loss = 0\n\n#     for batch_idx, (imgs, _) in enumerate(loader):\n#         print(\"Processing batch\", batch_idx+1)\n#         imgs = imgs.to(device)\n        \n#         patch_embed, global_embed, recon = model(imgs)\n        \n#         # Make sure target is on the same device\n#         target = split_patches(imgs).to(device)[:,:,0,:,:].mean(dim=1, keepdim=True)\n        \n#         loss_edge = edge_contrastive_loss(patch_embed)\n#         loss_recon = F.mse_loss(recon, target)\n#         loss = loss_edge + 0.1*loss_recon\n        \n#         optimizer.zero_grad()\n#         loss.backward()\n#         optimizer.step()\n        \n#         total_loss += loss.item()\n        \n#         if (batch_idx + 1) % 10 == 0:\n#             print(f\"Batch [{batch_idx+1}/{len(loader)}], Loss: {loss.item():.4f}\")\n    \n#     avg_loss = total_loss / len(loader)\n#     print(f\"Epoch [{epoch+1}/{epochs}], Average Loss: {avg_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T08:45:20.194837Z","iopub.execute_input":"2025-08-25T08:45:20.195111Z","iopub.status.idle":"2025-08-25T08:53:55.943346Z","shell.execute_reply.started":"2025-08-25T08:45:20.195091Z","shell.execute_reply":"2025-08-25T08:53:55.942308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9. Extract Embeddings\n# ======================\nembeddings_list = []\nwith torch.no_grad():\n    for imgs,_ in loader:\n        imgs = imgs.to(device)\n        _, global_embed, _ = model(imgs)\n        embeddings_list.append(global_embed.cpu())\nembeddings_all = torch.cat(embeddings_list).numpy()\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.280649Z","iopub.status.idle":"2025-08-24T21:26:56.280853Z","shell.execute_reply.started":"2025-08-24T21:26:56.280756Z","shell.execute_reply":"2025-08-24T21:26:56.280764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 10. Clustering + Evaluation\n# ======================\nnum_clusters = min(10, len(dataset))\nkmeans = KMeans(n_clusters=num_clusters, random_state=42)\npreds = kmeans.fit_predict(embeddings_all)\nsil_score = silhouette_score(embeddings_all, preds)\nprint(\"Silhouette Score:\", sil_score)\n\n# ======================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.281804Z","iopub.status.idle":"2025-08-24T21:26:56.282094Z","shell.execute_reply.started":"2025-08-24T21:26:56.281933Z","shell.execute_reply":"2025-08-24T21:26:56.281944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 11. Visualization\n# ======================\nplt.figure(figsize=(6,3))\nfor i in range(min(6, len(dataset))):\n    plt.subplot(2,6,i+1)\n    img,_ = dataset[i]\n    plt.imshow(img[0], cmap='gray')\n    plt.axis('off')\n    \n    plt.subplot(2,6,i+7)\n    with torch.no_grad():\n        recon_patch = model(img.unsqueeze(0).to(device))[2][0,0].cpu().numpy()\n    plt.imshow(recon_patch, cmap='gray')\n    plt.axis('off')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.282755Z","iopub.status.idle":"2025-08-24T21:26:56.283Z","shell.execute_reply.started":"2025-08-24T21:26:56.282868Z","shell.execute_reply":"2025-08-24T21:26:56.28288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# PCA Embedding Visualization\npca = PCA(n_components=2)\nemb_2d = pca.fit_transform(embeddings_all)\nplt.figure(figsize=(6,6))\nplt.scatter(emb_2d[:,0], emb_2d[:,1], c=preds, cmap='tab10', s=5)\nplt.title(\"Global Embeddings (PCA 2D)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T21:26:56.284056Z","iopub.status.idle":"2025-08-24T21:26:56.28428Z","shell.execute_reply.started":"2025-08-24T21:26:56.284176Z","shell.execute_reply":"2025-08-24T21:26:56.284186Z"}},"outputs":[],"execution_count":null}]}