# Iterate over dataloader for idx, (image, name) in enumerate(predict_loader): # Compute heatmap image = image.to(device) heatmap = gradCAM(image) image = image.cpu().squeeze(0).permute(1,2,0) heatmap = upsampleHeatmap(heatmap, image) # Plot images and heatmaps fig = plt.figure(figsize=(14,5)) fig.suptitle(f"nFile: {names[idx]}, Predicted label: {lbls[idx]}n", fontsize=24) plt.subplot(1, 2, 1) plt.imshow(image, cmap="gray") plt.title(f"Image", fontsize=14) plt.axis("off") plt.subplot(1, 2, 2) plt.imshow(heatmap) plt.title(f"Heatmap", fontsize=14) plt.tight_layout() plt.axis("off") plt.show() plt.close()