{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.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":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":8352421,"sourceType":"datasetVersion","datasetId":4962594},{"sourceId":8779646,"sourceType":"datasetVersion","datasetId":5277123}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torch_geometric torchlens -q\nimport torch\nimport torchvision\nimport torch_geometric as tg\nimport torchlens as tl\nimport mne\nimport cv2\nimport os\nimport networkx as nx\nimport numpy as np\nfrom matplotlib import pyplot as plt\nfrom torch import nn, optim, utils\nfrom torch_geometric import nn as gnn\nfrom torch.utils import data\nfrom torch_geometric.utils import from_networkx, to_networkx\nfrom torchvision.transforms import v2\nfrom torchvision import models as vis_models\nfrom sklearn.model_selection import train_test_split\nfrom torch_geometric.loader import DataLoader\nfrom torch_geometric.utils import sort_edge_index\nfrom torch.nn import functional as F\nfrom tqdm import tqdm\ntorch.cuda.is_available()","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:36:35.618252Z","iopub.execute_input":"2024-07-02T05:36:35.619051Z","iopub.status.idle":"2024-07-02T05:36:57.752026Z","shell.execute_reply.started":"2024-07-02T05:36:35.61901Z","shell.execute_reply":"2024-07-02T05:36:57.751022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nswaps = {\n    'Fpz': 'AFz',\n    'Iz': 'FCz',\n    'I1': 'O9',\n    'I2': 'O10'\n}\nChannels swapped manually\n\"\"\"\nchannels = ['Fp1','Fp2','F7','F3','Fz','F4','F8','FC5','FC1','FC2','FC6','T7','C3','Cz','C4','T8','TP9','CP5','CP1','CP2','CP6','TP10','P7','P3','Pz','P4','P8','PO9','O1','Oz','O2','PO10','AF7','AF3','AF4','AF8','F5','F1','F2','F6','FT9','FT7','FC3','FC4','FT8','FT10','C5','C1','C2','C6','TP7','CP3','CPz','CP4','TP8','P5','P1','P2','P6','PO7','PO3','POz','PO4','PO8','AFz','F9','AFF5h','AFF1h','AFF2h','AFF6h','F10','FTT9h','FTT7h','FCC5h','FCC3h','FCC1h','FCC2h','FCC4h','FCC6h','FTT8h','FTT10h','TPP9h','TPP7h','CPP5h','CPP3h','CPP1h','CPP2h','CPP4h','CPP6h','TPP8h','TPP10h','POO9h','POO1','POO2','POO10h','FCz','AFp1','AFp2','FFT9h','FFT7h','FFC5h','FFC3h','FFC1h','FFC2h','FFC4h','FFC6h','FFT8h','FFT10h','TTP7h','CCP5h','CCP3h','CCP1h','CCP2h','CCP4h','CCP6h','TTP8h','P9','PPO9h','PPO5h','PPO1h','PPO2h','PPO6h','PPO10h','P10','O9','OI1h','OI2h','O10'] \nmontage = mne.channels.read_custom_montage('/kaggle/input/acticap/montage.bvef')\ndataset = torch.load('/kaggle/input/main-dla-mae-saliency/eeg_5_95_std.pth')\ntrain_dataset, test_dataset = train_test_split(\n    dataset['dataset'], test_size=0.1\n)\n\ntrain_dataset = {\n    'dataset':train_dataset,\n    'images':dataset['images'],\n    'labels':dataset['labels']\n}\n\ntest_dataset = {\n    'dataset':test_dataset,\n    'images':dataset['images'],\n    'labels':dataset['labels']\n}\n\ndel dataset","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:36:57.753699Z","iopub.execute_input":"2024-07-02T05:36:57.75422Z","iopub.status.idle":"2024-07-02T05:37:26.528691Z","shell.execute_reply.started":"2024-07-02T05:36:57.754194Z","shell.execute_reply":"2024-07-02T05:37:26.527892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GDNBlock(nn.Module):\n    def __init__(self, in_dim:int, out_dim:int, device:torch.device) -> None:\n        super().__init__()\n        self.gcn = gnn.GCNConv(in_channels=in_dim, out_channels=out_dim).to(device)\n        self.norm = gnn.BatchNorm(in_channels=out_dim).to(device)\n    \n    def forward(self, feature, edge_idx):\n        x = F.relu(self.gcn(feature, edge_idx))\n        x = self.norm(x)\n        return x\n\nclass GDN(nn.Module):\n    def __init__(self, config:list[int], device:torch.device):\n        super().__init__()\n        self.GCN = [GDNBlock(config[i], config[i+1], device).to(device) for i in range(len(config)-1)]\n        self.drop = nn.Dropout1d().to(device)\n        self.lin_out = nn.Linear(in_features=6400, out_features=40).to(device)\n        \n    def forward(self, feature, edge_index, batch):\n        for block in self.GCN:\n            feature = block(feature, edge_index)\n        bs = torch.unique(batch).shape[0]\n        feature = feature.view(bs, -1)\n        x = self.drop(feature)\n        cx = self.lin_out(x)\n        return cx, x","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.529758Z","iopub.execute_input":"2024-07-02T05:37:26.53007Z","iopub.status.idle":"2024-07-02T05:37:26.539752Z","shell.execute_reply.started":"2024-07-02T05:37:26.530045Z","shell.execute_reply":"2024-07-02T05:37:26.538837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import os\n# import cv2\n# import torch\n# from torchvision import transforms\n\n# model = 'DPT_Hybrid'\n# midas = torch.hub.load('intel-isl/MiDaS', model)\n\n# transform = transforms.Compose([\n#     transforms.ToPILImage(),\n#     transforms.Resize((256, 256)),\n#     transforms.ToTensor(),\n# ])\n# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# midas = midas.to(device)\n# save_dir = \"/kaggle/working/generated_images\"\n# os.makedirs(save_dir, exist_ok=True)\n# images_folder = '/kaggle/input/main-dla-mae-saliency/imageNet_images/imageNet_images'\n# images_images_dir = os.listdir(images_folder)\n\n# for i in images_images_dir:\n#     print(\"Reading\", i)\n#     save_dir_folder = save_dir+\"/\"+i\n#     os.makedirs(save_dir_folder, exist_ok=True)\n    \n#     folder_path = os.path.join(images_folder, i)\n#     images = os.listdir(folder_path)\n    \n#     prev_img = None \n    \n#     save_category_dir = os.path.join(save_dir, i)\n#     os.makedirs(save_category_dir, exist_ok=True)\n    \n#     for j in images:\n#         img_path = os.path.join(folder_path, j)\n#         #print(img_path)\n#         try:\n#             img = cv2.imread(img_path)\n#             #print(\"Here\")\n#             img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n#             #print(\"Done\")\n#         except Exception as e:\n#             print(f\"Error reading image {j}: {e}\")\n#             if prev_img is None:\n#                 print(\"Skipping...\")\n#                 continue\n#             else:\n#                 print(\"Using previous image instead.\")\n#                 img = prev_img.copy()\n\n#         img = cv2.resize(img, (256, 256))\n#         input_batch = transform(img).unsqueeze(0).to(device) \n\n#         #PREDICT\n#         with torch.no_grad():\n#             prediction = midas(input_batch)\n\n#             prediction = torch.nn.functional.interpolate(\n#                 prediction.unsqueeze(1),\n#                 size=img.shape[:2],\n#                 mode=\"bicubic\",\n#                 align_corners=False,\n#             ).squeeze()\n\n#         output = prediction.cpu().numpy()\n#         output = cv2.normalize(output, None, 0, 255, cv2.NORM_MINMAX, cv2.CV_8U)\n\n#         #print(\"saveing sape\", output.shape)\n#         #print(\"save_path\", f'{save_dir_folder}/{j}')\n#         cv2.imwrite(f'{save_dir_folder}/{j}', output)\n        \n#         # Update incase of Error\n#         prev_img = img.copy()\n\n# print(\"Images saved in the directory:\", save_dir)","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.541785Z","iopub.execute_input":"2024-07-02T05:37:26.542113Z","iopub.status.idle":"2024-07-02T05:37:26.561864Z","shell.execute_reply.started":"2024-07-02T05:37:26.54209Z","shell.execute_reply":"2024-07-02T05:37:26.560984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Decoder(nn.Module):\n    def __init__(self, in_features)->None:\n        super().__init__()\n        self.in_lin = nn.Linear(in_features=in_features, out_features=1000)\n        self.lin1 = nn.Linear(in_features=1000, out_features=20000)\n        self.up_convT = nn.Sequential(\n            nn.ConvTranspose2d(8, 8, (4,4), stride=2, padding=1),\n            nn.LeakyReLU(0.2),\n            nn.ConvTranspose2d(8, 8, (4,4), stride=3, padding=1, output_padding=1),\n            nn.LeakyReLU(0.2),\n            nn.ConvTranspose2d(8, 8, (3,3), padding=1),\n            nn.LeakyReLU(0.2),\n            nn.ConvTranspose2d(8, 8, (3,3), padding=1),\n            nn.LeakyReLU(0.2),\n            nn.Conv2d(in_channels=8, out_channels=1, kernel_size=(3,3))\n        )\n    \n    def forward(self, x):\n        x = self.in_lin(x)\n        x = F.leaky_relu(self.lin1(x))\n        x = torch.reshape(x, (-1, 8, 50, 50))\n        return self.up_convT(x)","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.562791Z","iopub.execute_input":"2024-07-02T05:37:26.563069Z","iopub.status.idle":"2024-07-02T05:37:26.576111Z","shell.execute_reply.started":"2024-07-02T05:37:26.563044Z","shell.execute_reply":"2024-07-02T05:37:26.575208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Discriminator(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n        self.main = nn.Sequential(\n            nn.Conv2d(1, 32, kernel_size=(4, 4), stride=(2, 2), padding=(1, 1), bias=False),\n            nn.LeakyReLU(negative_slope=0.2, inplace=True),\n            nn.Conv2d(32, 64, kernel_size=(4, 4), stride=(2, 2), padding=(1, 1), bias=False),\n            nn.BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True),\n            nn.LeakyReLU(negative_slope=0.2, inplace=True),\n            nn.Conv2d(64, 128, kernel_size=(4, 4), stride=(2, 2), padding=(1, 1), bias=False),\n            nn.BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True),\n            nn.LeakyReLU(negative_slope=0.2, inplace=True),\n            nn.Conv2d(128, 256, kernel_size=(4, 4), stride=(2, 2), padding=(1, 1), bias=False),\n            nn.BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True),\n            nn.LeakyReLU(negative_slope=0.2, inplace=True),\n            nn.Conv2d(256, 1, kernel_size=(4, 4), stride=(1, 1), bias=False),\n            # nn.Sigmoid()\n        )\n    \n    def forward(self, img):\n        return torch.squeeze(self.main(img), dim=1)","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.57711Z","iopub.execute_input":"2024-07-02T05:37:26.577345Z","iopub.status.idle":"2024-07-02T05:37:26.591125Z","shell.execute_reply.started":"2024-07-02T05:37:26.577325Z","shell.execute_reply":"2024-07-02T05:37:26.590386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GANDataset(data.Dataset):\n    def __init__(self, dataset, IMAGE_PATH ='/kaggle/working/generated_images',ch_name=channels, montage='/kaggle/input/acticap/montage.bvef', img_tar_size=298, thresh:float=0.05) -> None:\n        super().__init__()\n        self.EEG = dataset['dataset']\n        self.IMAGES = dataset['images']\n        self.LABELS = dataset['labels']\n        self.montage_ch_names = ch_name\n        self.montage = mne.channels.read_custom_montage(montage)\n        self.IMAGE_PATH = IMAGE_PATH\n        self.tar_size = img_tar_size\n        self.pos = self.montage.get_positions()['ch_pos']\n        self.thresh = thresh\n        \n    def __len__(self, ):\n        return len(self.EEG)\n    \n    def makeGraph(self, EEG):\n        G = nx.Graph()\n        for i, ch_name in enumerate(self.montage_ch_names):\n            G.add_node(i, name = ch_name, pos = self.pos[ch_name], feature = EEG[i][20:460])\n            \n        for i in range(len(self.montage_ch_names)):\n            for j in range(i+1, 128):\n                distance = np.linalg.norm(np.array(self.pos[self.montage_ch_names[i]]) - np.array(self.pos[self.montage_ch_names[j]]))\n                if distance < self.thresh:\n                    G.add_edge(i, j)\n        return from_networkx(G)\n    \n    def readImage(self, label, image):\n        label_path = os.path.join(self.IMAGE_PATH, label)\n        image_path = os.path.join(label_path, image)\n        img = cv2.imread(image_path+'.JPEG')\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        img = cv2.resize(img, (self.tar_size, self.tar_size))\n        img = torch.from_numpy(img)/255.0\n        return img\n    \n    def __getitem__(self, idx):\n        try:\n            eeg = self.EEG[idx]\n            image = self.readImage(self.LABELS[eeg['label']], self.IMAGES[eeg['image']])\n        except:\n            eeg = self.EEG[idx+1]\n            image = self.readImage(self.LABELS[eeg['label']], self.IMAGES[eeg['image']])\n        graph = self.makeGraph(eeg['eeg'])\n        return image, graph, torch.tensor(eeg['label'])","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.592261Z","iopub.execute_input":"2024-07-02T05:37:26.592547Z","iopub.status.idle":"2024-07-02T05:37:26.606683Z","shell.execute_reply.started":"2024-07-02T05:37:26.592525Z","shell.execute_reply":"2024-07-02T05:37:26.605914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dataset), len(test_dataset)","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.607799Z","iopub.execute_input":"2024-07-02T05:37:26.608196Z","iopub.status.idle":"2024-07-02T05:37:26.621016Z","shell.execute_reply.started":"2024-07-02T05:37:26.608167Z","shell.execute_reply":"2024-07-02T05:37:26.620234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nclassification_loss_func = nn.CrossEntropyLoss()\nGAN_loss_func = nn.BCEWithLogitsLoss()\n\ngraph_enc = torch.load('/kaggle/working/checkpoint_graph_enc.pth', map_location=device)\ndecoder = torch.load('/kaggle/working/checkpoint_decoder.pth', map_location=device)\ndiscriminator = torch.load('/kaggle/working/checkpoint_discriminator.pth', map_location=device)\n\nlr=0.0001\n\noptimEnc = optim.SGD(graph_enc.parameters(), 0.001, weight_decay=4e-5)\noptimDec = optim.Adam(decoder.parameters(), lr)\noptimDis = optim.Adam(discriminator.parameters(), lr)\n\ntrain_data = DataLoader(GANDataset(dataset=train_dataset), 32, True)\ntest_data = DataLoader(GANDataset(dataset=test_dataset), 32, True)","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.621968Z","iopub.execute_input":"2024-07-02T05:37:26.622216Z","iopub.status.idle":"2024-07-02T05:37:26.89499Z","shell.execute_reply.started":"2024-07-02T05:37:26.622194Z","shell.execute_reply":"2024-07-02T05:37:26.894123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_loss_list = []\ndis_loss_list = []\ngen_loss_list = []\nval_class_loss_list = []\nval_dis_loss_list = []\nval_gen_loss_list = []\nn_epochs = 1000\nfor epoch in range(n_epochs):\n    step=0\n    class_loss = 0\n    dis_loss = 0\n    gen_loss = 0\n    progress_bar = tqdm(train_data, desc=f\"EPOCH: {epoch+34}\")\n    test_progress_bar = tqdm(test_data, desc=f\"Val: {epoch+34}\")\n    graph_enc.train()\n    decoder.train()\n    discriminator.train()\n    for image, graph, label in progress_bar:\n        graph = graph.to(device)\n        image = image.to(device)\n        label = label.to(device)\n        image = torch.unsqueeze(image, dim=1)\n        optimDec.zero_grad()\n        optimDis.zero_grad()\n        optimEnc.zero_grad()\n        \n        \"\"\"Training Classification\"\"\"\n        \n        label_pred, ltntVec = graph_enc.forward(graph.feature, graph.edge_index, graph.batch)\n        class_loss_step = classification_loss_func(label_pred, label)\n        class_loss_step.backward(retain_graph=True)\n        optimEnc.step()\n        class_loss+=class_loss_step.cpu().item()\n        \n        \"\"\"Training Discriminator\"\"\"\n        \n        real_labels = torch.ones([label.size(0), 15, 15]).to(device)\n        fake_labels = torch.zeros([label.size(0), 15, 15]).to(device)\n\n        fake_images = decoder(ltntVec)\n        fake_out = discriminator(fake_images.detach())\n        real_out = discriminator(image)\n        \n        d_loss_fake = GAN_loss_func(fake_out, fake_labels)\n        d_loss_real = GAN_loss_func(real_out, real_labels)\n        d_loss = d_loss_fake+d_loss_real\n        d_loss.backward()\n        optimDis.step()\n        dis_loss+=d_loss.cpu().item()\n        \n        \"\"\"Training Generator\"\"\"\n        \n        g_fake_out = discriminator(fake_images)\n        g_loss = GAN_loss_func(g_fake_out, real_labels)\n        g_loss.backward()\n        optimDec.step()\n        gen_loss +=g_loss.cpu().item()\n        \n        progress_bar.set_postfix({\n            \"Class Loss\":class_loss_step.cpu().item(),\n            \"Discriminator\":d_loss.cpu().item(),\n            \"Generator\":g_loss.cpu().item()\n        })\n        step+=1\n        \n        if step%100==0:\n            torch.save(graph_enc, 'checkpoint_graph_enc.pth')\n            torch.save(decoder, 'checkpoint_decoder.pth')\n            torch.save(discriminator, 'checkpoint_discriminator.pth')\n    class_loss_list.append(class_loss_step.cpu().item())\n    dis_loss_list.append(d_loss.cpu().item())\n    gen_loss_list.append(g_loss.cpu().item())\n    print(f\"Train: Class Loss {class_loss/step} || Discriminator Loss {dis_loss/step} || Generator Loss {gen_loss/step}\")\n    \n    graph_enc.eval()\n    decoder.eval()\n    discriminator.eval()\n    img_count=0\n    step=0\n    class_loss = 0\n    dis_loss = 0\n    gen_loss = 0                    \n    with torch.no_grad():\n        for image, graph, label in test_progress_bar:\n            step+=1\n            graph = graph.to(device)\n            image = image.to(device)\n            label = label.to(device)\n            image = torch.unsqueeze(image, dim=1)\n            real_labels = torch.ones([label.size(0), 15, 15]).to(device)\n            fake_labels = torch.zeros([label.size(0), 15, 15]).to(device)\n            label_pred, ltntVec = graph_enc.forward(graph.feature, graph.edge_index, graph.batch)\n            fake_image = decoder(ltntVec)\n            fake_out = discriminator(fake_image)\n            real_out = discriminator(image)\n            \n            class_loss_step = classification_loss_func(label_pred, label)\n            d_loss_fake = GAN_loss_func(fake_out, fake_labels)\n            d_loss_real = GAN_loss_func(real_out, real_labels)\n            d_loss = d_loss_fake+d_loss_real\n            g_loss = GAN_loss_func(fake_out, real_labels)\n            \n            class_loss+=class_loss_step.cpu().item()\n            dis_loss += d_loss.cpu().item()\n            gen_loss+=g_loss.cpu().item()\n                         \n            test_progress_bar.set_postfix({\n                \"Class Loss\":class_loss_step.cpu().item(),\n                \"Discriminator\":d_loss.cpu().item(),\n                \"Generator\":g_loss.cpu().item()\n            })\n            \n    plt.figure()\n    fig, axes = plt.subplots(nrows=2, ncols=1, figsize=(10, 8))\n    axes[0].imshow(image[0, 0].cpu().numpy(), cmap='gray')\n    axes[0].set_title(\"Real Image\")\n    axes[0].axis('off')\n    axes[1].imshow(fake_image[0, 0].cpu().numpy(), cmap='gray')\n    axes[1].set_title(\"Generated Image\")\n    axes[1].axis('off')\n    plt.show()            \n    val_class_loss_list.append(class_loss/step)\n    val_dis_loss_list.append(dis_loss/step)\n    val_gen_loss_list.append(gen_loss/step)\n    print(f\"Validation: Class Loss {class_loss/step} || Discriminator Loss {dis_loss/step} || Generator Loss {gen_loss/step}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-02T07:08:35.452192Z","iopub.execute_input":"2024-07-02T07:08:35.452541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(graph_enc, 'graph_enc.pth')\ntorch.save(decoder, 'decoder.pth')\ntorch.save(discriminator, 'discriminator.pth')","metadata":{"execution":{"iopub.status.busy":"2024-07-02T05:37:26.916881Z","iopub.status.idle":"2024-07-02T05:37:26.917326Z","shell.execute_reply.started":"2024-07-02T05:37:26.917091Z","shell.execute_reply":"2024-07-02T05:37:26.91711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}