from __future__ import annotations import json import os from collections.abc import Callable, Mapping from dataclasses import dataclass from pathlib import Path, PurePosixPath, PureWindowsPath from typing import TypeAlias import numpy as np import pycolmap from numpy.typing import NDArray from PIL import Image @dataclass(frozen=True, slots=True) class CalibrationView: image_id: int name: str rgb_path: Path mask_path: Path camera: pycolmap.Camera cam_from_world: "pycolmap.Rigid3d" @dataclass(frozen=True, slots=True) class CalibrationDataset: reconstruction: pycolmap.Reconstruction views: tuple[CalibrationView, ...] @dataclass(frozen=True, slots=True) class CalibrationResult: corrected_images: Mapping[int, NDArray[np.uint8]] report: Mapping[str, object] @dataclass(frozen=True, slots=True) class CalibrationOutputs: image_paths: tuple[Path, ...] report_path: Path Calibrator: TypeAlias = Callable[[CalibrationDataset], CalibrationResult] def _relative_image_path(name: str) -> Path: posix_path = PurePosixPath(name) windows_path = PureWindowsPath(name) if ( not posix_path.parts or posix_path.is_absolute() or ".." in posix_path.parts or windows_path.is_absolute() or windows_path.drive or ".." in windows_path.parts ): raise ValueError(f"unsafe COLMAP image name: {name}") return Path(*posix_path.parts) def _validate_output_target( output_root: Path, resolved_output_root: Path, relative: Path, seen_targets: set[str], seen_ancestors: set[str], ) -> Path: target = output_root / relative resolved_target = target.resolve() if not resolved_target.is_relative_to(resolved_output_root): raise ValueError(f"output path escapes output_dir: {target}") ancestor = target.parent while True: if (ancestor.exists() or ancestor.is_symlink()) and not ancestor.is_dir(): raise NotADirectoryError(f"output ancestor is not a directory: {ancestor}") if ancestor == output_root: break ancestor = ancestor.parent target_key = os.path.normcase(str(resolved_target)) if target_key in seen_targets: raise ValueError(f"duplicate output path: {target}") if target_key in seen_ancestors or any( os.path.normcase(str(parent)) in seen_targets for parent in resolved_target.parents ): raise ValueError(f"conflicting output path: {target}") seen_targets.add(target_key) seen_ancestors.update( os.path.normcase(str(parent)) for parent in resolved_target.parents ) if target.exists() or target.is_symlink(): raise FileExistsError(f"output already exists: {target}") return target def load_dataset( model_dir: str | os.PathLike[str], rgb_dir: str | os.PathLike[str], mask_dir: str | os.PathLike[str], ) -> CalibrationDataset: # ✅ pycolmap 4.x: 直接用 Reconstruction 加载模型目录 reconstruction = pycolmap.Reconstruction(str(Path(model_dir))) # 验证加载成功 if reconstruction.num_images() == 0: raise RuntimeError( f"Failed to load COLMAP model from {model_dir}. " f"Expected cameras.bin/images.bin/points3D.bin or .txt files." ) print(f"Loaded COLMAP model: {reconstruction.num_images()} images, " f"{reconstruction.num_cameras()} cameras, " f"{reconstruction.num_points3D()} 3D points") rgb_root = Path(rgb_dir) mask_root = Path(mask_dir) seen_names: set[str] = set() views: list[CalibrationView] = [] # ✅ pycolmap 4.x: images 是一个 dict-like 对象 for image_id, image in reconstruction.images.items(): if not image.has_pose: raise ValueError(f"COLMAP image has no pose: {image.name}") relative = _relative_image_path(image.name) normalized_name = os.path.normcase(str(relative)) if normalized_name in seen_names: raise ValueError(f"duplicate COLMAP image name: {image.name}") seen_names.add(normalized_name) rgb_path = rgb_root / relative mask_path = mask_root / relative.with_suffix(".png") if not rgb_path.is_file(): raise FileNotFoundError(f"RGB image not found: {rgb_path}") if not mask_path.is_file(): raise FileNotFoundError(f"mask image not found: {mask_path}") camera = image.camera if camera is None: raise ValueError(f"COLMAP image has no camera: {image.name}") with Image.open(rgb_path) as rgb_image: rgb_size = rgb_image.size with Image.open(mask_path) as mask_image: mask_size = mask_image.size camera_size = (camera.width, camera.height) if rgb_size != camera_size or mask_size != camera_size: raise ValueError( f"size mismatch for {image.name}: camera={camera_size}, " f"rgb={rgb_size}, mask={mask_size}" ) views.append( CalibrationView( image_id=image_id, name=image.name, rgb_path=rgb_path, mask_path=mask_path, camera=camera, # ✅ pycolmap 4.x: cam_from_world 是属性,不是方法 cam_from_world=image.cam_from_world, ) ) return CalibrationDataset(reconstruction=reconstruction, views=tuple(views)) def read_rgb(view: CalibrationView) -> NDArray[np.uint8]: with Image.open(view.rgb_path) as image: return np.array(image.convert("RGB"), dtype=np.uint8, copy=True) def read_mask(view: CalibrationView) -> NDArray[np.bool_]: with Image.open(view.mask_path) as image: return np.array(image.convert("L"), dtype=np.uint8, copy=True) != 0 def run_calibration( model_dir: str | os.PathLike[str], rgb_dir: str | os.PathLike[str], mask_dir: str | os.PathLike[str], output_dir: str | os.PathLike[str], *, calibrator: Calibrator, ) -> CalibrationOutputs: dataset = load_dataset(model_dir, rgb_dir, mask_dir) output_root = Path(output_dir) resolved_output_root = output_root.resolve() if resolved_output_root == Path(rgb_dir).resolve(): raise ValueError("output_dir must differ from rgb_dir") result = calibrator(dataset) if not isinstance(result, CalibrationResult): raise ValueError("calibrator must return CalibrationResult") if not isinstance(result.corrected_images, Mapping): raise ValueError("calibrator corrected_images must be a mapping") try: report_json = json.dumps( result.report, ensure_ascii=False, indent=2, sort_keys=True, allow_nan=False, ) except (TypeError, ValueError) as error: raise ValueError("calibration report is not valid JSON") from error expected_ids = {view.image_id for view in dataset.views} returned_ids = set(result.corrected_images) if returned_ids != expected_ids: missing = sorted(expected_ids - returned_ids) extra = sorted(returned_ids - expected_ids, key=repr) raise ValueError( f"calibrator returned wrong image ids: missing={missing} extra={extra}" ) Image.init() output_formats = Image.registered_extensions() seen_targets: set[str] = set() seen_ancestors: set[str] = set() outputs: list[tuple[Path, NDArray[np.uint8]]] = [] for view in dataset.views: rgb = result.corrected_images[view.image_id] expected_shape = (view.camera.height, view.camera.width, 3) if ( not isinstance(rgb, np.ndarray) or rgb.dtype != np.uint8 or rgb.shape != expected_shape ): raise ValueError( f"invalid output for image {view.name}: expected shape " f"{expected_shape} and dtype uint8, got shape " f"{getattr(rgb, 'shape', None)} and dtype " f"{getattr(rgb, 'dtype', None)}" ) target = _validate_output_target( output_root, resolved_output_root, _relative_image_path(view.name), seen_targets, seen_ancestors, ) output_format = output_formats.get(target.suffix.lower()) if output_format not in Image.SAVE: raise ValueError(f"unsupported output image extension: {target.suffix}") outputs.append((target, rgb)) report_path = _validate_output_target( output_root, resolved_output_root, Path("calibration_report.json"), seen_targets, seen_ancestors, ) for target, rgb in outputs: target.parent.mkdir(parents=True, exist_ok=True) Image.fromarray(rgb).save(target) report_path.parent.mkdir(parents=True, exist_ok=True) report_path.write_text(f"{report_json}\n", encoding="utf-8") return CalibrationOutputs( image_paths=tuple(target for target, _ in outputs), report_path=report_path, )