{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils import data\nfrom torchvision.models import vgg19\nfrom torchvision import transforms\nfrom torchvision import datasets\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n\n# use the ImageNet transformation\ntransform = transforms.Compose([transforms.Resize((224, 224)), \n                                transforms.ToTensor(),\n                                transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n\n# define a 1 image dataset\ndataset = datasets.ImageFolder(root='../input/siim-isic-melanoma-classification/jpeg', transform=transform)\n\n\n# define the dataloader to load that single image\ndataloader = data.DataLoader(dataset=dataset, shuffle=False, batch_size=1)\nclass VGG(nn.Module):\n    def __init__(self):\n        super(VGG, self).__init__()\n        \n        # get the pretrained VGG19 network\n        self.vgg = vgg19(pretrained=True)\n        \n        # disect the network to access its last convolutional layer\n        self.features_conv = self.vgg.features[:36]\n        \n        # get the max pool of the features stem\n        self.max_pool = nn.MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)\n        \n        # get the classifier of the vgg19\n        self.classifier = self.vgg.classifier\n        \n        # placeholder for the gradients\n        self.gradients = None\n    \n    # hook for the gradients of the activations\n    def activations_hook(self, grad):\n        self.gradients = grad\n        \n    def forward(self, x):\n        x = self.features_conv(x)\n        \n        # register the hook\n        h = x.register_hook(self.activations_hook)\n        \n        # apply the remaining pooling\n        x = self.max_pool(x)\n        x = x.view((1, -1))\n        x = self.classifier(x)\n        return x\n    \n    # method for the gradient extraction\n    def get_activations_gradient(self):\n        return self.gradients\n    \n    # method for the activation exctraction\n    def get_activations(self, x):\n        return self.features_conv(x)\n    \n\n# initialize the VGG model\nvgg = VGG()\n\n# set the evaluation mode\nvgg.eval()\n\n# get the image from the dataloader\nimg, _ = next(iter(dataloader))\n\n# get the most likely prediction of the model\npred = vgg(img)\n#print(pred)\n\n\n# get the gradient of the output with respect to the parameters of the model\npred[:, 386].backward()\n\n# pull the gradients out of the model\ngradients = vgg.get_activations_gradient()\n\n# pool the gradients across the channels\npooled_gradients = torch.mean(gradients, dim=[0, 2, 3])\n\n# get the activations of the last convolutional layer\nactivations = vgg.get_activations(img).detach()\n\n# weight the channels by corresponding gradients\nfor i in range(512):\n    activations[:, i, :, :] *= pooled_gradients[i]\n    \n# average the channels of the activations\nheatmap = torch.mean(activations, dim=1).squeeze()\n\n# relu on top of the heatmap\n# expression (2) in https://arxiv.org/pdf/1610.02391.pdf\nheatmap = np.maximum(heatmap, 0)\n\n# normalize the heatmap\nheatmap /= torch.max(heatmap)\n\n# draw the heatmap\nplt.matshow(heatmap.squeeze())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import cv2\nimg = cv2.imread('../input/siim-isic-melanoma-classification/jpeg/train/ISIC_0074268.jpg')\n#heatmap = cv2.resize(cv2.UMat(heatmap), (img.shape[1], img.shape[0]))\nheatmap = cv2.resize(np.array(heatmap), (img.shape[1], img.shape[0]))\nheatmap = np.uint8(255 * heatmap)\nheatmap = cv2.applyColorMap(heatmap, cv2.COLORMAP_JET)\nsuperimposed_img = heatmap * 0.4 + img\ncv2.imwrite('./map.jpg', superimposed_img)","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}