{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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":12004827,"sourceType":"datasetVersion","datasetId":7550133},{"sourceId":239894020,"sourceType":"kernelVersion"},{"sourceId":242601569,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# SiLK Submission\n\n使用SiLK模型进行关键点配对，完成聚类工作，3d重建部分未完成。","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/k/dgying/imc2022-dependencies-silk/wheels/hydra_core-1.3.2-py3-none-any.whl\n!pip install /kaggle/input/k/dgying/imc2022-dependencies-silk/wheels/loguru-0.7.3-py3-none-any.whl\n!pip uninstall -y torchtext torchaudio fastai allennlp\n# !pip uninstall -y etils","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T17:07:32.411964Z","iopub.execute_input":"2025-05-31T17:07:32.412596Z","iopub.status.idle":"2025-05-31T17:07:41.861143Z","shell.execute_reply.started":"2025-05-31T17:07:32.412570Z","shell.execute_reply":"2025-05-31T17:07:41.860424Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport torch\nimport h5py\nimport matplotlib.pyplot as plt\nfrom glob import glob\nimport csv\nimport cv2\nfrom functools import partial\nimport random\nfrom torchvision.transforms.functional import resize, InterpolationMode\n\nsys.path.append(os.path.join('/kaggle/input/silkmod/silk/silk'))\nsys.path.append(os.path.join('/kaggle/input/silkmod/silk/silk/scripts/examples'))\n\nimport silk\nimport common\nfrom silk.backbones.silk.silk import from_feature_coords_to_image_coords\n\nconf = {\n    \"common\": {\n        \"paths\": {\n            \"images\": \"/kaggle/input/image-matching-challenge-2025/test/\",\n        },\n        \"nms\": 0, # 0 = disabled\n        \"device\": torch.device('cuda' if torch.cuda.is_available() else 'cpu'),\n        \"topk\": 10_000,\n        \"ransac\": {\n            \"max_iter\": 200_000,\n            \"confidence\": 0.99999,\n            \"reproj_threshold\": 0.25,\n        }\n    },\n    \"model\": {\n        \"checkpoint\": \"/kaggle/input/silkmod/coco-rgb-aug.ckpt\",\n        \"image_max_side\": 480,\n        \"matcher\": {\n            \"postprocessing\": \"double-softmax\",\n            \"threshold\": 0.99,\n            \"temperature\": 0.1,\n        }\n    },\n}\n\ncommon.DEVICE = conf[\"common\"][\"device\"] # hacky\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T17:07:43.900401Z","iopub.execute_input":"2025-05-31T17:07:43.901172Z","iopub.status.idle":"2025-05-31T17:08:01.161823Z","shell.execute_reply.started":"2025-05-31T17:07:43.901147Z","shell.execute_reply":"2025-05-31T17:08:01.161280Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_silk(model, image, max_size = 1024):\n    assert image.shape[0] == 1\n    \n    # get scaling factor\n    if max_size is not None:\n        scale = max(max(image.shape[-2:]) / max_size, 1.)\n    else:\n        scale = 1.\n\n    # resize image if necessary\n    if scale != 1.:\n        image = resize(image, size = (int(image.shape[2] / scale), int(image.shape[3] / scale)), interpolation=InterpolationMode.BILINEAR)\n\n    # to greyscale if necessary\n    if image.shape[1] == 3:\n        image = torchvision.transforms.functional.rgb_to_grayscale(image)\n\n    keypoints, descriptors, prob = model(image)\n    keypoints = from_feature_coords_to_image_coords(model, keypoints)\n    descriptors = descriptors.reshape(1, 128, -1).permute(0, 2, 1)\n\n    keypoints = keypoints[0] * scale\n    descriptors = descriptors[0]\n    prob = prob[0]\n    \n    return keypoints, descriptors, prob\n\ndef get_top_k(keypoints, descriptors, k):\n    positions, scores = keypoints[:,:2], keypoints[:,2]\n\n    # top-k selection\n    idxs = scores.argsort()[-k:]\n\n    return positions[idxs], descriptors[idxs] / 1.41, scores[idxs]\n\ndef extract(model, max_side, top_k, *images):   \n    all_positions = []\n    all_descriptors = []\n    for image in images:\n        positions, descriptors, _ = run_silk(model, image, max_size = max_side)\n        positions, descriptors, scores = get_top_k(positions, descriptors, k = top_k)\n        \n        positions = positions[:,[1,0]]\n        descriptors = descriptors * 1.41\n        \n        all_positions.append(positions)\n        all_descriptors.append(descriptors)\n        \n    return all_positions, all_descriptors\n\ndef inlier_matches(inlier_mask, matches):\n    if inlier_mask is not None:\n        matches_after_ransac = np.array([match for match, is_inlier in zip(matches, inlier_mask) if is_inlier])\n    else:\n        matches_after_ransac = np.array([])\n        \n    return matches_after_ransac    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T17:08:10.379879Z","iopub.execute_input":"2025-05-31T17:08:10.380355Z","iopub.status.idle":"2025-05-31T17:08:10.389955Z","shell.execute_reply.started":"2025-05-31T17:08:10.380333Z","shell.execute_reply":"2025-05-31T17:08:10.389084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# load model\nmodel = common.get_model(\n    checkpoint=conf[\"model\"][\"checkpoint\"],\n    nms=conf[\"common\"][\"nms\"],\n    device=conf[\"common\"][\"device\"],\n)\n\n\n# create matcher\nmatcher = silk.models.silk.matcher(\n    postprocessing=conf[\"model\"][\"matcher\"][\"postprocessing\"],\n    threshold=conf[\"model\"][\"matcher\"][\"threshold\"],\n    temperature=conf[\"model\"][\"matcher\"][\"temperature\"],\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T17:08:13.586769Z","iopub.execute_input":"2025-05-31T17:08:13.587484Z","iopub.status.idle":"2025-05-31T17:08:14.007209Z","shell.execute_reply.started":"2025-05-31T17:08:13.587456Z","shell.execute_reply":"2025-05-31T17:08:14.006486Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport itertools\nimport csv\nimport random\nimport cv2\nimport numpy as np\nfrom typing import List, Tuple, Callable, Any, Dict, Set\n\ndef extract_sift_features(img_path: str) -> Tuple[np.ndarray, np.ndarray]:\n    \"\"\"提取图像的SIFT特征点和描述符\"\"\"\n    try:\n        img = cv2.imread(img_path)\n        if img is None:\n            print(f\"无法读取图像: {img_path}\")\n            return None, None\n        \n        # 转换为灰度图\n        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        \n        # 初始化SIFT检测器\n        sift = cv2.SIFT_create()\n        \n        # 检测关键点并计算描述符\n        kp, des = sift.detectAndCompute(gray, None)\n        \n        return kp, des\n    except Exception as e:\n        print(f\"提取SIFT特征失败: {img_path}, 错误: {str(e)}\")\n        return None, None\n\ndef match_sift_features(des1: np.ndarray, des2: np.ndarray) -> int:\n    \"\"\"匹配两个图像的SIFT描述符，返回匹配数量\"\"\"\n    if des1 is None or des2 is None or len(des1) < 2 or len(des2) < 2:\n        return 0\n    \n    # BFMatcher with default params\n    bf = cv2.BFMatcher()\n    matches = bf.knnMatch(des1, des2, k=2)\n    \n    # 应用比率测试 (Lowe's ratio test)\n    good_matches = []\n    for m, n in matches:\n        if m.distance < 0.75 * n.distance:\n            good_matches.append(m)\n    \n    return len(good_matches)\n\ndef process_image_pairs(\n    folder_path: str,\n    process_func: Callable[[str, str], Any],\n    extensions: List[str] = None,\n    clustering_threshold: float = 100.0,\n    sift_match_threshold: int = 10  # SIFT匹配点数量阈值\n) -> Dict[str, Set[str]]:\n    \"\"\"\n    处理单个文件夹中的图片对并进行聚类，使用SIFT特征预筛选提高效率\n    \n    参数:\n        folder_path: 包含图片的文件夹路径\n        process_func: 处理图片对的函数，接受两个图片路径作为参数\n        extensions: 允许的图片文件扩展名，默认为常见图片格式\n        clustering_threshold: 聚类阈值，默认为100.0\n        sift_match_threshold: SIFT特征匹配点数量阈值，用于预筛选\n    \"\"\"\n    # 设置默认的图片扩展名\n    if extensions is None:\n        extensions = ['.jpg', '.jpeg', '.png', '.bmp', '.gif']\n    \n    # 获取所有图片文件\n    image_files = []\n    for file in os.listdir(folder_path):\n        file_ext = os.path.splitext(file)[1].lower()\n        if file_ext in extensions:\n            image_files.append(os.path.join(folder_path, file))\n    \n    # 提取所有图像的SIFT特征\n    print(f\"提取 {len(image_files)} 张图片的SIFT特征...\")\n    image_features = {}\n    for img_path in image_files:\n        kp, des = extract_sift_features(img_path)\n        if des is not None:\n            image_features[img_path] = des\n    \n    # 使用SIFT特征预筛选生成候选图片对\n    print(\"生成候选图片对...\")\n    candidate_pairs = []\n    total_pairs = len(image_files) * (len(image_files) - 1) // 2\n    \n    # 优化：按描述符数量排序，优先处理特征点多的图像\n    sorted_images = sorted(image_features.items(), key=lambda x: len(x[1]), reverse=True)\n    \n    for i, (img1, des1) in enumerate(sorted_images):\n        for img2, des2 in sorted_images[i+1:]:\n            # 计算SIFT特征匹配数量\n            match_count = match_sift_features(des1, des2)\n            \n            if match_count >= sift_match_threshold:\n                candidate_pairs.append((img1, img2, match_count))\n    \n    # 按匹配数量排序，优先处理匹配点多的图像对\n    candidate_pairs.sort(key=lambda x: x[2], reverse=True)\n    candidate_pairs = [(img1, img2) for img1, img2, _ in candidate_pairs]\n    \n    print(f\"原始图片对数量: {total_pairs}\")\n    print(f\"预筛选后候选对数量: {len(candidate_pairs)}\")\n    print(f\"匹配次数减少: {100 - len(candidate_pairs)/total_pairs*100:.2f}%\")\n    \n    # 处理候选图片对（后续代码保持不变）\n    valid_pairs = []\n    for img1, img2 in candidate_pairs:\n        try:\n            result = process_func(img1, img2)\n            if isinstance(result, (int, float)) and result >= clustering_threshold:\n                valid_pairs.append((img1, img2, result))\n        except Exception as e:\n            print(f\"处理失败: {os.path.basename(img1)} 和 {os.path.basename(img2)}, 错误: {str(e)}\")\n    \n    # 构建聚类（代码保持不变）\n    clusters = []\n    for img1, img2, _ in valid_pairs:\n        cluster1_idx = next((i for i, c in enumerate(clusters) if img1 in c), None)\n        cluster2_idx = next((i for i, c in enumerate(clusters) if img2 in c), None)\n        \n        if cluster1_idx is not None and cluster2_idx is not None:\n            if cluster1_idx != cluster2_idx:\n                clusters[cluster1_idx].update(clusters[cluster2_idx])\n                del clusters[cluster2_idx]\n        elif cluster1_idx is not None:\n            clusters[cluster1_idx].add(img2)\n        elif cluster2_idx is not None:\n            clusters[cluster2_idx].add(img1)\n        else:\n            clusters.append({img1, img2})\n    \n    # 生成聚类名称（代码保持不变）\n    cluster_dict = {}\n    for i, cluster in enumerate(clusters, 1):\n        cluster_name = f\"cluster{i:02d}\"\n        cluster_dict[cluster_name] = cluster\n    \n    # 找出未被聚类的图片（代码保持不变）\n    clustered_images = set()\n    for cluster in clusters:\n        clustered_images.update(cluster)\n    \n    outliers = set(image_files) - clustered_images\n    if outliers:\n        cluster_dict[\"outliers\"] = outliers\n    \n    return cluster_dict\ndef generate_random_rotation_matrix() -> str:\n    \"\"\"生成随机旋转矩阵（9个0-1之间的数字，用;分隔）\"\"\"\n    return ';'.join([f\"{random.random():.6f}\" for _ in range(9)])\n\ndef generate_random_translation_vector() -> str:\n    \"\"\"生成随机位移向量（3个0-1之间的数字，用;分隔）\"\"\"\n    return ';'.join([f\"{random.random():.6f}\" for _ in range(3)])\n\ndef process_multiple_datasets(\n    main_folder: str,\n    process_func: Callable[[str, str], Any],\n    output_file: str = \"clusters.csv\",\n    extensions: List[str] = None,\n    clustering_threshold: float = 100.0\n) -> None:\n    \"\"\"\n    处理主文件夹下的所有子文件夹，并将结果汇总到CSV文件\n    \n    参数:\n        main_folder: 包含多个数据集文件夹的主文件夹\n        process_func: 处理图片对的函数\n        output_file: 输出CSV文件路径\n        extensions: 允许的图片扩展名\n        clustering_threshold: 聚类阈值\n    \"\"\"\n    # 检查主文件夹是否存在\n    if not os.path.isdir(main_folder):\n        print(f\"错误: 主文件夹 '{main_folder}' 不存在\")\n        return\n    \n    # 获取所有子文件夹（数据集）\n    datasets = [\n        d for d in os.listdir(main_folder) \n        if os.path.isdir(os.path.join(main_folder, d))\n    ]\n    \n    if not datasets:\n        print(f\"错误: 主文件夹 '{main_folder}' 中没有子文件夹\")\n        return\n    \n    print(f\"找到 {len(datasets)} 个数据集文件夹\")\n    \n    # 写入CSV文件\n    with open(output_file, 'w', newline='', encoding='utf-8') as csvfile:\n        fieldnames = ['dataset', 'scene', 'image', 'rotation_matrix', 'translation_vector']\n        writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n        writer.writeheader()  # 写入表头\n        \n        # 处理每个数据集\n        for dataset_idx, dataset in enumerate(datasets, 1):\n            dataset_path = os.path.join(main_folder, dataset)\n            print(f\"\\n处理数据集 {dataset_idx}/{len(datasets)}: {dataset}\")\n            \n            # 处理当前数据集的图片\n            clusters = process_image_pairs(\n                dataset_path,\n                process_func=process_func,\n                extensions=extensions,\n                clustering_threshold=clustering_threshold\n            )\n            \n            # 将结果写入CSV\n            for scene, images in clusters.items():\n                for img_path in images:\n                    image_name = os.path.basename(img_path)\n                    writer.writerow({\n                        'dataset': dataset,\n                        'scene': scene,\n                        'image': image_name,\n                        'rotation_matrix': generate_random_rotation_matrix(),\n                        'translation_vector': generate_random_translation_vector()\n                    })\n            \n            print(f\"数据集 {dataset} 处理完成，生成 {len(clusters)} 个聚类\")\n    \n    print(f\"\\n所有数据集处理完成，结果已保存到 {output_file}\")\n\n#推理函数\ndef inference(img_path1: str, img_path2: str) -> float:\n    # load\n    images_1 = common.load_images(img_path1)\n    images_2 = common.load_images(img_path2)\n\n    # extract\n    (positions_1, positions_2), (descriptors_1, descriptors_2) = extract(\n        model,\n        conf[\"model\"][\"image_max_side\"],\n        conf[\"common\"][\"topk\"],\n        images_1,\n        images_2,\n    )\n\n    # match\n    matches = matcher(descriptors_1, descriptors_2)\n    matches = matches.cpu().numpy()\n    print(matches.shape)\n\n    return len(matches)\n\nif __name__ == \"__main__\":\n    # 设置主文件夹路径\n    main_folder = '/kaggle/input/image-matching-challenge-2025/test'\n    \n    # 调用处理函数\n    process_multiple_datasets(\n        main_folder,\n        process_func=inference,  # 替换为你自己定义的函数\n        output_file=\"submission.csv\",\n        clustering_threshold=30.0  # 根据实际需要调整阈值\n    )    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T17:08:16.290420Z","iopub.execute_input":"2025-05-31T17:08:16.290803Z"}},"outputs":[],"execution_count":null}]}