{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":6799,"databundleVersionId":4225553}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Task 01: SAE Feature Discovery & Interpretability\nEvaluates a custom Sparse Autoencoder on LLaVA-1.5-7B's vision tower to discover\nand visualize monosemantic features for specific concepts.","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport gc\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom PIL import Image\nfrom pathlib import Path\nfrom huggingface_hub import hf_hub_download\nfrom transformers import LlavaForConditionalGeneration, AutoProcessor\nfrom datasets import load_dataset\nimport numpy as np\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\nCFG = {\n    # Replace with your actual Hugging Face repository ID\n    \"hf_repo_id\": \"AKG2/llava-micro-sae\", \n    \"sae_filename\": \"micro_sae_1024d.pt\",\n    \n    \"model_id\": \"llava-hf/llava-1.5-7b-hf\",\n    \"images_per_concept\": 100, # We only need 100 images per concept to test the SAE\n    \"patch_size\": 14,\n    \"image_size\": 336,\n    \"grid_size\": 24, # 336 / 14 = 24 patches per dimension\n    \n    \"output_dir\": \"/kaggle/working/feature_discovery\",\n}\n\nos.makedirs(CFG[\"output_dir\"], exist_ok=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# SAE Architecture Defination\nclass SparseAutoencoder(nn.Module):\n    \"\"\"Must match the architecture used during training exactly.\"\"\"\n    def __init__(self, input_dim: int, dict_size: int):\n        super().__init__()\n        self.encoder = nn.Linear(input_dim, dict_size)\n        self.decoder = nn.Linear(dict_size, input_dim, bias=True)\n\n    def encode(self, x: torch.Tensor) -> torch.Tensor:\n        return F.relu(self.encoder(x))\n\n    def forward(self, x: torch.Tensor):\n        z = self.encode(x)\n        return self.decoder(z), z","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# DATA LOADING\ndef load_evaluation_images():\n    print(\"=\" * 70)\n    print(\"STEP 1: Data Setup (Strict Training Distribution Match)\")\n    print(\"=\" * 70)\n    \n    datasets_dict = {\"baseline\": [], \"zebra\": [], \"fire_truck\": [], \"redness\": []}\n    \n    # --- 1. Fetch Baseline (COCO) ---\n    print(\"[...] Fetching Baseline (COCO) images...\")\n    coco = load_dataset(\"detection-datasets/coco\", split=\"train\", streaming=True)\n    for i, ex in enumerate(coco):\n        if i >= CFG[\"images_per_concept\"]: break\n        img = ex[\"image\"].convert(\"RGB\").resize((CFG[\"image_size\"], CFG[\"image_size\"]), Image.Resampling.LANCZOS)\n        datasets_dict[\"baseline\"].append(img)\n        \n    # --- 2. Fetch Kaggle Native ImageNet Concepts ---\n    def fetch_local_imagenet(wnid, max_count):\n        source_dir = Path(f\"/kaggle/input/competitions/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train/{wnid}\")\n        # Fallback path just in case the Kaggle mount point is slightly different\n        if not source_dir.exists():\n            source_dir = Path(f\"/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train/{wnid}\")\n            \n        if not source_dir.exists():\n            raise FileNotFoundError(f\"Missing Kaggle dataset: {source_dir}\")\n            \n        jpeg_files = list(source_dir.glob(\"*.JPEG\"))\n        imgs = []\n        for file_path in jpeg_files:\n            if len(imgs) >= max_count: break\n            try:\n                with Image.open(file_path) as img:\n                    img_resized = img.convert(\"RGB\").resize((CFG[\"image_size\"], CFG[\"image_size\"]), Image.Resampling.LANCZOS)\n                    imgs.append(img_resized)\n            except Exception:\n                continue\n        return imgs\n\n    print(\"[...] Fetching Zebra images (from ImageNet)...\")\n    datasets_dict[\"zebra\"] = fetch_local_imagenet(\"n02391049\", CFG[\"images_per_concept\"])\n    \n    print(\"[...] Fetching Fire Truck images (from ImageNet)...\")\n    datasets_dict[\"fire_truck\"] = fetch_local_imagenet(\"n03345487\", CFG[\"images_per_concept\"])\n    \n    # --- 3. Fetch Redness (CIFAR-100 Red Classes to bypass DDG rate limits) ---\n    print(\"[...] Fetching Redness images (from CIFAR-100 red classes: apple, rose, tulip)...\")\n    try:\n        cifar_stream = load_dataset(\"cifar100\", split=\"train\", streaming=True)\n        # Class IDs for overwhelmingly red objects\n        red_classes = [0, 70, 83, 84] # apple, rose, tulip, sweet_pepper\n        \n        for ex in cifar_stream:\n            if len(datasets_dict[\"redness\"]) >= CFG[\"images_per_concept\"]: break\n            if ex[\"fine_label\"] in red_classes:\n                try:\n                    img = ex[\"img\"]\n                    if not isinstance(img, Image.Image):\n                        img = Image.fromarray(np.array(img))\n                    img_resized = img.convert(\"RGB\").resize((CFG[\"image_size\"], CFG[\"image_size\"]), Image.Resampling.LANCZOS)\n                    datasets_dict[\"redness\"].append(img_resized)\n                except Exception:\n                    continue\n    except Exception as e:\n        print(f\"[FAIL] CIFAR-100 redness fetch failed: {e}\")\n            \n    for k, v in datasets_dict.items():\n        print(f\"  [OK] {k}: {len(v)} images loaded.\")\n        \n    return datasets_dict","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MODEL & SAE LOADING\ndef load_models():\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 2: Loading LLaVA and SAE\")\n    print(\"=\" * 70)\n    \n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    print(\"[...] Loading LLaVA vision tower (FP16)...\")\n    processor = AutoProcessor.from_pretrained(CFG[\"model_id\"], use_fast=False)\n    model = LlavaForConditionalGeneration.from_pretrained(\n        CFG[\"model_id\"], torch_dtype=torch.float16, low_cpu_mem_usage=True\n    ).to(device)\n    model.eval()\n    \n    print(f\"[...] Downloading SAE from {CFG['hf_repo_id']}...\")\n    sae_path = hf_hub_download(repo_id=CFG[\"hf_repo_id\"], filename=CFG[\"sae_filename\"])\n    \n    # Init and load weights\n    sae = SparseAutoencoder(input_dim=1024, dict_size=4096).to(device)\n    sae.load_state_dict(torch.load(sae_path, map_location=device))\n    sae.eval()\n    \n    print(\"[OK] Models loaded and ready.\")\n    return model, processor, sae, device","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# SELECTIVITY SCORING\ndef compute_selectivity(model, processor, sae, datasets_dict, device):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 3: Extracting Latents & Scoring Features\")\n    print(\"=\" * 70)\n    \n    # Store all sparse latents for plotting later\n    # Shape per concept: [N_images * 576, 4096]\n    all_latents = {}\n    \n    vision_tower = model.model.vision_tower\n    \n    with torch.no_grad():\n        for concept, imgs in datasets_dict.items():\n            print(f\"  Processing {concept}...\")\n            concept_latents = []\n            \n            for img in imgs:\n                # Preprocess and extract patch tokens\n                # FIX: Call processor.image_processor directly to avoid the text requirement\n                pixel_values = processor.image_processor(images=img, return_tensors=\"pt\")[\"pixel_values\"].to(device, dtype=torch.float16)\n                vision_outputs = vision_tower(pixel_values, output_hidden_states=True)\n                \n                # Layer 23, drop CLS -> [1, 576, 1024]\n                patch_tokens = vision_outputs.hidden_states[-2][:, 1:, :] \n                \n                # Push through SAE to get sparse features -> [576, 4096]\n                flat_tokens = patch_tokens.squeeze(0).float()\n                sparse_acts = sae.encode(flat_tokens) \n                \n                concept_latents.append(sparse_acts.cpu())\n            \n            # Combine all patches for this concept\n            all_latents[concept] = torch.cat(concept_latents, dim=0) \n            \n    # Calculate Mean Activations\n    baseline_mean = all_latents[\"baseline\"].mean(dim=0) # [4096]\n    \n    concept_features = {}\n    top_features_per_concept = {}\n    \n    for concept in [\"zebra\", \"fire_truck\", \"redness\"]:\n        concept_mean = all_latents[concept].mean(dim=0) # [4096]\n        \n        # Selectivity Score = Concept Mean - Baseline Mean\n        selectivity_scores = concept_mean - baseline_mean\n        \n        # Rank features\n        top_scores, top_indices = selectivity_scores.topk(20)\n        \n        concept_features[concept] = top_indices.tolist()\n        top_features_per_concept[concept] = top_indices.tolist()[:5] # Keep top 5 for plotting\n        \n        print(f\"\\n[OK] Top 5 Features for {concept.upper()}:\")\n        for rank, (idx, score) in enumerate(zip(top_indices[:5], top_scores[:5])):\n            print(f\"     Rank {rank+1}: Feature #{idx} (Score: {score:.4f})\")\n            \n    # Save to JSON\n    json_path = Path(CFG[\"output_dir\"]) / \"concept_features.json\"\n    with open(json_path, \"w\") as f:\n        json.dump(concept_features, f, indent=2)\n    print(f\"\\n[SAVE] Saved feature rankings to {json_path}\")\n    \n    return all_latents, top_features_per_concept","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# VISUAL INSPECTION (Patch Plotting)\n\ndef plot_max_activating_patches(datasets_dict, all_latents, top_features):\n    print(\"\\n\" + \"=\" * 70)\n    print(\"STEP 4: Visual Inspection (Plotting Max Patches)\")\n    print(\"=\" * 70)\n    \n    for concept, features in top_features.items():\n        concept_imgs = datasets_dict[concept]\n        latents = all_latents[concept] # [N_images * 576, 4096]\n        \n        for feat_idx in features:\n            # Find the top 16 highest activating patches for this specific feature\n            feat_activations = latents[:, feat_idx]\n            top_16_vals, top_16_indices = feat_activations.topk(16)\n            \n            fig, axes = plt.subplots(4, 4, figsize=(10, 10))\n            fig.suptitle(f\"Top 16 Patches for {concept.upper()} - Feature #{feat_idx}\", fontsize=16)\n            \n            for ax, flat_idx, act_val in zip(axes.flatten(), top_16_indices, top_16_vals):\n                flat_idx = flat_idx.item()\n                img_idx = flat_idx // (CFG[\"grid_size\"] ** 2)\n                patch_idx = flat_idx % (CFG[\"grid_size\"] ** 2)\n                \n                # Math to find pixel coordinates of the patch\n                row = patch_idx // CFG[\"grid_size\"]\n                col = patch_idx % CFG[\"grid_size\"]\n                y_start = row * CFG[\"patch_size\"]\n                x_start = col * CFG[\"patch_size\"]\n                \n                # We crop a slightly larger context window (42x42) around the 14x14 patch\n                context = 14 \n                img = concept_imgs[img_idx]\n                \n                left = max(0, x_start - context)\n                upper = max(0, y_start - context)\n                right = min(img.width, x_start + CFG[\"patch_size\"] + context)\n                lower = min(img.height, y_start + CFG[\"patch_size\"] + context)\n                \n                crop = img.crop((left, upper, right, lower))\n                ax.imshow(crop)\n                \n                # Draw a red bounding box exactly where the 14x14 patch is\n                rect_x = x_start - left\n                rect_y = y_start - upper\n                rect = patches.Rectangle((rect_x, rect_y), CFG[\"patch_size\"], CFG[\"patch_size\"], \n                                         linewidth=2, edgecolor='r', facecolor='none')\n                ax.add_patch(rect)\n                \n                ax.set_title(f\"Act: {act_val:.2f}\", fontsize=10)\n                ax.axis('off')\n                \n            plt.tight_layout()\n            save_path = Path(CFG[\"output_dir\"]) / f\"{concept}_feature_{feat_idx}.png\"\n            plt.savefig(save_path, dpi=150, bbox_inches='tight')\n            plt.close()\n            \n        print(f\"[SAVE] Saved 5 visualization grids for {concept}.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"datasets = load_evaluation_images()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model, processor, sae, device = load_models()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"latents, top_feats = compute_selectivity(model, processor, sae, datasets, device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del model, sae\ngc.collect()\ntorch.cuda.empty_cache()\n\nplot_max_activating_patches(datasets, latents, top_feats)\nprint(\"\\nTask 01 Complete!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}