{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"},{"sourceId":11968461,"sourceType":"datasetVersion","datasetId":7525964}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Preliminaries\nDownload libraries, set up device, load model (ViT Base)","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport random\nimport matplotlib.pyplot as plt\nimport matplotlib.cm as cm\nimport shutil\nimport numpy as np\nimport pandas as pd\nfrom matplotlib.colors import Normalize\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader, Dataset, ConcatDataset, Subset\nfrom torchvision.datasets import ImageFolder\nfrom transformers import ViTImageProcessor, ViTModel\nfrom concurrent.futures import ThreadPoolExecutor\nfrom datetime import date\nfrom PIL import Image","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# init random generator\nrandom.seed(42)\n\n# Set device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Load feature extractor and model\nmodel_name = \"google/vit-base-patch16-224-in21k\"\nprocessor = ViTImageProcessor.from_pretrained(model_name)\nmodel = ViTModel.from_pretrained(model_name).to(device).eval()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset\nLoad dataset (IMAGENET), then create dataloader, and define transformation","metadata":{}},{"cell_type":"code","source":"imagenet_path = \"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train\"\n\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Resize to 224x224 for models like ViT\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])  # Standard normalization\n])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CustomImageDataset(Dataset):\n    def __init__(self, root, label, transform=None):\n        self.root = root\n        self.transform = transform\n        self.image_paths = [os.path.join(root, fname) for fname in os.listdir(root)]\n        self.label = label\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        image = Image.open(img_path).convert(\"RGB\")\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"zebra_id = \"n02391049\" # zebras\n\nzebra_subset = CustomImageDataset(root=f\"{imagenet_path}/{zebra_id}\", label=340, transform=transform)\nzebra_loader = DataLoader(zebra_subset, batch_size=64, shuffle=False, num_workers=4, pin_memory=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Functions\nDefine functions to extract patch embeddings, and visualize them","metadata":{}},{"cell_type":"code","source":"def extract_patch_embeddings_after(image_tensor, model):\n    if isinstance(image_tensor, np.ndarray):\n        image_tensor = torch.from_numpy(image_tensor).float()\n    device = next(model.parameters()).device\n    image_tensor = image_tensor.to(device)\n    \n    if len(image_tensor.shape) == 3:\n        if image_tensor.shape[-1] == 3:\n            image_tensor = image_tensor.permute(2, 0, 1)\n        image_tensor = image_tensor.unsqueeze(0)\n    elif len(image_tensor.shape) == 4:\n        if image_tensor.shape[-1] == 3: \n            image_tensor = image_tensor.permute(0, 3, 1, 2) \n            \n    with torch.no_grad():\n        outputs = model(image_tensor)\n        cls = outputs.last_hidden_state[:, 0, :]\n        patch_embeddings = outputs.last_hidden_state[:, 1:, :] \n    return patch_embeddings, cls\n\ndef numpy_image(image):\n    if isinstance(image, torch.Tensor):\n        if len(image.shape) == 4: \n            image_np = image.cpu().squeeze(0).permute(1, 2, 0).numpy() \n        elif len(image.shape) == 3:\n            image_np = image.cpu().permute(1, 2, 0).numpy()  # Convert to (H, W, C)\n        else:\n            image_np = image.cpu().numpy()\n    else:\n        if len(image.shape) == 4:\n            image_np = image.squeeze(0)\n        elif len(image.shape) == 3 and image.shape[0] == 3:\n            image_np = np.transpose(image, (1, 2, 0))\n        else:\n            image_np = image \n    \n    if image_np.max() > 1.0 or image_np.min() < 0.0:\n        image_np = (image_np - image_np.min()) / (image_np.max() - image_np.min())\n\n    return image_np\n\ndef visualize_patches_overlay(image, patch_embeddings, original_img=None):\n       \n        image_np = numpy_image(image)\n        if original_img is not None:\n            original_img = numpy_image(original_img)\n    \n        h, w = image_np.shape[:2]\n        \n        if isinstance(patch_embeddings, torch.Tensor):\n            patch_embeddings = patch_embeddings.cpu().numpy()\n        \n        patch_activations = patch_embeddings.mean(axis=-1)\n        patch_activations = (patch_activations - patch_activations.min()) / (patch_activations.max() - patch_activations.min())\n        \n        if len(patch_activations.shape) == 2:  # If shape is (B, num_patches)\n            patch_activations = patch_activations.reshape(-1)  # Convert to flattened\n        \n        patch_activations = patch_activations.flatten()[:196].reshape(14, 14)\n        \n        if original_img is not None:\n            fig, axes = plt.subplots(1, 4, figsize=(18, 6))\n            axes[3].imshow(original_img)\n            axes[3].set_title(\"Unshuffled Image\")\n            axes[3].axis(\"off\")\n        else:\n            fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n        \n        # 1. Original image\n        axes[0].imshow(image_np)\n        axes[0].set_title(\"Original Image\")\n        axes[0].axis(\"off\")\n        \n        # 2. Patch activation heatmap\n        im = axes[1].imshow(patch_activations, cmap=\"RdBu_r\", interpolation=\"nearest\")\n        axes[1].set_title(\"Patch Activations\")\n        axes[1].axis(\"off\")\n        plt.colorbar(im, ax=axes[1], fraction=0.046, pad=0.04)\n        \n        # 3. Overlay visualization\n        axes[2].imshow(image_np)\n        \n        # Create upsampled heatmap to match image size\n        patch_size_h, patch_size_w = h // 14, w // 14\n        heatmap = np.zeros((h, w))\n        \n        # Resize patch activations to match image size\n        for i in range(14):\n            for j in range(14):\n                activation_value = patch_activations[i, j]\n                heatmap[i*patch_size_h:(i+1)*patch_size_h, j*patch_size_w:(j+1)*patch_size_w] = activation_value\n        \n        # Normalize heatmap for colormap\n        norm = Normalize(vmin=heatmap.min(), vmax=heatmap.max())\n        heatmap_colored = cm.magma(norm(heatmap))\n        \n        # Set alpha based on activation intensity\n        alpha = 0.85  # Adjust overlay transparency\n        heatmap_colored[..., 3] = alpha * norm(heatmap)\n        \n        # Add heatmap overlay\n        axes[2].imshow(heatmap_colored, alpha=0.7)\n        axes[2].set_title(\"Overlay Visualization\")\n        axes[2].axis(\"off\")\n        \n        plt.tight_layout()\n        plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Tests\nTesting the code using 5 randomly chosen sample image","metadata":{}},{"cell_type":"code","source":"# Select a sample image from the test loader\ncount = 0\nfor inputs in zebra_loader:\n    inputs = inputs.to(device)\n\n    patch_embeddings_after, cls = extract_patch_embeddings_after(inputs, model)\n\n    for i in range(inputs.size(0)):\n        randnum = random.randint(0, inputs.size(0))\n        visualize_patches_overlay(inputs[randnum], patch_embeddings_after[randnum])  # Visualize overlayed patches\n        # print(cls[randnum])\n        count += 1\n        if count == 5:\n            break  # Stop after 5 images, done only to reduce computation time\n    \n    if count == 5:\n        break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Compare Embeddings\nLoad datasets of concepts, thanks to: https://captum.ai/tutorials/TCAV_Image","metadata":{}},{"cell_type":"code","source":"# courtesy of ChatGPT, but very fast!\n\n# Load index once\nindex = pd.read_csv(\"/kaggle/input/broden-concepts/broden1_224/index.csv\")\nbroden_base_path = \"/kaggle/input/broden-concepts/broden1_224\"\nimage_base_path = os.path.join(broden_base_path, \"images\")\noutput_base_path = \"/kaggle/working/concepts\"\n\nconcept_files = {\n    \"texture\": \"c_texture.csv\",\n    \"color\": \"c_color.csv\",\n    \"material\": \"c_material.csv\",\n    \"object\": \"c_object.csv\",\n    \"part\": \"c_part.csv\",\n    \"scene\": \"c_scene.csv\"\n}\n\nconcept_loader = {}\n\n# Preload index columns as strings to avoid repeated conversion\nindex_str_cols = {col: index[col].astype(str) for col in index.columns}\n\ndef copy_images(rows, concept_dir):\n    os.makedirs(concept_dir, exist_ok=True)\n    tasks = []\n    for rel_path in rows['image']:\n        src = os.path.join(image_base_path, rel_path)\n        dst = os.path.join(concept_dir, os.path.basename(src))\n        if not os.path.exists(dst):\n            tasks.append((src, dst))\n    with ThreadPoolExecutor(max_workers=8) as executor:\n        executor.map(lambda x: shutil.copy(*x), tasks)\n\nfor concept_type, csv_file in concept_files.items():\n    df = pd.read_csv(os.path.join(broden_base_path, csv_file))\n    names_and_codes = list(zip(df['name'].astype(str), df['number'].astype(int)))\n    concept_loader[concept_type] = []\n\n    if concept_type not in index.columns:\n        continue\n\n    column_str = index_str_cols[concept_type]\n\n    for name, code in names_and_codes:\n        concept_dir = os.path.join(output_base_path, concept_type, name)\n\n        # Use faster string contains filtering\n        matching_rows = index[column_str.str.contains(str(code), na=False)]\n\n        if not matching_rows.empty:\n            copy_images(matching_rows, concept_dir)\n\n            subset = CustomImageDataset(root=concept_dir, label=code, transform=transform)\n            loader = DataLoader(subset, batch_size=len(subset), shuffle=False,\n                                num_workers=4, pin_memory=True)\n            concept_loader[concept_type].append(loader)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_dataloaders_grid(dataloaders, columns=5):\n    num_dataloaders = len(dataloaders)\n    rows = (num_dataloaders + columns - 1) // columns\n\n    fig, axs = plt.subplots(rows, columns, figsize=(columns * 3, rows * 3))\n    axs = axs.flatten()\n    images_to_plot = []\n    titles = []\n\n    for i, dl in enumerate(dataloaders):\n        for topic in dataloaders[dl]:\n            for image in topic:\n                randnum = random.randint(0, image.size(0) - 1)\n                img = image[randnum]\n                img_np = img.permute(1, 2, 0).cpu().numpy()\n                images_to_plot.append(img_np)\n                # titles.append(f\"{concept_type}/{concept_names[i]}\")\n                titles.append(f\"{concept_type}\")\n                break\n\n    for i, (img_np, title) in enumerate(zip(images_to_plot, titles)):\n        axs[i].imshow(img_np)\n        axs[i].axis('off')\n        axs[i].set_title(title)\n\n    for j in range(i + 1, len(axs)):\n        axs[j].axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot_dataloaders_grid(concept_loader)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Extract patch embeddings of known concepts. Compute centroids of patches\n\n![cosine similarity](https://www.timescale.com/_next/image?url=https%3A%2F%2Fimages.ctfassets.net%2Fnpizagvkn99r%2F4N5sVZo7nNVf7OD2C9IoXq%2Ff8be1c88ff13ff6e27be2ae06650edba%2FCosine_similarity_python_formula.png%3Ffm%3Djpg%26fl%3Dprogressive&w=1920&q=75)","metadata":{}},{"cell_type":"code","source":"def centroid(embeddings):\n    return embeddings.mean(dim=0)\n\ncentroid_dict = {}\n\ni = 0\n\nfor concept_type, csv_file in concept_files.items():\n    df = pd.read_csv(os.path.join(broden_base_path, csv_file))\n    names_and_codes = list(zip(df['name'].astype(str), df['number'].astype(int)))\n\n    temp_dict = {}\n    \n    for dataset in concept_loader:\n        for dl in concept_loader[dataset]:\n            for inputs in dl:\n                inputs = inputs.to(device)\n                patch_embeddings_after, cls = extract_patch_embeddings_after(inputs, model)\n                temp_dict[names_and_codes[i]] = centroid(patch_embeddings_after)\n                i += 1\n    \n    print(temp_dict)\n    centroid_dict[concept_type] = temp_dict","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def cos_sim(target, centroid):\n\n    t_cpu = target.cpu().numpy().reshape(-1)\n    c_cpu = centroid.cpu().numpy().reshape(-1)\n    \n    dot_product = np.dot(t_cpu, c_cpu)\n    magnitude_A = np.linalg.norm(t_cpu)\n    magnitude_B = np.linalg.norm(c_cpu)\n    cosine_similarity = dot_product / (magnitude_A * magnitude_B)\n    \n    return cosine_similarity","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Compare centroids with target image\nImage-wise","metadata":{}},{"cell_type":"code","source":"def pick_samples(dataloader, k=15):\n    all_embeddings = []\n\n    for inputs in dataloader:\n        inputs = inputs.to(device)\n        embeddings, cls = extract_patch_embeddings_after(inputs, model)\n        all_embeddings.append(embeddings.detach().cpu())\n\n    all_embeddings = torch.cat(all_embeddings, dim=0)\n\n    if len(all_embeddings) < k:\n        k = len(all_embeddings)\n    indices = torch.randperm(len(all_embeddings))[:k]\n    samples = all_embeddings[indices]\n\n    return samples","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_similarity(samples, concept_type=\"texture\"):\n    similarity_scores = {}\n\n    relevant_concepts = centroid_dict.get(concept_type, {})\n\n    for (name, code), concept_centroid in relevant_concepts.items():\n        similarities = []\n        for sample in samples:\n            embedding = sample.reshape(-1)\n            similarity = cos_sim(concept_centroid, embedding)\n            similarities.append(similarity.item() if hasattr(similarity, \"item\") else similarity)\n        \n        similarity_scores[name] = np.mean(similarities)\n\n    return similarity_scores","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_similarities(similarity):\n    plt.figure(figsize=(8, 5))\n    plt.bar(similarity.keys(), similarity.values(), color='skyblue')\n    plt.xlabel('Concepts')\n    plt.ylabel('Cosine Similarity')\n    plt.title('Similarity Scores per texture')\n    plt.xticks(rotation=90)\n    plt.grid(axis='y', linestyle='--', alpha=0.7)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"zebra_samples = pick_samples(zebra_loader)\n\nsimilarities = {}\n\nfor concept_type in concept_files.items():\n    similarities[concept_type] = compute_similarity(zebra_samples)\n    show_similarities(similarities[concept_type])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Analyze the attention scores\nExtract attention scores from\n* Last layer\n* Every other layer\n  \nTo understand where the model is focusing more, based on the level of abstraction of the layers.\nhttps://arxiv.org/abs/2012.14913","metadata":{}},{"cell_type":"code","source":"# print(model.get_submodule(\"encoder.layer\"))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Reference for the code used below: [here](http://https://medium.com/@nivonl/exploring-visual-attention-in-transformer-models-ab538c06083a), [here](https://www.kaggle.com/code/piantic/vision-transformer-vit-visualize-attention-map) and chatgpt","metadata":{}},{"cell_type":"code","source":"def plot_attention_map(original_img_tensor, att_2d_map, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]):\n    img = original_img_tensor.clone()\n    for t, m, s in zip(img, mean, std):\n        t.mul_(s).add_(m)\n    img = img.permute(1, 2, 0).cpu().numpy()\n\n    if not isinstance(att_2d_map, torch.Tensor):\n        att_2d_map = torch.tensor(att_2d_map)\n\n    att_resized = T.functional.resize(\n        att_2d_map.unsqueeze(0).unsqueeze(0),\n        img.shape[:2],\n        interpolation=T.InterpolationMode.BILINEAR\n    ).squeeze().numpy()\n\n    # Plot\n    fig, (ax1, ax2) = plt.subplots(ncols=2, figsize=(12, 6))\n    ax1.set_title('Original Image')\n    ax1.imshow(img)\n    ax1.axis('off')\n\n    ax2.set_title('Attention Map Overlay')\n    ax2.imshow(img)\n    ax2.imshow(att_resized, cmap='inferno', alpha=0.5)\n    ax2.axis('off')\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# compute & plot rolling attention\ndef attention_rollout(attentions):\n    rollout = torch.eye(attentions[0].size(-1)).to(attentions[0].device)\n\n    for attention in attentions:\n        attention_heads_fused = attention.mean(dim=1) # Average attention across heads\n        attention_heads_fused += torch.eye(attention_heads_fused.size(-1)).to(attention_heads_fused.device) # A + I\n        attention_heads_fused /= attention_heads_fused.sum(dim=-1, keepdim=True) # Normalizing A\n        rollout = torch.matmul(rollout, attention_heads_fused) # Multiplication\n\n    return rollout","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"count = 0\n\n# feeding batches of zebras\nfor inputs in zebra_loader:\n    inputs = inputs.to(device)\n    with torch.no_grad():\n        outputs = model(inputs, output_attentions=True)\n        attentions = outputs.attentions\n        rollout = attention_rollout(attentions)  # Shape: [B, N, N]\n\n        for i in range(0, 5):\n            cls_attention = rollout[i, 1:, 0]\n            side = int(np.sqrt(cls_attention.shape[0]))\n            cls_attention = 1 - cls_attention.reshape(side, side)\n\n            plot_attention_map(inputs[i].cpu(), cls_attention.cpu())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot raw attentions of the first image\nig, axs = plt.subplots(3, 4, figsize=(20, 20))\nfor i, ax in enumerate(axs.flatten()):\n    ax.imshow(attention_rollout(attentions[-1][0, i, :, :].detach().cpu().numpy()))\n    ax.axis('off')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}