Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .gitmodules
Original file line number Diff line number Diff line change
@@ -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
Expand Down
5 changes: 3 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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/):

Expand Down
2 changes: 1 addition & 1 deletion third_party/RoMaV2
214 changes: 214 additions & 0 deletions vidmap/frontend/models/compiled_graph.py
Original file line number Diff line number Diff line change
@@ -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)
32 changes: 20 additions & 12 deletions vidmap/frontend/models/romav2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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",
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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,
Expand Down
65 changes: 65 additions & 0 deletions vidmap/frontend/models/romav2_compile_cache.py
Original file line number Diff line number Diff line change
@@ -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
4 changes: 2 additions & 2 deletions vidmap/frontend/models/romav2_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading