{"metadata":{"kernelspec":{"display_name":"Python (tf-gpu)","language":"python","name":"tf-gpu"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.16"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91498,"databundleVersionId":11655853,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"44a8aa2d-97c0-4604-be36-c2b2fa853206","cell_type":"markdown","source":"# Image Matching Challenge 2025\n### Phase 1: Image Loading & Preprocessing\nThis phase loads training and test images into normalized grayscale tensors.\nIt supports PyTorch GPU acceleration (CUDA/MPS) and prepares images for feature extraction.","metadata":{}},{"id":"51da82ed-e73d-41e6-bc86-1db0eb838957","cell_type":"code","source":"import torch\n\ndef get_device():\n    if torch.cuda.is_available():\n        return torch.device(\"cuda\")\n    elif torch.backends.mps.is_available():\n        return torch.device(\"mps\")\n    else:\n        return torch.device(\"cpu\")\n\ndevice = get_device()\nprint(f\" Using device: {device}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"3b966185-b043-49ed-92bb-0521cb30bbfb","cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nfrom collections import defaultdict\n\n# Dataset config\nDATA_ROOT = \"./\"\nIMAGE_EXT = \".png\"\n","metadata":{},"outputs":[],"execution_count":null},{"id":"6e7200b0-8a6a-43d8-9816-f0e569a7ae11","cell_type":"code","source":"def load_images(root_dir=\"train\", device=torch.device(\"cpu\")):\n    image_data = {}\n    folders = sorted([os.path.join(root_dir, d) for d in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, d))])\n\n    for folder in tqdm(folders, desc=f\"📂 Loading {root_dir}\"):\n        image_files = sorted(glob(os.path.join(folder, f\"*{IMAGE_EXT}\")))\n        for img_path in image_files:\n            img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            if img is not None:\n                img_tensor = torch.tensor(img, dtype=torch.float32, device=device) / 255.0\n                rel_path = os.path.relpath(img_path, start=DATA_ROOT)\n                image_data[rel_path] = img_tensor\n            else:\n                print(f\"⚠️ Could not load: {img_path}\")\n    return image_data\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c0b289f1-9762-4e6b-95f7-ed510c5df29d","cell_type":"code","source":"def show_sample_images(image_dict, title, num_samples=9):\n    keys = list(image_dict.keys())\n    if len(keys) < num_samples:\n        print(f\"Only {len(keys)} images found.\")\n        return\n\n    sample_keys = np.random.choice(keys, num_samples, replace=False)\n    fig, axs = plt.subplots(3, 3, figsize=(10, 10))\n    fig.suptitle(title, fontsize=14)\n\n    for i, key in enumerate(sample_keys):\n        ax = axs[i // 3, i % 3]\n        img = image_dict[key].cpu().numpy()\n        ax.imshow(img, cmap='gray')\n        ax.set_title(os.path.basename(key), fontsize=8)\n        ax.axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"19888b92-e8ef-4485-a85b-a915dbeef644","cell_type":"code","source":"train_images = load_images(\"train\", device=device)\ntest_images = load_images(\"test\", device=device)\n\nprint(f\"✅ Loaded {len(train_images)} training images\")\nprint(f\"✅ Loaded {len(test_images)} testing images\")\n\n# Preview sample scenes\nshow_sample_images({k: v for k, v in train_images.items() if \"amy_gardens\" in k}, \"Train Sample: amy_gardens\")\nshow_sample_images({k: v for k, v in test_images.items() if \"ETs\" in k}, \"Test Sample: ETs\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c941c946-8b33-40f7-bac0-85eb4f0107ca","cell_type":"markdown","source":"## Phase 2: Harris Corner Detection (from scratch)\nThis phase computes Harris corners per image using Sobel gradients, structure tensors, and non-maximum suppression. Keypoints are sorted by response score, and the top N are cached per image.\n","metadata":{}},{"id":"e4264508-6a59-4af2-80a8-4d3fa46b4d93","cell_type":"code","source":"import torch.nn.functional as F\nimport os\n\nN_KEYPOINTS = 500\nCACHE_DIR = \"./cache/keypoints\"\nos.makedirs(CACHE_DIR, exist_ok=True)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"4fd8372b-249d-456d-b6eb-f34f8c9c86f6","cell_type":"code","source":"def harris_corners(img, device, k=0.04, window_size=3):\n    assert img.ndim == 2  # grayscale\n    img = img.unsqueeze(0).unsqueeze(0)  # shape: (1, 1, H, W)\n\n    # Sobel kernels\n    sobel_x = torch.tensor([[1, 0, -1],\n                            [2, 0, -2],\n                            [1, 0, -1]], dtype=torch.float32, device=device).view(1, 1, 3, 3)\n    sobel_y = torch.tensor([[1, 2, 1],\n                            [0, 0, 0],\n                            [-1, -2, -1]], dtype=torch.float32, device=device).view(1, 1, 3, 3)\n\n    Ix = F.conv2d(img, sobel_x, padding=1)[0, 0]\n    Iy = F.conv2d(img, sobel_y, padding=1)[0, 0]\n\n    Ixx = Ix * Ix\n    Iyy = Iy * Iy\n    Ixy = Ix * Iy\n\n    # Window smoothing\n    window = torch.ones((1, 1, window_size, window_size), device=device) / (window_size ** 2)\n    Sxx = F.conv2d(Ixx.unsqueeze(0).unsqueeze(0), window, padding=window_size//2)[0, 0]\n    Syy = F.conv2d(Iyy.unsqueeze(0).unsqueeze(0), window, padding=window_size//2)[0, 0]\n    Sxy = F.conv2d(Ixy.unsqueeze(0).unsqueeze(0), window, padding=window_size//2)[0, 0]\n\n    detM = Sxx * Syy - Sxy**2\n    traceM = Sxx + Syy\n    R = detM - k * (traceM ** 2)\n    return R\n","metadata":{},"outputs":[],"execution_count":null},{"id":"ca4622b9-be63-460c-84a1-07f021018c07","cell_type":"code","source":"def extract_keypoints(image_dict, dataset_name, cache_dir=CACHE_DIR, top_k=N_KEYPOINTS):\n    keypoints = {}\n\n    dataset_cache_dir = os.path.join(cache_dir, dataset_name)\n    os.makedirs(dataset_cache_dir, exist_ok=True)\n\n    for path, img in tqdm(image_dict.items(), desc=f\"🧠 Detecting Harris corners: {dataset_name}\"):\n        R = harris_corners(img, device)\n        R_np = R.detach().cpu().numpy()\n\n        # Flatten and sort by corner strength\n        flat_indices = np.argpartition(R_np.ravel(), -top_k)[-top_k:]\n        y, x = np.unravel_index(flat_indices, R_np.shape)\n        scores = R_np[y, x]\n        sorted_idx = np.argsort(-scores)\n\n        keypoints_xy = np.stack([x[sorted_idx], y[sorted_idx]], axis=-1)\n\n        # Cache as .pt file\n        cache_path = os.path.join(dataset_cache_dir, os.path.basename(path).replace('.png', '.pt'))\n        torch.save(torch.tensor(keypoints_xy, dtype=torch.int32), cache_path)\n        keypoints[path] = keypoints_xy\n\n    return keypoints\n","metadata":{},"outputs":[],"execution_count":null},{"id":"2520c3a4-80e9-485c-bed3-4f96ef0352f5","cell_type":"code","source":"def show_keypoints(img_tensor, keypoints, title=\"Keypoints\", max_k=100):\n    import matplotlib.pyplot as plt\n\n    img = img_tensor.detach().cpu().numpy()\n    pts = keypoints[:max_k]  # (x, y)\n\n    plt.figure(figsize=(6, 6))\n    plt.imshow(img, cmap='gray')\n    plt.scatter(pts[:, 0], pts[:, 1], c='red', s=10)\n    plt.title(title)\n    plt.axis('off')\n    plt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"8dbce3fc-0223-4380-850f-4f7c1caf900d","cell_type":"code","source":"# Run detection on one folder for now\nscene_name = \"ETs\"\ntest_subset = {k: v for k, v in test_images.items() if scene_name in k}\nkeypoints_test = extract_keypoints(test_subset, dataset_name=scene_name)\n\n# Preview one image's keypoints\nfirst_img = list(test_subset.keys())[0]\nkp = keypoints_test[first_img]\nshow_keypoints(test_subset[first_img], kp, title=f\"Keypoints: {os.path.basename(first_img)}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b5654d07-d699-4417-98ec-8f3d0d8bd27a","cell_type":"markdown","source":"## Phase 3: Descriptor Extraction + Keypoint Matching\nWe extract local grayscale patches around Harris keypoints and compare descriptors using cosine similarity with mutual match filtering.\n","metadata":{}},{"id":"2d8a2605-2ea8-425c-ac8f-12cf54d91e4c","cell_type":"code","source":"def extract_descriptors(image_dict, keypoints_dict, patch_size=15):\n    descriptors = {}\n\n    half = patch_size // 2\n\n    for path, keypoints in tqdm(keypoints_dict.items(), desc=\" Extracting descriptors\"):\n        img = image_dict[path]\n        H, W = img.shape\n        padded_img = F.pad(img.unsqueeze(0).unsqueeze(0), (half, half, half, half), mode='reflect')[0, 0]\n\n        desc_list = []\n        for x, y in keypoints:\n            x, y = int(x.item()), int(y.item())\n            patch = padded_img[y:y+patch_size, x:x+patch_size]\n            if patch.shape == (patch_size, patch_size):\n                desc_list.append(patch.flatten())\n\n        if desc_list:\n            descriptors[path] = torch.stack(desc_list)\n        else:\n            descriptors[path] = torch.empty((0, patch_size * patch_size))\n\n    return descriptors\n","metadata":{},"outputs":[],"execution_count":null},{"id":"cbaf5dcd-47e3-4348-b5aa-cb8f5782e245","cell_type":"code","source":"def match_descriptors(desc1, desc2, top_k=1, ratio_thresh=0.75):\n    if desc1.size(0) == 0 or desc2.size(0) == 0:\n        return []\n\n    # Normalize\n    d1 = F.normalize(desc1, p=2, dim=1)\n    d2 = F.normalize(desc2, p=2, dim=1)\n\n    # Cosine similarity\n    sim = torch.matmul(d1, d2.T)  # shape: (N1, N2)\n    scores, indices = sim.topk(k=top_k, dim=1)\n\n    matches = []\n    for i, (score, idx) in enumerate(zip(scores, indices)):\n        if top_k == 1 or score[0] > score[1] * ratio_thresh:\n            matches.append((i, idx[0].item(), score[0].item()))\n\n    return matches\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f658dd00-8259-4ec5-885b-ecde50c89811","cell_type":"code","source":"def show_matches(img1, img2, kp1, kp2, matches, max_lines=50):\n    import matplotlib.pyplot as plt\n\n    img1 = img1.cpu().numpy()\n    img2 = img2.cpu().numpy()\n    h1, w1 = img1.shape\n    h2, w2 = img2.shape\n\n    canvas = np.zeros((max(h1, h2), w1 + w2), dtype=np.float32)\n    canvas[:h1, :w1] = img1\n    canvas[:h2, w1:] = img2\n\n    plt.figure(figsize=(10, 6))\n    plt.imshow(canvas, cmap='gray')\n\n    for i, j, _ in matches[:max_lines]:\n        pt1 = kp1[i]\n        pt2 = kp2[j]\n        x1, y1 = pt1[0], pt1[1]\n        x2, y2 = pt2[0] + w1, pt2[1]\n        plt.plot([x1, x2], [y1, y2], 'r-', linewidth=1)\n        plt.scatter([x1, x2], [y1, y2], c='lime', s=10)\n\n    plt.axis(\"off\")\n    plt.title(f\"{len(matches)} matches\")\n    plt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c40028d9-5d82-4e60-984e-9b13803b4a7b","cell_type":"code","source":"# Extract descriptors for ETs subset\ndesc_ETs = extract_descriptors(test_subset, keypoints_test)\n\n# Pick two ET images\npaths = list(desc_ETs.keys())[:2]\ndesc1 = desc_ETs[paths[0]]\ndesc2 = desc_ETs[paths[1]]\n\nmatches = match_descriptors(desc1, desc2)\nprint(f\"Found {len(matches)} matches\")\n\n# Visualize\nshow_matches(\n    test_subset[paths[0]], test_subset[paths[1]],\n    keypoints_test[paths[0]], keypoints_test[paths[1]],\n    matches\n)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f2d95127-1748-4466-9d89-d95ae5a86104","cell_type":"markdown","source":"## Phase 4.0: Scene Graph Construction\nWe build a graph where:\n- Each node = image\n- Each edge = strong match connection\n- We cluster connected components to discover sub-scenes\nThis provides a clean base for novel edge weighting and pose-aware refinement later.\n","metadata":{}},{"id":"3b9c647a-93fc-4160-beed-617c24ad934a","cell_type":"code","source":"import networkx as nx\nfrom itertools import combinations\n\ndef build_match_graph(image_dict, kp_dict, desc_dict, match_threshold=20):\n    print(f\" Building image match graph from {len(image_dict)} images...\")\n    G = nx.Graph()\n    paths = list(image_dict.keys())\n\n    for img1, img2 in tqdm(combinations(paths, 2), total=len(paths)*(len(paths)-1)//2):\n        desc1 = desc_dict[img1]\n        desc2 = desc_dict[img2]\n\n        matches = match_descriptors(desc1, desc2)\n        if len(matches) >= match_threshold:\n            G.add_edge(img1, img2, weight=len(matches))\n\n    return G\n","metadata":{},"outputs":[],"execution_count":null},{"id":"d0e85ff9-db0c-4913-930d-e15ac4a9b345","cell_type":"code","source":"def cluster_graph(graph):\n    clusters = list(nx.connected_components(graph))\n    cluster_map = {}\n\n    for i, group in enumerate(clusters):\n        for img_path in group:\n            cluster_map[img_path] = f\"cluster_{i+1}\"\n\n    return cluster_map\n","metadata":{},"outputs":[],"execution_count":null},{"id":"2764fc7c-d83e-40e1-82e3-94b64c8c6306","cell_type":"code","source":"def visualize_graph(graph, cluster_map=None):\n    import matplotlib.pyplot as plt\n\n    pos = nx.spring_layout(graph, seed=42)\n    plt.figure(figsize=(10, 7))\n    if cluster_map:\n        colors = [hash(cluster_map[node]) % 20 for node in graph.nodes()]\n    else:\n        colors = 'lightblue'\n\n    nx.draw_networkx_nodes(graph, pos, node_size=300, node_color=colors, cmap='tab20')\n    nx.draw_networkx_edges(graph, pos, alpha=0.5)\n    nx.draw_networkx_labels(graph, pos, font_size=7)\n    plt.axis('off')\n    plt.title(\"Scene Match Graph\")\n    plt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"9e9696d3-99f3-44d5-a9bd-2035a5d7c805","cell_type":"code","source":"scene_name = \"ETs\"\ntest_subset = {k: v for k, v in test_images.items() if scene_name in k}\ndesc_ETs = extract_descriptors(test_subset, keypoints_test)\n\nG = build_match_graph(test_subset, keypoints_test, desc_ETs, match_threshold=20)\ncluster_map = cluster_graph(G)\n\nvisualize_graph(G, cluster_map)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"394c5596-626e-4945-8f48-7e1a0d71f47d","cell_type":"markdown","source":"## Phase 4.5: Novel Graph Refinement (Average Match Quality)\nWe enhance the match graph by weighting edges using the average similarity score from matched descriptor pairs.\nThis improves graph clustering by favoring high-quality visual links over noisy ones.\n","metadata":{}},{"id":"16838392-d4c4-4976-8cd9-ff2112278e72","cell_type":"code","source":"def build_quality_weighted_graph(image_dict, keypoints, descriptors, match_threshold=20, min_avg_sim=0.5):\n    print(f\" Building enhanced match graph using average descriptor similarity...\")\n    G = nx.Graph()\n    paths = list(image_dict.keys())\n\n    for img1, img2 in tqdm(combinations(paths, 2), total=len(paths)*(len(paths)-1)//2):\n        desc1 = descriptors[img1]\n        desc2 = descriptors[img2]\n\n        matches = match_descriptors(desc1, desc2)\n        if len(matches) < match_threshold:\n            continue\n\n        # Compute mean cosine similarity from valid matches\n        match_sims = []\n        for i, j, score in matches:\n            match_sims.append(score)\n        mean_sim = np.mean(match_sims)\n\n        if mean_sim >= min_avg_sim:\n            G.add_edge(img1, img2, weight=mean_sim)\n\n    return G\n","metadata":{},"outputs":[],"execution_count":null},{"id":"b6941db6-98e2-4660-9c66-b7d20fd3135a","cell_type":"code","source":"def visualize_weighted_graph(graph, cluster_map=None):\n    import matplotlib.pyplot as plt\n\n    pos = nx.spring_layout(graph, seed=42, weight='weight')\n    plt.figure(figsize=(10, 7))\n    if cluster_map:\n        colors = [hash(cluster_map[node]) % 20 for node in graph.nodes()]\n    else:\n        colors = 'lightblue'\n\n    edges = graph.edges(data=True)\n    weights = [d['weight'] for (_, _, d) in edges]\n\n    nx.draw_networkx_nodes(graph, pos, node_size=300, node_color=colors, cmap='tab20')\n    nx.draw_networkx_edges(graph, pos, width=[w * 2 for w in weights], alpha=0.5)\n    nx.draw_networkx_labels(graph, pos, font_size=7)\n    plt.axis('off')\n    plt.title(\"Enhanced Match Graph (Weighted by Similarity)\")\n    plt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"6a6567a8-0eb8-4f87-b9a5-6ebe06c1a750","cell_type":"code","source":"G_quality = build_quality_weighted_graph(\n    test_subset,\n    keypoints_test,\n    desc_ETs,\n    match_threshold=20,\n    min_avg_sim=0.5\n)\n\ncluster_map_quality = cluster_graph(G_quality)\nvisualize_weighted_graph(G_quality, cluster_map_quality)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"de830052-8718-452b-95af-735645175ab6","cell_type":"markdown","source":"## Phase 4.6: Cluster Inspection & Summary\nThis step visually and numerically inspects the clusters created from the refined match graph.\nWe verify whether images were grouped correctly before moving to pose estimation.\n","metadata":{}},{"id":"198dbd18-41d2-4d94-b049-a0f9f4097a53","cell_type":"code","source":"def summarize_clusters(cluster_map):\n    from collections import defaultdict\n\n    cluster_counts = defaultdict(list)\n    for path, cluster_id in cluster_map.items():\n        cluster_counts[cluster_id].append(path)\n\n    print(\" Cluster Summary:\")\n    for cid, images in sorted(cluster_counts.items(), key=lambda x: -len(x[1])):\n        print(f\"{cid}: {len(images)} images\")\n    \n    return cluster_counts\n","metadata":{},"outputs":[],"execution_count":null},{"id":"0c67fba3-4de9-4eea-a261-dcc494f3fda7","cell_type":"code","source":"def show_cluster_images(cluster_id, cluster_dict, image_dict, max_images=9):\n    import matplotlib.pyplot as plt\n    images = cluster_dict[cluster_id][:max_images]\n    \n    n = len(images)\n    cols = min(3, n)\n    rows = (n + cols - 1) // cols\n    fig, axs = plt.subplots(rows, cols, figsize=(4*cols, 4*rows))\n\n    if rows == 1: axs = [axs]\n    axs = np.array(axs).flatten()\n\n    for i, path in enumerate(images):\n        img = image_dict[path].cpu().numpy()\n        axs[i].imshow(img, cmap='gray')\n        axs[i].set_title(os.path.basename(path), fontsize=8)\n        axs[i].axis('off')\n\n    for j in range(i+1, len(axs)):\n        axs[j].axis('off')\n\n    fig.suptitle(f\" Cluster: {cluster_id}\", fontsize=14)\n    plt.tight_layout()\n    plt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c53fc8d0-2485-47d1-82c0-19ca6a8f8d7b","cell_type":"code","source":"# Get cluster map from earlier\ncluster_dict = summarize_clusters(cluster_map_quality)\n\n# Visualize each cluster\nfor cid in cluster_dict:\n    show_cluster_images(cid, cluster_dict, test_subset)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"0289635a-f3bb-498e-8969-0b3280d4c033","cell_type":"markdown","source":"## Phase 5: Camera Pose Estimation (Essential Matrix)\nWe use matched keypoints across image pairs to estimate relative camera poses.\nFrom these, we register each image into a global scene frame.\n","metadata":{}},{"id":"3d8484f9-e96d-48af-a912-bcf461fa78bf","cell_type":"code","source":"import cv2\nimport numpy as np\n\ndef estimate_relative_pose(kp1, kp2, matches, img_shape, focal_length=1.0):\n    pts1 = np.float32([kp1[i] for i, j, _ in matches])\n    pts2 = np.float32([kp2[j] for i, j, _ in matches])\n    \n    # Camera intrinsics: assume focal length = 1, centered\n    h, w = img_shape\n    K = np.array([[focal_length, 0, w/2],\n                  [0, focal_length, h/2],\n                  [0, 0, 1]], dtype=np.float64)\n\n    E, mask = cv2.findEssentialMat(pts1, pts2, K, method=cv2.RANSAC, prob=0.999, threshold=1.0)\n    if E is None or mask is None:\n        return None, None, None\n\n    _, R, t, inliers = cv2.recoverPose(E, pts1, pts2, K)\n    return R, t, mask.ravel()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"483b63c3-54d3-440a-877f-e11824764e9d","cell_type":"code","source":"def initialize_pose_graph(cluster_paths, kp_dict, desc_dict, img_tensor_dict):\n    pose_graph = {cluster_paths[0]: (np.eye(3), np.zeros((3, 1)))}\n    visited = set([cluster_paths[0]])\n    to_visit = [cluster_paths[0]]\n\n    while to_visit:\n        current = to_visit.pop()\n        R_current, t_current = pose_graph[current]\n        desc1 = desc_dict[current]\n        kp1 = kp_dict[current]\n        shape1 = img_tensor_dict[current].shape\n\n        for neighbor in cluster_paths:\n            if neighbor in visited or neighbor == current:\n                continue\n\n            desc2 = desc_dict[neighbor]\n            kp2 = kp_dict[neighbor]\n            shape2 = img_tensor_dict[neighbor].shape\n\n            matches = match_descriptors(desc1, desc2)\n            if len(matches) < 20:\n                continue\n\n            R_rel, t_rel, inlier_mask = estimate_relative_pose(kp1, kp2, matches, shape1)\n            if R_rel is None:\n                continue\n\n            # Chain transformations\n            R_global = R_rel @ R_current\n            t_global = R_rel @ t_current + t_rel\n\n            pose_graph[neighbor] = (R_global, t_global)\n            visited.add(neighbor)\n            to_visit.append(neighbor)\n\n    return pose_graph\n","metadata":{},"outputs":[],"execution_count":null},{"id":"8295cbfd-7b37-4665-ba6d-0c831df746d8","cell_type":"code","source":"def format_pose_output(pose_graph, dataset_name, cluster_id):\n    rows = []\n    for path, (R, t) in pose_graph.items():\n        R_flat = \";\".join([f\"{v:.6f}\" for v in R.flatten()])\n        t_flat = \";\".join([f\"{v:.6f}\" for v in t.flatten()])\n        rows.append({\n            \"dataset\": dataset_name,\n            \"scene\": cluster_id,\n            \"image\": os.path.basename(path),\n            \"rotation_matrix\": R_flat,\n            \"translation_vector\": t_flat\n        })\n    return rows\n","metadata":{},"outputs":[],"execution_count":null},{"id":"ac2d97a4-f917-4708-854e-1d1dfa3c01cf","cell_type":"code","source":"cluster_id = list(cluster_dict.keys())[0]\ncluster_paths = cluster_dict[cluster_id]\n\npose_graph = initialize_pose_graph(cluster_paths, keypoints_test, desc_ETs, test_subset)\n\npose_rows = format_pose_output(pose_graph, dataset_name=\"testETs\", cluster_id=cluster_id)\nimport pandas as pd\npose_df = pd.DataFrame(pose_rows)\npose_df.head()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"1ece0664-ea13-4a09-b744-da28ce865ed6","cell_type":"markdown","source":"## Phase 7: Pose-Consistency-Aware Graph Pruning\n\nEven visually similar images can result in bad 3D reconstructions if their underlying geometry is incompatible.\nIn this phase, we filter the image match graph using a geometric verification step:\n\n- For every edge (image pair), we estimate an **Essential Matrix** and count **RANSAC inliers**\n- Edges with insufficient geometric support (e.g., < 20 inliers) are pruned\n- The result is a cleaner match graph with higher-quality clusters and more reliable pose chains\n\nThis improves both:\n- Mean Average Accuracy (mAA) — by avoiding pose failures\n- Clustering Score — by preventing false-positive scene merges\n","metadata":{}},{"id":"91170def-448b-4315-9ccd-de98d5986a26","cell_type":"code","source":"def pose_consistency_filter(image_dict, kp_dict, desc_dict, edge_list, threshold_inliers=20):\n    pruned_edges = []\n    failed_edges = []\n\n    for img1, img2 in tqdm(edge_list, desc=\" Checking pose consistency\"):\n        desc1 = desc_dict[img1]\n        desc2 = desc_dict[img2]\n        kp1 = kp_dict[img1]\n        kp2 = kp_dict[img2]\n        shape = image_dict[img1].shape\n\n        matches = match_descriptors(desc1, desc2)\n        if len(matches) < threshold_inliers:\n            failed_edges.append((img1, img2, 0))\n            continue\n\n        R, t, inlier_mask = estimate_relative_pose(kp1, kp2, matches, shape)\n        if R is None or inlier_mask is None or np.sum(inlier_mask) < threshold_inliers:\n            failed_edges.append((img1, img2, int(np.sum(inlier_mask)) if inlier_mask is not None else 0))\n        else:\n            pruned_edges.append((img1, img2, float(np.mean(inlier_mask))))\n\n    return pruned_edges, failed_edges\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f73dabe1-8949-4952-979b-a929838396a2","cell_type":"code","source":"def build_pose_pruned_graph(valid_edges):\n    G = nx.Graph()\n    for img1, img2, weight in valid_edges:\n        G.add_edge(img1, img2, weight=weight)\n    return G\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f46cfc6a-3809-4e33-be71-bf55c1da8d1f","cell_type":"code","source":"original_edges = list(G_quality.edges())\n\npose_valid_edges, pose_failed_edges = pose_consistency_filter(\n    test_subset, keypoints_test, desc_ETs, original_edges, threshold_inliers=20\n)\n\nG_pose_pruned = build_pose_pruned_graph(pose_valid_edges)\ncluster_map_pose = cluster_graph(G_pose_pruned)\n\nvisualize_weighted_graph(G_pose_pruned, cluster_map_pose)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"366eab56-baa9-4bcf-b5c8-9ee950fc955e","cell_type":"code","source":"print(\"✅ Retained edges:\", len(pose_valid_edges))\nprint(\"❌ Pruned edges:\", len(pose_failed_edges))\nprint(\"🧠 Clusters after pruning:\", len(set(cluster_map_pose.values())))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"35e72919-88dd-43c5-ab76-f676b63d83e7","cell_type":"markdown","source":"## Notebook Cells Phase 6","metadata":{}},{"id":"a7403552-721d-43b9-9cc8-614e0edafffc","cell_type":"code","source":"def add_outlier_rows(all_image_paths, pose_df, dataset_name=\"testETs\"):\n    pose_images = set(pose_df['image'])\n    outlier_rows = []\n\n    for path in all_image_paths:\n        fname = os.path.basename(path)\n        if fname not in pose_images:\n            outlier_rows.append({\n                \"dataset\": dataset_name,\n                \"scene\": \"outliers\",\n                \"image\": fname,\n                \"rotation_matrix\": \"nan;nan;nan;nan;nan;nan;nan;nan;nan\",\n                \"translation_vector\": \"nan;nan;nan\"\n            })\n\n    return pd.DataFrame(outlier_rows)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"67974793-b96a-42fc-986c-b7c37f527a2a","cell_type":"code","source":"# Get all test image paths (from your image dict)\nall_test_paths = list(test_images.keys())\n\n# Build outlier entries\noutlier_df = add_outlier_rows(all_test_paths, pose_df)\n\n# Combine poses + outliers\nfinal_df = pd.concat([pose_df, outlier_df], ignore_index=True)\n\n# Sort by dataset + image for clean output\nfinal_df = final_df.sort_values(by=[\"dataset\", \"image\"])\n\n# Save\nfinal_df.to_csv(\"submission.csv\", index=False)\nfinal_df.head()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"27799bf1-fbe8-4ad7-8221-9d0d567650b9","cell_type":"code","source":"print(\"submission.csv shape:\", final_df.shape)\n# print(\"Saved to:\", os.getcwd())\n","metadata":{},"outputs":[],"execution_count":null},{"id":"3a4a4ab9-2651-4236-9e32-596659de847f","cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}