Skip to content
Open
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
9 changes: 6 additions & 3 deletions demo_colmap.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,11 +278,14 @@ def rename_colmap_recons_and_rescale_camera(
pycamera.height = real_image_size[1]

if shift_point2d_to_original_res:
# Also shift the point2D to original resolution
top_left = original_coords[pyimageid - 1, :2]
x1, y1, x2, y2, w, h = original_coords[pyimageid - 1]
if x1 == 0:
shift = np.array([x1, y1]) * (w / x2)
else:
shift = np.array([x1, y1]) * (h / y2)

for point2D in pyimage.points2D:
point2D.xy = (point2D.xy - top_left) * resize_ratio
point2D.xy = point2D.xy * resize_ratio - shift

if shared_camera:
# If shared_camera, all images share the same camera
Expand Down
2 changes: 1 addition & 1 deletion vggt/dependency/np_to_pycolmap.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,8 +55,8 @@ def batch_np_matrix_to_pycolmap(

if max_reproj_error is not None:
projected_points_2d, projected_points_cam = project_3D_points_np(points3d, extrinsics, intrinsics)
projected_diff = np.linalg.norm(projected_points_2d - tracks, axis=-1)
projected_points_2d[projected_points_cam[:, -1] <= 0] = 1e6
projected_diff = np.linalg.norm(projected_points_2d - tracks, axis=-1)
reproj_mask = projected_diff < max_reproj_error

if masks is not None and reproj_mask is not None:
Expand Down
75 changes: 49 additions & 26 deletions vggt/dependency/track_predict.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,16 @@
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.

import torch
import math
import numpy as np
from .vggsfm_utils import *
from .projection import project_3D_points_np


def predict_tracks(
images,
extrinsics,
intrinsics,
conf=None,
points_3d=None,
masks=None,
Expand All @@ -20,6 +23,9 @@ def predict_tracks(
max_points_num=163840,
fine_tracking=True,
complete_non_vis=True,
track_vis_thresh=0.1,
reproj_error_thresh=8,
min_inlier_per_frame=64
):
"""
Predict tracks for the given images and masks.
Expand All @@ -33,6 +39,8 @@ def predict_tracks(

Args:
images: Tensor of shape [S, 3, H, W] containing the input images.
extrinsics: Array of shape [S, 3, 4] containing the extrinsic parameters for each frame.
intrinsics: Array of shape [S, 3, 3] containing the intrinsic parameters for each frame.
conf: Tensor of shape [S, 1, H, W] containing the confidence scores. Default is None.
points_3d: Tensor containing 3D points. Default is None.
masks: Optional tensor of shape [S, 1, H, W] containing masks. Default is None.
Expand All @@ -42,6 +50,9 @@ def predict_tracks(
max_points_num: Maximum number of points to process at once. Default is 163840.
fine_tracking: Whether to use fine tracking. Default is True.
complete_non_vis: Whether to augment non-visible frames. Default is True.
track_vis_thresh: Visibility threshold for track filtering
reproj_error_thresh: Reprojection error threshold for track filtering
min_inlier_per_frame: Minimum number of inliers per frame

Returns:
pred_tracks: Numpy array containing the predicted tracks.
Expand Down Expand Up @@ -108,16 +119,19 @@ def predict_tracks(
pred_points_3d,
pred_colors,
images,
extrinsics,
intrinsics,
conf,
points_3d,
fmaps_for_tracker,
keypoint_extractors,
tracker,
max_points_num,
fine_tracking,
min_vis=500,
non_vis_thresh=0.1,
device=device,
track_vis_thresh=track_vis_thresh,
reproj_error_thresh=reproj_error_thresh,
min_inlier_per_frame=min_inlier_per_frame,
device=device
)

pred_tracks = np.concatenate(pred_tracks, axis=1)
Expand Down Expand Up @@ -209,16 +223,17 @@ def _forward_on_query(
fmaps_feed = fmaps_feed[None] # add batch dimension

all_points_num = images_feed.shape[1] * query_points.shape[1]
gpu_count = torch.cuda.device_count()
gpu_points_num = math.ceil(all_points_num / gpu_count)

# Don't need to be scared, this is just chunking to make GPU happy
if all_points_num > max_points_num:
num_splits = (all_points_num + max_points_num - 1) // max_points_num
query_points = torch.chunk(query_points, num_splits, dim=1)
else:
if gpu_points_num <= max_points_num:
query_points = [query_points]
else:
chunk_num = math.ceil(all_points_num / (max_points_num * gpu_count))
query_points = torch.chunk(query_points, chunk_num, dim=1)

pred_track, pred_vis, _ = predict_tracks_in_chunks(
tracker, images_feed, query_points, fmaps_feed, fine_tracking=fine_tracking
tracker, images_feed, query_points, fmaps_feed, fine_tracking=fine_tracking, parallel=gpu_count > 1
)

pred_track, pred_vis = switch_tensor_order([pred_track, pred_vis], reorder_index, dim=1)
Expand All @@ -236,16 +251,18 @@ def _augment_non_visible_frames(
pred_points_3d: list, # ← running list of np.ndarrays for 3D points
pred_colors: list, # ← running list of np.ndarrays for colors
images: torch.Tensor,
extrinsics,
intrinsics,
conf,
points_3d,
fmaps_for_tracker,
keypoint_extractors,
tracker,
max_points_num: int,
fine_tracking: bool,
*,
min_vis: int = 500,
non_vis_thresh: float = 0.1,
track_vis_thresh: float = 0.1,
reproj_error_thresh: float = 8.0,
min_inlier_per_frame: int = 64,
device: torch.device = None,
):
"""
Expand All @@ -258,15 +275,18 @@ def _augment_non_visible_frames(
pred_points_3d: List of numpy arrays containing 3D points.
pred_colors: List of numpy arrays containing point colors.
images: Tensor of shape [S, 3, H, W] containing the input images.
extrinsics: Array of shape [S, 3, 4] containing the extrinsic parameters for each frame.
intrinsics: Array of shape [S, 3, 3] containing the intrinsic parameters for each frame.
conf: Tensor of shape [S, 1, H, W] containing confidence scores
points_3d: Tensor containing 3D points
fmaps_for_tracker: Feature maps for the tracker
keypoint_extractors: Initialized feature extractors
tracker: VGG-SFM tracker
max_points_num: Maximum number of points to process at once
fine_tracking: Whether to use fine tracking
min_vis: Minimum visibility threshold
non_vis_thresh: Non-visibility threshold
track_vis_thresh: Visibility threshold for tracks filtering
reproj_error_thresh: Reprojection error threshold for track filtering
min_inlier_per_frame: Minimum number of inliers per frame
device: Device to use for computation

Returns:
Expand All @@ -277,28 +297,31 @@ def _augment_non_visible_frames(
cur_extractors = keypoint_extractors # may be replaced on the final trial

while True:
# Visibility per frame
vis_array = np.concatenate(pred_vis_scores, axis=1)

# Count frames with sufficient visibility using numpy
sufficient_vis_count = (vis_array > non_vis_thresh).sum(axis=-1)
non_vis_frames = np.where(sufficient_vis_count < min_vis)[0].tolist()
vis_mask = np.concatenate(pred_vis_scores, axis=-1) > track_vis_thresh

projected_points_2d, projected_points_cam = project_3D_points_np(np.concatenate(pred_points_3d, axis=0), extrinsics, intrinsics)
projected_points_2d[projected_points_cam[:, -1] <= 0] = 1e6
projected_diff = np.linalg.norm(projected_points_2d - np.concatenate(pred_tracks, axis=1), axis=-1)
reproj_mask = projected_diff < reproj_error_thresh
mask = np.logical_and(vis_mask, reproj_mask)
non_inlier_frames = np.where(mask.sum(axis=1) < min_inlier_per_frame)[0].tolist()

if len(non_vis_frames) == 0:
if len(non_inlier_frames) == 0:
break

print("Processing non visible frames:", non_vis_frames)
print("Processing non enough inlier frames:", non_inlier_frames)

# Decide the frames & extractor for this round
if non_vis_frames[0] == last_query:
if non_inlier_frames[0] == last_query:
# Same frame failed twice - final "all-in" attempt
final_trial = True
cur_extractors = initialize_feature_extractors(2048, extractor_method="sp+sift+aliked", device=device)
query_frame_list = non_vis_frames # blast them all at once
query_frame_list = non_inlier_frames # blast them all at once
else:
query_frame_list = [non_vis_frames[0]] # Process one at a time
query_frame_list = [non_inlier_frames[0]] # Process one at a time

last_query = non_vis_frames[0]
last_query = non_inlier_frames[0]

# Run the tracker for every selected frame
for query_index in query_frame_list:
Expand Down
45 changes: 33 additions & 12 deletions vggt/dependency/vggsfm_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,13 @@
# LICENSE file in the root directory of this source tree.

import logging
import math
import warnings
from typing import Dict, List, Optional, Tuple, Union

import numpy as np
import pycolmap
import torch
import torch.nn.functional as F
from lightglue import ALIKED, SIFT, SuperPoint
from torch.nn import DataParallel

from .vggsfm_tracker import TrackerPredictor

Expand Down Expand Up @@ -95,17 +94,19 @@ def generate_rank_by_dino(
frame_feat_norm = F.normalize(frame_feat, p=2, dim=1)
similarity_matrix = torch.mm(frame_feat_norm, frame_feat_norm.transpose(-1, -2))

distance_matrix = 100 - similarity_matrix.clone()
with torch.amp.autocast("cuda", enabled=False):
similarity_matrix = similarity_matrix.float()
distance_matrix = 100 - similarity_matrix

# Ignore self-pairing
similarity_matrix.fill_diagonal_(-100)
similarity_sum = similarity_matrix.sum(dim=1)
# Ignore self-pairing
similarity_matrix.fill_diagonal_(-100)
similarity_sum = similarity_matrix.sum(dim=1)

# Find the most common frame
most_common_frame_index = torch.argmax(similarity_sum).item()
# Find the most common frame
most_common_frame_index = torch.argmax(similarity_sum).item()

# Conduct FPS sampling starting from the most common frame
fps_idx = farthest_point_sampling(distance_matrix, query_frame_num, most_common_frame_index)
# Conduct FPS sampling starting from the most common frame
fps_idx = farthest_point_sampling(distance_matrix, query_frame_num, most_common_frame_index)

# Clean up all tensors and models to free memory
del frame_feat, frame_feat_norm, similarity_matrix, distance_matrix
Expand Down Expand Up @@ -253,7 +254,7 @@ def extract_keypoints(query_image, extractors, round_keypoints=True):


def predict_tracks_in_chunks(
track_predictor, images_feed, query_points_list, fmaps_feed, fine_tracking, num_splits=None, fine_chunk=40960
track_predictor, images_feed, query_points_list, fmaps_feed, fine_tracking, num_splits=None, fine_chunk=40960, parallel=True,
):
"""
Process a list of query points to avoid memory issues.
Expand All @@ -265,6 +266,7 @@ def predict_tracks_in_chunks(
fmaps_feed (torch.Tensor): A tensor of feature maps for the tracker.
fine_tracking (bool): Whether to perform fine tracking.
num_splits (int, optional): Ignored when query_points_list is provided. Kept for backward compatibility.
parallel: Using multiple GPUs to speed up predicting tracks. Default is True.

Returns:
tuple: A tuple containing the concatenated predicted tracks, visibility, and scores.
Expand All @@ -284,11 +286,30 @@ def predict_tracks_in_chunks(
pred_vis_list = []
pred_score_list = []

track_predictor = DataParallel(track_predictor) if parallel else track_predictor

for split_points in query_points_list:
# Feed into track predictor for each split
if parallel:
gpu_count = torch.cuda.device_count()
points_num = split_points.shape[1]
batch_points_num = math.ceil(points_num / gpu_count)
padding_num = batch_points_num * gpu_count - points_num
mask = torch.arange(batch_points_num * gpu_count) < points_num
split_points = F.pad(split_points, pad=(0, 0, 0, padding_num))
split_points = split_points.reshape(gpu_count, batch_points_num, -1)
images_feed = images_feed.expand(gpu_count, *images_feed.shape[1:])
fmaps_feed = fmaps_feed.expand(gpu_count, *fmaps_feed.shape[1:])

fine_pred_track, _, pred_vis, pred_score = track_predictor(
images_feed, split_points, fmaps=fmaps_feed, fine_tracking=fine_tracking, fine_chunk=fine_chunk
)

if parallel:
B, S, N, D = fine_pred_track.shape
fine_pred_track = fine_pred_track.permute(1, 0, 2, 3).reshape(1, S, -1, D)[:, :, mask, :]
pred_vis = pred_vis.permute(1, 0, 2).reshape(1, S, -1)[:, :, mask]

fine_pred_track_list.append(fine_pred_track)
pred_vis_list.append(pred_vis)
pred_score_list.append(pred_score)
Expand Down