{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":5048,"databundleVersionId":868335}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(\"Using device:\", device)\n\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:46:25.636757Z","iopub.execute_input":"2026-03-02T15:46:25.637030Z","iopub.status.idle":"2026-03-02T15:46:29.631176Z","shell.execute_reply.started":"2026-03-02T15:46:25.636996Z","shell.execute_reply":"2026-03-02T15:46:29.630409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import datasets, transforms\nfrom torch.utils.data import DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:46:34.954938Z","iopub.execute_input":"2026-03-02T15:46:34.955631Z","iopub.status.idle":"2026-03-02T15:46:38.083851Z","shell.execute_reply.started":"2026-03-02T15:46:34.955600Z","shell.execute_reply":"2026-03-02T15:46:38.083264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_path = \"/kaggle/input/competitions/state-farm-distracted-driver-detection/imgs/train\"\n\nprint(\"Classes:\", os.listdir(train_path))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:47:06.224968Z","iopub.execute_input":"2026-03-02T15:47:06.225628Z","iopub.status.idle":"2026-03-02T15:47:06.240485Z","shell.execute_reply.started":"2026-03-02T15:47:06.225600Z","shell.execute_reply":"2026-03-02T15:47:06.239863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:46:51.059648Z","iopub.execute_input":"2026-03-02T15:46:51.060392Z","iopub.status.idle":"2026-03-02T15:46:51.063903Z","shell.execute_reply.started":"2026-03-02T15:46:51.060334Z","shell.execute_reply":"2026-03-02T15:46:51.063099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = datasets.ImageFolder(\n    root=train_path,\n    transform=transform\n)\n\nprint(\"Total Images:\", len(train_dataset))\nprint(\"Number of Classes:\", len(train_dataset.classes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:47:10.120103Z","iopub.execute_input":"2026-03-02T15:47:10.120437Z","iopub.status.idle":"2026-03-02T15:47:57.851924Z","shell.execute_reply.started":"2026-03-02T15:47:10.120398Z","shell.execute_reply":"2026-03-02T15:47:57.851132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_dataset,\n    batch_size=32,\n    shuffle=True,\n    num_workers=2\n)\n\nprint(\"DataLoader Ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:49:14.104592Z","iopub.execute_input":"2026-03-02T15:49:14.105398Z","iopub.status.idle":"2026-03-02T15:49:14.110011Z","shell.execute_reply.started":"2026-03-02T15:49:14.105337Z","shell.execute_reply":"2026-03-02T15:49:14.109325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Subset\nimport numpy as np\n\n# --- IMPORTANT: create GAN dataset with 64x64 ---\ngan_transform = transforms.Compose([\n    transforms.Resize((64, 64)),\n    transforms.ToTensor(),\n    transforms.Normalize((0.5,), (0.5,))\n])\n\ngan_dataset_full = datasets.ImageFolder(\n    root=train_path,\n    transform=gan_transform\n)\n\nclass_datasets = {}\n\ntargets = np.array(gan_dataset_full.targets)\n\nfor class_idx, class_name in enumerate(gan_dataset_full.classes):\n    indices = np.where(targets == class_idx)[0]\n    class_subset = Subset(gan_dataset_full, indices)\n    class_datasets[class_name] = class_subset\n\nprint(\"Class-wise GAN datasets ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:49:19.110921Z","iopub.execute_input":"2026-03-02T15:49:19.111207Z","iopub.status.idle":"2026-03-02T15:49:26.852033Z","shell.execute_reply.started":"2026-03-02T15:49:19.111184Z","shell.execute_reply":"2026-03-02T15:49:26.851424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Generator(nn.Module):\n    def __init__(self, nz=100):\n        super().__init__()\n        self.main = nn.Sequential(\n            nn.ConvTranspose2d(nz, 256, 4, 1, 0),   # 4x4\n            nn.BatchNorm2d(256),\n            nn.ReLU(True),\n\n            nn.ConvTranspose2d(256, 128, 4, 2, 1),  # 8x8\n            nn.BatchNorm2d(128),\n            nn.ReLU(True),\n\n            nn.ConvTranspose2d(128, 64, 4, 2, 1),   # 16x16\n            nn.BatchNorm2d(64),\n            nn.ReLU(True),\n\n            nn.ConvTranspose2d(64, 32, 4, 2, 1),    # 32x32\n            nn.BatchNorm2d(32),\n            nn.ReLU(True),\n\n            nn.ConvTranspose2d(32, 3, 4, 2, 1),     # 64x64\n            nn.Tanh()\n        )\n\n    def forward(self, x):\n        return self.main(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:49:38.964756Z","iopub.execute_input":"2026-03-02T15:49:38.965089Z","iopub.status.idle":"2026-03-02T15:49:38.971126Z","shell.execute_reply.started":"2026-03-02T15:49:38.965063Z","shell.execute_reply":"2026-03-02T15:49:38.970410Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Discriminator(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.main = nn.Sequential(\n            nn.Conv2d(3, 64, 4, 2, 1),     # 64x32x32\n            nn.LeakyReLU(0.2),\n\n            nn.Conv2d(64, 128, 4, 2, 1),   # 128x16x16\n            nn.BatchNorm2d(128),\n            nn.LeakyReLU(0.2),\n\n            nn.Conv2d(128, 256, 4, 2, 1),  # 256x8x8\n            nn.BatchNorm2d(256),\n            nn.LeakyReLU(0.2),\n\n            nn.Conv2d(256, 1, 8, 1, 0),    # 1x1\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        return self.main(x).view(-1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:49:44.344535Z","iopub.execute_input":"2026-03-02T15:49:44.345217Z","iopub.status.idle":"2026-03-02T15:49:44.350000Z","shell.execute_reply.started":"2026-03-02T15:49:44.345188Z","shell.execute_reply":"2026-03-02T15:49:44.349412Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torchvision.utils as vutils\n\naug_path = \"/kaggle/working/gan_augmented\"\nos.makedirs(aug_path, exist_ok=True)\n\nnz = 100\nepochs = 3          # keep small first\nnum_fake_per_class = 500\n\nfor class_name, dataset_subset in class_datasets.items():\n\n    print(f\"\\nTraining GAN for class: {class_name}\")\n\n    loader = DataLoader(dataset_subset, batch_size=64, shuffle=True)\n\n    netG = Generator().to(device)\n    netD = Discriminator().to(device)\n\n    criterion = nn.BCELoss()\n    optimizerG = torch.optim.Adam(netG.parameters(), lr=0.0002)\n    optimizerD = torch.optim.Adam(netD.parameters(), lr=0.0002)\n\n    for epoch in range(epochs):\n        for real_imgs, _ in loader:\n            real_imgs = real_imgs.to(device)\n            b_size = real_imgs.size(0)\n\n            real_labels = torch.ones(b_size).to(device)\n            fake_labels = torch.zeros(b_size).to(device)\n\n            # Train D\n            netD.zero_grad()\n            output_real = netD(real_imgs)\n            loss_real = criterion(output_real, real_labels)\n\n            noise = torch.randn(b_size, nz, 1, 1).to(device)\n            fake_imgs = netG(noise)\n\n            output_fake = netD(fake_imgs.detach())\n            loss_fake = criterion(output_fake, fake_labels)\n\n            loss_D = loss_real + loss_fake\n            loss_D.backward()\n            optimizerD.step()\n\n            # Train G\n            netG.zero_grad()\n            output = netD(fake_imgs)\n            loss_G = criterion(output, real_labels)\n            loss_G.backward()\n            optimizerG.step()\n\n        print(f\"{class_name} GAN Epoch [{epoch+1}/{epochs}] Done\")\n\n    # Generate synthetic images\n    class_folder = os.path.join(aug_path, class_name)\n    os.makedirs(class_folder, exist_ok=True)\n\n    noise = torch.randn(num_fake_per_class, nz, 1, 1).to(device)\n    fake_imgs = netG(noise).detach().cpu()\n\n    for i in range(num_fake_per_class):\n        vutils.save_image(\n            fake_imgs[i],\n            f\"{class_folder}/fake_{i}.png\",\n            normalize=True\n        )\n\nprint(\"All class-wise GAN images generated\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:49:48.195136Z","iopub.execute_input":"2026-03-02T15:49:48.195464Z","iopub.status.idle":"2026-03-02T15:57:00.616032Z","shell.execute_reply.started":"2026-03-02T15:49:48.195436Z","shell.execute_reply":"2026-03-02T15:57:00.615416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import ConcatDataset\n\nreal_dataset = datasets.ImageFolder(\n    root=train_path,\n    transform=transform  # 224x224 transform\n)\n\ngan_dataset = datasets.ImageFolder(\n    root=\"/kaggle/working/gan_augmented\",\n    transform=transform\n)\n\ncombined_dataset = ConcatDataset([real_dataset, gan_dataset])\n\ntrain_loader = DataLoader(\n    combined_dataset,\n    batch_size=32,\n    shuffle=True,\n    num_workers=2\n)\n\nprint(\"Combined Dataset Size:\", len(combined_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:57:15.230245Z","iopub.execute_input":"2026-03-02T15:57:15.230948Z","iopub.status.idle":"2026-03-02T15:57:24.372455Z","shell.execute_reply.started":"2026-03-02T15:57:15.230918Z","shell.execute_reply":"2026-03-02T15:57:24.371821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_dataset = datasets.ImageFolder(\n    root=train_path,\n    transform=transform\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=32,\n    shuffle=False,\n    num_workers=2\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:57:30.820109Z","iopub.execute_input":"2026-03-02T15:57:30.820851Z","iopub.status.idle":"2026-03-02T15:57:31.188789Z","shell.execute_reply.started":"2026-03-02T15:57:30.820821Z","shell.execute_reply":"2026-03-02T15:57:31.188235Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\nfrom torchvision.models import resnet18, ResNet18_Weights\n\n\n# ================================\n# 1️⃣ CNN + TCN + Swin\n# ================================\nclass CNN_TCN_Swin(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.tcn = nn.Sequential(\n            nn.Conv1d(512, 256, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.Conv1d(256, 256, kernel_size=3, padding=1),\n            nn.ReLU()\n        )\n\n        self.transformer = timm.create_model(\n            \"swin_base_patch4_window7_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(256 + 1024, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n\n        tcn_input = cnn_feat.unsqueeze(2)\n        tcn_out = self.tcn(tcn_input)\n        tcn_feat = tcn_out.mean(dim=2)\n\n        swin_feat = self.transformer(x)\n\n        fused = torch.cat((tcn_feat, swin_feat), dim=1)\n        return self.classifier(fused)\n\n\n# ================================\n# 2️⃣ CNN + TCN + ViT\n# ================================\nclass CNN_TCN_ViT(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.tcn = nn.Sequential(\n            nn.Conv1d(512, 256, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.Conv1d(256, 256, kernel_size=3, padding=1),\n            nn.ReLU()\n        )\n\n        self.transformer = timm.create_model(\n            \"vit_base_patch16_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(256 + 768, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n\n        tcn_input = cnn_feat.unsqueeze(2)\n        tcn_out = self.tcn(tcn_input)\n        tcn_feat = tcn_out.mean(dim=2)\n\n        vit_feat = self.transformer(x)\n\n        fused = torch.cat((tcn_feat, vit_feat), dim=1)\n        return self.classifier(fused)\n\n\n# ================================\n# 3️⃣ CNN + LSTM + ViT\n# ================================\nclass CNN_LSTM_ViT(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.lstm = nn.LSTM(512, 256, batch_first=True)\n\n        self.transformer = timm.create_model(\n            \"vit_base_patch16_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(256 + 768, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n\n        lstm_input = cnn_feat.unsqueeze(1)\n        lstm_out, _ = self.lstm(lstm_input)\n        lstm_feat = lstm_out[:, -1, :]\n\n        vit_feat = self.transformer(x)\n\n        fused = torch.cat((lstm_feat, vit_feat), dim=1)\n        return self.classifier(fused)\n\n\n# ================================\n# 4️⃣ CNN + GNN + LSTM + ViT\n# ================================\nclass CNN_GNN_LSTM_ViT(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.gnn = nn.Linear(512, 256)\n        self.lstm = nn.LSTM(256, 256, batch_first=True)\n\n        self.transformer = timm.create_model(\n            \"vit_base_patch16_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(256 + 768, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n        gnn_feat = self.gnn(cnn_feat)\n\n        lstm_out, _ = self.lstm(gnn_feat.unsqueeze(1))\n        lstm_feat = lstm_out[:, -1, :]\n\n        vit_feat = self.transformer(x)\n\n        fused = torch.cat((lstm_feat, vit_feat), dim=1)\n        return self.classifier(fused)\n\n\n# ================================\n# 5️⃣ CNN + GNN + Swin\n# ================================\nclass CNN_GNN_Swin(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.gnn = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU()\n        )\n\n        self.transformer = timm.create_model(\n            \"swin_base_patch4_window7_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(256 + 1024, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n        gnn_feat = self.gnn(cnn_feat)\n        swin_feat = self.transformer(x)\n\n        fused = torch.cat((gnn_feat, swin_feat), dim=1)\n        return self.classifier(fused)\n\n\n# ================================\n# 6️⃣ CNN + GNN + ViT\n# ================================\nclass CNN_GNN_ViT(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.gnn = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU()\n        )\n\n        self.transformer = timm.create_model(\n            \"vit_base_patch16_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(256 + 768, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n        gnn_feat = self.gnn(cnn_feat)\n        vit_feat = self.transformer(x)\n\n        fused = torch.cat((gnn_feat, vit_feat), dim=1)\n        return self.classifier(fused)\n\n\n# ================================\n# 7️⃣ CNN + GNN + LSTM + Swin\n# ================================\nclass CNN_GNN_LSTM_SWIN(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        self.gnn = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU()\n        )\n\n        self.lstm = nn.LSTM(256, 128, batch_first=True)\n\n        self.transformer = timm.create_model(\n            \"swin_tiny_patch4_window7_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        self.classifier = nn.Linear(128 + 768, num_classes)\n\n    def forward(self, x):\n        cnn_feat = self.cnn(x)\n        gnn_feat = self.gnn(cnn_feat)\n\n        lstm_out, _ = self.lstm(gnn_feat.unsqueeze(1))\n        lstm_feat = lstm_out[:, -1, :]\n\n        trans_feat = self.transformer(x)\n\n        fused = torch.cat((lstm_feat, trans_feat), dim=1)\n        return self.classifier(fused)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:57:36.524444Z","iopub.execute_input":"2026-03-02T15:57:36.525067Z","iopub.status.idle":"2026-03-02T15:57:40.316219Z","shell.execute_reply.started":"2026-03-02T15:57:36.525038Z","shell.execute_reply":"2026-03-02T15:57:40.315653Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, model_name, epochs=5):\n\n    model = model.to(device)\n\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.Adam(\n        model.parameters(),\n        lr=1e-4,\n        weight_decay=1e-4\n    )\n\n    for epoch in range(epochs):\n\n        # ===== TRAIN =====\n        model.train()\n        train_correct = 0\n        train_total = 0\n        train_loss = 0\n\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            train_total += labels.size(0)\n            train_correct += (preds == labels).sum().item()\n\n        train_acc = train_correct / train_total\n\n        # ===== VALIDATION =====\n        model.eval()\n        val_correct = 0\n        val_total = 0\n\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs, labels = imgs.to(device), labels.to(device)\n                outputs = model(imgs)\n                _, preds = torch.max(outputs, 1)\n\n                val_total += labels.size(0)\n                val_correct += (preds == labels).sum().item()\n\n        val_acc = val_correct / val_total\n\n        print(f\"{model_name} | Epoch [{epoch+1}/{epochs}]\")\n        print(f\"Train Loss: {train_loss/len(train_loader):.4f}\")\n        print(f\"Train Acc: {train_acc:.4f}\")\n        print(f\"Val Acc: {val_acc:.4f}\")\n        print(\"-\"*40)\n\n    return val_acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:57:49.365726Z","iopub.execute_input":"2026-03-02T15:57:49.366567Z","iopub.status.idle":"2026-03-02T15:57:49.373862Z","shell.execute_reply.started":"2026-03-02T15:57:49.366536Z","shell.execute_reply":"2026-03-02T15:57:49.373115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"after_gan_results = {}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:57:58.054941Z","iopub.execute_input":"2026-03-02T15:57:58.055649Z","iopub.status.idle":"2026-03-02T15:57:58.058782Z","shell.execute_reply.started":"2026-03-02T15:57:58.055617Z","shell.execute_reply":"2026-03-02T15:57:58.058103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_TCN_Swin()\nacc = train_model(model, \"CNN_TCN_Swin\", epochs=5)\nafter_gan_results[\"CNN_TCN_Swin\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_TCN_Swin_after_gan.pth\")\nprint(\"Saved CNN_TCN_Swin\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T06:35:55.983107Z","iopub.execute_input":"2026-02-25T06:35:55.983844Z","iopub.status.idle":"2026-02-25T08:32:58.999756Z","shell.execute_reply.started":"2026-02-25T06:35:55.983787Z","shell.execute_reply":"2026-02-25T08:32:58.998888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_LSTM_ViT()\nacc = train_model(model, \"CNN_LSTM_ViT\", epochs=5)\nafter_gan_results[\"CNN_LSTM_ViT\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_LSTM_ViT_after_gan.pth\")\nprint(\"Saved CNN_LSTM_ViT\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T09:36:08.242853Z","iopub.execute_input":"2026-02-25T09:36:08.243766Z","iopub.status.idle":"2026-02-25T11:28:17.956804Z","shell.execute_reply.started":"2026-02-25T09:36:08.243715Z","shell.execute_reply":"2026-02-25T11:28:17.955842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_GNN_LSTM_ViT()\nacc = train_model(model, \"CNN_GNN_LSTM_ViT\", epochs=5)\nafter_gan_results[\"CNN_GNN_LSTM_ViT\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_GNN_LSTM_ViT_after_gan.pth\")\nprint(\"Saved CNN_GNN_LSTM_ViT\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T11:31:34.308082Z","iopub.execute_input":"2026-02-25T11:31:34.308608Z","iopub.status.idle":"2026-02-25T13:24:14.381569Z","shell.execute_reply.started":"2026-02-25T11:31:34.308571Z","shell.execute_reply":"2026-02-25T13:24:14.380703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_GNN_LSTM_SWIN()\nacc = train_model(model, \"CNN_GNN_LSTM_SWIN\", epochs=5)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T16:01:24.983378Z","iopub.execute_input":"2026-02-25T16:01:24.984077Z","iopub.status.idle":"2026-02-25T16:54:41.413995Z","shell.execute_reply.started":"2026-02-25T16:01:24.984045Z","shell.execute_reply":"2026-02-25T16:54:41.412809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"after_gan_results[\"CNN_GNN_LSTM_SWIN\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_GNN_LSTM_SWIN_after_gan.pth\")\nprint(\"Saved CNN_GNN_LSTM_SWIN\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T16:57:35.716655Z","iopub.execute_input":"2026-02-25T16:57:35.717506Z","iopub.status.idle":"2026-02-25T16:57:35.961899Z","shell.execute_reply.started":"2026-02-25T16:57:35.717470Z","shell.execute_reply":"2026-02-25T16:57:35.961084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_GNN_Swin()\nacc = train_model(model, \"CNN_GNN_Swin\", epochs=5)\nafter_gan_results[\"CNN_GNN_Swin\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_GNN_Swin_after_gan.pth\")\nprint(\"Saved CNN_GNN_Swin\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-25T16:58:25.984428Z","iopub.execute_input":"2026-02-25T16:58:25.985226Z","iopub.status.idle":"2026-02-25T19:04:12.029347Z","shell.execute_reply.started":"2026-02-25T16:58:25.985193Z","shell.execute_reply":"2026-02-25T19:04:12.028303Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_GNN_ViT()\nacc = train_model(model, \"CNN_GNN_ViT\", epochs=5)\nafter_gan_results[\"CNN_GNN_ViT\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_GNN_ViT_after_gan.pth\")\nprint(\"Saved CNN_GNN_ViT\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T15:58:05.279211Z","iopub.execute_input":"2026-03-02T15:58:05.280022Z","iopub.status.idle":"2026-03-02T17:47:37.375089Z","shell.execute_reply.started":"2026-03-02T15:58:05.279989Z","shell.execute_reply":"2026-03-02T17:47:37.374432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_TCN_ViT()\nacc = train_model(model, \"CNN_TCN_ViT\", epochs=5)\nafter_gan_results[\"CNN_TCN_ViT\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_TCN_ViT_after_gan.pth\")\nprint(\"Saved CNN_TCN_ViT\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T17:50:22.631265Z","iopub.execute_input":"2026-03-02T17:50:22.632062Z","iopub.status.idle":"2026-03-02T19:40:37.104288Z","shell.execute_reply.started":"2026-03-02T17:50:22.632030Z","shell.execute_reply":"2026-03-02T19:40:37.103383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CNN_LSTM_SWIN(nn.Module):\n    def __init__(self, num_classes=10):\n        super().__init__()\n\n        # CNN Backbone\n        self.cnn = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)\n        self.cnn.fc = nn.Identity()\n\n        # LSTM\n        self.lstm = nn.LSTM(512, 256, batch_first=True)\n\n        # Swin Transformer\n        self.transformer = timm.create_model(\n            \"swin_base_patch4_window7_224\",\n            pretrained=True,\n            num_classes=0\n        )\n\n        # Final classifier\n        self.classifier = nn.Linear(256 + 1024, num_classes)\n\n    def forward(self, x):\n        # CNN features\n        cnn_feat = self.cnn(x)\n\n        # LSTM expects sequence, so add dimension\n        lstm_input = cnn_feat.unsqueeze(1)\n        lstm_out, _ = self.lstm(lstm_input)\n        lstm_feat = lstm_out[:, -1, :]\n\n        # Swin features\n        swin_feat = self.transformer(x)\n\n        # Fusion\n        fused = torch.cat((lstm_feat, swin_feat), dim=1)\n\n        return self.classifier(fused)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T19:41:15.565826Z","iopub.execute_input":"2026-03-02T19:41:15.566282Z","iopub.status.idle":"2026-03-02T19:41:15.572938Z","shell.execute_reply.started":"2026-03-02T19:41:15.566201Z","shell.execute_reply":"2026-03-02T19:41:15.572283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CNN_LSTM_SWIN()\nacc = train_model(model, \"CNN_LSTM_SWIN\", epochs=5)\nafter_gan_results[\"CNN_LSTM_SWIN\"] = acc\n\ntorch.save(model.state_dict(), \"CNN_LSTM_SWIN_after_gan.pth\")\nprint(\"Saved CNN_LSTM_SWIN\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T19:41:30.470107Z","iopub.execute_input":"2026-03-02T19:41:30.470446Z","iopub.status.idle":"2026-03-02T21:36:35.964689Z","shell.execute_reply.started":"2026-03-02T19:41:30.470416Z","shell.execute_reply":"2026-03-02T21:36:35.963808Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\n\nwith open(\"after_gan_results.json\", \"w\") as f:\n    json.dump(after_gan_results, f, indent=4)\n\nprint(\"After GAN results saved successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:37:19.477178Z","iopub.execute_input":"2026-03-02T21:37:19.477602Z","iopub.status.idle":"2026-03-02T21:37:19.482631Z","shell.execute_reply.started":"2026-03-02T21:37:19.477567Z","shell.execute_reply":"2026-03-02T21:37:19.482057Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"before_gan_results = {\n    \"CNN_TCN_Swin\": 0.9955,\n    \"CNN_TCN_ViT\": 0.9946,\n    \"CNN_LSTM_ViT\": 0.9967,\n    \"CNN_GNN_LSTM_ViT\": 0.9911,\n    \"CNN_GNN_Swin\": 0.9978,\n    \"CNN_GNN_ViT\": 0.9942,\n    \"CNN_GNN_LSTM_SWIN\": 0.9944,\n    \"CNN_LSTM_SWIN\":0.9912\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:43:30.705027Z","iopub.execute_input":"2026-03-02T21:43:30.705666Z","iopub.status.idle":"2026-03-02T21:43:30.709368Z","shell.execute_reply.started":"2026-03-02T21:43:30.705638Z","shell.execute_reply":"2026-03-02T21:43:30.708694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"before_gan_results.json\", \"w\") as f:\n    json.dump(before_gan_results, f, indent=4)\n\nprint(\"Before GAN results saved successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:43:49.860476Z","iopub.execute_input":"2026-03-02T21:43:49.860745Z","iopub.status.idle":"2026-03-02T21:43:49.865873Z","shell.execute_reply.started":"2026-03-02T21:43:49.860722Z","shell.execute_reply":"2026-03-02T21:43:49.865172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"after_gan_results.update({\n    \"CNN_TCN_Swin\": 0.9971,\n    \"CNN_LSTM_ViT\": 0.998,\n    \"CNN_GNN_LSTM_ViT\": 0.9977,\n    \"CNN_GNN_LSTM_SWIN\": 0.9980,\n    \"CNN_GNN_Swin\": 0.9983\n})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:50:49.924466Z","iopub.execute_input":"2026-03-02T21:50:49.925077Z","iopub.status.idle":"2026-03-02T21:50:49.928549Z","shell.execute_reply.started":"2026-03-02T21:50:49.925046Z","shell.execute_reply":"2026-03-02T21:50:49.927768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"After GAN models available:\")\nprint(after_gan_results.keys())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:51:02.482706Z","iopub.execute_input":"2026-03-02T21:51:02.482996Z","iopub.status.idle":"2026-03-02T21:51:02.487247Z","shell.execute_reply.started":"2026-03-02T21:51:02.482971Z","shell.execute_reply":"2026-03-02T21:51:02.486580Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ncomparison = pd.DataFrame({\n    \"Model\": list(before_gan_results.keys()),\n    \"Before GAN\": [before_gan_results[m] for m in before_gan_results],\n    \"After GAN\": [after_gan_results.get(m, None) for m in before_gan_results]\n})\n\ncomparison[\"Improvement\"] = comparison[\"After GAN\"] - comparison[\"Before GAN\"]\n\ncomparison","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:51:23.187674Z","iopub.execute_input":"2026-03-02T21:51:23.188245Z","iopub.status.idle":"2026-03-02T21:51:23.218086Z","shell.execute_reply.started":"2026-03-02T21:51:23.188219Z","shell.execute_reply":"2026-03-02T21:51:23.217279Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"comparison.to_csv(\"final_gan_comparison.csv\", index=False)\nprint(\"Comparison CSV saved successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:53:52.257008Z","iopub.execute_input":"2026-03-02T21:53:52.257746Z","iopub.status.idle":"2026-03-02T21:53:52.267565Z","shell.execute_reply.started":"2026-03-02T21:53:52.257716Z","shell.execute_reply":"2026-03-02T21:53:52.266956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\n\nwith open(\"before_gan_results.json\", \"w\") as f:\n    json.dump(before_gan_results, f, indent=4)\n\nwith open(\"after_gan_results.json\", \"w\") as f:\n    json.dump(after_gan_results, f, indent=4)\n\nprint(\"All results saved successfully\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:54:09.862170Z","iopub.execute_input":"2026-03-02T21:54:09.862778Z","iopub.status.idle":"2026-03-02T21:54:09.868125Z","shell.execute_reply.started":"2026-03-02T21:54:09.862748Z","shell.execute_reply":"2026-03-02T21:54:09.867503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.listdir(\"/kaggle/working\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T21:54:24.237619Z","iopub.execute_input":"2026-03-02T21:54:24.238406Z","iopub.status.idle":"2026-03-02T21:54:24.243176Z","shell.execute_reply.started":"2026-03-02T21:54:24.238370Z","shell.execute_reply":"2026-03-02T21:54:24.242641Z"}},"outputs":[],"execution_count":null}]}