diff --git a/.gitmodules b/.gitmodules index 317b1eb..2da6058 100644 --- a/.gitmodules +++ b/.gitmodules @@ -1,6 +1,6 @@ [submodule "third_party/RoMaV2"] path = third_party/RoMaV2 - url = https://github.com/Parskatt/RoMaV2.git + url = https://github.com/Zador-Pataki/RoMaV2.git [submodule "third_party/Depth-Anything-3"] path = third_party/Depth-Anything-3 url = https://github.com/ByteDance-Seed/Depth-Anything-3.git diff --git a/README.md b/README.md index 3892ba6..bc717b8 100644 --- a/README.md +++ b/README.md @@ -48,10 +48,11 @@ pip install -e . ``` VidMap was last tested on Linux x86-64 with Python 3.10, COLMAP and PyCOLMAP -4.1, PyTorch 2.7.1, TorchVision 0.22.1, and xFormers 0.0.31 on NVIDIA GPUs. +4.1, PyTorch 2.14.0, TorchVision 0.29.0 and xFormers 0.0.35 on NVIDIA GPUs. +Fast loading of cached compiled models requires PyTorch 2.10+. The first frontend run automatically downloads approximately 9 GB of model -checkpoints. +checkpoints and caches approximately 1.2 GB of compiled models. **Optional visualization setup.** Install [Rerun](https://rerun.io/): diff --git a/third_party/RoMaV2 b/third_party/RoMaV2 index 95c9968..f23bab4 160000 --- a/third_party/RoMaV2 +++ b/third_party/RoMaV2 @@ -1 +1 @@ -Subproject commit 95c9968145c8906b7b59383258e9f73b02853d89 +Subproject commit f23bab45a53ffb3f3c3cdda0566d0de5365e7339 diff --git a/vidmap/frontend/models/compiled_graph.py b/vidmap/frontend/models/compiled_graph.py new file mode 100644 index 0000000..a14e57e --- /dev/null +++ b/vidmap/frontend/models/compiled_graph.py @@ -0,0 +1,214 @@ +"""Checksummed, atomic storage around PyTorch's compiled-function serializer.""" + +import fcntl +import hashlib +import inspect +import json +import logging +import os +import platform +from contextlib import contextmanager +from pathlib import Path +from tempfile import TemporaryDirectory +from unittest.mock import patch + +import torch +from torch.utils._pytree import tree_flatten, treespec_dumps + +from vidmap.frontend.cache import file_fingerprint + +logger = logging.getLogger(__name__) + + +def tensor_signature(value): + return {"shape": list(value.shape), "stride": list(value.stride()), "dtype": str(value.dtype)} + + +def digest(value): + return hashlib.sha256(json.dumps(value, sort_keys=True).encode()).hexdigest() + + +def runtime_identity(): + import triton + from torch._inductor import config + + device = torch.cuda.current_device() + return { + "torch": torch.__version__, + "cuda": torch.version.cuda, + "triton": triton.__version__, + "python": platform.python_version(), + "gpu": torch.cuda.get_device_name(device), + "capability": list(torch.cuda.get_device_capability(device)), + "device": device, + "inductor": hashlib.sha256(config.save_config()).hexdigest(), + "cudnn": torch.backends.cudnn.version(), + "cudnn_tf32": torch.backends.cudnn.allow_tf32, + "cudnn_benchmark": torch.backends.cudnn.benchmark, + "cudnn_deterministic": torch.backends.cudnn.deterministic, + "deterministic": torch.are_deterministic_algorithms_enabled(), + "fp16_reduction": torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, + "bf16_reduction": torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, + "cublas_workspace": os.environ["CUBLAS_WORKSPACE_CONFIG"] if "CUBLAS_WORKSPACE_CONFIG" in os.environ else None, + "flash_sdp": torch.backends.cuda.flash_sdp_enabled(), + "math_sdp": torch.backends.cuda.math_sdp_enabled(), + "memory_efficient_sdp": torch.backends.cuda.mem_efficient_sdp_enabled(), + "cudnn_sdp": torch.backends.cuda.cudnn_sdp_enabled(), + } + + +@contextmanager +def graph_entry(root, key): + """Lock an entry and publish a successful capture atomically.""" + root.mkdir(parents=True, exist_ok=True) + directory = root / key + with (root / f"{key}.lock").open("a") as lock: + fcntl.flock(lock, fcntl.LOCK_EX) + if directory.exists(): + logger.info(f"Loading saved graph: {directory}") + yield directory, False + else: + logger.info(f"Capturing graph: {directory}") + with TemporaryDirectory(prefix=f".{key}-", dir=root) as temporary: + candidate = Path(temporary) / "entry" + candidate.mkdir() + yield candidate, True + candidate.rename(directory) + + +def _manifest(directory, identity, payload, *, write=False): + path = directory / "manifest.json" + if write: + path.write_text(json.dumps({"identity": identity, "sha256": file_fingerprint(directory / payload)}) + "\n") + else: + manifest = json.loads(path.read_text()) + if manifest["identity"] != identity or manifest["sha256"] != file_fingerprint(directory / payload): + raise ValueError(f"Compiled artifact identity/checksum mismatch: {directory}; remove it to rebuild") + + +class CachedGraph(torch.nn.Module): + """Share model weights across shapes; delegate executable persistence to PyTorch.""" + + def __init__(self, *, net=None, factory=None, namespace, component, model_identity, sources, extra_identity=None): + super().__init__() + assert (net is None) != (factory is None) + self.net = net + self._factory = factory + self.graphs = {} + self.persistent = ( + "load_compiled_function" in vars(torch.compiler) + and "f_globals" in inspect.signature(torch.compiler.load_compiled_function).parameters + ) + if not self.persistent: + logger.info("Compiled-function persistence is unavailable; using ordinary torch.compile") + self.runtime = runtime_identity() + self.identity = { + **({} if extra_identity is None else extra_identity), + "format": "torch-aot-function-v1", + "component": component, + "model": model_identity, + "source_sha256": [file_fingerprint(path) for path in [*sources, Path(__file__)]], + "runtime": self.runtime, + } + variable = f"VIDMAP_{namespace.upper()}_CACHE_DIR" + root = ( + Path(os.environ[variable]).expanduser() + if variable in os.environ + else Path.home() / ".cache/vidmap" / namespace + ) + self.root = root / digest(self.identity) + if net is not None: + self.train(net.training) + + def prepare_module(self, net): + return net + + def load_module(self, path): + return torch.load(path, map_location=f"cuda:{self.runtime['device']}", weights_only=False) + + def _ensure_module(self): + if self.net is not None: + return + if not self.persistent: + self.net = self.prepare_module(self._factory()) + return + with graph_entry(self.root, "module") as (directory, build): + if build: + self.net = self.prepare_module(self._factory()) + torch.save(self.net, directory / "module.pt") + _manifest(directory, self.identity, "module.pt", write=True) + else: + _manifest(directory, self.identity, "module.pt") + self.net = self.prepare_module(self.load_module(directory / "module.pt")) + + def _apply(self, fn, recurse=True): + self.graphs.clear() + if self._factory is not None and self.net is not None: + self.net.cpu() + self.net = None + return super()._apply(fn, recurse=recurse) + + def forward(self, *args, **kwargs): + if self.training or not torch.is_inference_mode_enabled(): + raise ValueError("Saved graphs require eval and inference mode") + if torch.get_float32_matmul_precision() != "highest" or runtime_identity() != self.runtime: + raise ValueError("Compiler, device or precision settings changed") + return self._run(args, kwargs) + + def _run(self, args, kwargs): + if not self.persistent: + self._ensure_module() + if "jit" not in self.graphs: + self.graphs["jit"] = torch.compile(type(self.net).forward, fullgraph=True, dynamic=False) + return self.graphs["jit"](self.net, *args, **kwargs) + leaves, structure = tree_flatten((args, kwargs)) + assert all( + isinstance(value, torch.Tensor) or type(value) in (str, int, float, bool, type(None)) for value in leaves + ) + tensors = [value for value in leaves if isinstance(value, torch.Tensor)] + assert tensors + identity = { + **self.identity, + "structure": treespec_dumps(structure), + "inputs": [tensor_signature(value) if isinstance(value, torch.Tensor) else value for value in leaves], + "aliases": [ + [ + i + for i, candidate in enumerate(tensors) + if value is candidate or value.data_ptr() == candidate.data_ptr() + ] + for value in tensors + ], + "autocast": torch.is_autocast_enabled("cuda"), + "autocast_dtype": str(torch.get_autocast_dtype("cuda")), + } + key = digest(identity) + if key not in self.graphs: + self._ensure_module() + with graph_entry(self.root, key) as (directory, build): + if build: + from torch._functorch._aot_autograd.autograd_cache import AOTAutogradCache + + # Seeded AOT nonce keys can collide. Isolate their lookup, retaining Triton tuning. + with ( + torch._dynamo.convert_frame.compile_lock, + TemporaryDirectory(dir=directory) as temporary, + patch.object(AOTAutogradCache, "_get_tmp_dir", return_value=temporary), + torch._functorch.config.patch(enable_remote_autograd_cache=False), + ): + graph = torch.compile(type(self.net).forward, fullgraph=True, dynamic=False).aot_compile( + ((self.net, *args), kwargs) + ) + graph.save_compiled_function(str(directory / "compiled.pt")) + _manifest(directory, identity, "compiled.pt", write=True) + else: + _manifest(directory, identity, "compiled.pt") + with (directory / "compiled.pt").open("rb") as stream: + graph = torch.compiler.load_compiled_function( + stream, f_globals=type(self.net).forward.__globals__ + ) + # Validate execution before publishing a newly built program. + result = graph(self.net, *args, **kwargs) + self.graphs[key] = graph + return result + return self.graphs[key](self.net, *args, **kwargs) diff --git a/vidmap/frontend/models/romav2.py b/vidmap/frontend/models/romav2.py index aad7bea..6dae3a9 100644 --- a/vidmap/frontend/models/romav2.py +++ b/vidmap/frontend/models/romav2.py @@ -9,6 +9,7 @@ import torch from vidmap.frontend.cache import file_fingerprint +from vidmap.frontend.models.romav2_compile_cache import CachedRoMaGraph from vidmap.frontend.models.romav2_inference import ( match_lowres_batch, match_true_highres_pair, @@ -17,7 +18,7 @@ from vidmap.frontend.options.matching import RoMaV2Options from vidmap.model_sources import import_model_package, model_package_root -ROMAV2_SOURCE_REVISION = "95c9968145c8906b7b59383258e9f73b02853d89" +ROMAV2_SOURCE_REVISION = "f23bab45a53ffb3f3c3cdda0566d0de5365e7339" ROMAV2_PACKAGE_ROOT = model_package_root( "romav2", "third_party/RoMaV2/src/romav2", @@ -71,18 +72,24 @@ def __init__(self, conf: RoMaV2Options): module = import_model_package("romav2", ROMAV2_PACKAGE_ROOT) _configure_romav2_logging() - cfg = module.RoMaV2.Cfg( - setting="precise", - compile=conf.compile, - ) use_native_local_correlation() - # RoMaV2 downloads its release checkpoint on first construction. - self._net = module.RoMaV2(cfg) - _verify_romav2_checkpoint() - self._net.bidirectional = False - self._net.eval() - for parameter in self.parameters(): - parameter.requires_grad = False + cached = conf.compile and torch.cuda.is_available() + + def create_net(): + # RoMaV2 downloads its release checkpoint on first construction. + cfg = module.RoMaV2.Cfg(setting="precise", compile=conf.compile and not cached) + net = module.RoMaV2(cfg) + _verify_romav2_checkpoint() + net.bidirectional = False + net.return_intermediates = False + return net.eval().requires_grad_(False) + + if cached: + self._net = CachedRoMaGraph( + create_net, romav2_cache_identity(conf), ROMAV2_PACKAGE_ROOT, component="whole_model" + ).eval() + else: + self._net = create_net() def forward(self, data): raise NotImplementedError("Use the explicit RoMaV2 inference operations") @@ -216,6 +223,7 @@ def romav2_cache_identity(conf: RoMaV2Options): "bidirectional": False, }, "version": ROMAV2_VERSION, + "torch": torch.__version__, "source_revision": ROMAV2_SOURCE_REVISION, "checkpoint_sha256": ROMAV2_CHECKPOINT_SHA256, "true_highres": True, diff --git a/vidmap/frontend/models/romav2_compile_cache.py b/vidmap/frontend/models/romav2_compile_cache.py new file mode 100644 index 0000000..4c18d47 --- /dev/null +++ b/vidmap/frontend/models/romav2_compile_cache.py @@ -0,0 +1,65 @@ +"""RoMa saved-program identity and image/feature input contract.""" + +import sys +from collections.abc import Callable +from dataclasses import is_dataclass +from pathlib import Path +from types import SimpleNamespace + +import torch + +from vidmap.frontend.models.compiled_graph import CachedGraph + + +class CachedRoMaGraph(CachedGraph): + """Lazily capture or load a tensor-only RoMa component.""" + + bidirectional = False + threshold = None + + def __init__( + self, factory: Callable[[], torch.nn.Module], model_identity: dict, source_root: Path, *, component: str + ): + sources = [ + *sorted(source_root.rglob("*.py")), + Path(__file__), + Path(__file__).with_name("romav2.py"), + Path(__file__).with_name("romav2_inference.py"), + ] + super().__init__( + factory=factory, namespace="romav2", component=component, model_identity=model_identity, sources=sources + ) + + def forward(self, *inputs: torch.Tensor) -> dict: + if torch.is_autocast_enabled("cuda"): + raise ValueError("Saved RoMa graphs require outer autocast to be disabled") + if not inputs or any( + value.device != torch.device("cuda", self.runtime["device"]) + or value.dtype not in (torch.float32, torch.bfloat16) + or not value.is_contiguous() + or not value.is_inference() + for value in inputs + ): + raise ValueError("RoMa components require contiguous float32 or bfloat16 CUDA inference tensors") + if self.identity["component"] == "whole_model" and ( + len(inputs) not in (2, 4) + or any(value.dtype != torch.float32 or value.ndim != 4 or value.shape[1] != 3 for value in inputs) + ): + raise ValueError("RoMa requires two or four float32 image tensors") + return super().forward(*inputs) + + def load_module(self, path): + descriptor_source = ( + Path(torch.hub.get_dir()) / "facebookresearch_dinov3_adc254450203739c8149213a7a69d8d905b4fcfa" + ) + if not descriptor_source.is_dir(): + raise FileNotFoundError(f"Cached RoMa requires its pinned DINOv3 source: {descriptor_source}") + sys.path.insert(0, str(descriptor_source)) + return super().load_module(path) + + def prepare_module(self, net): + # PyTorch's portable type guards require globally importable config types. + for module in net.modules(): + if "cfg" in module.__dict__ and is_dataclass(module.cfg): + module.cfg = SimpleNamespace(**vars(module.cfg)) + return net diff --git a/vidmap/frontend/models/romav2_inference.py b/vidmap/frontend/models/romav2_inference.py index 1fbd4db..c829ac3 100644 --- a/vidmap/frontend/models/romav2_inference.py +++ b/vidmap/frontend/models/romav2_inference.py @@ -81,7 +81,7 @@ def match_true_highres_pair( predictions = model( _resize(image_a_lowres, lowres_size), _resize(image_b_lowres, lowres_size), - img_A_hr=_resize(image_a_highres, highres_size), - img_B_hr=_resize(image_b_highres, highres_size), + _resize(image_a_highres, highres_size), + _resize(image_b_highres, highres_size), ) return _finalize_predictions(model, predictions)