{"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":"none","dataSources":[{"sourceType":"competition","sourceId":4117,"databundleVersionId":46665},{"sourceType":"datasetVersion","sourceId":3337316,"datasetId":2015227,"databundleVersionId":3388378}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import shutil\n\nshutil.rmtree(\"/kaggle/working/\", ignore_errors=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:28.406303Z","iopub.execute_input":"2026-05-11T17:41:28.406680Z","iopub.status.idle":"2026-05-11T17:41:28.417953Z","shell.execute_reply.started":"2026-05-11T17:41:28.406586Z","shell.execute_reply":"2026-05-11T17:41:28.416822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:28.419809Z","iopub.execute_input":"2026-05-11T17:41:28.420870Z","iopub.status.idle":"2026-05-11T17:41:28.853245Z","shell.execute_reply.started":"2026-05-11T17:41:28.420825Z","shell.execute_reply":"2026-05-11T17:41:28.852176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torchvision import datasets, transforms\nfrom torch.utils.data import DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:28.854447Z","iopub.execute_input":"2026-05-11T17:41:28.854969Z","iopub.status.idle":"2026-05-11T17:41:36.009042Z","shell.execute_reply.started":"2026-05-11T17:41:28.854936Z","shell.execute_reply":"2026-05-11T17:41:36.007930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.Grayscale(num_output_channels=3),  # convert to 3-channel\n    transforms.ToTensor(),\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:36.010322Z","iopub.execute_input":"2026-05-11T17:41:36.010933Z","iopub.status.idle":"2026-05-11T17:41:36.018538Z","shell.execute_reply.started":"2026-05-11T17:41:36.010890Z","shell.execute_reply":"2026-05-11T17:41:36.017514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nDATA_PATH = \"/kaggle/input/datasets/manmandes/malimg/malimg_dataset\"\n\nos.listdir(DATA_PATH)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:36.021309Z","iopub.execute_input":"2026-05-11T17:41:36.021631Z","iopub.status.idle":"2026-05-11T17:41:36.050613Z","shell.execute_reply.started":"2026-05-11T17:41:36.021601Z","shell.execute_reply":"2026-05-11T17:41:36.049724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BASE_PATH = \"/kaggle/input/datasets/manmandes/malimg/malimg_dataset\"\n\ntrain_data = datasets.ImageFolder(os.path.join(BASE_PATH, \"train\"), transform=transform)\nval_data = datasets.ImageFolder(os.path.join(BASE_PATH, \"val\"), transform=transform)\ntest_data = datasets.ImageFolder(os.path.join(BASE_PATH, \"test\"), transform=transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:36.051828Z","iopub.execute_input":"2026-05-11T17:41:36.052111Z","iopub.status.idle":"2026-05-11T17:41:39.917782Z","shell.execute_reply.started":"2026-05-11T17:41:36.052084Z","shell.execute_reply":"2026-05-11T17:41:39.916852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(train_data, batch_size=32, shuffle=True)\nval_loader = DataLoader(val_data, batch_size=32, shuffle=False)\ntest_loader = DataLoader(test_data, batch_size=32, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:39.918992Z","iopub.execute_input":"2026-05-11T17:41:39.919823Z","iopub.status.idle":"2026-05-11T17:41:39.925265Z","shell.execute_reply.started":"2026-05-11T17:41:39.919786Z","shell.execute_reply":"2026-05-11T17:41:39.924090Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_data.classes)\nprint(\"Number of classes:\", len(train_data.classes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:39.926548Z","iopub.execute_input":"2026-05-11T17:41:39.926864Z","iopub.status.idle":"2026-05-11T17:41:39.951504Z","shell.execute_reply.started":"2026-05-11T17:41:39.926836Z","shell.execute_reply":"2026-05-11T17:41:39.950416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import models","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:39.953139Z","iopub.execute_input":"2026-05-11T17:41:39.953441Z","iopub.status.idle":"2026-05-11T17:41:39.979293Z","shell.execute_reply.started":"2026-05-11T17:41:39.953413Z","shell.execute_reply":"2026-05-11T17:41:39.978266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:39.980426Z","iopub.execute_input":"2026-05-11T17:41:39.980811Z","iopub.status.idle":"2026-05-11T17:41:40.001971Z","shell.execute_reply.started":"2026-05-11T17:41:39.980783Z","shell.execute_reply":"2026-05-11T17:41:40.000821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pretrained resnet\n\nmodel = models.resnet18(pretrained=True)\n\n# Modify final layer\nnum_classes = 25\nmodel.fc = nn.Linear(model.fc.in_features, num_classes)\n\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:40.003372Z","iopub.execute_input":"2026-05-11T17:41:40.003853Z","iopub.status.idle":"2026-05-11T17:41:40.761341Z","shell.execute_reply.started":"2026-05-11T17:41:40.003796Z","shell.execute_reply":"2026-05-11T17:41:40.760272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.001)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:40.762877Z","iopub.execute_input":"2026-05-11T17:41:40.763771Z","iopub.status.idle":"2026-05-11T17:41:40.769374Z","shell.execute_reply.started":"2026-05-11T17:41:40.763734Z","shell.execute_reply":"2026-05-11T17:41:40.768529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, train_loader, val_loader, epochs=15):\n    for epoch in range(epochs):\n        model.train()\n        total_loss = 0\n        correct = 0\n\n        for images, labels in train_loader:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            correct += (preds == labels).sum().item()\n\n        train_acc = correct / len(train_loader.dataset)\n\n        print(f\"Epoch {epoch+1}/{epochs}, Loss: {total_loss:.4f}, Train Acc: {train_acc:.4f}\")\n\n        evaluate(model, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:40.770697Z","iopub.execute_input":"2026-05-11T17:41:40.771065Z","iopub.status.idle":"2026-05-11T17:41:40.791391Z","shell.execute_reply.started":"2026-05-11T17:41:40.771024Z","shell.execute_reply":"2026-05-11T17:41:40.790329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(model, loader):\n    model.eval()\n    correct = 0\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n\n            outputs = model(images)\n            _, preds = torch.max(outputs, 1)\n\n            correct += (preds == labels).sum().item()\n\n    acc = correct / len(loader.dataset)\n    print(f\"Test Accuracy: {acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:40.795766Z","iopub.execute_input":"2026-05-11T17:41:40.796100Z","iopub.status.idle":"2026-05-11T17:41:40.816123Z","shell.execute_reply.started":"2026-05-11T17:41:40.796073Z","shell.execute_reply":"2026-05-11T17:41:40.815205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_model(model, train_loader, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T17:41:40.817376Z","iopub.execute_input":"2026-05-11T17:41:40.817835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#baseline CNN accuracy\n\nevaluate(model, test_loader)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, confusion_matrix\nimport numpy as np\n\ndef get_preds(model, loader):\n    model.eval()\n    y_true, y_pred = [], []\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(device)\n            out = model(x)\n            preds = out.argmax(1).cpu().numpy()\n            y_pred.extend(preds)\n            y_true.extend(y.numpy())\n    return np.array(y_true), np.array(y_pred)\n\ny_true, y_pred = get_preds(model, test_loader)\n\nprint(classification_report(y_true, y_pred, target_names=train_data.classes))\n\ncm = confusion_matrix(y_true, y_pred)\nprint(cm)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#SE Attention Block\nimport torch\nimport torch.nn as nn\n\nclass SEBlock(nn.Module):\n    def __init__(self, channels, reduction=16):\n        super(SEBlock, self).__init__()\n        self.fc1 = nn.Linear(channels, channels // reduction)\n        self.fc2 = nn.Linear(channels // reduction, channels)\n\n    def forward(self, x):\n        b, c, h, w = x.size()\n\n        y = x.view(b, c, -1).mean(dim=2)  # global avg pool\n        y = torch.relu(self.fc1(y))\n        y = torch.sigmoid(self.fc2(y))\n        y = y.view(b, c, 1, 1)\n\n        return x * y","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Modify ResNet18\n\nfrom torchvision import models\n\nclass ResNetWithSE(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n\n        self.base = models.resnet18(pretrained=True)\n\n        self.se = SEBlock(512)\n\n        self.base.fc = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        x = self.base.conv1(x)\n        x = self.base.bn1(x)\n        x = self.base.relu(x)\n        x = self.base.maxpool(x)\n\n        x = self.base.layer1(x)\n        x = self.base.layer2(x)\n        x = self.base.layer3(x)\n        x = self.base.layer4(x)\n\n        x = self.se(x)  # 🔥 ATTENTION HERE\n\n        x = self.base.avgpool(x)\n        x = torch.flatten(x, 1)\n        x = self.base.fc(x)\n\n        return x","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#intialise new model \nmodel_att = ResNetWithSE(num_classes=25).to(device)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model_att.parameters(), lr=0.0003, weight_decay=1e-4)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, train_loader, val_loader, epochs=15):\n    train_acc_list = []\n    val_acc_list = []\n    train_loss_list = []\n\n    for epoch in range(epochs):\n        model.train()\n        total_loss = 0\n        correct = 0\n\n        for images, labels in train_loader:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            loss.backward()\n            optimizer.step()\n\n            total_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            correct += (preds == labels).sum().item()\n\n        # ✅ Fix here\n        avg_loss = total_loss / len(train_loader)\n\n        train_acc = correct / len(train_loader.dataset)\n        val_acc = evaluate_return(model, val_loader)\n\n        train_acc_list.append(train_acc)\n        val_acc_list.append(val_acc)\n        train_loss_list.append(avg_loss)\n\n        print(f\"Epoch {epoch+1}: Train Acc={train_acc:.4f}, Val Acc={val_acc:.4f}, Loss={avg_loss:.4f}\")\n\n    return train_acc_list, val_acc_list, train_loss_list","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_return(model, loader):\n    model.eval()\n    correct = 0\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images, labels = images.to(device), labels.to(device)\n\n            outputs = model(images)\n            _, preds = torch.max(outputs, 1)\n            correct += (preds == labels).sum().item()\n\n    return correct / len(loader.dataset)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_acc, val_acc, train_loss = train_model(model_att, train_loader, val_loader)\n\nplt.figure()\nplt.plot(train_acc, label='Train Accuracy')\nplt.plot(val_acc, label='Validation Accuracy')\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\nplt.legend()\nplt.title('Training vs Validation Accuracy')\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#attention cnn accuracy \nprint(f\"Test Accuracy: {evaluate_return(model_att, test_loader):.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure()\nplt.plot(train_loss, label='Training Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\nplt.title('Training Loss Curve')\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, confusion_matrix\nimport numpy as np\n\ndef get_preds(model, loader):\n    model.eval()\n    y_true, y_pred = [], []\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(device)\n            out = model(x)\n            preds = out.argmax(1).cpu().numpy()\n            y_pred.extend(preds)\n            y_true.extend(y.numpy())\n    return np.array(y_true), np.array(y_pred)\n\ny_true, y_pred = get_preds(model_att, test_loader)\n\nprint(classification_report(y_true, y_pred, target_names=train_data.classes))\n\ncm = confusion_matrix(y_true, y_pred)\nprint(cm)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns\nfrom sklearn.metrics import confusion_matrix\n\ncm = confusion_matrix(y_true, y_pred)\n\nplt.figure(figsize=(12,10))\nsns.heatmap(cm, cmap=\"Blues\")\nplt.title(\"Confusion Matrix\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sns.heatmap(cm, xticklabels=train_data.classes, yticklabels=train_data.classes)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = ['Baseline CNN', 'CNN + Attention']\naccuracy = [0.9655, 0.9937]\n\nplt.figure()\nplt.bar(models, accuracy)\nplt.ylabel('Accuracy')\nplt.title('Model Comparison')\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#grad-came {Explainable-AI}\n\n!pip install grad-cam","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pytorch_grad_cam import GradCAM\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\nimport cv2\nimport numpy as np","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_layer = model_att.base.layer4[-1]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cam = GradCAM(model=model_att, target_layers=[target_layer])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = test_data\nimg, label = dataset[0]\n\ninput_tensor = img.unsqueeze(0).to(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"grayscale_cam = cam(input_tensor=input_tensor)[0]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_np = img.permute(1, 2, 0).cpu().numpy()\nimg_np = img_np / img_np.max()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualization = show_cam_on_image(img_np, grayscale_cam, use_rgb=True)\n\nimport matplotlib.pyplot as plt\n\nplt.imshow(visualization)\nplt.title(f\"True: {train_data.classes[label]}\")\nplt.axis('off')\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Comparision of plain cnn vs att cnn","metadata":{}},{"cell_type":"code","source":"model.eval()\nmodel_att.eval()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_layer_base = model.layer4[-1]          # baseline\ntarget_layer_att = model_att.base.layer4[-1] # attention","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cam_base = GradCAM(model=model, target_layers=[target_layer_base])\ncam_att = GradCAM(model=model_att, target_layers=[target_layer_att])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img, label = test_data[10]   # change index for different cases\ninput_tensor = img.unsqueeze(0).to(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cam_base_out = cam_base(input_tensor=input_tensor)[0]\ncam_att_out = cam_att(input_tensor=input_tensor)[0]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_np = img.permute(1, 2, 0).cpu().numpy()\nimg_np = img_np / img_np.max()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"vis_base = show_cam_on_image(img_np, cam_base_out, use_rgb=True)\nvis_att = show_cam_on_image(img_np, cam_att_out, use_rgb=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12,4))\n\nplt.subplot(1,3,1)\nplt.imshow(img_np)\nplt.title(\"Original\")\nplt.axis('off')\n\nplt.subplot(1,3,2)\nplt.imshow(vis_base)\nplt.title(\"Baseline CNN\")\nplt.axis('off')\n\nplt.subplot(1,3,3)\nplt.imshow(vis_att)\nplt.title(\"Attention CNN\")\nplt.axis('off')\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef show_comparison(idx):\n    img, label = test_data[idx]\n    input_tensor = img.unsqueeze(0).to(device)\n\n    cam_base_out = cam_base(input_tensor=input_tensor)[0]\n    cam_att_out = cam_att(input_tensor=input_tensor)[0]\n\n    img_np = img.permute(1,2,0).cpu().numpy()\n    img_np = img_np / img_np.max()\n\n    vis_base = show_cam_on_image(img_np, cam_base_out, use_rgb=True)\n    vis_att = show_cam_on_image(img_np, cam_att_out, use_rgb=True)\n\n    plt.figure(figsize=(10,3))\n\n    plt.subplot(1,3,1)\n    plt.imshow(img_np)\n    plt.title(\"Original\")\n    plt.axis('off')\n\n    plt.subplot(1,3,2)\n    plt.imshow(vis_base)\n    plt.title(\"Baseline\")\n    plt.axis('off')\n\n    plt.subplot(1,3,3)\n    plt.imshow(vis_att)\n    plt.title(\"Attention\")\n    plt.axis('off')\n\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(10, 30):\n    show_comparison(i)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_class_index(class_name):\n    for i in range(len(test_data)):\n        _, label = test_data[i]\n        if train_data.classes[label] == class_name:\n            return i\n    return None\n\nidx_autorun = find_class_index(\"Autorun.K\")\nprint(\"Autorun.K index:\", idx_autorun)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx_swizzor = find_class_index(\"Swizzor.gen!E\")  # or \"Swizzor.gen!I\"\nprint(\"Swizzor index:\", idx_swizzor)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_gradcam_comparison(idx):\n    img, label = test_data[idx]\n    input_tensor = img.unsqueeze(0).to(device)\n\n    # CAMs\n    cam_base_out = cam_base(input_tensor=input_tensor)[0]\n    cam_att_out = cam_att(input_tensor=input_tensor)[0]\n\n    # Image prep\n    img_np = img.permute(1,2,0).cpu().numpy()\n    img_np = img_np / img_np.max()\n\n    vis_base = show_cam_on_image(img_np, cam_base_out, use_rgb=True)\n    vis_att = show_cam_on_image(img_np, cam_att_out, use_rgb=True)\n\n    class_name = train_data.classes[label]\n\n    import matplotlib.pyplot as plt\n    plt.figure(figsize=(12,4))\n\n    plt.subplot(1,3,1)\n    plt.imshow(img_np)\n    plt.title(f\"Original\\n{class_name}\")\n    plt.axis('off')\n\n    plt.subplot(1,3,2)\n    plt.imshow(vis_base)\n    plt.title(\"Baseline CNN\")\n    plt.axis('off')\n\n    plt.subplot(1,3,3)\n    plt.imshow(vis_att)\n    plt.title(\"Attention CNN\")\n    plt.axis('off')\n\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_gradcam_comparison(idx_autorun)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"show_gradcam_comparison(idx_swizzor)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}