{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":91498,"databundleVersionId":11655853,"sourceType":"competition"},{"sourceId":7884485,"sourceType":"datasetVersion","datasetId":4628051},{"sourceId":11217117,"sourceType":"datasetVersion","datasetId":6988459},{"sourceId":4534,"sourceType":"modelInstanceVersion","modelInstanceId":3326,"modelId":986},{"sourceId":17191,"sourceType":"modelInstanceVersion","modelInstanceId":14317,"modelId":21716},{"sourceId":17555,"sourceType":"modelInstanceVersion","modelInstanceId":14611,"modelId":22086}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# IMPORTANT: Install dependencies and copy model weights\n!pip install --no-index /kaggle/input/imc2024-packages-lightglue-rerun-kornia/* --no-deps\n!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp /kaggle/input/aliked/pytorch/aliked-n16/1/aliked-n16.pth /root/.cache/torch/hub/checkpoints/\n!cp /kaggle/input/lightglue/pytorch/aliked/1/aliked_lightglue.pth /root/.cache/torch/hub/checkpoints/\n!cp /kaggle/input/lightglue/pytorch/aliked/1/aliked_lightglue.pth /root/.cache/torch/hub/checkpoints/aliked_lightglue_v0-1_arxiv-pth","metadata":{"_uuid":"f3cf6c24-c3f8-41db-aca0-1df6cfd5a2f2","_cell_guid":"0bad243a-71e2-4eab-b0fa-2c3833fe6bdf","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-04-01T08:55:23.780908Z","iopub.execute_input":"2025-04-01T08:55:23.781261Z","iopub.status.idle":"2025-04-01T08:55:28.928298Z","shell.execute_reply.started":"2025-04-01T08:55:23.781228Z","shell.execute_reply":"2025-04-01T08:55:28.926942Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nimport os\nfrom tqdm import tqdm\nfrom time import time, sleep\nimport gc\nimport numpy as np\nimport h5py\nimport dataclasses\nimport pandas as pd\nfrom IPython.display import clear_output\nfrom collections import defaultdict, deque\nfrom copy import deepcopy\nfrom PIL import Image\nimport cv2\nimport torch\nimport torch.nn.functional as F\nimport kornia as K\nimport kornia.feature as KF\nfrom lightglue import ALIKED, LightGlue\nfrom lightglue.utils import load_image, rbd\nfrom transformers import AutoImageProcessor, AutoModel\nimport pycolmap\nfrom kornia_moons.feature import draw_LAF_matches\nimport traceback\n\n# IMPORTANT Utilities\nsys.path.append('/kaggle/input/imc25-utils')\nfrom database import *\nfrom h5_to_db import *\nimport metric","metadata":{"_uuid":"52a940c3-a991-4d14-bdef-ff1721b13c2d","_cell_guid":"e0e91669-a4e4-4a79-8b85-15e61e5b3620","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configuration class for easy parameter tuning\n@dataclasses.dataclass\nclass Config:\n    # Feature extraction\n    num_features: int = 4096\n    resize_to: int = 1024\n    feature_extractor: str = 'aliked'  # 'aliked', 'disk'\n    \n    # Matching\n    matcher: str = 'lightglue'  # 'lightglue', 'superglue'\n    min_matches: int = 20\n    match_threshold: float = 0.9\n    \n    # Pair selection\n    sim_th: float = 0.3\n    min_pairs: int = 20\n    exhaustive_if_less: int = 10\n    use_reciprocal_matches: bool = True\n    \n    # Geometric verification\n    geometric_verification: bool = True\n    ransac_thresh: float = 2.0\n    min_inliers: int = 15\n    confidence: float = 0.9999\n    \n    # Reconstruction\n    min_model_size: int = 2\n    max_num_models: int = 25\n    retry_reconstruction: bool = True\n    retry_with_fewer_images: bool = True\n    \n    # Debugging\n    visualize_matches: bool = False\n    debug_dir: str = '/kaggle/working/debug'\n    save_debug_images: bool = False\n    \nconfig = Config()","metadata":{"_uuid":"bfa6f1c4-22ac-4db9-8620-908079ade79d","_cell_guid":"cb04784c-8cbc-4f74-acb3-95faf543b5e1","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Device setup\ndevice = K.utils.get_cuda_device_if_available(0)\nprint(f'{device=}')","metadata":{"_uuid":"fba93293-e56e-4f4b-97da-b59786493a9e","_cell_guid":"595da94c-baad-4314-bd79-e4bb7ff055b6","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_torch_image(fname, device=torch.device('cpu')):\n    \"\"\"Load image as torch tensor with normalization.\"\"\"\n    img = K.io.load_image(fname, K.io.ImageLoadType.RGB32, device=device)[None, ...]\n    return img\n\ndef get_global_desc(fnames, device=torch.device('cpu')):\n    \"\"\"Compute global descriptors using DINOv2 with batching.\"\"\"\n    processor = AutoImageProcessor.from_pretrained('/kaggle/input/dinov2/pytorch/base/1')\n    model = AutoModel.from_pretrained('/kaggle/input/dinov2/pytorch/base/1').eval().to(device)\n    \n    global_descs = []\n    batch_size = 8  # Adjust based on GPU memory\n    \n    for i in range(0, len(fnames), batch_size):\n        batch_fnames = fnames[i:i+batch_size]\n        batch_images = []\n        \n        for img_fname in batch_fnames:\n            timg = load_torch_image(img_fname, device)\n            batch_images.append(timg)\n        \n        batch_images = torch.cat(batch_images, dim=0)\n        with torch.inference_mode():\n            inputs = processor(images=batch_images, return_tensors=\"pt\", do_rescale=False).to(device)\n            outputs = model(**inputs)\n            dino_mac = F.normalize(outputs.last_hidden_state[:,1:].max(dim=1)[0], dim=1, p=2)\n            global_descs.append(dino_mac.detach().cpu())\n    \n    global_descs = torch.cat(global_descs, dim=0)\n    return global_descs\n\ndef select_diverse_pairs(scores, num_pairs):\n    \"\"\"Select diverse pairs using non-maximum suppression.\"\"\"\n    pairs = []\n    scores_copy = scores.copy()\n    np.fill_diagonal(scores_copy, -np.inf)\n    \n    for _ in range(num_pairs):\n        if np.all(scores_copy == -np.inf):\n            break\n        idx = np.unravel_index(np.argmax(scores_copy), scores_copy.shape)\n        pairs.append(idx)\n        # Suppress nearby pairs\n        scores_copy[idx[0], :] = -np.inf\n        scores_copy[:, idx[1]] = -np.inf\n    \n    return pairs\n\ndef get_image_pairs_shortlist(fnames, config=Config(), device=torch.device('cpu')):\n    \"\"\"Improved pair selection with diversity and reciprocal checks.\"\"\"\n    num_imgs = len(fnames)\n    if num_imgs <= config.exhaustive_if_less:\n        return [(i, j) for i in range(num_imgs) for j in range(i+1, num_imgs)]\n    \n    # Compute global descriptors\n    descs = get_global_desc(fnames, device=device)\n    dm = torch.cdist(descs, descs, p=2).detach().cpu().numpy()\n    \n    # Convert distances to similarities\n    sim_matrix = 1 - dm / dm.max()\n    \n    # Select pairs using diversity\n    pairs = set()\n    for i in range(num_imgs):\n        # Get top candidates for this image\n        candidates = np.argsort(sim_matrix[i])[-config.min_pairs*2:][::-1]\n        candidates = [j for j in candidates if j != i]\n        \n        # Add reciprocal pairs if enabled\n        if config.use_reciprocal_matches:\n            for j in candidates:\n                if sim_matrix[j, i] > config.sim_th:\n                    pairs.add(tuple(sorted((i, j))))\n        \n        # Add top pairs for this image\n        for j in candidates[:config.min_pairs]:\n            pairs.add(tuple(sorted((i, j))))\n    \n    # Convert to list and sort\n    pairs = sorted(list(pairs))\n    \n    # If we still don't have enough pairs, add more diverse ones\n    if len(pairs) < config.min_pairs * num_imgs // 2:\n        additional_pairs = select_diverse_pairs(sim_matrix, config.min_pairs * num_imgs)\n        pairs.extend([tuple(sorted(p)) for p in additional_pairs])\n        pairs = sorted(list(set(pairs)))\n    \n    return pairs\n\ndef detect_features(img_fnames, feature_dir='.featureout', config=Config(), device=torch.device('cpu')):\n    \"\"\"Enhanced feature detection with fallback\"\"\"\n    try:\n        dtype = torch.float32\n        if config.feature_extractor == 'aliked':\n            extractor = ALIKED(\n                max_num_keypoints=config.num_features,\n                detection_threshold=0.01,\n                resize=config.resize_to\n            ).eval().to(device, dtype)\n        else:\n            raise ValueError(f\"Unsupported feature extractor: {config.feature_extractor}\")\n        \n        os.makedirs(feature_dir, exist_ok=True)\n        \n        with h5py.File(f'{feature_dir}/keypoints.h5', mode='w') as f_kp, \\\n             h5py.File(f'{feature_dir}/descriptors.h5', mode='w') as f_desc:\n            \n            for img_path in tqdm(img_fnames, desc='Extracting features'):\n                try:\n                    img_fname = os.path.basename(img_path)\n                    key = img_fname\n                    \n                    with torch.inference_mode():\n                        image = load_torch_image(img_path, device).to(dtype)\n                        feats = extractor.extract(image)\n                        \n                        kpts = feats['keypoints'].reshape(-1, 2).detach().cpu().numpy()\n                        descs = feats['descriptors'].reshape(len(kpts), -1).detach().cpu().numpy()\n                        \n                        f_kp[key] = kpts\n                        f_desc[key] = descs\n                except Exception as e:\n                    print(f\"Error processing {img_path}: {str(e)}\")\n                    continue\n    except Exception as e:\n        print(f\"Feature extraction failed: {str(e)}\")\n        traceback.print_exc()\n        raise\n\ndef geometric_verification(kpts1, kpts2, matches, config=Config()):\n    \"\"\"Perform geometric verification using RANSAC.\"\"\"\n    if len(matches) < config.min_inliers:\n        return np.zeros(0, dtype=bool)\n    \n    src_pts = kpts1[matches[:, 0]]\n    dst_pts = kpts2[matches[:, 1]]\n    \n    # Use fundamental matrix for uncalibrated images\n    F, mask = cv2.findFundamentalMat(\n        src_pts, dst_pts,\n        method=cv2.USAC_MAGSAC,\n        ransacReprojThreshold=config.ransac_thresh,\n        confidence=config.confidence,\n        maxIters=10000\n    )\n    \n    if mask is None:\n        return np.zeros(len(matches), dtype=bool)\n    \n    return mask.ravel().astype(bool)\n\ndef visualize_match_pair(img1, img2, kpts1, kpts2, matches, inliers=None, save_path=None):\n    \"\"\"Visualize matches between two images.\"\"\"\n    img1 = cv2.cvtColor(cv2.imread(img1), cv2.COLOR_BGR2RGB)\n    img2 = cv2.cvtColor(cv2.imread(img2), cv2.COLOR_BGR2RGB)\n    \n    if inliers is not None:\n        good_matches = matches[inliers]\n        bad_matches = matches[~inliers]\n    else:\n        good_matches = matches\n        bad_matches = np.zeros((0, 2), dtype=int)\n    \n    # Draw all matches\n    display = cv2.drawMatches(\n        img1, kpts1, img2, kpts2, \n        [cv2.DMatch(_[0], _[1], 0) for _ in good_matches],\n        None,\n        matchColor=(0, 255, 0),  # Green for good matches\n        singlePointColor=(255, 0, 0),\n        flags=cv2.DrawMatchesFlags_NOT_DRAW_SINGLE_POINTS\n    )\n    \n    # Draw bad matches if any\n    if len(bad_matches) > 0:\n        display = cv2.drawMatches(\n            img1, kpts1, img2, kpts2, \n            [cv2.DMatch(_[0], _[1], 0) for _ in bad_matches],\n            display,\n            matchColor=(255, 0, 0),  # Red for bad matches\n            singlePointColor=(255, 0, 0),\n            flags=cv2.DrawMatchesFlags_NOT_DRAW_SINGLE_POINTS | cv2.DrawMatchesFlags_DRAW_OVER_OUTIMG\n        )\n    \n    if save_path:\n        os.makedirs(os.path.dirname(save_path), exist_ok=True)\n        cv2.imwrite(save_path, cv2.cvtColor(display, cv2.COLOR_RGB2BGR))\n    \n    return display\n\ndef match_features(img_fnames, index_pairs, feature_dir='.featureout', config=Config(), device=torch.device('cpu')):\n    \"\"\"Enhanced matching with multiple fallbacks\"\"\"\n    try:\n        if config.matcher == 'lightglue':\n            matcher = KF.LightGlueMatcher(\n                \"aliked\", \n                {\"width_confidence\": -1, \"depth_confidence\": -1, \"mp\": True if 'cuda' in str(device) else False}\n            ).eval().to(device)\n        else:\n            raise ValueError(f\"Unsupported matcher: {config.matcher}\")\n        \n        os.makedirs(config.debug_dir, exist_ok=True)\n        \n        with h5py.File(f'{feature_dir}/keypoints.h5', mode='r') as f_kp, \\\n             h5py.File(f'{feature_dir}/descriptors.h5', mode='r') as f_desc, \\\n             h5py.File(f'{feature_dir}/matches.h5', mode='w') as f_match:\n            \n            for pair_idx in tqdm(index_pairs, desc='Matching pairs'):\n                try:\n                    idx1, idx2 = pair_idx\n                    fname1, fname2 = img_fnames[idx1], img_fnames[idx2]\n                    key1, key2 = os.path.basename(fname1), os.path.basename(fname2)\n                    \n                    # Load features\n                    kp1 = torch.from_numpy(f_kp[key1][...]).to(device)\n                    kp2 = torch.from_numpy(f_kp[key2][...]).to(device)\n                    desc1 = torch.from_numpy(f_desc[key1][...]).to(device)\n                    desc2 = torch.from_numpy(f_desc[key2][...]).to(device)\n                    \n                    # Match features\n                    with torch.inference_mode():\n                        dists, idxs = matcher(\n                            desc1, desc2,\n                            KF.laf_from_center_scale_ori(kp1[None]),\n                            KF.laf_from_center_scale_ori(kp2[None])\n                        )\n                        matches = idxs.detach().cpu().numpy().reshape(-1, 2)\n                    \n                    # Skip if no matches\n                    if len(matches) == 0:\n                        continue\n                        \n                    # Geometric verification\n                    if config.geometric_verification and len(matches) >= config.min_inliers:\n                        kp1_np = kp1.cpu().numpy()\n                        kp2_np = kp2.cpu().numpy()\n                        inliers = geometric_verification(kp1_np, kp2_np, matches, config)\n                        \n                        if np.sum(inliers) < config.min_inliers:\n                            continue\n                        \n                        matches = matches[inliers]\n                        \n                        if config.visualize_matches and config.save_debug_images:\n                            debug_path = os.path.join(config.debug_dir, f'{key1}_{key2}.jpg')\n                            visualize_match_pair(\n                                fname1, fname2, \n                                kp1_np.astype(np.int32), \n                                kp2_np.astype(np.int32),\n                                matches,\n                                save_path=debug_path\n                            )\n                    \n                    # Save matches\n                    if len(matches) >= config.min_matches:\n                        group = f_match.require_group(key1)\n                        group.create_dataset(key2, data=matches)\n                except Exception as e:\n                    print(f\"Error matching pair {pair_idx}: {str(e)}\")\n                    continue\n    except Exception as e:\n        print(f\"Matching failed: {str(e)}\")\n        traceback.print_exc()\n        raise\n\ndef run_reconstruction(images_dir, feature_dir, database_path, output_path, config):\n    \"\"\"Enhanced reconstruction with retries\"\"\"\n    try:\n        # First try with standard parameters\n        mapper_options = pycolmap.IncrementalPipelineOptions()\n        mapper_options.min_model_size = config.min_model_size\n        mapper_options.max_num_models = config.max_num_models\n        \n        maps = pycolmap.incremental_mapping(\n            database_path=database_path,\n            image_path=images_dir,\n            output_path=output_path,\n            options=mapper_options\n        )\n        \n        # If failed but we have matches, try with relaxed parameters\n        if not maps and config.retry_reconstruction:\n            print(\"First reconstruction attempt failed, retrying with relaxed parameters...\")\n            mapper_options.min_model_size = max(2, config.min_model_size - 1)\n            maps = pycolmap.incremental_mapping(\n                database_path=database_path,\n                image_path=images_dir,\n                output_path=output_path + \"_retry\",\n                options=mapper_options\n            )\n        \n        return maps\n    except Exception as e:\n        print(f\"Reconstruction failed: {str(e)}\")\n        traceback.print_exc()\n        return None\n\ndef import_into_colmap(img_dir, feature_dir='.featureout', database_path='colmap.db'):\n    \"\"\"Import features and matches into COLMAP database.\"\"\"\n    if os.path.isfile(database_path):\n        os.remove(database_path)\n    \n    db = COLMAPDatabase.connect(database_path)\n    db.create_tables()\n    single_camera = False\n    fname_to_id = add_keypoints(db, feature_dir, img_dir, '', 'simple-pinhole', single_camera)\n    add_matches(db, feature_dir, fname_to_id)\n    db.commit()\n    return\n\n@dataclasses.dataclass\nclass Prediction:\n    image_id: str | None\n    dataset: str\n    filename: str\n    cluster_index: int | None = None\n    rotation: np.ndarray | None = None\n    translation: np.ndarray | None = None","metadata":{"_uuid":"09f97351-9ef2-4f2a-9450-db443ade29cc","_cell_guid":"37ac179d-459f-4c45-9cf5-ab7d756b6282","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Main execution\ndef main():\n    # Configuration\n    is_train = False\n    data_dir = '/kaggle/input/image-matching-challenge-2025'\n    workdir = '/kaggle/working/result/'\n    os.makedirs(workdir, exist_ok=True)\n    \n    if is_train:\n        sample_submission_csv = os.path.join(data_dir, 'train_labels.csv')\n    else:\n        sample_submission_csv = os.path.join(data_dir, 'sample_submission.csv')\n    \n    # Load dataset\n    samples = {}\n    competition_data = pd.read_csv(sample_submission_csv)\n    for _, row in competition_data.iterrows():\n        if row.dataset not in samples:\n            samples[row.dataset] = []\n        samples[row.dataset].append(\n            Prediction(\n                image_id=None if is_train else row.image_id,\n                dataset=row.dataset,\n                filename=row.image\n            )\n        )\n    \n    # Debug: print dataset info\n    for dataset in samples:\n        print(f'Dataset \"{dataset}\" -> num_images={len(samples[dataset])}')\n    \n    # Process each dataset\n    max_images = None  # For debugging\n    datasets_to_process = None  # None means all datasets\n    \n    if is_train:\n        # Example for training (will hit time limit)\n        # max_images = 5\n        datasets_to_process = [\n            'amy_gardens',\n            'ETs',\n            'fbk_vineyard',\n            'stairs',\n        ]\n    \n    timings = {\n        \"shortlisting\": [],\n        \"feature_detection\": [],\n        \"feature_matching\": [],\n        \"RANSAC\": [],\n        \"Reconstruction\": [],\n    }\n    mapping_result_strs = []\n    \n    print(f\"Running on device {device}\")\n    for dataset, predictions in samples.items():\n        if datasets_to_process and dataset not in datasets_to_process:\n            print(f'Skipping \"{dataset}\"')\n            continue\n        \n        images_dir = os.path.join(data_dir, 'train' if is_train else 'test', dataset)\n        images = [os.path.join(images_dir, p.filename) for p in predictions]\n        if max_images is not None:\n            images = images[:max_images]\n        \n        print(f'\\nProcessing dataset \"{dataset}\": {len(images)} images')\n        filename_to_index = {p.filename: idx for idx, p in enumerate(predictions)}\n        feature_dir = os.path.join(workdir, 'featureout', dataset)\n        os.makedirs(feature_dir, exist_ok=True)\n        \n        try:\n            # 1. Pair selection\n            t = time()\n            index_pairs = get_image_pairs_shortlist(\n                images,\n                config=config,\n                device=device\n            )\n            timings['shortlisting'].append(time() - t)\n            print(f'Shortlisting. Number of pairs to match: {len(index_pairs)}. Done in {time() - t:.4f} sec')\n            gc.collect()\n            \n            # 2. Feature detection\n            t = time()\n            detect_features(images, feature_dir, config=config, device=device)\n            gc.collect()\n            timings['feature_detection'].append(time() - t)\n            print(f'Features detected in {time() - t:.4f} sec')\n            \n            # 3. Feature matching\n            t = time()\n            match_features(images, index_pairs, feature_dir=feature_dir, config=config, device=device)\n            timings['feature_matching'].append(time() - t)\n            print(f'Features matched in {time() - t:.4f} sec')\n            \n            # 4. Reconstruction\n            database_path = os.path.join(feature_dir, 'colmap.db')\n            if os.path.isfile(database_path):\n                os.remove(database_path)\n            gc.collect()\n            sleep(1)\n            \n            import_into_colmap(images_dir, feature_dir=feature_dir, database_path=database_path)\n            output_path = f'{feature_dir}/colmap_rec_aliked'\n            \n            # RANSAC\n            t = time()\n            pycolmap.match_exhaustive(database_path)\n            timings['RANSAC'].append(time() - t)\n            print(f'Ran RANSAC in {time() - t:.4f} sec')\n            \n            # Reconstruction\n            t = time()\n            maps = run_reconstruction(\n                images_dir,\n                feature_dir,\n                database_path,\n                output_path,\n                config\n            )\n            timings['Reconstruction'].append(time() - t)\n            print(f'Reconstruction done in {time() - t:.4f} sec')\n            \n            if maps:\n                print(maps)\n                clear_output(wait=False)\n                \n                # Process results\n                registered = 0\n                for map_index, cur_map in maps.items():\n                    for index, image in cur_map.images.items():\n                        prediction_index = filename_to_index[image.name]\n                        predictions[prediction_index].cluster_index = map_index\n                        predictions[prediction_index].rotation = deepcopy(image.cam_from_world.rotation.matrix())\n                        predictions[prediction_index].translation = deepcopy(image.cam_from_world.translation)\n                        registered += 1\n                \n                mapping_result_str = f'Dataset \"{dataset}\" -> Registered {registered} / {len(images)} images with {len(maps)} clusters'\n            else:\n                mapping_result_str = f'Dataset \"{dataset}\" -> Reconstruction failed!'\n            \n            mapping_result_strs.append(mapping_result_str)\n            print(mapping_result_str)\n            gc.collect()\n        \n        except Exception as e:\n            print(f\"Error processing {dataset}: {str(e)}\")\n            traceback.print_exc()\n            mapping_result_str = f'Dataset \"{dataset}\" -> Failed!'\n            mapping_result_strs.append(mapping_result_str)\n            print(mapping_result_str)\n    \n    # Print summary\n    print('\\nResults')\n    for s in mapping_result_strs:\n        print(s)\n    \n    print('\\nTimings')\n    for k, v in timings.items():\n        print(f'{k} -> total={sum(v):.02f} sec.')\n    \n    # Create submission file\n    array_to_str = lambda array: ';'.join([f\"{x:.09f}\" for x in array])\n    none_to_str = lambda n: ';'.join(['nan'] * n)\n    \n    submission_file = '/kaggle/working/submission.csv'\n    with open(submission_file, 'w') as f:\n        if is_train:\n            f.write('dataset,scene,image,rotation_matrix,translation_vector\\n')\n            for dataset in samples:\n                for prediction in samples[dataset]:\n                    cluster_name = 'outliers' if prediction.cluster_index is None else f'cluster{prediction.cluster_index}'\n                    rotation = none_to_str(9) if prediction.rotation is None else array_to_str(prediction.rotation.flatten())\n                    translation = none_to_str(3) if prediction.translation is None else array_to_str(prediction.translation)\n                    f.write(f'{prediction.dataset},{cluster_name},{prediction.filename},{rotation},{translation}\\n')\n        else:\n            f.write('image_id,dataset,scene,image,rotation_matrix,translation_vector\\n')\n            for dataset in samples:\n                for prediction in samples[dataset]:\n                    cluster_name = 'outliers' if prediction.cluster_index is None else f'cluster{prediction.cluster_index}'\n                    rotation = none_to_str(9) if prediction.rotation is None else array_to_str(prediction.rotation.flatten())\n                    translation = none_to_str(3) if prediction.translation is None else array_to_str(prediction.translation)\n                    f.write(f'{prediction.image_id},{prediction.dataset},{cluster_name},{prediction.filename},{rotation},{translation}\\n')\n    \n    print(f'\\nSubmission file created at {submission_file}')\n    !head {submission_file}\n    \n    # Compute metric if running on training set\n    if is_train:\n        t = time()\n        final_score, dataset_scores = metric.score(\n            gt_csv='/kaggle/input/image-matching-challenge-2025/train_labels.csv',\n            user_csv=submission_file,\n            thresholds_csv='/kaggle/input/image-matching-challenge-2025/train_thresholds.csv',\n            mask_csv=None if is_train else os.path.join(data_dir, 'mask.csv'),\n            inl_cf=0,\n            strict_cf=-1,\n            verbose=True,\n        )\n        print(f'Computed metric in: {time() - t:.02f} sec.')\n\nif __name__ == '__main__':\n    main()","metadata":{"_uuid":"bcc1bcca-f25a-468b-acb6-a74d6277e90a","_cell_guid":"7a610f08-ea19-4b78-9dfb-e46dff53815f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}