{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":91498,"databundleVersionId":11655853,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **DINOv2 Bike Image Similarity Search**","metadata":{}},{"cell_type":"markdown","source":"https://huggingface.co/facebook/dinov2-base","metadata":{}},{"cell_type":"code","source":"!pip install protobuf==3.20.3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-13T09:02:03.908792Z","iopub.execute_input":"2025-11-13T09:02:03.908993Z","iopub.status.idle":"2025-11-13T09:02:09.077188Z","shell.execute_reply.started":"2025-11-13T09:02:03.908975Z","shell.execute_reply":"2025-11-13T09:02:09.076209Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom transformers import AutoImageProcessor, AutoModel\nimport torch\nfrom pathlib import Path\nfrom sklearn.metrics.pairwise import cosine_similarity\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"=\"*80)\nprint(\"          DINOv2 IMAGE SIMILARITY SEARCH - QUERY MODE\")\nprint(\"=\"*80)\nprint(\"\\nPIPELINE:\")\nprint(\"1. Load query image and extract DINOv2 features\")\nprint(\"2. Load all images from folder and extract features\")\nprint(\"3. Calculate cosine similarity between query and all images\")\nprint(\"4. Display Top 5 most similar images\")\nprint(\"=\"*80)\nprint()\n\n\nclass DINOv2SimilaritySearch:\n    def __init__(self, model_name='facebook/dinov2-base'):\n        \"\"\"Initialize DINOv2 model for feature extraction\"\"\"\n        print(f\"Loading DINOv2 model: {model_name}...\")\n        self.processor = AutoImageProcessor.from_pretrained(model_name, use_fast=True)\n        self.model = AutoModel.from_pretrained(model_name)\n        self.model.eval()\n        print(\"✓ Model loaded successfully\\n\")\n        \n    def extract_feature(self, image_path):\n        \"\"\"\n        Extract 768-dimensional feature vector from image using [CLS] token\n        \"\"\"\n        try:\n            image = Image.open(image_path).convert('RGB')\n            inputs = self.processor(images=image, return_tensors=\"pt\")\n            \n            with torch.no_grad():\n                outputs = self.model(**inputs)\n            \n            # Extract [CLS] token (global image representation)\n            cls_token = outputs.last_hidden_state[0, 0, :].numpy()\n            return cls_token, image\n        except Exception as e:\n            print(f\"✗ Error processing {image_path}: {e}\")\n            return None, None\n    \n    def search_similar_images(self, query_image_path, folder_path, top_k=5):\n        \"\"\"\n        Find top-k images most similar to query image\n        \n        Args:\n            query_image_path: Path to query image\n            folder_path: Folder containing database images\n            top_k: Number of top similar images to return\n        \"\"\"\n        print(\"=\"*80)\n        print(f\"QUERY IMAGE: {Path(query_image_path).name}\")\n        print(\"=\"*80)\n        \n        # Extract query features\n        print(\"\\n[1/4] Extracting query image features...\")\n        query_feature, query_image = self.extract_feature(query_image_path)\n        if query_feature is None:\n            print(\"Error: Could not process query image\")\n            return\n        print(f\"✓ Query feature shape: {query_feature.shape}\")\n        \n        # Load database images\n        print(f\"\\n[2/4] Loading images from: {folder_path}\")\n        folder = Path(folder_path)\n        image_extensions = {'.jpg', '.jpeg', '.png', '.bmp', '.gif', '.webp'}\n        \n        image_paths = []\n        for ext in image_extensions:\n            image_paths.extend(folder.glob(f'*{ext}'))\n            image_paths.extend(folder.glob(f'*{ext.upper()}'))\n        \n        image_paths = sorted(set(image_paths))\n        \n        # Remove query image from database if it exists in the folder\n        query_path = Path(query_image_path)\n        image_paths = [p for p in image_paths if p.resolve() != query_path.resolve()]\n        \n        if len(image_paths) == 0:\n            print(\"Error: No images found in folder\")\n            return\n        \n        print(f\"✓ Found {len(image_paths)} images in database\")\n        \n        # Extract features from all database images\n        print(\"\\n[3/4] Extracting features from database images...\")\n        db_features = []\n        db_images = []\n        valid_paths = []\n        \n        for i, img_path in enumerate(image_paths):\n            print(f\"  Processing [{i+1}/{len(image_paths)}]: {img_path.name}\", end='\\r')\n            feature, image = self.extract_feature(img_path)\n            \n            if feature is not None:\n                db_features.append(feature)\n                db_images.append(image)\n                valid_paths.append(img_path)\n        \n        print(f\"\\n✓ Successfully processed {len(db_features)} database images\")\n        \n        if len(db_features) == 0:\n            print(\"Error: No valid images found in database\")\n            return\n        \n        db_features = np.array(db_features)\n        \n        # Calculate similarities\n        print(\"\\n[4/4] Calculating cosine similarities...\")\n        query_feature = query_feature.reshape(1, -1)\n        similarities = cosine_similarity(query_feature, db_features)[0]\n        \n        # Get top-k most similar\n        top_k = min(top_k, len(similarities))\n        top_indices = np.argsort(similarities)[-top_k:][::-1]\n        \n        print(f\"✓ Found top {top_k} similar images\\n\")\n        \n        # Display results\n        self._visualize_results(query_image, query_path, \n                               db_images, valid_paths, \n                               similarities, top_indices, top_k)\n        \n        # Print detailed results\n        self._print_results(valid_paths, similarities, top_indices)\n        \n        return top_indices, similarities[top_indices], valid_paths\n    \n    def _visualize_results(self, query_image, query_path, \n                          db_images, db_paths, similarities, top_indices, top_k):\n        \"\"\"Visualize query image and top-k results\"\"\"\n        \n        # Create figure with query + top-k results\n        fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n        fig.suptitle('DINOv2 Image Similarity Search - Top 5 Results', \n                     fontsize=18, fontweight='bold', y=0.98)\n        \n        # Flatten axes for easier indexing\n        axes = axes.flatten()\n        \n        # Display query image in first position\n        axes[0].imshow(query_image)\n        axes[0].set_title('QUERY IMAGE\\n' + query_path.name, \n                         fontsize=13, fontweight='bold', \n                         color='white', backgroundcolor='red', pad=10)\n        axes[0].axis('off')\n        axes[0].set_facecolor('#f0f0f0')\n        \n        # Display top-k results\n        for i, idx in enumerate(top_indices):\n            ax = axes[i+1]\n            ax.imshow(db_images[idx])\n            \n            # Color code by similarity\n            sim = similarities[idx]\n            if sim >= 0.9:\n                color = 'green'\n                label = 'Very Similar'\n            elif sim >= 0.8:\n                color = 'yellowgreen'\n                label = 'Similar'\n            elif sim >= 0.7:\n                color = 'orange'\n                label = 'Moderately Similar'\n            else:\n                color = 'orangered'\n                label = 'Less Similar'\n            \n            title = f'RANK #{i+1} - {label}\\n{db_paths[idx].name}\\nSimilarity: {sim:.4f}'\n            ax.set_title(title, fontsize=11, fontweight='bold', \n                        color='white', backgroundcolor=color, pad=10)\n            ax.axis('off')\n            ax.set_facecolor('#f0f0f0')\n        \n        plt.tight_layout()\n        \n        # Save result\n        output_path = 'similarity_search_results.png'\n        plt.savefig(output_path, dpi=150, bbox_inches='tight', facecolor='white')\n        print(f\"✓ Results saved to: {output_path}\")\n        plt.show()\n    \n    def _print_results(self, db_paths, similarities, top_indices):\n        \"\"\"Print detailed similarity scores\"\"\"\n        print(\"\\n\" + \"=\"*80)\n        print(\"                         SIMILARITY RANKING\")\n        print(\"=\"*80)\n        print(f\"{'Rank':<6} {'Similarity':<12} {'Image Name':<50}\")\n        print(\"-\"*80)\n        \n        for i, idx in enumerate(top_indices):\n            sim = similarities[idx]\n            name = db_paths[idx].name\n            \n            # Add visual indicator\n            if sim >= 0.9:\n                indicator = \"★★★★★\"\n            elif sim >= 0.8:\n                indicator = \"★★★★☆\"\n            elif sim >= 0.7:\n                indicator = \"★★★☆☆\"\n            elif sim >= 0.6:\n                indicator = \"★★☆☆☆\"\n            else:\n                indicator = \"★☆☆☆☆\"\n            \n            print(f\"#{i+1:<5} {sim:<12.6f} {name:<50} {indicator}\")\n        \n        print(\"=\"*80)\n        print(f\"\\nSimilarity Statistics:\")\n        print(f\"  • Highest similarity: {similarities[top_indices[0]]:.6f}\")\n        print(f\"  • Lowest similarity:  {similarities[top_indices[-1]]:.6f}\")\n        print(f\"  • Average similarity: {np.mean(similarities[top_indices]):.6f}\")\n        print(f\"  • Total database size: {len(similarities)} images\")\n        print(\"=\"*80)\n\n\n# ==================== MAIN USAGE ====================\n\ndef main():\n    \"\"\"\n    Main function to run similarity search\n    \"\"\"\n    # ========== CONFIGURATION ==========\n    # Set your paths here\n    DATABASE_FOLDER = \"/kaggle/input/image-matching-challenge-2025/train/imc2023_haiper\"   \n    QUERY_IMAGE = F\"{DATABASE_FOLDER}/bike_image_004.png\"        \n   \n    TOP_K = 5                         # Number of similar images to find\n    # ===================================\n    \n    # Check if paths exist\n    if not os.path.exists(QUERY_IMAGE):\n        print(f\"❌ Error: Query image not found: {QUERY_IMAGE}\")\n        print(\"\\nPlease update QUERY_IMAGE path in the code.\")\n        print(\"Example: QUERY_IMAGE = '/path/to/your/query.jpg'\")\n        return\n    \n    if not os.path.exists(DATABASE_FOLDER):\n        print(f\"❌ Error: Database folder not found: {DATABASE_FOLDER}\")\n        print(\"\\nPlease update DATABASE_FOLDER path in the code.\")\n        print(\"Example: DATABASE_FOLDER = '/path/to/your/images'\")\n        return\n    \n    # Run similarity search\n    searcher = DINOv2SimilaritySearch()\n    results = searcher.search_similar_images(\n        query_image_path=QUERY_IMAGE,\n        folder_path=DATABASE_FOLDER,\n        top_k=TOP_K\n    )\n    \n    if results:\n        print(\"\\n✓ Search completed successfully!\")\n\n\n# ==================== ALTERNATIVE: DIRECT FUNCTION ====================\n\ndef find_similar_images(query_image_path, folder_path, top_k=5):\n    \"\"\"\n    Quick function to find similar images\n    \n    Usage:\n        find_similar_images(\"my_cat.jpg\", \"./animal_images\", top_k=5)\n    \"\"\"\n    searcher = DINOv2SimilaritySearch()\n    return searcher.search_similar_images(query_image_path, folder_path, top_k)\n\n\nif __name__ == \"__main__\":\n    main()\n    \n    # Or use the quick function:\n    # find_similar_images(\"query.jpg\", \"./images\", top_k=5)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}