{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n!pip install git+https://github.com/openai/CLIP.git","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-14T14:53:12.360019Z","iopub.execute_input":"2022-08-14T14:53:12.361074Z","iopub.status.idle":"2022-08-14T14:53:30.150848Z","shell.execute_reply.started":"2022-08-14T14:53:12.360964Z","shell.execute_reply":"2022-08-14T14:53:30.149278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision.transforms import InterpolationMode\nimport torchvision.transforms as T\nimport matplotlib.pyplot as plt\nimport clip\nfrom scipy.ndimage import filters\nfrom PIL import Image\nimport requests","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:49:16.765928Z","iopub.execute_input":"2022-08-14T15:49:16.766449Z","iopub.status.idle":"2022-08-14T15:49:16.773892Z","shell.execute_reply.started":"2022-08-14T15:49:16.766411Z","shell.execute_reply":"2022-08-14T15:49:16.772906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CLIP GradCAM Visualizations\n\nThis code is an example of GradCAM applied to CLIP model.\n\nThe code is adapted from this original link:\nhttps://colab.research.google.com/github/kevinzakka/clip_playground/blob/main/CLIP_GradCAM_Visualization.ipynb","metadata":{}},{"cell_type":"markdown","source":"# Utility functions","metadata":{}},{"cell_type":"code","source":"def _convert_image_to_rgb(image):\n    return image.convert(\"RGB\")\n\ndef preprocess(n_px):\n    return T.Compose([\n        T.Resize(n_px, interpolation=T.InterpolationMode.BICUBIC),\n        T.CenterCrop(n_px),\n        _convert_image_to_rgb,\n        T.ToTensor(),\n        T.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)),\n    ])\n\ndef pil_preprocess(n_px):\n    return T.Compose([\n        T.Resize(n_px, interpolation=T.InterpolationMode.BICUBIC),\n        T.CenterCrop(n_px)\n    ])\n\ndef normalize(x: np.ndarray) -> np.ndarray:\n    # Normalize to [0, 1].\n    x = x - x.min()\n    if x.max() > 0:\n        x = x / x.max()\n    return x\n\ndef getAttMap(img, attn_map, blur=True):\n    if blur:\n        attn_map = filters.gaussian_filter(attn_map, 0.02*max(img.shape[:2]))\n    attn_map = normalize(attn_map)\n    cmap = plt.get_cmap('jet')\n    attn_map_c = np.delete(cmap(attn_map), 3, 2)\n    attn_map = 1*(1-attn_map**0.7).reshape(attn_map.shape + (1,))*img + \\\n            (attn_map**0.7).reshape(attn_map.shape+(1,)) * attn_map_c\n    return attn_map\n\ndef viz_attn(img, attn_map, image_caption, blur=True):\n    fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n    axes[0].imshow(img)\n    axes[1].imshow(getAttMap(img, attn_map, blur))\n    plt.gca().set_title(image_caption)\n    for ax in axes:\n        ax.axis(\"off\")\n    plt.show()\n\ndef load_image(img_path, clip_input_size, resize=None):\n    image = Image.open(img_path).convert(\"RGB\")\n    image = pil_preprocess(clip_input_size)(image)\n    if resize is not None:\n        image = image.resize((resize, resize))\n    return np.asarray(image).astype(np.float32) / 255.\n\ndef get_image_from_url(url):\n    img_data = requests.get(url).content\n    with open('image.jpg', 'wb') as handler:\n        handler.write(img_data)","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:49:40.142007Z","iopub.execute_input":"2022-08-14T15:49:40.142720Z","iopub.status.idle":"2022-08-14T15:49:40.178655Z","shell.execute_reply.started":"2022-08-14T15:49:40.142670Z","shell.execute_reply":"2022-08-14T15:49:40.177337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GradCAM method implementation","metadata":{}},{"cell_type":"code","source":"\nclass Hook:\n    \"\"\"Attaches to a module and records its activations and gradients.\"\"\"\n\n    def __init__(self, module: nn.Module):\n        self.data = None\n        self.hook = module.register_forward_hook(self.save_grad)\n        \n    def save_grad(self, module, input, output):\n        self.data = output\n        output.requires_grad_(True)\n        output.retain_grad()\n        \n    def __enter__(self):\n        return self\n    \n    def __exit__(self, exc_type, exc_value, exc_traceback):\n        self.hook.remove()\n        \n    @property\n    def activation(self) -> torch.Tensor:\n        return self.data\n    \n    @property\n    def gradient(self) -> torch.Tensor:\n        return self.data.grad\n\n\n# Reference: https://arxiv.org/abs/1610.02391\ndef gradCAM(\n    model: nn.Module,\n    input: torch.Tensor,\n    target: torch.Tensor,\n    layer: nn.Module\n) -> torch.Tensor:\n    # Zero out any gradients at the input.\n    if input.grad is not None:\n        input.grad.data.zero_()\n        \n    # Disable gradient settings.\n    requires_grad = {}\n    for name, param in model.named_parameters():\n        requires_grad[name] = param.requires_grad\n        param.requires_grad_(False)\n        \n    # Attach a hook to the model at the desired layer.\n    assert isinstance(layer, nn.Module)\n    with Hook(layer) as hook:        \n        # Do a forward and backward pass.\n        output = model(input)\n        output.backward(target)\n\n        grad = hook.gradient.float()\n        act = hook.activation.float()\n    \n        # Global average pool gradient across spatial dimension\n        # to obtain importance weights.\n        alpha = grad.mean(dim=(2, 3), keepdim=True)\n        # Weighted combination of activation maps over channel\n        # dimension.\n        gradcam = torch.sum(act * alpha, dim=1, keepdim=True)\n        # We only want neurons with positive influence so we\n        # clamp any negative ones.\n        gradcam = torch.clamp(gradcam, min=0)\n\n    # Resize gradcam to input resolution.\n    gradcam = F.interpolate(\n        gradcam,\n        input.shape[2:],\n        mode='bicubic',\n        align_corners=False)\n    \n    # Restore gradient settings.\n    for name, param in model.named_parameters():\n        param.requires_grad_(requires_grad[name])\n        \n    return gradcam","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:25:15.386302Z","iopub.execute_input":"2022-08-14T15:25:15.386761Z","iopub.status.idle":"2022-08-14T15:25:15.402649Z","shell.execute_reply.started":"2022-08-14T15:25:15.386728Z","shell.execute_reply":"2022-08-14T15:25:15.401263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizations","metadata":{}},{"cell_type":"code","source":"# Take care, this code does not work with ViT models, it must be adapted\nclip_model = \"RN50\"\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel, _ = clip.load(clip_model, device=device, jit=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-14T16:00:57.091667Z","iopub.execute_input":"2022-08-14T16:00:57.092146Z","iopub.status.idle":"2022-08-14T16:01:01.231020Z","shell.execute_reply.started":"2022-08-14T16:00:57.092110Z","shell.execute_reply":"2022-08-14T16:01:01.229998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_gradcam(model, image_path, image_caption):\n    # Download the image from the web.\n    image_input = preprocess(model.visual.input_resolution)(Image.open(image_path)).unsqueeze(0).to(device)\n    image_np = load_image(image_path, model.visual.input_resolution, model.visual.input_resolution)\n    text_input = clip.tokenize([image_caption]).to(device)\n\n    attn_map = gradCAM(\n        model.visual,\n        image_input,\n        model.encode_text(text_input).float(),\n        getattr(model.visual, \"layer4\")\n    )\n    attn_map = attn_map.squeeze().detach().cpu().numpy()\n\n    viz_attn(image_np, attn_map, image_caption, True)","metadata":{"execution":{"iopub.status.busy":"2022-08-14T16:01:01.233875Z","iopub.execute_input":"2022-08-14T16:01:01.234867Z","iopub.status.idle":"2022-08-14T16:01:01.241785Z","shell.execute_reply.started":"2022-08-14T16:01:01.234817Z","shell.execute_reply":"2022-08-14T16:01:01.240910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# The cats","metadata":{}},{"cell_type":"code","source":"image_path = '../input/clip-test-images/cat_color.jpg'\n\nvisualize_gradcam(model, image_path, 'the cat')\nvisualize_gradcam(model, image_path, 'the white cat')\nvisualize_gradcam(model, image_path, 'the black cat')","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:57:14.811746Z","iopub.execute_input":"2022-08-14T15:57:14.813063Z","iopub.status.idle":"2022-08-14T15:57:16.734095Z","shell.execute_reply.started":"2022-08-14T15:57:14.813027Z","shell.execute_reply":"2022-08-14T15:57:16.733231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# A photo of a store","metadata":{}},{"cell_type":"code","source":"image_path = '../input/clip-test-images/apple_store.jpg'\n\nvisualize_gradcam(model, image_path, 'apple')\nvisualize_gradcam(model, image_path, 'person')\nvisualize_gradcam(model, image_path, 'railing')","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:57:16.736298Z","iopub.execute_input":"2022-08-14T15:57:16.737259Z","iopub.status.idle":"2022-08-14T15:57:18.541382Z","shell.execute_reply.started":"2022-08-14T15:57:16.737219Z","shell.execute_reply":"2022-08-14T15:57:18.540102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Example with an oil bottle","metadata":{}},{"cell_type":"code","source":"image_path = '../input/clip-test-images/product_oil.jpg'\nvisualize_gradcam(model, image_path, 'the bottle')\nvisualize_gradcam(model, image_path, 'cap')\nvisualize_gradcam(model, image_path, 'the oil bottle')","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:57:18.543145Z","iopub.execute_input":"2022-08-14T15:57:18.543941Z","iopub.status.idle":"2022-08-14T15:57:20.161680Z","shell.execute_reply.started":"2022-08-14T15:57:18.543899Z","shell.execute_reply":"2022-08-14T15:57:20.160829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# A photo of a landmark","metadata":{}},{"cell_type":"code","source":"image_path = '../input/clip-test-images/colosseum_outside.jpg'\nvisualize_gradcam(model, image_path, 'colosseum')\nvisualize_gradcam(model, image_path, 'the arc')\nvisualize_gradcam(model, image_path, 'sky')","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:57:20.162927Z","iopub.execute_input":"2022-08-14T15:57:20.163925Z","iopub.status.idle":"2022-08-14T15:57:21.988608Z","shell.execute_reply.started":"2022-08-14T15:57:20.163880Z","shell.execute_reply":"2022-08-14T15:57:21.987530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Nutella time!","metadata":{}},{"cell_type":"code","source":"image_path = '../input/clip-test-images/kinder.jpg'\nvisualize_gradcam(model, image_path, 'nutella')\nvisualize_gradcam(model, image_path, 'candy')\nvisualize_gradcam(model, image_path, 'the nutella text')","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:57:21.990585Z","iopub.execute_input":"2022-08-14T15:57:21.991309Z","iopub.status.idle":"2022-08-14T15:57:23.889019Z","shell.execute_reply.started":"2022-08-14T15:57:21.991266Z","shell.execute_reply":"2022-08-14T15:57:23.887938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# You can download an image directly from the internet and have fun with it!","metadata":{}},{"cell_type":"code","source":"get_image_from_url(\"https://www.americanpost.news/wp-content/uploads/2022/03/Dragon-Ball-Super-Super-Hero-confirms-all-its-characters-in.jpg\")","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:49:45.711118Z","iopub.execute_input":"2022-08-14T15:49:45.712470Z","iopub.status.idle":"2022-08-14T15:49:45.953261Z","shell.execute_reply.started":"2022-08-14T15:49:45.712417Z","shell.execute_reply":"2022-08-14T15:49:45.952012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_gradcam(model, \"image.jpg\", 'goku')\nvisualize_gradcam(model, \"image.jpg\", 'gohan')\nvisualize_gradcam(model, \"image.jpg\", 'manga')","metadata":{"execution":{"iopub.status.busy":"2022-08-14T15:57:29.545184Z","iopub.execute_input":"2022-08-14T15:57:29.545739Z","iopub.status.idle":"2022-08-14T15:57:31.605706Z","shell.execute_reply.started":"2022-08-14T15:57:29.545697Z","shell.execute_reply":"2022-08-14T15:57:31.604803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}